[Disagg] GPU staging buffer with dynamic ring allocator for heterogeneous TP KV transfer (#19890)
This commit is contained in:
@@ -148,6 +148,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
# These timeout requests should be aborted to release the tree cache.
|
# These timeout requests should be aborted to release the tree cache.
|
||||||
self.bootstrap_timeout = envs.SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT.get()
|
self.bootstrap_timeout = envs.SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT.get()
|
||||||
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||||
|
self.enable_staging: bool = False
|
||||||
self.connection_pool: Dict[str, Dict[str, Union[str, int]]] = {}
|
self.connection_pool: Dict[str, Dict[str, Union[str, int]]] = {}
|
||||||
self.connection_lock = threading.Lock()
|
self.connection_lock = threading.Lock()
|
||||||
self.required_prefill_response_num_table: Dict[int, int] = {}
|
self.required_prefill_response_num_table: Dict[int, int] = {}
|
||||||
@@ -501,6 +502,7 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
self.bootstrap_addr = bootstrap_addr
|
self.bootstrap_addr = bootstrap_addr
|
||||||
self.kv_mgr = mgr
|
self.kv_mgr = mgr
|
||||||
self.conclude_state: Optional[KVPoll] = None
|
self.conclude_state: Optional[KVPoll] = None
|
||||||
|
self.require_staging: bool = False
|
||||||
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room)
|
self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room)
|
||||||
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping)
|
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping)
|
||||||
|
|
||||||
@@ -529,6 +531,12 @@ class CommonKVReceiver(BaseKVReceiver):
|
|||||||
self.required_prefill_response_num
|
self.required_prefill_response_num
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if self.kv_mgr.enable_staging:
|
||||||
|
self.require_staging = (
|
||||||
|
self.prefill_info.attn_tp_size != 0
|
||||||
|
and self.prefill_info.attn_tp_size != self.kv_mgr.attn_tp_size
|
||||||
|
)
|
||||||
|
|
||||||
self.prefill_dp_rank = prefill_dp_rank
|
self.prefill_dp_rank = prefill_dp_rank
|
||||||
self._setup_bootstrap_infos()
|
self._setup_bootstrap_infos()
|
||||||
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput)
|
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput)
|
||||||
|
|||||||
@@ -0,0 +1,768 @@
|
|||||||
|
"""
|
||||||
|
GPU Staging Buffer for heterogeneous TP KV cache transfer.
|
||||||
|
|
||||||
|
When prefill attn_tp_size != decode attn_tp_size, the per-token RDMA approach
|
||||||
|
generates O(tokens * layers) small RDMA requests. This module provides a staging
|
||||||
|
buffer mechanism that gathers scattered head slices into contiguous GPU memory,
|
||||||
|
enabling bulk RDMA transfers that reduce request count to O(layers) or O(1).
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
Activated by setting SGLANG_DISAGG_STAGING_BUFFER=1.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import threading
|
||||||
|
from typing import List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# TODO(yangminl): remove torch fallback implementations once the Triton kernels
|
||||||
|
# have been validated in production across all configurations.
|
||||||
|
_USE_TRITON_STAGING = not bool(os.environ.get("SGLANG_STAGING_USE_TORCH", ""))
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_gather_to_staging_kernel(
|
||||||
|
layer_ptrs,
|
||||||
|
page_indices,
|
||||||
|
staging,
|
||||||
|
num_tokens,
|
||||||
|
stride_pool_token,
|
||||||
|
head_offset,
|
||||||
|
per_layer_elems,
|
||||||
|
ELEMS_PER_TOKEN: tl.constexpr,
|
||||||
|
PAGE_SIZE: tl.constexpr,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
layer_id = tl.program_id(0)
|
||||||
|
block_id = tl.program_id(1)
|
||||||
|
|
||||||
|
layer_ptr = tl.load(layer_ptrs + layer_id).to(staging.dtype)
|
||||||
|
|
||||||
|
offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = offsets < per_layer_elems
|
||||||
|
|
||||||
|
t_idx = offsets // ELEMS_PER_TOKEN
|
||||||
|
e_idx = offsets % ELEMS_PER_TOKEN
|
||||||
|
|
||||||
|
page_id = t_idx // PAGE_SIZE
|
||||||
|
intra_page = t_idx % PAGE_SIZE
|
||||||
|
page_val = tl.load(page_indices + page_id, mask=mask, other=0)
|
||||||
|
pool_token = page_val * PAGE_SIZE + intra_page
|
||||||
|
|
||||||
|
src_offsets = (
|
||||||
|
pool_token * stride_pool_token.to(tl.int64) + head_offset.to(tl.int64) + e_idx
|
||||||
|
)
|
||||||
|
vals = tl.load(layer_ptr + src_offsets, mask=mask)
|
||||||
|
|
||||||
|
dst_offsets = tl.program_id(0).to(tl.int64) * per_layer_elems.to(tl.int64) + offsets
|
||||||
|
tl.store(staging + dst_offsets, vals, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def _fused_scatter_from_staging_kernel(
|
||||||
|
layer_ptrs,
|
||||||
|
page_indices,
|
||||||
|
staging,
|
||||||
|
writer_head_offsets,
|
||||||
|
num_tokens,
|
||||||
|
stride_pool_token,
|
||||||
|
per_layer_elems,
|
||||||
|
ELEMS_PER_TOKEN: tl.constexpr,
|
||||||
|
PAGE_SIZE: tl.constexpr,
|
||||||
|
NUM_LAYERS_X2: tl.constexpr,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
):
|
||||||
|
prog_id = tl.program_id(0)
|
||||||
|
block_id = tl.program_id(1)
|
||||||
|
|
||||||
|
writer_id = prog_id // NUM_LAYERS_X2
|
||||||
|
layer_kv_id = prog_id % NUM_LAYERS_X2
|
||||||
|
|
||||||
|
layer_ptr = tl.load(layer_ptrs + layer_kv_id).to(staging.dtype)
|
||||||
|
head_offset = tl.load(writer_head_offsets + writer_id)
|
||||||
|
|
||||||
|
offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = offsets < per_layer_elems
|
||||||
|
|
||||||
|
t_idx = offsets // ELEMS_PER_TOKEN
|
||||||
|
e_idx = offsets % ELEMS_PER_TOKEN
|
||||||
|
|
||||||
|
page_id = t_idx // PAGE_SIZE
|
||||||
|
intra_page = t_idx % PAGE_SIZE
|
||||||
|
page_val = tl.load(page_indices + page_id, mask=mask, other=0)
|
||||||
|
pool_token = page_val * PAGE_SIZE + intra_page
|
||||||
|
|
||||||
|
per_rank_elems = per_layer_elems.to(tl.int64) * NUM_LAYERS_X2
|
||||||
|
src_offsets = (
|
||||||
|
writer_id.to(tl.int64) * per_rank_elems
|
||||||
|
+ layer_kv_id.to(tl.int64) * per_layer_elems.to(tl.int64)
|
||||||
|
+ offsets
|
||||||
|
)
|
||||||
|
vals = tl.load(staging + src_offsets, mask=mask)
|
||||||
|
|
||||||
|
dst_offsets = (
|
||||||
|
pool_token * stride_pool_token.to(tl.int64) + head_offset.to(tl.int64) + e_idx
|
||||||
|
)
|
||||||
|
tl.store(layer_ptr + dst_offsets, vals, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
class StagingBuffer:
|
||||||
|
"""Pre-allocated GPU staging buffer for bulk KV transfer.
|
||||||
|
|
||||||
|
When a custom_mem_pool is provided (e.g., mooncake NVLink allocator),
|
||||||
|
the buffer is allocated within that pool so it's compatible with
|
||||||
|
NVLink/MNNVL transport (requires cuMemCreate-backed memory).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
size_bytes: int,
|
||||||
|
device: str,
|
||||||
|
gpu_id: int,
|
||||||
|
custom_mem_pool=None,
|
||||||
|
):
|
||||||
|
self.size_bytes = size_bytes
|
||||||
|
self.device = device
|
||||||
|
self.gpu_id = gpu_id
|
||||||
|
|
||||||
|
torch.cuda.set_device(gpu_id)
|
||||||
|
if custom_mem_pool is not None:
|
||||||
|
with torch.cuda.use_mem_pool(custom_mem_pool):
|
||||||
|
self.buffer = torch.empty(size_bytes, dtype=torch.uint8, device=device)
|
||||||
|
alloc_method = "custom_mem_pool (cuMemCreate)"
|
||||||
|
else:
|
||||||
|
self.buffer = torch.empty(size_bytes, dtype=torch.uint8, device=device)
|
||||||
|
alloc_method = "cudaMalloc (NVLink incompatible!)"
|
||||||
|
self.data_ptr = self.buffer.data_ptr()
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"StagingBuffer allocated: {size_bytes / (1024*1024):.1f} MB "
|
||||||
|
f"on {device}, method={alloc_method}, ptr=0x{self.data_ptr:x}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_ptr(self) -> int:
|
||||||
|
return self.data_ptr
|
||||||
|
|
||||||
|
def get_size(self) -> int:
|
||||||
|
return self.size_bytes
|
||||||
|
|
||||||
|
def fits(self, required_bytes: int) -> bool:
|
||||||
|
return required_bytes <= self.size_bytes
|
||||||
|
|
||||||
|
|
||||||
|
class StagingAllocator:
|
||||||
|
"""Decode-side dynamic staging ring buffer allocator with overcommit.
|
||||||
|
|
||||||
|
One large pre-allocated GPU buffer used as a ring buffer. Each request
|
||||||
|
gets a (alloc_id, offset, round) triple based on its actual byte
|
||||||
|
requirement. Allocation (assign) is overcommit — it always succeeds
|
||||||
|
as long as the request fits in the buffer. Overlap safety is enforced
|
||||||
|
on the prefill side before RDMA, using a watermark that tracks the
|
||||||
|
oldest un-freed allocation.
|
||||||
|
|
||||||
|
The watermark (round, tail_offset) is periodically sent to prefill.
|
||||||
|
Prefill transfer workers wait before writing if their target region
|
||||||
|
overlaps with not-yet-freed data from a previous round.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Permanent alloc failure: chunk exceeds ring buffer total size.
|
||||||
|
ALLOC_OVERSIZED = -2
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
total_size_bytes: int,
|
||||||
|
device: str,
|
||||||
|
gpu_id: int,
|
||||||
|
custom_mem_pool=None,
|
||||||
|
):
|
||||||
|
self.buffer = StagingBuffer(total_size_bytes, device, gpu_id, custom_mem_pool)
|
||||||
|
self.total_size = total_size_bytes
|
||||||
|
self.base_ptr = self.buffer.data_ptr
|
||||||
|
self.head = 0
|
||||||
|
self.round = 0
|
||||||
|
self.allocations: dict = {} # alloc_id -> (offset, size, round)
|
||||||
|
self.alloc_order: List[int] = []
|
||||||
|
self.next_alloc_id = 0
|
||||||
|
self.watermark_round = 0
|
||||||
|
self.watermark_tail = 0
|
||||||
|
self.lock = threading.Lock()
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"StagingAllocator (ring+overcommit): "
|
||||||
|
f"{total_size_bytes / (1024*1024):.1f} MB "
|
||||||
|
f"on {device}, ptr=0x{self.base_ptr:x}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def assign(self, required_bytes: int) -> Optional[Tuple[int, int, int]]:
|
||||||
|
"""Allocate a region. Returns (alloc_id, offset, round) or None."""
|
||||||
|
with self.lock:
|
||||||
|
if required_bytes > self.total_size:
|
||||||
|
return None
|
||||||
|
|
||||||
|
space_at_end = self.total_size - self.head
|
||||||
|
if required_bytes <= space_at_end:
|
||||||
|
offset = self.head
|
||||||
|
self.head += required_bytes
|
||||||
|
else:
|
||||||
|
self.round += 1
|
||||||
|
offset = 0
|
||||||
|
self.head = required_bytes
|
||||||
|
|
||||||
|
alloc_id = self.next_alloc_id
|
||||||
|
self.next_alloc_id += 1
|
||||||
|
self.allocations[alloc_id] = (offset, required_bytes, self.round)
|
||||||
|
self.alloc_order.append(alloc_id)
|
||||||
|
return (alloc_id, offset, self.round)
|
||||||
|
|
||||||
|
def free(self, alloc_id: int):
|
||||||
|
"""Free an allocation and advance watermark past consecutive freed entries."""
|
||||||
|
with self.lock:
|
||||||
|
if alloc_id not in self.allocations:
|
||||||
|
return
|
||||||
|
self.allocations.pop(alloc_id)
|
||||||
|
|
||||||
|
while self.alloc_order and self.alloc_order[0] not in self.allocations:
|
||||||
|
self.alloc_order.pop(0)
|
||||||
|
|
||||||
|
if not self.allocations:
|
||||||
|
self.watermark_round = self.round
|
||||||
|
self.watermark_tail = self.head
|
||||||
|
elif self.alloc_order:
|
||||||
|
off, _, rnd = self.allocations[self.alloc_order[0]]
|
||||||
|
self.watermark_round = rnd
|
||||||
|
self.watermark_tail = off
|
||||||
|
|
||||||
|
def get_watermark(self) -> Tuple[int, int]:
|
||||||
|
"""Return (round, tail_offset). Everything before this is safe to write."""
|
||||||
|
with self.lock:
|
||||||
|
return (self.watermark_round, self.watermark_tail)
|
||||||
|
|
||||||
|
def get_ptr(self, alloc_id: int) -> int:
|
||||||
|
offset, _, _ = self.allocations[alloc_id]
|
||||||
|
return self.base_ptr + offset
|
||||||
|
|
||||||
|
def get_offset(self, alloc_id: int) -> int:
|
||||||
|
offset, _, _ = self.allocations[alloc_id]
|
||||||
|
return offset
|
||||||
|
|
||||||
|
def get_round(self, alloc_id: int) -> int:
|
||||||
|
_, _, rnd = self.allocations[alloc_id]
|
||||||
|
return rnd
|
||||||
|
|
||||||
|
def get_base_ptr(self) -> int:
|
||||||
|
return self.base_ptr
|
||||||
|
|
||||||
|
def get_total_size(self) -> int:
|
||||||
|
return self.total_size
|
||||||
|
|
||||||
|
|
||||||
|
def gather_kv_head_slices(
|
||||||
|
kv_buffer_tensor: torch.Tensor,
|
||||||
|
gather_idx: torch.Tensor,
|
||||||
|
head_start: int,
|
||||||
|
num_heads: int,
|
||||||
|
staging_tensor: torch.Tensor,
|
||||||
|
):
|
||||||
|
"""Gather KV head slices from scattered pages into contiguous staging buffer.
|
||||||
|
|
||||||
|
Uses torch.gather(out=) to write directly into staging_tensor without
|
||||||
|
allocating temporary tensors (avoids CUDA caching allocator stalls).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
kv_buffer_tensor: [pool_size, head_num, head_dim], one layer.
|
||||||
|
gather_idx: [num_tokens, num_heads, head_dim] int64, pre-computed
|
||||||
|
token indices expanded for gather on dim=0.
|
||||||
|
head_start: Starting head index for the slice.
|
||||||
|
num_heads: Number of heads to gather.
|
||||||
|
staging_tensor: Output tensor, shape [num_tokens, num_heads, head_dim].
|
||||||
|
"""
|
||||||
|
src = kv_buffer_tensor[:, head_start : head_start + num_heads, :]
|
||||||
|
torch.gather(src, 0, gather_idx, out=staging_tensor)
|
||||||
|
|
||||||
|
|
||||||
|
def scatter_kv_head_slices(
|
||||||
|
staging_tensor: torch.Tensor,
|
||||||
|
kv_buffer_tensor: torch.Tensor,
|
||||||
|
page_indices: torch.Tensor,
|
||||||
|
head_start: int,
|
||||||
|
num_heads: int,
|
||||||
|
page_size: int = 1,
|
||||||
|
):
|
||||||
|
"""Scatter KV head slices from contiguous staging buffer to KV cache.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
staging_tensor: Input tensor from staging buffer (contiguous packed data).
|
||||||
|
kv_buffer_tensor: The KV buffer for one layer, shape [pool_size, head_num, head_dim].
|
||||||
|
page_indices: [num_pages] int32/int64 tensor of page indices.
|
||||||
|
head_start: Starting head index for the slice.
|
||||||
|
num_heads: Number of heads to scatter.
|
||||||
|
page_size: Number of tokens per page.
|
||||||
|
"""
|
||||||
|
head_dim = kv_buffer_tensor.shape[-1]
|
||||||
|
if page_size == 1:
|
||||||
|
num_tokens = page_indices.shape[0]
|
||||||
|
data = staging_tensor.reshape(num_tokens, num_heads, head_dim)
|
||||||
|
kv_buffer_tensor[page_indices, head_start : head_start + num_heads, :] = data
|
||||||
|
else:
|
||||||
|
num_tokens = page_indices.shape[0] * page_size
|
||||||
|
offsets = torch.arange(page_size, device=page_indices.device)
|
||||||
|
token_indices = (page_indices.unsqueeze(1) * page_size + offsets).reshape(-1)
|
||||||
|
data = staging_tensor.reshape(num_tokens, num_heads, head_dim)
|
||||||
|
kv_buffer_tensor[token_indices, head_start : head_start + num_heads, :] = data
|
||||||
|
|
||||||
|
|
||||||
|
def _gather_all_layers_torch(
|
||||||
|
k_buffers: list,
|
||||||
|
v_buffers: list,
|
||||||
|
page_indices_np,
|
||||||
|
staging_buffer: StagingBuffer,
|
||||||
|
src_head_start: int,
|
||||||
|
num_heads: int,
|
||||||
|
page_size: int,
|
||||||
|
gpu_id: int,
|
||||||
|
) -> int:
|
||||||
|
"""torch.gather path: zero per-layer allocation, one kernel per layer."""
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
num_layers = len(k_buffers)
|
||||||
|
head_dim = k_buffers[0].shape[-1]
|
||||||
|
dtype_size = k_buffers[0].element_size()
|
||||||
|
num_tokens = len(page_indices_np) * page_size
|
||||||
|
per_layer_bytes = num_tokens * num_heads * head_dim * dtype_size
|
||||||
|
|
||||||
|
device = f"cuda:{gpu_id}"
|
||||||
|
torch.cuda.set_device(gpu_id)
|
||||||
|
page_idx_tensor = torch.from_numpy(page_indices_np.astype(np.int64)).to(device)
|
||||||
|
|
||||||
|
if page_size == 1:
|
||||||
|
token_indices = page_idx_tensor
|
||||||
|
else:
|
||||||
|
offsets = torch.arange(page_size, device=device)
|
||||||
|
token_indices = (page_idx_tensor.unsqueeze(1) * page_size + offsets).reshape(-1)
|
||||||
|
|
||||||
|
gather_idx = token_indices.view(-1, 1, 1).expand(num_tokens, num_heads, head_dim)
|
||||||
|
|
||||||
|
if not hasattr(staging_buffer, "_gather_stream"):
|
||||||
|
staging_buffer._gather_stream = torch.cuda.Stream(device=device)
|
||||||
|
|
||||||
|
staging_buffer._gather_stream.wait_stream(
|
||||||
|
torch.cuda.default_stream(torch.device(device))
|
||||||
|
)
|
||||||
|
|
||||||
|
staging_view = staging_buffer.buffer
|
||||||
|
offset = 0
|
||||||
|
with torch.cuda.stream(staging_buffer._gather_stream):
|
||||||
|
for layer_id in range(num_layers):
|
||||||
|
dst = (
|
||||||
|
staging_view[offset : offset + per_layer_bytes]
|
||||||
|
.view(k_buffers[layer_id].dtype)
|
||||||
|
.reshape(num_tokens, num_heads, head_dim)
|
||||||
|
)
|
||||||
|
gather_kv_head_slices(
|
||||||
|
k_buffers[layer_id],
|
||||||
|
gather_idx,
|
||||||
|
src_head_start,
|
||||||
|
num_heads,
|
||||||
|
dst,
|
||||||
|
)
|
||||||
|
offset += per_layer_bytes
|
||||||
|
for layer_id in range(num_layers):
|
||||||
|
dst = (
|
||||||
|
staging_view[offset : offset + per_layer_bytes]
|
||||||
|
.view(v_buffers[layer_id].dtype)
|
||||||
|
.reshape(num_tokens, num_heads, head_dim)
|
||||||
|
)
|
||||||
|
gather_kv_head_slices(
|
||||||
|
v_buffers[layer_id],
|
||||||
|
gather_idx,
|
||||||
|
src_head_start,
|
||||||
|
num_heads,
|
||||||
|
dst,
|
||||||
|
)
|
||||||
|
offset += per_layer_bytes
|
||||||
|
|
||||||
|
staging_buffer._gather_stream.synchronize()
|
||||||
|
return offset
|
||||||
|
|
||||||
|
|
||||||
|
def _gather_all_layers_triton(
|
||||||
|
k_buffers: list,
|
||||||
|
v_buffers: list,
|
||||||
|
page_indices_np,
|
||||||
|
staging_buffer: StagingBuffer,
|
||||||
|
src_head_start: int,
|
||||||
|
num_heads: int,
|
||||||
|
page_size: int,
|
||||||
|
gpu_id: int,
|
||||||
|
) -> int:
|
||||||
|
"""Triton fused kernel path: single kernel launch for all layers."""
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
num_layers = len(k_buffers)
|
||||||
|
head_dim = k_buffers[0].shape[-1]
|
||||||
|
total_heads = k_buffers[0].shape[1]
|
||||||
|
dtype_size = k_buffers[0].element_size()
|
||||||
|
num_tokens = len(page_indices_np) * page_size
|
||||||
|
elems_per_token = num_heads * head_dim
|
||||||
|
per_layer_elems = num_tokens * elems_per_token
|
||||||
|
per_layer_bytes = per_layer_elems * dtype_size
|
||||||
|
total_bytes = per_layer_bytes * num_layers * 2
|
||||||
|
|
||||||
|
device = f"cuda:{gpu_id}"
|
||||||
|
torch.cuda.set_device(gpu_id)
|
||||||
|
page_idx_tensor = torch.from_numpy(page_indices_np.astype(np.int64)).to(device)
|
||||||
|
|
||||||
|
layer_ptrs = torch.tensor(
|
||||||
|
[buf.data_ptr() for buf in k_buffers] + [buf.data_ptr() for buf in v_buffers],
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
# Use integer dtype matching element size for bit-preserving copy
|
||||||
|
int_dtype_map = {1: torch.int8, 2: torch.int16, 4: torch.int32}
|
||||||
|
int_dtype = int_dtype_map.get(dtype_size, torch.int16)
|
||||||
|
staging_typed = staging_buffer.buffer[:total_bytes].view(int_dtype)
|
||||||
|
|
||||||
|
if not hasattr(staging_buffer, "_gather_stream"):
|
||||||
|
staging_buffer._gather_stream = torch.cuda.Stream(device=device)
|
||||||
|
|
||||||
|
staging_buffer._gather_stream.wait_stream(
|
||||||
|
torch.cuda.default_stream(torch.device(device))
|
||||||
|
)
|
||||||
|
|
||||||
|
BLOCK_SIZE = 1024
|
||||||
|
grid = (2 * num_layers, triton.cdiv(per_layer_elems, BLOCK_SIZE))
|
||||||
|
|
||||||
|
with torch.cuda.stream(staging_buffer._gather_stream):
|
||||||
|
_fused_gather_to_staging_kernel[grid](
|
||||||
|
layer_ptrs,
|
||||||
|
page_idx_tensor,
|
||||||
|
staging_typed,
|
||||||
|
num_tokens,
|
||||||
|
total_heads * head_dim,
|
||||||
|
src_head_start * head_dim,
|
||||||
|
per_layer_elems,
|
||||||
|
elems_per_token,
|
||||||
|
page_size,
|
||||||
|
BLOCK_SIZE,
|
||||||
|
)
|
||||||
|
|
||||||
|
staging_buffer._gather_stream.synchronize()
|
||||||
|
return total_bytes
|
||||||
|
|
||||||
|
|
||||||
|
def gather_all_layers_to_staging(
|
||||||
|
k_buffers: list,
|
||||||
|
v_buffers: list,
|
||||||
|
page_indices_np,
|
||||||
|
staging_buffer: StagingBuffer,
|
||||||
|
src_head_start: int,
|
||||||
|
num_heads: int,
|
||||||
|
page_size: int,
|
||||||
|
gpu_id: int,
|
||||||
|
) -> int:
|
||||||
|
"""Gather all layers' K and V head slices into a staging buffer.
|
||||||
|
|
||||||
|
Returns total bytes written.
|
||||||
|
Dispatches to Triton fused kernel when available, falls back to torch.gather.
|
||||||
|
"""
|
||||||
|
if _USE_TRITON_STAGING:
|
||||||
|
return _gather_all_layers_triton(
|
||||||
|
k_buffers,
|
||||||
|
v_buffers,
|
||||||
|
page_indices_np,
|
||||||
|
staging_buffer,
|
||||||
|
src_head_start,
|
||||||
|
num_heads,
|
||||||
|
page_size,
|
||||||
|
gpu_id,
|
||||||
|
)
|
||||||
|
return _gather_all_layers_torch(
|
||||||
|
k_buffers,
|
||||||
|
v_buffers,
|
||||||
|
page_indices_np,
|
||||||
|
staging_buffer,
|
||||||
|
src_head_start,
|
||||||
|
num_heads,
|
||||||
|
page_size,
|
||||||
|
gpu_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _scatter_staging_to_kv_torch(
|
||||||
|
staging_buffer_view: torch.Tensor,
|
||||||
|
k_buffers: list,
|
||||||
|
v_buffers: list,
|
||||||
|
page_idx_tensor: torch.Tensor,
|
||||||
|
page_size: int,
|
||||||
|
prefill_attn_tp_size: int,
|
||||||
|
decode_attn_tp_size: int,
|
||||||
|
dst_tp_rank: int,
|
||||||
|
total_kv_heads: int,
|
||||||
|
) -> None:
|
||||||
|
"""torch path for scatter."""
|
||||||
|
num_layers = len(k_buffers)
|
||||||
|
head_dim = k_buffers[0].shape[-1]
|
||||||
|
dtype_size = k_buffers[0].element_size()
|
||||||
|
num_tokens = page_idx_tensor.shape[0] * page_size
|
||||||
|
|
||||||
|
if prefill_attn_tp_size > decode_attn_tp_size:
|
||||||
|
num_writers = prefill_attn_tp_size // max(1, decode_attn_tp_size)
|
||||||
|
else:
|
||||||
|
num_writers = 1
|
||||||
|
|
||||||
|
for writer_rank in range(num_writers):
|
||||||
|
_, num_heads, dst_head_start, _ = compute_head_slice_params(
|
||||||
|
prefill_attn_tp_size,
|
||||||
|
decode_attn_tp_size,
|
||||||
|
writer_rank,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
)
|
||||||
|
per_layer_bytes = num_tokens * num_heads * head_dim * dtype_size
|
||||||
|
per_rank_bytes = per_layer_bytes * num_layers * 2
|
||||||
|
rank_base = writer_rank * per_rank_bytes
|
||||||
|
|
||||||
|
offset = rank_base
|
||||||
|
for layer_id in range(num_layers):
|
||||||
|
layer_data = (
|
||||||
|
staging_buffer_view[offset : offset + per_layer_bytes]
|
||||||
|
.view(k_buffers[layer_id].dtype)
|
||||||
|
.reshape(num_tokens, num_heads, head_dim)
|
||||||
|
)
|
||||||
|
scatter_kv_head_slices(
|
||||||
|
layer_data,
|
||||||
|
k_buffers[layer_id],
|
||||||
|
page_idx_tensor,
|
||||||
|
dst_head_start,
|
||||||
|
num_heads,
|
||||||
|
page_size,
|
||||||
|
)
|
||||||
|
offset += per_layer_bytes
|
||||||
|
for layer_id in range(num_layers):
|
||||||
|
layer_data = (
|
||||||
|
staging_buffer_view[offset : offset + per_layer_bytes]
|
||||||
|
.view(v_buffers[layer_id].dtype)
|
||||||
|
.reshape(num_tokens, num_heads, head_dim)
|
||||||
|
)
|
||||||
|
scatter_kv_head_slices(
|
||||||
|
layer_data,
|
||||||
|
v_buffers[layer_id],
|
||||||
|
page_idx_tensor,
|
||||||
|
dst_head_start,
|
||||||
|
num_heads,
|
||||||
|
page_size,
|
||||||
|
)
|
||||||
|
offset += per_layer_bytes
|
||||||
|
|
||||||
|
|
||||||
|
def _scatter_staging_to_kv_triton(
|
||||||
|
staging_buffer_view: torch.Tensor,
|
||||||
|
k_buffers: list,
|
||||||
|
v_buffers: list,
|
||||||
|
page_idx_tensor: torch.Tensor,
|
||||||
|
page_size: int,
|
||||||
|
prefill_attn_tp_size: int,
|
||||||
|
decode_attn_tp_size: int,
|
||||||
|
dst_tp_rank: int,
|
||||||
|
total_kv_heads: int,
|
||||||
|
) -> None:
|
||||||
|
"""Triton fused kernel path for scatter."""
|
||||||
|
num_layers = len(k_buffers)
|
||||||
|
head_dim = k_buffers[0].shape[-1]
|
||||||
|
total_heads = k_buffers[0].shape[1]
|
||||||
|
dtype_size = k_buffers[0].element_size()
|
||||||
|
num_tokens = page_idx_tensor.shape[0] * page_size
|
||||||
|
device = page_idx_tensor.device
|
||||||
|
|
||||||
|
if prefill_attn_tp_size > decode_attn_tp_size:
|
||||||
|
num_writers = prefill_attn_tp_size // max(1, decode_attn_tp_size)
|
||||||
|
else:
|
||||||
|
num_writers = 1
|
||||||
|
|
||||||
|
# All writers share the same num_heads; only dst_head_start differs
|
||||||
|
_, num_heads, _, _ = compute_head_slice_params(
|
||||||
|
prefill_attn_tp_size,
|
||||||
|
decode_attn_tp_size,
|
||||||
|
0,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
)
|
||||||
|
elems_per_token = num_heads * head_dim
|
||||||
|
per_layer_elems = num_tokens * elems_per_token
|
||||||
|
|
||||||
|
layer_ptrs = torch.tensor(
|
||||||
|
[buf.data_ptr() for buf in k_buffers] + [buf.data_ptr() for buf in v_buffers],
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
writer_head_offsets = torch.tensor(
|
||||||
|
[
|
||||||
|
compute_head_slice_params(
|
||||||
|
prefill_attn_tp_size,
|
||||||
|
decode_attn_tp_size,
|
||||||
|
wr,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
)[2]
|
||||||
|
* head_dim
|
||||||
|
for wr in range(num_writers)
|
||||||
|
],
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
int_dtype_map = {1: torch.int8, 2: torch.int16, 4: torch.int32}
|
||||||
|
int_dtype = int_dtype_map.get(dtype_size, torch.int16)
|
||||||
|
total_staging_bytes = (
|
||||||
|
num_tokens * elems_per_token * dtype_size * num_layers * 2 * num_writers
|
||||||
|
)
|
||||||
|
staging_typed = staging_buffer_view[:total_staging_bytes].view(int_dtype)
|
||||||
|
|
||||||
|
BLOCK_SIZE = 1024
|
||||||
|
num_layers_x2 = 2 * num_layers
|
||||||
|
grid = (num_writers * num_layers_x2, triton.cdiv(per_layer_elems, BLOCK_SIZE))
|
||||||
|
|
||||||
|
_fused_scatter_from_staging_kernel[grid](
|
||||||
|
layer_ptrs,
|
||||||
|
page_idx_tensor,
|
||||||
|
staging_typed,
|
||||||
|
writer_head_offsets,
|
||||||
|
num_tokens,
|
||||||
|
total_heads * head_dim,
|
||||||
|
per_layer_elems,
|
||||||
|
elems_per_token,
|
||||||
|
page_size,
|
||||||
|
num_layers_x2,
|
||||||
|
BLOCK_SIZE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def scatter_staging_to_kv(
|
||||||
|
staging_buffer_view: torch.Tensor,
|
||||||
|
k_buffers: list,
|
||||||
|
v_buffers: list,
|
||||||
|
page_idx_tensor: torch.Tensor,
|
||||||
|
page_size: int,
|
||||||
|
prefill_attn_tp_size: int,
|
||||||
|
decode_attn_tp_size: int,
|
||||||
|
dst_tp_rank: int,
|
||||||
|
total_kv_heads: int,
|
||||||
|
) -> None:
|
||||||
|
"""Scatter data from a contiguous staging region into KV cache buffers."""
|
||||||
|
if _USE_TRITON_STAGING:
|
||||||
|
return _scatter_staging_to_kv_triton(
|
||||||
|
staging_buffer_view,
|
||||||
|
k_buffers,
|
||||||
|
v_buffers,
|
||||||
|
page_idx_tensor,
|
||||||
|
page_size,
|
||||||
|
prefill_attn_tp_size,
|
||||||
|
decode_attn_tp_size,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
)
|
||||||
|
return _scatter_staging_to_kv_torch(
|
||||||
|
staging_buffer_view,
|
||||||
|
k_buffers,
|
||||||
|
v_buffers,
|
||||||
|
page_idx_tensor,
|
||||||
|
page_size,
|
||||||
|
prefill_attn_tp_size,
|
||||||
|
decode_attn_tp_size,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def compute_head_slice_params(
|
||||||
|
src_attn_tp_size: int,
|
||||||
|
dst_attn_tp_size: int,
|
||||||
|
src_tp_rank: int,
|
||||||
|
dst_tp_rank: int,
|
||||||
|
total_kv_heads: int,
|
||||||
|
) -> Tuple[int, int, int, int]:
|
||||||
|
"""Compute head slicing parameters for heterogeneous TP transfer.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(src_head_start, num_heads_to_send, dst_head_start, num_heads_to_send)
|
||||||
|
"""
|
||||||
|
src_heads_per_rank = max(1, total_kv_heads // src_attn_tp_size)
|
||||||
|
dst_heads_per_rank = max(1, total_kv_heads // dst_attn_tp_size)
|
||||||
|
|
||||||
|
local_tp_rank = src_tp_rank % src_attn_tp_size
|
||||||
|
dst_tp_rank_in_group = dst_tp_rank % dst_attn_tp_size
|
||||||
|
|
||||||
|
if src_attn_tp_size > dst_attn_tp_size:
|
||||||
|
src_head_start = 0
|
||||||
|
num_heads_to_send = src_heads_per_rank
|
||||||
|
src_replication = max(1, src_attn_tp_size // total_kv_heads)
|
||||||
|
unique_head_idx = local_tp_rank // src_replication
|
||||||
|
dst_head_start = (unique_head_idx * src_heads_per_rank) % dst_heads_per_rank
|
||||||
|
else:
|
||||||
|
src_head_start = (
|
||||||
|
dst_tp_rank_in_group * dst_heads_per_rank
|
||||||
|
) % src_heads_per_rank
|
||||||
|
num_heads_to_send = dst_heads_per_rank
|
||||||
|
dst_head_start = 0
|
||||||
|
|
||||||
|
return src_head_start, num_heads_to_send, dst_head_start, num_heads_to_send
|
||||||
|
|
||||||
|
|
||||||
|
def compute_staging_layout(
|
||||||
|
src_attn_tp_size: int,
|
||||||
|
dst_attn_tp_size: int,
|
||||||
|
dst_tp_rank: int,
|
||||||
|
total_kv_heads: int,
|
||||||
|
num_tokens: int,
|
||||||
|
bytes_per_head_token: int,
|
||||||
|
num_layers: int,
|
||||||
|
) -> Tuple[int, List[int], int]:
|
||||||
|
"""Compute per-writer byte layout for a staging region.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(num_writers, writer_bytes_list, total_bytes)
|
||||||
|
where writer_bytes_list[i] = bytes for writer i covering all layers (K+V).
|
||||||
|
"""
|
||||||
|
if src_attn_tp_size > dst_attn_tp_size:
|
||||||
|
num_writers = src_attn_tp_size // max(1, dst_attn_tp_size)
|
||||||
|
else:
|
||||||
|
num_writers = 1
|
||||||
|
|
||||||
|
writer_bytes = []
|
||||||
|
for wr in range(num_writers):
|
||||||
|
_, nh, _, _ = compute_head_slice_params(
|
||||||
|
src_attn_tp_size,
|
||||||
|
dst_attn_tp_size,
|
||||||
|
wr,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
)
|
||||||
|
writer_bytes.append(num_tokens * nh * bytes_per_head_token * num_layers * 2)
|
||||||
|
return num_writers, writer_bytes, sum(writer_bytes)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_total_kv_heads(
|
||||||
|
kv_args,
|
||||||
|
attn_tp_size: int,
|
||||||
|
) -> int:
|
||||||
|
"""Resolve the global total KV head count from kv_args metadata."""
|
||||||
|
total = getattr(kv_args, "total_kv_head_num", 0)
|
||||||
|
if total > 0:
|
||||||
|
return total
|
||||||
|
per_rank = getattr(kv_args, "kv_head_num", 0)
|
||||||
|
if per_rank > 0:
|
||||||
|
return per_rank * attn_tp_size
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot resolve total_kv_heads: kv_args has neither total_kv_head_num "
|
||||||
|
"nor kv_head_num. "
|
||||||
|
"Ensure DecodePreallocQueue._init_kv_manager sets kv_args.kv_head_num."
|
||||||
|
)
|
||||||
@@ -0,0 +1,732 @@
|
|||||||
|
"""
|
||||||
|
Staging handler for heterogeneous TP KV cache transfer.
|
||||||
|
|
||||||
|
Isolates staging scatter lifecycle from decode.py and conn.py.
|
||||||
|
Generic (backend-agnostic) code is at the top; mooncake-specific
|
||||||
|
protocol code is at the bottom.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
|
import logging
|
||||||
|
import struct
|
||||||
|
import threading
|
||||||
|
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.disaggregation.decode import DecodeRequest
|
||||||
|
|
||||||
|
|
||||||
|
# ======================================================================
|
||||||
|
# Generic staging state and handler (backend-agnostic)
|
||||||
|
# ======================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class DecodeStagingContext:
|
||||||
|
"""Staging-specific context for decode mode."""
|
||||||
|
|
||||||
|
allocator: object = None
|
||||||
|
room_bootstrap: dict = dataclasses.field(default_factory=dict)
|
||||||
|
room_receivers: dict = dataclasses.field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class PrefillStagingContext:
|
||||||
|
"""Staging-specific context for prefill mode."""
|
||||||
|
|
||||||
|
buffers: list = dataclasses.field(default_factory=list)
|
||||||
|
remote_watermarks: dict = dataclasses.field(default_factory=dict)
|
||||||
|
watermark_cv: threading.Condition = dataclasses.field(
|
||||||
|
default_factory=threading.Condition
|
||||||
|
)
|
||||||
|
prefetch_requested: set = dataclasses.field(default_factory=set)
|
||||||
|
prefetch_sockets: dict = dataclasses.field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
class DecodeStagingHandler:
|
||||||
|
"""Decode-side staging scatter lifecycle manager.
|
||||||
|
|
||||||
|
Scatter submission can be called from the decode_thread (background) as
|
||||||
|
soon as all writers/ranks have arrived, while event checking and freeing
|
||||||
|
always run on the scheduler main thread.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
kv_manager,
|
||||||
|
staging_allocator,
|
||||||
|
kv_buffer_info: dict,
|
||||||
|
decode_tp: int,
|
||||||
|
total_kv_heads: int,
|
||||||
|
tp_rank: int,
|
||||||
|
scheduler,
|
||||||
|
):
|
||||||
|
self.kv_manager = kv_manager
|
||||||
|
self.staging_allocator = staging_allocator
|
||||||
|
self.kv_buffer_info = kv_buffer_info
|
||||||
|
self.decode_tp = decode_tp
|
||||||
|
self.total_kv_heads = total_kv_heads
|
||||||
|
self.tp_rank = tp_rank
|
||||||
|
self.scheduler = scheduler
|
||||||
|
self._room_to_decode_req: dict = {}
|
||||||
|
self._wm_subscribers: dict = {}
|
||||||
|
|
||||||
|
def register_wm_subscriber(self, receiver, session_id: str) -> None:
|
||||||
|
"""Register a prefill's bootstrap connection for watermark broadcasts."""
|
||||||
|
if receiver is None or not getattr(receiver, "bootstrap_infos", None):
|
||||||
|
return
|
||||||
|
key = tuple(str(bi) for bi in receiver.bootstrap_infos)
|
||||||
|
if key not in self._wm_subscribers:
|
||||||
|
self._wm_subscribers[key] = (receiver, session_id)
|
||||||
|
|
||||||
|
def num_writers_for(self, decode_req) -> int:
|
||||||
|
"""Compute num_writers for a specific request based on its prefill TP."""
|
||||||
|
prefill_tp = decode_req.kv_receiver.prefill_info.attn_tp_size
|
||||||
|
if prefill_tp > self.decode_tp:
|
||||||
|
return prefill_tp // max(1, self.decode_tp)
|
||||||
|
return 1
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, kv_manager, scheduler, tp_rank: int) -> "DecodeStagingHandler":
|
||||||
|
"""Factory: create handler. Raises if staging infra is missing."""
|
||||||
|
staging_allocator = kv_manager._staging_ctx.allocator
|
||||||
|
if staging_allocator is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Staging is enabled but kv_manager._staging_ctx.allocator is None. "
|
||||||
|
"Check that the transfer backend correctly initializes the staging allocator."
|
||||||
|
)
|
||||||
|
kv_buffer_info = kv_manager.kv_buffer_tensors
|
||||||
|
if kv_buffer_info is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Staging is enabled but kv_manager.kv_buffer_tensors is None. "
|
||||||
|
"Check that set_kv_buffer_tensors() was called during kv_manager init."
|
||||||
|
)
|
||||||
|
decode_tp = kv_manager.attn_tp_size
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import (
|
||||||
|
resolve_total_kv_heads,
|
||||||
|
)
|
||||||
|
|
||||||
|
total_kv_heads = resolve_total_kv_heads(kv_manager.kv_args, decode_tp)
|
||||||
|
return cls(
|
||||||
|
kv_manager=kv_manager,
|
||||||
|
staging_allocator=staging_allocator,
|
||||||
|
kv_buffer_info=kv_buffer_info,
|
||||||
|
decode_tp=decode_tp,
|
||||||
|
total_kv_heads=total_kv_heads,
|
||||||
|
tp_rank=tp_rank,
|
||||||
|
scheduler=scheduler,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Registration: called from main thread (DecodeTransferQueue)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def register_decode_req(self, room: int, decode_req: "DecodeRequest") -> None:
|
||||||
|
self._room_to_decode_req[room] = decode_req
|
||||||
|
|
||||||
|
def unregister_decode_req(self, room: int) -> None:
|
||||||
|
self._room_to_decode_req.pop(room, None)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Scatter submission: called from decode_thread (background)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def submit_chunk_scatter(
|
||||||
|
self, room: int, chunk_idx: int, page_start: int, num_pages: int
|
||||||
|
) -> bool:
|
||||||
|
"""Submit scatter for an intermediate chunk whose writers all arrived.
|
||||||
|
|
||||||
|
Called from decode_thread. Records a CUDA event on decode_req so
|
||||||
|
the main thread can later check completion and free the allocation.
|
||||||
|
"""
|
||||||
|
decode_req = self._room_to_decode_req.get(room)
|
||||||
|
if decode_req is None:
|
||||||
|
logger.warning(
|
||||||
|
"[STAGING] submit_chunk_scatter: room=%s not registered, "
|
||||||
|
"chunk_idx=%s. This should not happen if register_decode_req "
|
||||||
|
"is called at kv_receiver.init() time.",
|
||||||
|
room,
|
||||||
|
chunk_idx,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
chunk_infos = getattr(decode_req.kv_receiver, "chunk_staging_infos", [])
|
||||||
|
if chunk_idx >= len(chunk_infos):
|
||||||
|
return False
|
||||||
|
alloc_id, staging_offset, _, _, _ = chunk_infos[chunk_idx]
|
||||||
|
if staging_offset < 0 or alloc_id < 0:
|
||||||
|
return False
|
||||||
|
|
||||||
|
ok = self._scatter_region(staging_offset, page_start, num_pages, decode_req)
|
||||||
|
if ok:
|
||||||
|
event = torch.cuda.Event()
|
||||||
|
event.record(self.staging_allocator._scatter_stream)
|
||||||
|
if not hasattr(decode_req, "_chunk_events"):
|
||||||
|
decode_req._chunk_events = []
|
||||||
|
decode_req._chunk_events.append((event, alloc_id))
|
||||||
|
chunk_infos[chunk_idx] = (-1, -1, 0, -1, 0)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"submit_chunk_scatter failed room=%s chunk_idx=%s tp_rank=%s",
|
||||||
|
room,
|
||||||
|
chunk_idx,
|
||||||
|
self.tp_rank,
|
||||||
|
)
|
||||||
|
return ok
|
||||||
|
|
||||||
|
def is_staging_room(self, room: int) -> bool:
|
||||||
|
"""Check if a room is registered for staging scatter."""
|
||||||
|
return room in self._room_to_decode_req
|
||||||
|
|
||||||
|
def submit_last_scatter_async(self, room: int) -> bool:
|
||||||
|
"""Submit scatter for the last chunk when all ranks report Success.
|
||||||
|
|
||||||
|
Called from decode_thread. Sets ``_scatter_event`` **before**
|
||||||
|
``_staging_last_scatter_submitted`` so the main thread sees the
|
||||||
|
event when it checks the flag (CPython GIL guarantees ordering).
|
||||||
|
"""
|
||||||
|
decode_req = self._room_to_decode_req.get(room)
|
||||||
|
if decode_req is None:
|
||||||
|
logger.warning(
|
||||||
|
"[STAGING] submit_last_scatter_async: room=%s not registered. "
|
||||||
|
"This should not happen if register_decode_req is called at "
|
||||||
|
"kv_receiver.init() time.",
|
||||||
|
room,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
alloc_id = self._submit_last_scatter(decode_req)
|
||||||
|
if alloc_id >= 0:
|
||||||
|
event = torch.cuda.Event()
|
||||||
|
event.record(self.staging_allocator._scatter_stream)
|
||||||
|
decode_req._scatter_event = event
|
||||||
|
decode_req._scatter_alloc_id = alloc_id
|
||||||
|
decode_req._staging_last_scatter_submitted = True
|
||||||
|
else:
|
||||||
|
decode_req._staging_scatter_done = True
|
||||||
|
return True
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Event check + free: called from main thread (pop_transferred)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def is_done(self, decode_req: "DecodeRequest") -> bool:
|
||||||
|
"""Return True if staging scatter is complete for this request."""
|
||||||
|
if not getattr(decode_req, "_staging_scatter_done", False):
|
||||||
|
return False
|
||||||
|
return not getattr(decode_req, "_chunk_events", None)
|
||||||
|
|
||||||
|
def advance_scatter(self, decode_req: "DecodeRequest") -> None:
|
||||||
|
"""Check CUDA events and free completed staging allocations.
|
||||||
|
|
||||||
|
Scatter kernels have already been submitted by the decode_thread
|
||||||
|
(via submit_chunk_scatter / submit_last_scatter_async). This
|
||||||
|
method only polls the recorded events and releases staging memory.
|
||||||
|
"""
|
||||||
|
room = decode_req.req.bootstrap_room
|
||||||
|
chunk_events = getattr(decode_req, "_chunk_events", None)
|
||||||
|
if chunk_events:
|
||||||
|
for i in range(len(chunk_events) - 1, -1, -1):
|
||||||
|
event, alloc_id = chunk_events[i]
|
||||||
|
if event.query():
|
||||||
|
chunk_events.pop(i)
|
||||||
|
self._free_and_send_watermark(alloc_id, decode_req)
|
||||||
|
|
||||||
|
if not getattr(decode_req, "_staging_last_scatter_submitted", False):
|
||||||
|
return
|
||||||
|
|
||||||
|
event = getattr(decode_req, "_scatter_event", None)
|
||||||
|
if event is not None and event.query():
|
||||||
|
self._free_and_send_watermark(decode_req._scatter_alloc_id, decode_req)
|
||||||
|
decode_req._scatter_event = None
|
||||||
|
decode_req._scatter_alloc_id = -1
|
||||||
|
decode_req._staging_scatter_done = True
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Internal methods
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _scatter_region(
|
||||||
|
self,
|
||||||
|
staging_offset: int,
|
||||||
|
page_start: int,
|
||||||
|
num_pages: int,
|
||||||
|
decode_req: "DecodeRequest",
|
||||||
|
) -> bool:
|
||||||
|
"""Submit scatter kernels for a staging region to scatter_stream.
|
||||||
|
|
||||||
|
May be called from the decode_thread (background). All GPU work
|
||||||
|
runs on scatter_stream so that the decode_thread never blocks on
|
||||||
|
the default stream (which carries the main-thread forward pass).
|
||||||
|
"""
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import (
|
||||||
|
scatter_staging_to_kv,
|
||||||
|
)
|
||||||
|
|
||||||
|
k_buffers = self.kv_buffer_info["k_buffers"]
|
||||||
|
v_buffers = self.kv_buffer_info["v_buffers"]
|
||||||
|
page_size = self.kv_buffer_info["page_size"]
|
||||||
|
dst_tp_rank = self.kv_manager.kv_args.engine_rank % self.decode_tp
|
||||||
|
|
||||||
|
device = k_buffers[0].device
|
||||||
|
torch.cuda.set_device(device)
|
||||||
|
|
||||||
|
if not hasattr(self.staging_allocator, "_scatter_stream"):
|
||||||
|
self.staging_allocator._scatter_stream = torch.cuda.Stream(device=device)
|
||||||
|
|
||||||
|
scatter_stream = self.staging_allocator._scatter_stream
|
||||||
|
|
||||||
|
staging_view = self.staging_allocator.buffer.buffer[staging_offset:]
|
||||||
|
|
||||||
|
req_pool_idx = decode_req.req.req_pool_idx
|
||||||
|
token_start = page_start * page_size
|
||||||
|
token_end = token_start + num_pages * page_size
|
||||||
|
prefill_tp = decode_req.kv_receiver.prefill_info.attn_tp_size
|
||||||
|
|
||||||
|
with torch.cuda.stream(scatter_stream):
|
||||||
|
kv_indices = self.scheduler.req_to_token_pool.req_to_token[
|
||||||
|
req_pool_idx, token_start:token_end
|
||||||
|
]
|
||||||
|
if page_size > 1:
|
||||||
|
page_idx_tensor = kv_indices[::page_size] // page_size
|
||||||
|
else:
|
||||||
|
page_idx_tensor = kv_indices
|
||||||
|
|
||||||
|
scatter_staging_to_kv(
|
||||||
|
staging_view,
|
||||||
|
k_buffers,
|
||||||
|
v_buffers,
|
||||||
|
page_idx_tensor,
|
||||||
|
page_size,
|
||||||
|
prefill_tp,
|
||||||
|
self.decode_tp,
|
||||||
|
dst_tp_rank,
|
||||||
|
self.total_kv_heads,
|
||||||
|
)
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _submit_last_scatter(self, decode_req: "DecodeRequest") -> int:
|
||||||
|
"""Submit scatter for the last chunk. Returns alloc_id >= 0, or -1."""
|
||||||
|
receiver = decode_req.kv_receiver
|
||||||
|
chunk_infos = getattr(receiver, "chunk_staging_infos", [])
|
||||||
|
if not chunk_infos:
|
||||||
|
return -1
|
||||||
|
|
||||||
|
last_info = chunk_infos[-1]
|
||||||
|
alloc_id, staging_offset, _, _, last_num_pages = last_info
|
||||||
|
if staging_offset < 0 or alloc_id < 0:
|
||||||
|
return -1
|
||||||
|
|
||||||
|
seq_len = len(decode_req.req.origin_input_ids)
|
||||||
|
ps = self.scheduler.token_to_kv_pool_allocator.page_size
|
||||||
|
total_pages = (seq_len + ps - 1) // ps
|
||||||
|
page_start = total_pages - last_num_pages
|
||||||
|
|
||||||
|
ok = self._scatter_region(
|
||||||
|
staging_offset, page_start, last_num_pages, decode_req
|
||||||
|
)
|
||||||
|
return alloc_id if ok else -1
|
||||||
|
|
||||||
|
def _free_and_send_watermark(
|
||||||
|
self, alloc_id: int, decode_req: "DecodeRequest"
|
||||||
|
) -> None:
|
||||||
|
"""Free a staging allocation and broadcast watermark to all prefills."""
|
||||||
|
self.staging_allocator.free(alloc_id)
|
||||||
|
post_wm = self.staging_allocator.get_watermark()
|
||||||
|
room = decode_req.req.bootstrap_room
|
||||||
|
wm_round, wm_tail = post_wm
|
||||||
|
wm_round_b = str(wm_round).encode("ascii")
|
||||||
|
wm_tail_b = str(wm_tail).encode("ascii")
|
||||||
|
for _key, (receiver, session_id) in list(self._wm_subscribers.items()):
|
||||||
|
sid_b = session_id.encode("ascii")
|
||||||
|
for bootstrap_info in receiver.bootstrap_infos:
|
||||||
|
try:
|
||||||
|
sock, lock = receiver._connect_to_bootstrap_server(bootstrap_info)
|
||||||
|
with lock:
|
||||||
|
sock.send_multipart(
|
||||||
|
[b"WATERMARK", wm_round_b, wm_tail_b, sid_b]
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def is_watermark_ready(
|
||||||
|
staging_state, session_id: str, alloc_round: int, alloc_end: int
|
||||||
|
) -> bool:
|
||||||
|
"""Non-blocking check: is the staging region safe to write?"""
|
||||||
|
if alloc_round <= 0:
|
||||||
|
return True
|
||||||
|
prev_round = alloc_round - 1
|
||||||
|
wm_round, wm_tail = staging_state.remote_watermarks.get(session_id, (0, 0))
|
||||||
|
return prev_round < wm_round or (prev_round == wm_round and alloc_end <= wm_tail)
|
||||||
|
|
||||||
|
|
||||||
|
# ======================================================================
|
||||||
|
# Mooncake-specific staging protocol and utilities
|
||||||
|
# ======================================================================
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class StagingTransferInfo:
|
||||||
|
"""Per-chunk staging allocation info attached to a TransferInfo."""
|
||||||
|
|
||||||
|
offsets: List[int] = dataclasses.field(default_factory=lambda: [-1])
|
||||||
|
rounds: List[int] = dataclasses.field(default_factory=lambda: [0])
|
||||||
|
ends: List[int] = dataclasses.field(default_factory=lambda: [-1])
|
||||||
|
|
||||||
|
def set_chunk(self, idx: int, offset: int, rnd: int, end: int):
|
||||||
|
while len(self.offsets) <= idx:
|
||||||
|
self.offsets.append(-1)
|
||||||
|
self.rounds.append(0)
|
||||||
|
self.ends.append(-1)
|
||||||
|
self.offsets[idx] = offset
|
||||||
|
self.rounds[idx] = rnd
|
||||||
|
self.ends[idx] = end
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class StagingRegisterInfo:
|
||||||
|
"""Staging buffer registration info attached to a KVArgsRegisterInfo."""
|
||||||
|
|
||||||
|
base_ptr: int = 0
|
||||||
|
total_size: int = 0
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_zmq_fields(
|
||||||
|
cls, msg: list, msg_start_offset: int
|
||||||
|
) -> Optional["StagingRegisterInfo"]:
|
||||||
|
i = msg_start_offset
|
||||||
|
base_ptr = (
|
||||||
|
struct.unpack("Q", msg[i])[0] if len(msg) > i and len(msg[i]) == 8 else 0
|
||||||
|
)
|
||||||
|
total_size = (
|
||||||
|
int(msg[i + 1].decode("ascii"))
|
||||||
|
if len(msg) > i + 1 and len(msg[i + 1]) > 0
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
if base_ptr == 0 and total_size == 0:
|
||||||
|
return None
|
||||||
|
return cls(base_ptr=base_ptr, total_size=total_size)
|
||||||
|
|
||||||
|
|
||||||
|
class PrefillStagingStrategy:
|
||||||
|
"""Prefill-side staging transfer: readiness check + gather-RDMA execution.
|
||||||
|
|
||||||
|
Encapsulates the decision logic (chunk index calculation, staging offset
|
||||||
|
lookup, watermark readiness) and delegates actual RDMA to the kv_manager.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, kv_manager, staging_buffer):
|
||||||
|
self.kv_manager = kv_manager
|
||||||
|
self.staging_buffer = staging_buffer
|
||||||
|
page_size = kv_manager.kv_buffer_tensors["page_size"]
|
||||||
|
cps = kv_manager.server_args.chunked_prefill_size or 8192
|
||||||
|
self.full_chunk_pages = max(1, cps // page_size)
|
||||||
|
|
||||||
|
def check_ready(
|
||||||
|
self,
|
||||||
|
req,
|
||||||
|
kv_chunk_index_start: int,
|
||||||
|
num_chunk_pages: int,
|
||||||
|
) -> Tuple[bool, int, int, int, int]:
|
||||||
|
"""Check if staging offset and watermark are ready for this chunk.
|
||||||
|
|
||||||
|
Returns (ready, chunk_idx, offset, round, end).
|
||||||
|
offset == ALLOC_OVERSIZED means permanent failure (fall back to slice).
|
||||||
|
offset == -1 means allocation pending (re-enqueue).
|
||||||
|
"""
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator
|
||||||
|
|
||||||
|
chunk_idx = (
|
||||||
|
kv_chunk_index_start // self.full_chunk_pages
|
||||||
|
if self.full_chunk_pages > 0
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
|
||||||
|
stg = req.staging
|
||||||
|
if stg is None or chunk_idx >= len(stg.offsets):
|
||||||
|
return (False, chunk_idx, -1, 0, -1)
|
||||||
|
|
||||||
|
c_offset = stg.offsets[chunk_idx]
|
||||||
|
if c_offset == StagingAllocator.ALLOC_OVERSIZED:
|
||||||
|
return (False, chunk_idx, StagingAllocator.ALLOC_OVERSIZED, 0, -1)
|
||||||
|
if c_offset < 0:
|
||||||
|
return (False, chunk_idx, -1, 0, -1)
|
||||||
|
|
||||||
|
c_round = stg.rounds[chunk_idx]
|
||||||
|
c_end = stg.ends[chunk_idx]
|
||||||
|
|
||||||
|
if not self.kv_manager._is_watermark_ready(
|
||||||
|
req.mooncake_session_id, c_round, c_end
|
||||||
|
):
|
||||||
|
return (False, chunk_idx, c_offset, c_round, c_end)
|
||||||
|
|
||||||
|
return (True, chunk_idx, c_offset, c_round, c_end)
|
||||||
|
|
||||||
|
def transfer(
|
||||||
|
self,
|
||||||
|
session_id: str,
|
||||||
|
prefill_kv_indices,
|
||||||
|
dst_staging_ptr: int,
|
||||||
|
dst_staging_size: int,
|
||||||
|
target_info,
|
||||||
|
) -> int:
|
||||||
|
"""Execute staged transfer (gather + RDMA).
|
||||||
|
|
||||||
|
Returns 0 on success, -1 to signal fallback to slice path.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
return self.kv_manager.send_kvcache_staged(
|
||||||
|
session_id,
|
||||||
|
prefill_kv_indices,
|
||||||
|
dst_staging_ptr,
|
||||||
|
dst_staging_size,
|
||||||
|
target_info.dst_tp_rank,
|
||||||
|
target_info.dst_attn_tp_size,
|
||||||
|
target_info.dst_kv_item_len,
|
||||||
|
staging_buffer=self.staging_buffer,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"[Staging] KV transfer via staging buffer failed: {e}. "
|
||||||
|
f"session={session_id}"
|
||||||
|
) from e
|
||||||
|
|
||||||
|
|
||||||
|
def init_staging_buffers(engine, kv_args, count: int) -> list:
|
||||||
|
"""Create prefill-side staging buffers and register them with the engine.
|
||||||
|
|
||||||
|
Returns list of StagingBuffer instances.
|
||||||
|
"""
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import StagingBuffer
|
||||||
|
from sglang.srt.disaggregation.mooncake.utils import (
|
||||||
|
init_mooncake_custom_mem_pool,
|
||||||
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
|
size_mb = envs.SGLANG_DISAGG_STAGING_BUFFER_SIZE_MB.get()
|
||||||
|
size_bytes = size_mb * 1024 * 1024
|
||||||
|
gpu_id = kv_args.gpu_id
|
||||||
|
device = f"cuda:{gpu_id}"
|
||||||
|
|
||||||
|
_, custom_mem_pool, pool_type = init_mooncake_custom_mem_pool(device)
|
||||||
|
if custom_mem_pool is None:
|
||||||
|
logger.warning(
|
||||||
|
"No mooncake custom mem pool available for staging buffer. "
|
||||||
|
"NVLink transport will NOT work. Set SGLANG_MOONCAKE_CUSTOM_MEM_POOL."
|
||||||
|
)
|
||||||
|
|
||||||
|
buffers = []
|
||||||
|
for _ in range(count):
|
||||||
|
buf = StagingBuffer(size_bytes, device, gpu_id, custom_mem_pool=custom_mem_pool)
|
||||||
|
engine.batch_register([buf.get_ptr()], [buf.get_size()])
|
||||||
|
buffers.append(buf)
|
||||||
|
return buffers
|
||||||
|
|
||||||
|
|
||||||
|
def init_staging_allocator(engine, kv_args):
|
||||||
|
"""Create decode-side staging ring-buffer allocator and register with engine.
|
||||||
|
|
||||||
|
Returns a StagingAllocator instance.
|
||||||
|
"""
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator
|
||||||
|
from sglang.srt.disaggregation.mooncake.utils import (
|
||||||
|
init_mooncake_custom_mem_pool,
|
||||||
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
|
pool_size_mb = envs.SGLANG_DISAGG_STAGING_POOL_SIZE_MB.get()
|
||||||
|
pool_size_bytes = pool_size_mb * 1024 * 1024
|
||||||
|
gpu_id = kv_args.gpu_id
|
||||||
|
device = f"cuda:{gpu_id}"
|
||||||
|
|
||||||
|
_, custom_mem_pool, _ = init_mooncake_custom_mem_pool(device)
|
||||||
|
allocator = StagingAllocator(pool_size_bytes, device, gpu_id, custom_mem_pool)
|
||||||
|
engine.batch_register([allocator.get_base_ptr()], [allocator.get_total_size()])
|
||||||
|
return allocator
|
||||||
|
|
||||||
|
|
||||||
|
def handle_staging_req(
|
||||||
|
msg,
|
||||||
|
staging_allocator,
|
||||||
|
kv_args,
|
||||||
|
attn_tp_size: int,
|
||||||
|
prefill_attn_tp_size: int,
|
||||||
|
kv_buffer_tensors,
|
||||||
|
room_receivers: dict,
|
||||||
|
room_bootstrap: dict,
|
||||||
|
):
|
||||||
|
"""Allocate staging for a chunk on-demand and send STAGING_RSP to prefill.
|
||||||
|
|
||||||
|
Deduplicates: multiple prefill TP ranks requesting the same (room, chunk_idx)
|
||||||
|
only allocate once. Sends ALLOC_OVERSIZED on permanent failure.
|
||||||
|
"""
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator
|
||||||
|
|
||||||
|
room = int(msg[1].decode("ascii"))
|
||||||
|
chunk_idx = int(msg[2].decode("ascii"))
|
||||||
|
chunk_num_pages = int(msg[3].decode("ascii"))
|
||||||
|
session_id = msg[4].decode("ascii")
|
||||||
|
|
||||||
|
if staging_allocator is None:
|
||||||
|
logger.warning(
|
||||||
|
"STAGING_REQ ignored: allocator is None room=%s chunk=%s",
|
||||||
|
room,
|
||||||
|
chunk_idx,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
receiver = room_receivers.get(room)
|
||||||
|
if receiver is None:
|
||||||
|
logger.warning(
|
||||||
|
"STAGING_REQ dropped: no receiver for room=%s chunk=%s session=%s",
|
||||||
|
room,
|
||||||
|
chunk_idx,
|
||||||
|
session_id,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
infos = getattr(receiver, "chunk_staging_infos", [])
|
||||||
|
|
||||||
|
if chunk_idx < len(infos) and infos[chunk_idx][0] >= 0:
|
||||||
|
_, offset, rnd, end, _ = infos[chunk_idx]
|
||||||
|
elif (
|
||||||
|
chunk_idx < len(infos)
|
||||||
|
and infos[chunk_idx][1] == StagingAllocator.ALLOC_OVERSIZED
|
||||||
|
):
|
||||||
|
offset, rnd, end = StagingAllocator.ALLOC_OVERSIZED, 0, -1
|
||||||
|
else:
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import (
|
||||||
|
compute_staging_layout,
|
||||||
|
resolve_total_kv_heads,
|
||||||
|
)
|
||||||
|
|
||||||
|
page_size = kv_args.page_size
|
||||||
|
kv_item_lens = kv_args.kv_item_lens
|
||||||
|
num_kv_layers = len(kv_item_lens) // 2
|
||||||
|
decode_bytes_per_token = kv_item_lens[0] // page_size
|
||||||
|
total_kv_heads = resolve_total_kv_heads(kv_args, attn_tp_size)
|
||||||
|
dst_heads_per_rank = max(1, total_kv_heads // max(1, attn_tp_size))
|
||||||
|
bytes_per_head_per_token = decode_bytes_per_token // dst_heads_per_rank
|
||||||
|
dst_tp_rank = kv_args.engine_rank % max(1, attn_tp_size)
|
||||||
|
|
||||||
|
chunk_tokens = chunk_num_pages * page_size
|
||||||
|
_, _, required = compute_staging_layout(
|
||||||
|
prefill_attn_tp_size,
|
||||||
|
attn_tp_size,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
chunk_tokens,
|
||||||
|
bytes_per_head_per_token,
|
||||||
|
num_kv_layers,
|
||||||
|
)
|
||||||
|
result = staging_allocator.assign(required)
|
||||||
|
if result is None:
|
||||||
|
logger.error(
|
||||||
|
"[STAGING_REQ] alloc failed room=%s chunk=%d (need %d bytes, "
|
||||||
|
"buffer total=%d bytes). Increase SGLANG_DISAGG_STAGING_POOL_SIZE_MB.",
|
||||||
|
room,
|
||||||
|
chunk_idx,
|
||||||
|
required,
|
||||||
|
staging_allocator.total_size,
|
||||||
|
)
|
||||||
|
offset, rnd, end = StagingAllocator.ALLOC_OVERSIZED, 0, -1
|
||||||
|
while len(infos) <= chunk_idx:
|
||||||
|
infos.append((-1, -1, 0, -1, 0))
|
||||||
|
infos[chunk_idx] = (
|
||||||
|
-1,
|
||||||
|
StagingAllocator.ALLOC_OVERSIZED,
|
||||||
|
0,
|
||||||
|
-1,
|
||||||
|
chunk_num_pages,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
alloc_id, offset, rnd = result
|
||||||
|
end = offset + required
|
||||||
|
while len(infos) <= chunk_idx:
|
||||||
|
infos.append((-1, -1, 0, -1, 0))
|
||||||
|
infos[chunk_idx] = (alloc_id, offset, rnd, end, chunk_num_pages)
|
||||||
|
|
||||||
|
bootstrap_infos = room_bootstrap.get(room)
|
||||||
|
if bootstrap_infos:
|
||||||
|
for bi in bootstrap_infos:
|
||||||
|
try:
|
||||||
|
sock, lock = receiver._connect_to_bootstrap_server(bi)
|
||||||
|
with lock:
|
||||||
|
sock.send_multipart(
|
||||||
|
[
|
||||||
|
b"STAGING_RSP",
|
||||||
|
str(room).encode("ascii"),
|
||||||
|
str(chunk_idx).encode("ascii"),
|
||||||
|
str(offset).encode("ascii"),
|
||||||
|
str(rnd).encode("ascii"),
|
||||||
|
str(end).encode("ascii"),
|
||||||
|
session_id.encode("ascii"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def prefetch_staging_reqs(
|
||||||
|
room: int,
|
||||||
|
transfer_infos: dict,
|
||||||
|
kv_buffer_tensors: dict,
|
||||||
|
chunked_prefill_size: int,
|
||||||
|
staging_requested: set,
|
||||||
|
prefetch_sockets: dict,
|
||||||
|
) -> None:
|
||||||
|
"""Send STAGING_REQ for all chunks before the prefill forward starts.
|
||||||
|
|
||||||
|
Called from the scheduler right after batch formation, so that decode
|
||||||
|
allocates staging during the GPU forward pass.
|
||||||
|
"""
|
||||||
|
import zmq
|
||||||
|
|
||||||
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
|
||||||
|
page_size = kv_buffer_tensors["page_size"]
|
||||||
|
cps = chunked_prefill_size or 8192
|
||||||
|
full_chunk_pages = max(1, cps // page_size)
|
||||||
|
|
||||||
|
for session_id, tinfo in transfer_infos[room].items():
|
||||||
|
if tinfo.is_dummy:
|
||||||
|
continue
|
||||||
|
total_pages = len(tinfo.dst_kv_indices)
|
||||||
|
if total_pages == 0:
|
||||||
|
continue
|
||||||
|
num_chunks = (total_pages + full_chunk_pages - 1) // full_chunk_pages
|
||||||
|
|
||||||
|
for chunk_idx in range(num_chunks):
|
||||||
|
stg_key = (room, chunk_idx, session_id)
|
||||||
|
if stg_key in staging_requested:
|
||||||
|
continue
|
||||||
|
staging_requested.add(stg_key)
|
||||||
|
|
||||||
|
remaining = total_pages - chunk_idx * full_chunk_pages
|
||||||
|
chunk_pages = min(full_chunk_pages, remaining)
|
||||||
|
try:
|
||||||
|
na = NetworkAddress(tinfo.endpoint, tinfo.dst_port)
|
||||||
|
ep = na.to_tcp()
|
||||||
|
if ep not in prefetch_sockets:
|
||||||
|
sock = zmq.Context().socket(zmq.PUSH)
|
||||||
|
if na.is_ipv6:
|
||||||
|
sock.setsockopt(zmq.IPV6, 1)
|
||||||
|
sock.connect(ep)
|
||||||
|
prefetch_sockets[ep] = sock
|
||||||
|
prefetch_sockets[ep].send_multipart(
|
||||||
|
[
|
||||||
|
b"STAGING_REQ",
|
||||||
|
str(room).encode("ascii"),
|
||||||
|
str(chunk_idx).encode("ascii"),
|
||||||
|
str(chunk_pages).encode("ascii"),
|
||||||
|
session_id.encode("ascii"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
staging_requested.discard(stg_key)
|
||||||
@@ -45,6 +45,7 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
is_mla_backend,
|
is_mla_backend,
|
||||||
kv_to_page_indices,
|
kv_to_page_indices,
|
||||||
poll_and_all_reduce,
|
poll_and_all_reduce,
|
||||||
|
poll_and_all_reduce_with_staging,
|
||||||
prepare_abort,
|
prepare_abort,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -290,7 +291,10 @@ class DecodePreallocQueue:
|
|||||||
self._max_ensure_retries: int = 15 # scheduling cycles
|
self._max_ensure_retries: int = 15 # scheduling cycles
|
||||||
self._ensure_last_attempt_time: Dict[str, float] = {}
|
self._ensure_last_attempt_time: Dict[str, float] = {}
|
||||||
self._ensure_retry_interval: float = 1.0 # seconds
|
self._ensure_retry_interval: float = 1.0 # seconds
|
||||||
|
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
self.kv_manager = self._init_kv_manager()
|
self.kv_manager = self._init_kv_manager()
|
||||||
|
if self.enable_staging:
|
||||||
|
self.transfer_queue._init_staging_handler(self.kv_manager)
|
||||||
|
|
||||||
if self.scheduler.tp_worker.is_hybrid_swa:
|
if self.scheduler.tp_worker.is_hybrid_swa:
|
||||||
# FIXME: current SWA allocation allocate full kv cache size in prefill
|
# FIXME: current SWA allocation allocate full kv cache size in prefill
|
||||||
@@ -366,6 +370,21 @@ class DecodePreallocQueue:
|
|||||||
self.scheduler.server_args,
|
self.scheduler.server_args,
|
||||||
self.is_mla_backend,
|
self.is_mla_backend,
|
||||||
)
|
)
|
||||||
|
# Staging buffer setup (only when heterogeneous TP staging is enabled)
|
||||||
|
if self.enable_staging and not self.is_mla_backend:
|
||||||
|
kv_pool_for_heads = self.token_to_kv_pool
|
||||||
|
if hasattr(kv_pool_for_heads, "full_kv_pool"):
|
||||||
|
kv_pool_for_heads = kv_pool_for_heads.full_kv_pool
|
||||||
|
per_rank_kv_heads = getattr(kv_pool_for_heads, "head_num", 0)
|
||||||
|
if per_rank_kv_heads > 0:
|
||||||
|
kv_args.kv_head_num = per_rank_kv_heads
|
||||||
|
kv_args.total_kv_head_num = per_rank_kv_heads * attn_tp_size
|
||||||
|
if hasattr(kv_manager, "set_kv_buffer_tensors"):
|
||||||
|
kv_pool = kv_pool_for_heads
|
||||||
|
if hasattr(kv_pool, "k_buffer") and hasattr(kv_pool, "v_buffer"):
|
||||||
|
kv_manager.set_kv_buffer_tensors(
|
||||||
|
kv_pool.k_buffer, kv_pool.v_buffer, kv_pool.page_size
|
||||||
|
)
|
||||||
return kv_manager
|
return kv_manager
|
||||||
|
|
||||||
def add(self, req: Req, is_retracted: bool = False) -> None:
|
def add(self, req: Req, is_retracted: bool = False) -> None:
|
||||||
@@ -737,6 +756,13 @@ class DecodePreallocQueue:
|
|||||||
decode_req.kv_receiver.send_metadata(
|
decode_req.kv_receiver.send_metadata(
|
||||||
page_indices, decode_req.metadata_buffer_index, state_indices
|
page_indices, decode_req.metadata_buffer_index, state_indices
|
||||||
)
|
)
|
||||||
|
if (
|
||||||
|
self.transfer_queue.enable_staging
|
||||||
|
and decode_req.kv_receiver.require_staging
|
||||||
|
):
|
||||||
|
self.transfer_queue.staging_handler.register_decode_req(
|
||||||
|
decode_req.req.bootstrap_room, decode_req
|
||||||
|
)
|
||||||
preallocated_reqs.append(decode_req)
|
preallocated_reqs.append(decode_req)
|
||||||
indices_to_remove.add(i)
|
indices_to_remove.add(i)
|
||||||
decode_req.req.time_stats.set_decode_transfer_queue_entry_time()
|
decode_req.req.time_stats.set_decode_transfer_queue_entry_time()
|
||||||
@@ -863,12 +889,18 @@ class DecodeTransferQueue:
|
|||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
self.tree_cache = tree_cache
|
self.tree_cache = tree_cache
|
||||||
self.spec_algorithm = scheduler.spec_algorithm
|
self.spec_algorithm = scheduler.spec_algorithm
|
||||||
|
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
|
self.staging_handler = None
|
||||||
|
|
||||||
def add(self, decode_req: DecodeRequest) -> None:
|
def add(self, decode_req: DecodeRequest) -> None:
|
||||||
self.queue.append(decode_req)
|
self.queue.append(decode_req)
|
||||||
|
|
||||||
def extend(self, decode_reqs: List[DecodeRequest]) -> None:
|
def extend(self, decode_reqs: List[DecodeRequest]) -> None:
|
||||||
self.queue.extend(decode_reqs)
|
self.queue.extend(decode_reqs)
|
||||||
|
if self.enable_staging:
|
||||||
|
for dr in decode_reqs:
|
||||||
|
if dr.kv_receiver.require_staging:
|
||||||
|
self.staging_handler.register_decode_req(dr.req.bootstrap_room, dr)
|
||||||
|
|
||||||
def _commit_transfer_to_req(self, decode_req: DecodeRequest) -> bool:
|
def _commit_transfer_to_req(self, decode_req: DecodeRequest) -> bool:
|
||||||
"""
|
"""
|
||||||
@@ -951,18 +983,39 @@ class DecodeTransferQueue:
|
|||||||
decode_req.req.time_stats.set_wait_queue_entry_time()
|
decode_req.req.time_stats.set_wait_queue_entry_time()
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
def _poll_with_staging(self) -> list:
|
||||||
|
return poll_and_all_reduce_with_staging(
|
||||||
|
self.queue, self.staging_handler, self.gloo_group
|
||||||
|
)
|
||||||
|
|
||||||
|
def _init_staging_handler(self, kv_manager):
|
||||||
|
"""Create staging handler from kv_manager. Must be called exactly once."""
|
||||||
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
DecodeStagingHandler,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.staging_handler = DecodeStagingHandler.create(
|
||||||
|
kv_manager, self.scheduler, self.tp_rank
|
||||||
|
)
|
||||||
|
kv_manager._staging_handler = self.staging_handler
|
||||||
|
|
||||||
def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req]:
|
def pop_transferred(self, rids_to_check: Optional[List[str]] = None) -> List[Req]:
|
||||||
if not self.queue:
|
if not self.queue:
|
||||||
return []
|
return []
|
||||||
polls = poll_and_all_reduce(
|
|
||||||
[decode_req.kv_receiver for decode_req in self.queue], self.gloo_group
|
if self.enable_staging:
|
||||||
)
|
polls = self._poll_with_staging()
|
||||||
|
else:
|
||||||
|
polls = poll_and_all_reduce(
|
||||||
|
[dr.kv_receiver for dr in self.queue], self.gloo_group
|
||||||
|
)
|
||||||
|
|
||||||
transferred_reqs = []
|
transferred_reqs = []
|
||||||
indices_to_remove = set()
|
indices_to_remove = set()
|
||||||
for i, (decode_req, poll) in enumerate(zip(self.queue, polls)):
|
for i, (decode_req, poll) in enumerate(zip(self.queue, polls)):
|
||||||
if rids_to_check is not None and decode_req.req.rid not in rids_to_check:
|
if rids_to_check is not None and decode_req.req.rid not in rids_to_check:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if poll == KVPoll.Failed:
|
if poll == KVPoll.Failed:
|
||||||
error_message = f"Decode transfer failed for request rank={self.tp_rank} {decode_req.req.rid=} {decode_req.req.bootstrap_room=}"
|
error_message = f"Decode transfer failed for request rank={self.tp_rank} {decode_req.req.rid=} {decode_req.req.bootstrap_room=}"
|
||||||
try:
|
try:
|
||||||
@@ -1010,6 +1063,12 @@ class DecodeTransferQueue:
|
|||||||
raise ValueError(f"Unexpected poll case: {poll}")
|
raise ValueError(f"Unexpected poll case: {poll}")
|
||||||
|
|
||||||
for i in indices_to_remove:
|
for i in indices_to_remove:
|
||||||
|
if self.enable_staging and self.staging_handler.is_staging_room(
|
||||||
|
self.queue[i].req.bootstrap_room
|
||||||
|
):
|
||||||
|
self.staging_handler.unregister_decode_req(
|
||||||
|
self.queue[i].req.bootstrap_room
|
||||||
|
)
|
||||||
idx = self.queue[i].metadata_buffer_index
|
idx = self.queue[i].metadata_buffer_index
|
||||||
assert idx != -1
|
assert idx != -1
|
||||||
self.req_to_metadata_buffer_idx_allocator.free(idx)
|
self.req_to_metadata_buffer_idx_allocator.free(idx)
|
||||||
|
|||||||
@@ -61,6 +61,14 @@ class TransferKVChunk:
|
|||||||
state_indices: Optional[List[int]]
|
state_indices: Optional[List[int]]
|
||||||
|
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
DecodeStagingContext,
|
||||||
|
PrefillStagingContext,
|
||||||
|
StagingRegisterInfo,
|
||||||
|
StagingTransferInfo,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# decode
|
# decode
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class TransferInfo:
|
class TransferInfo:
|
||||||
@@ -73,6 +81,7 @@ class TransferInfo:
|
|||||||
dst_state_indices: List[int]
|
dst_state_indices: List[int]
|
||||||
required_dst_info_num: int
|
required_dst_info_num: int
|
||||||
is_dummy: bool
|
is_dummy: bool
|
||||||
|
staging: Optional[StagingTransferInfo] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_zmq(cls, msg: List[bytes]):
|
def from_zmq(cls, msg: List[bytes]):
|
||||||
@@ -118,6 +127,7 @@ class KVArgsRegisterInfo:
|
|||||||
# for mamba state different tp slice transfer
|
# for mamba state different tp slice transfer
|
||||||
dst_state_item_lens: list[int]
|
dst_state_item_lens: list[int]
|
||||||
dst_state_dim_per_tensor: list[int]
|
dst_state_dim_per_tensor: list[int]
|
||||||
|
staging: Optional[StagingRegisterInfo] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_zmq(cls, msg: List[bytes]):
|
def from_zmq(cls, msg: List[bytes]):
|
||||||
@@ -142,6 +152,7 @@ class KVArgsRegisterInfo:
|
|||||||
if len(msg) > 11 and len(msg[11]) > 0
|
if len(msg) > 11 and len(msg[11]) > 0
|
||||||
else []
|
else []
|
||||||
),
|
),
|
||||||
|
staging=StagingRegisterInfo.from_zmq_fields(msg, 12),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -178,6 +189,7 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
|
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
|
||||||
self.init_engine()
|
self.init_engine()
|
||||||
self.register_buffer_to_engine()
|
self.register_buffer_to_engine()
|
||||||
|
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||||
self.start_prefill_thread()
|
self.start_prefill_thread()
|
||||||
self.session_failures = defaultdict(int)
|
self.session_failures = defaultdict(int)
|
||||||
@@ -204,14 +216,34 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
for _ in range(transfer_queue_size)
|
for _ in range(transfer_queue_size)
|
||||||
]
|
]
|
||||||
for queue, executor in zip(self.transfer_queues, self.executors):
|
|
||||||
threading.Thread(
|
|
||||||
target=self.transfer_worker, args=(queue, executor), daemon=True
|
|
||||||
).start()
|
|
||||||
self.enable_custom_mem_pool, self.custom_mem_pool_type = (
|
self.enable_custom_mem_pool, self.custom_mem_pool_type = (
|
||||||
check_mooncake_custom_mem_pool_enabled()
|
check_mooncake_custom_mem_pool_enabled()
|
||||||
)
|
)
|
||||||
|
self._staging_ctx = PrefillStagingContext() if self.enable_staging else None
|
||||||
|
if self.enable_staging:
|
||||||
|
self._init_staging_buffers(len(self.transfer_queues))
|
||||||
|
for i, (queue, executor) in enumerate(
|
||||||
|
zip(self.transfer_queues, self.executors)
|
||||||
|
):
|
||||||
|
threading.Thread(
|
||||||
|
target=self.transfer_worker,
|
||||||
|
args=(
|
||||||
|
queue,
|
||||||
|
executor,
|
||||||
|
(
|
||||||
|
self._staging_ctx.buffers[i]
|
||||||
|
if self.enable_staging and self._staging_ctx.buffers
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
),
|
||||||
|
daemon=True,
|
||||||
|
).start()
|
||||||
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
elif self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||||
|
self._staging_ctx = DecodeStagingContext() if self.enable_staging else None
|
||||||
|
if self.enable_staging:
|
||||||
|
self._init_staging_allocator()
|
||||||
|
self._staging_handler = None
|
||||||
|
self._chunk_writer_counts: dict = defaultdict(lambda: defaultdict(list))
|
||||||
self.start_decode_thread()
|
self.start_decode_thread()
|
||||||
|
|
||||||
def init_engine(self):
|
def init_engine(self):
|
||||||
@@ -236,6 +268,297 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
self.kv_args.state_data_ptrs, self.kv_args.state_data_lens
|
self.kv_args.state_data_ptrs, self.kv_args.state_data_lens
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Staging buffer methods (all delegate to staging_handler.py)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def register_staging_room_bootstrap(self, room, bootstrap_infos, receiver):
|
||||||
|
self._staging_ctx.room_bootstrap[room] = bootstrap_infos
|
||||||
|
self._staging_ctx.room_receivers[room] = receiver
|
||||||
|
|
||||||
|
def set_kv_buffer_tensors(self, k_buffers: list, v_buffers: list, page_size: int):
|
||||||
|
self.kv_buffer_tensors = {
|
||||||
|
"k_buffers": k_buffers,
|
||||||
|
"v_buffers": v_buffers,
|
||||||
|
"page_size": page_size,
|
||||||
|
}
|
||||||
|
|
||||||
|
def _init_staging_buffers(self, count: int):
|
||||||
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
init_staging_buffers,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._staging_ctx.buffers = init_staging_buffers(
|
||||||
|
self.engine, self.kv_args, count
|
||||||
|
)
|
||||||
|
self.kv_buffer_tensors = None
|
||||||
|
|
||||||
|
def _init_staging_allocator(self):
|
||||||
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
init_staging_allocator,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._staging_ctx.allocator = init_staging_allocator(self.engine, self.kv_args)
|
||||||
|
self.kv_buffer_tensors = None
|
||||||
|
|
||||||
|
def _handle_staging_req(self, msg):
|
||||||
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
handle_staging_req,
|
||||||
|
)
|
||||||
|
|
||||||
|
room = int(msg[1].decode("ascii"))
|
||||||
|
session_id = msg[4].decode("ascii")
|
||||||
|
handler = self._staging_handler
|
||||||
|
assert (
|
||||||
|
handler is not None
|
||||||
|
), "STAGING_REQ received before staging handler initialized"
|
||||||
|
decode_req = handler._room_to_decode_req.get(room)
|
||||||
|
if decode_req is None:
|
||||||
|
logger.warning(
|
||||||
|
"STAGING_REQ received for unregistered room=%s, skipping",
|
||||||
|
room,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
prefill_tp = decode_req.kv_receiver.prefill_info.attn_tp_size
|
||||||
|
handle_staging_req(
|
||||||
|
msg,
|
||||||
|
self._staging_ctx.allocator,
|
||||||
|
self.kv_args,
|
||||||
|
self.attn_tp_size,
|
||||||
|
prefill_tp,
|
||||||
|
getattr(self, "kv_buffer_tensors", None),
|
||||||
|
self._staging_ctx.room_receivers,
|
||||||
|
self._staging_ctx.room_bootstrap,
|
||||||
|
)
|
||||||
|
|
||||||
|
receiver = self._staging_ctx.room_receivers.get(room)
|
||||||
|
if receiver is not None:
|
||||||
|
handler.register_wm_subscriber(receiver, session_id)
|
||||||
|
|
||||||
|
def _is_watermark_ready(
|
||||||
|
self, session_id: str, alloc_round: int, alloc_end: int
|
||||||
|
) -> bool:
|
||||||
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
is_watermark_ready,
|
||||||
|
)
|
||||||
|
|
||||||
|
return is_watermark_ready(self._staging_ctx, session_id, alloc_round, alloc_end)
|
||||||
|
|
||||||
|
def _try_create_staging_strategy(self, staging_buffer):
|
||||||
|
if not self.enable_staging or self.kv_buffer_tensors is None:
|
||||||
|
return None
|
||||||
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
PrefillStagingStrategy,
|
||||||
|
)
|
||||||
|
|
||||||
|
return PrefillStagingStrategy(self, staging_buffer)
|
||||||
|
|
||||||
|
def _send_chunk_ready(self, req, chunk_idx, kv_chunk, prefill_unique_rank):
|
||||||
|
"""Notify decode that a non-last staging chunk RDMA is complete."""
|
||||||
|
try:
|
||||||
|
na = NetworkAddress(req.endpoint, req.dst_port)
|
||||||
|
self._connect(
|
||||||
|
na.to_tcp(),
|
||||||
|
is_ipv6=na.is_ipv6,
|
||||||
|
).send_multipart(
|
||||||
|
[
|
||||||
|
b"CHUNK_READY",
|
||||||
|
str(req.room).encode("ascii"),
|
||||||
|
str(chunk_idx).encode("ascii"),
|
||||||
|
str(kv_chunk.index_slice.start).encode("ascii"),
|
||||||
|
str(len(kv_chunk.prefill_kv_indices)).encode("ascii"),
|
||||||
|
req.mooncake_session_id.encode("ascii"),
|
||||||
|
str(prefill_unique_rank).encode("ascii"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def _do_staging_transfer(
|
||||||
|
self,
|
||||||
|
staging_strategy,
|
||||||
|
kv_chunk,
|
||||||
|
req,
|
||||||
|
target_info,
|
||||||
|
chunked_dst_kv_indice,
|
||||||
|
executor,
|
||||||
|
queue,
|
||||||
|
prefill_unique_rank,
|
||||||
|
):
|
||||||
|
"""Execute staging transfer for one chunk. Returns (ret, deferred).
|
||||||
|
|
||||||
|
Handles readiness check, transfer, fallback, and CHUNK_READY notification.
|
||||||
|
deferred=True means caller should re-enqueue and break.
|
||||||
|
"""
|
||||||
|
_tp = self.attn_tp_rank
|
||||||
|
ready, chunk_idx, c_offset, _, _ = staging_strategy.check_ready(
|
||||||
|
req,
|
||||||
|
kv_chunk.index_slice.start,
|
||||||
|
len(kv_chunk.prefill_kv_indices),
|
||||||
|
)
|
||||||
|
if not ready:
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator
|
||||||
|
|
||||||
|
if c_offset == StagingAllocator.ALLOC_OVERSIZED:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"[Staging] Chunk staging allocation permanently failed: "
|
||||||
|
f"chunk exceeds ring buffer total size (room={kv_chunk.room}). "
|
||||||
|
f"Increase SGLANG_DISAGG_STAGING_POOL_SIZE_MB."
|
||||||
|
)
|
||||||
|
queue.put(kv_chunk)
|
||||||
|
return (-1, True)
|
||||||
|
|
||||||
|
ret = staging_strategy.transfer(
|
||||||
|
req.mooncake_session_id,
|
||||||
|
kv_chunk.prefill_kv_indices,
|
||||||
|
target_info.staging.base_ptr + c_offset,
|
||||||
|
target_info.staging.total_size - c_offset,
|
||||||
|
target_info,
|
||||||
|
)
|
||||||
|
if ret == -1:
|
||||||
|
logger.warning(
|
||||||
|
f"[Staging][tp{_tp}] Falling back to per-token slice path "
|
||||||
|
f"(room={kv_chunk.room})"
|
||||||
|
)
|
||||||
|
ret = self.send_kvcache_slice(
|
||||||
|
req.mooncake_session_id,
|
||||||
|
kv_chunk.prefill_kv_indices,
|
||||||
|
target_info.dst_kv_ptrs,
|
||||||
|
chunked_dst_kv_indice,
|
||||||
|
target_info.dst_tp_rank,
|
||||||
|
target_info.dst_attn_tp_size,
|
||||||
|
target_info.dst_kv_item_len,
|
||||||
|
executor,
|
||||||
|
)
|
||||||
|
elif ret == 0 and not kv_chunk.is_last_chunk:
|
||||||
|
self._send_chunk_ready(req, chunk_idx, kv_chunk, prefill_unique_rank)
|
||||||
|
return (ret, False)
|
||||||
|
|
||||||
|
def _prefetch_staging_reqs(self, room: int):
|
||||||
|
if not self.enable_staging or self.kv_buffer_tensors is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
room_infos = self.transfer_infos.get(room, {})
|
||||||
|
needs_staging = any(
|
||||||
|
not tinfo.is_dummy
|
||||||
|
and self.decode_kv_args_table.get(tinfo.mooncake_session_id) is not None
|
||||||
|
and self.decode_kv_args_table[tinfo.mooncake_session_id].dst_attn_tp_size
|
||||||
|
!= self.attn_tp_size
|
||||||
|
for tinfo in room_infos.values()
|
||||||
|
)
|
||||||
|
if not needs_staging:
|
||||||
|
return
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.common.staging_handler import (
|
||||||
|
prefetch_staging_reqs,
|
||||||
|
)
|
||||||
|
|
||||||
|
prefetch_staging_reqs(
|
||||||
|
room,
|
||||||
|
self.transfer_infos,
|
||||||
|
self.kv_buffer_tensors,
|
||||||
|
self.server_args.chunked_prefill_size,
|
||||||
|
self._staging_ctx.prefetch_requested,
|
||||||
|
self._staging_ctx.prefetch_sockets,
|
||||||
|
)
|
||||||
|
|
||||||
|
def send_kvcache_staged(
|
||||||
|
self,
|
||||||
|
mooncake_session_id: str,
|
||||||
|
prefill_kv_indices: npt.NDArray[np.int32],
|
||||||
|
dst_staging_ptr: int,
|
||||||
|
dst_staging_size: int,
|
||||||
|
dst_tp_rank: int,
|
||||||
|
dst_attn_tp_size: int,
|
||||||
|
dst_kv_item_len: int,
|
||||||
|
staging_buffer=None,
|
||||||
|
) -> int:
|
||||||
|
"""Transfer KV cache via staging buffers (gather -> bulk RDMA -> scatter on decode)."""
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import (
|
||||||
|
compute_head_slice_params,
|
||||||
|
compute_staging_layout,
|
||||||
|
resolve_total_kv_heads,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.kv_buffer_tensors is None or staging_buffer is None:
|
||||||
|
return -1
|
||||||
|
|
||||||
|
k_buffers = self.kv_buffer_tensors["k_buffers"]
|
||||||
|
v_buffers = self.kv_buffer_tensors["v_buffers"]
|
||||||
|
page_size = self.kv_buffer_tensors["page_size"]
|
||||||
|
num_layers = len(k_buffers)
|
||||||
|
head_dim = k_buffers[0].shape[-1]
|
||||||
|
dtype_size = k_buffers[0].element_size()
|
||||||
|
|
||||||
|
total_kv_heads = resolve_total_kv_heads(self.kv_args, self.attn_tp_size)
|
||||||
|
|
||||||
|
local_tp_rank = self.kv_args.engine_rank % self.attn_tp_size
|
||||||
|
src_head_start, num_heads_to_send, _, _ = compute_head_slice_params(
|
||||||
|
self.attn_tp_size,
|
||||||
|
dst_attn_tp_size,
|
||||||
|
local_tp_rank,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
)
|
||||||
|
|
||||||
|
num_tokens = len(prefill_kv_indices) * page_size
|
||||||
|
per_layer_bytes = num_tokens * num_heads_to_send * head_dim * dtype_size
|
||||||
|
per_rank_bytes = per_layer_bytes * num_layers * 2
|
||||||
|
|
||||||
|
num_writers, writer_rank_bytes, total_staging_needed = compute_staging_layout(
|
||||||
|
self.attn_tp_size,
|
||||||
|
dst_attn_tp_size,
|
||||||
|
dst_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
num_tokens,
|
||||||
|
head_dim * dtype_size,
|
||||||
|
num_layers,
|
||||||
|
)
|
||||||
|
writer_idx = local_tp_rank % num_writers if num_writers > 1 else 0
|
||||||
|
rank_offset = sum(writer_rank_bytes[:writer_idx])
|
||||||
|
|
||||||
|
if not staging_buffer.fits(per_rank_bytes):
|
||||||
|
logger.warning(
|
||||||
|
f"Prefill staging too small for {per_rank_bytes} bytes, falling back"
|
||||||
|
)
|
||||||
|
return -1
|
||||||
|
if dst_staging_size < total_staging_needed:
|
||||||
|
logger.warning(
|
||||||
|
f"Decode staging too small: need {total_staging_needed} bytes "
|
||||||
|
f"({num_writers if self.attn_tp_size > dst_attn_tp_size else 1} writers "
|
||||||
|
f"x {per_rank_bytes} bytes/rank), have {dst_staging_size}, falling back"
|
||||||
|
)
|
||||||
|
return -1
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import (
|
||||||
|
gather_all_layers_to_staging,
|
||||||
|
)
|
||||||
|
|
||||||
|
gather_all_layers_to_staging(
|
||||||
|
k_buffers,
|
||||||
|
v_buffers,
|
||||||
|
prefill_kv_indices,
|
||||||
|
staging_buffer,
|
||||||
|
src_head_start,
|
||||||
|
num_heads_to_send,
|
||||||
|
page_size,
|
||||||
|
self.kv_args.gpu_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
dst_write_ptr = dst_staging_ptr + rank_offset
|
||||||
|
ret = self._transfer_data(
|
||||||
|
mooncake_session_id,
|
||||||
|
[(staging_buffer.get_ptr(), dst_write_ptr, per_rank_bytes)],
|
||||||
|
)
|
||||||
|
if ret != 0:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"[Staging] Bulk RDMA transfer failed with ret={ret}. "
|
||||||
|
f"src_ptr=0x{staging_buffer.get_ptr():x}, "
|
||||||
|
f"dst_ptr=0x{dst_write_ptr:x}, size={per_rank_bytes}. "
|
||||||
|
f"The decode staging buffer may not be properly registered."
|
||||||
|
)
|
||||||
|
return ret
|
||||||
|
|
||||||
def _transfer_data(self, mooncake_session_id, transfer_blocks):
|
def _transfer_data(self, mooncake_session_id, transfer_blocks):
|
||||||
if not transfer_blocks:
|
if not transfer_blocks:
|
||||||
return 0
|
return 0
|
||||||
@@ -770,11 +1093,22 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def transfer_worker(
|
def transfer_worker(
|
||||||
self, queue: FastQueue, executor: concurrent.futures.ThreadPoolExecutor
|
self,
|
||||||
|
queue: FastQueue,
|
||||||
|
executor: concurrent.futures.ThreadPoolExecutor,
|
||||||
|
staging_buffer=None,
|
||||||
):
|
):
|
||||||
|
staging_strategy = None
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
kv_chunk: TransferKVChunk = queue.get()
|
kv_chunk: TransferKVChunk = queue.get()
|
||||||
|
if (
|
||||||
|
self.enable_staging
|
||||||
|
and staging_strategy is None
|
||||||
|
and staging_buffer is not None
|
||||||
|
):
|
||||||
|
staging_strategy = self._try_create_staging_strategy(staging_buffer)
|
||||||
reqs_to_be_processed = (
|
reqs_to_be_processed = (
|
||||||
self.transfer_infos[kv_chunk.room].values()
|
self.transfer_infos[kv_chunk.room].values()
|
||||||
if kv_chunk.room in self.transfer_infos
|
if kv_chunk.room in self.transfer_infos
|
||||||
@@ -788,6 +1122,9 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
+ self.pp_rank * self.attn_cp_size
|
+ self.pp_rank * self.attn_cp_size
|
||||||
+ self.attn_cp_rank
|
+ self.attn_cp_rank
|
||||||
)
|
)
|
||||||
|
# When staging transfer is not yet ready (watermark/allocation pending),
|
||||||
|
# the chunk is re-enqueued and we break out of the req loop to retry later.
|
||||||
|
staging_deferred = False
|
||||||
for req in reqs_to_be_processed:
|
for req in reqs_to_be_processed:
|
||||||
if not req.is_dummy:
|
if not req.is_dummy:
|
||||||
# Early exit if the request has failed
|
# Early exit if the request has failed
|
||||||
@@ -835,6 +1172,25 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
chunked_dst_kv_indice,
|
chunked_dst_kv_indice,
|
||||||
executor,
|
executor,
|
||||||
)
|
)
|
||||||
|
elif (
|
||||||
|
self.enable_staging
|
||||||
|
and staging_strategy is not None
|
||||||
|
and target_rank_registration_info.staging is not None
|
||||||
|
):
|
||||||
|
ret, deferred = self._do_staging_transfer(
|
||||||
|
staging_strategy,
|
||||||
|
kv_chunk,
|
||||||
|
req,
|
||||||
|
target_rank_registration_info,
|
||||||
|
chunked_dst_kv_indice,
|
||||||
|
executor,
|
||||||
|
queue,
|
||||||
|
prefill_unique_rank,
|
||||||
|
)
|
||||||
|
if deferred:
|
||||||
|
staging_deferred = True
|
||||||
|
# Chunk re-enqueued; stop processing remaining reqs for this chunk
|
||||||
|
break
|
||||||
else:
|
else:
|
||||||
ret = self.send_kvcache_slice(
|
ret = self.send_kvcache_slice(
|
||||||
req.mooncake_session_id,
|
req.mooncake_session_id,
|
||||||
@@ -909,6 +1265,9 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
if kv_chunk.is_last_chunk and req.room in self.request_status:
|
if kv_chunk.is_last_chunk and req.room in self.request_status:
|
||||||
self.update_status(req.room, KVPoll.Success)
|
self.update_status(req.room, KVPoll.Success)
|
||||||
|
|
||||||
|
if staging_deferred:
|
||||||
|
continue
|
||||||
|
|
||||||
if (
|
if (
|
||||||
kv_chunk.room not in self.request_status
|
kv_chunk.room not in self.request_status
|
||||||
or self.check_status(kv_chunk.room) == KVPoll.Success
|
or self.check_status(kv_chunk.room) == KVPoll.Success
|
||||||
@@ -929,6 +1288,50 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
while True:
|
while True:
|
||||||
waiting_req_bytes = self.server_socket.recv_multipart()
|
waiting_req_bytes = self.server_socket.recv_multipart()
|
||||||
room = waiting_req_bytes[0].decode("ascii")
|
room = waiting_req_bytes[0].decode("ascii")
|
||||||
|
# Staging: decode reports consumption watermark back to prefill
|
||||||
|
if room == "WATERMARK":
|
||||||
|
wm_round = int(waiting_req_bytes[1].decode("ascii"))
|
||||||
|
wm_tail = int(waiting_req_bytes[2].decode("ascii"))
|
||||||
|
wm_session = (
|
||||||
|
waiting_req_bytes[3].decode("ascii")
|
||||||
|
if len(waiting_req_bytes) > 3
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
with self._staging_ctx.watermark_cv:
|
||||||
|
prev = self._staging_ctx.remote_watermarks.get(
|
||||||
|
wm_session, (0, 0)
|
||||||
|
)
|
||||||
|
if (wm_round, wm_tail) > prev:
|
||||||
|
self._staging_ctx.remote_watermarks[wm_session] = (
|
||||||
|
wm_round,
|
||||||
|
wm_tail,
|
||||||
|
)
|
||||||
|
self._staging_ctx.watermark_cv.notify_all()
|
||||||
|
continue
|
||||||
|
# Staging: decode replies with allocated staging offset
|
||||||
|
if room == "STAGING_RSP":
|
||||||
|
stg_room = int(waiting_req_bytes[1].decode("ascii"))
|
||||||
|
stg_chunk_idx = int(waiting_req_bytes[2].decode("ascii"))
|
||||||
|
stg_offset = int(waiting_req_bytes[3].decode("ascii"))
|
||||||
|
stg_round = int(waiting_req_bytes[4].decode("ascii"))
|
||||||
|
stg_end = int(waiting_req_bytes[5].decode("ascii"))
|
||||||
|
stg_session = waiting_req_bytes[6].decode("ascii")
|
||||||
|
room_infos = self.transfer_infos.get(stg_room, {})
|
||||||
|
tinfo = room_infos.get(stg_session)
|
||||||
|
if tinfo is not None:
|
||||||
|
if tinfo.staging is None:
|
||||||
|
tinfo.staging = StagingTransferInfo()
|
||||||
|
tinfo.staging.set_chunk(
|
||||||
|
stg_chunk_idx, stg_offset, stg_round, stg_end
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"STAGING_RSP RECV but tinfo=None room=%s chunk=%d session=%s",
|
||||||
|
stg_room,
|
||||||
|
stg_chunk_idx,
|
||||||
|
stg_session,
|
||||||
|
)
|
||||||
|
continue
|
||||||
mooncake_session_id = waiting_req_bytes[3].decode("ascii")
|
mooncake_session_id = waiting_req_bytes[3].decode("ascii")
|
||||||
if room == "None":
|
if room == "None":
|
||||||
self.decode_kv_args_table[mooncake_session_id] = (
|
self.decode_kv_args_table[mooncake_session_id] = (
|
||||||
@@ -966,6 +1369,42 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
self._handle_aux_data(msg)
|
self._handle_aux_data(msg)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
# Staging: prefill notifies a chunk written to staging buffer
|
||||||
|
if msg[0] == b"CHUNK_READY":
|
||||||
|
room = int(msg[1].decode("ascii"))
|
||||||
|
chunk_idx = int(msg[2].decode("ascii"))
|
||||||
|
page_start = int(msg[3].decode("ascii"))
|
||||||
|
num_pages = int(msg[4].decode("ascii"))
|
||||||
|
session_id = msg[5].decode("ascii")
|
||||||
|
self._chunk_writer_counts[room][chunk_idx].append(
|
||||||
|
(page_start, num_pages, session_id)
|
||||||
|
)
|
||||||
|
handler = self._staging_handler
|
||||||
|
assert (
|
||||||
|
handler is not None
|
||||||
|
), "CHUNK_READY received before staging handler initialized"
|
||||||
|
writers_arrived = len(self._chunk_writer_counts[room][chunk_idx])
|
||||||
|
decode_req = handler._room_to_decode_req.get(room)
|
||||||
|
if decode_req is None:
|
||||||
|
logger.warning(
|
||||||
|
"CHUNK_READY received for unregistered room=%s chunk=%d, skipping",
|
||||||
|
room,
|
||||||
|
chunk_idx,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
num_writers = handler.num_writers_for(decode_req)
|
||||||
|
if writers_arrived >= num_writers:
|
||||||
|
handler.submit_chunk_scatter(
|
||||||
|
room, chunk_idx, page_start, num_pages
|
||||||
|
)
|
||||||
|
del self._chunk_writer_counts[room][chunk_idx]
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Staging: prefill pre-requests staging allocation before forward
|
||||||
|
if msg[0] == b"STAGING_REQ":
|
||||||
|
self._handle_staging_req(msg)
|
||||||
|
continue
|
||||||
|
|
||||||
bootstrap_room, status, prefill_rank = msg
|
bootstrap_room, status, prefill_rank = msg
|
||||||
status = int(status.decode("ascii"))
|
status = int(status.decode("ascii"))
|
||||||
bootstrap_room = int(bootstrap_room.decode("ascii"))
|
bootstrap_room = int(bootstrap_room.decode("ascii"))
|
||||||
@@ -981,6 +1420,11 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
self.prefill_response_tracker[bootstrap_room]
|
self.prefill_response_tracker[bootstrap_room]
|
||||||
)
|
)
|
||||||
if arrived_response_num == expected_response_num:
|
if arrived_response_num == expected_response_num:
|
||||||
|
if self.enable_staging:
|
||||||
|
handler = self._staging_handler
|
||||||
|
if handler.is_staging_room(bootstrap_room):
|
||||||
|
handler.submit_last_scatter_async(bootstrap_room)
|
||||||
|
self._chunk_writer_counts.pop(bootstrap_room, None)
|
||||||
self.update_status(bootstrap_room, KVPoll.Success)
|
self.update_status(bootstrap_room, KVPoll.Success)
|
||||||
elif status == KVPoll.Failed:
|
elif status == KVPoll.Failed:
|
||||||
self.record_failure(
|
self.record_failure(
|
||||||
@@ -1272,6 +1716,17 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
|||||||
dst_attn_tp_size = str(self.kv_mgr.attn_tp_size).encode("ascii")
|
dst_attn_tp_size = str(self.kv_mgr.attn_tp_size).encode("ascii")
|
||||||
dst_kv_item_len = str(kv_item_len).encode("ascii")
|
dst_kv_item_len = str(kv_item_len).encode("ascii")
|
||||||
|
|
||||||
|
if (
|
||||||
|
self.kv_mgr.enable_staging
|
||||||
|
and self.kv_mgr._staging_ctx.allocator is not None
|
||||||
|
):
|
||||||
|
_alloc = self.kv_mgr._staging_ctx.allocator
|
||||||
|
packed_staging_base_ptr = struct.pack("Q", _alloc.get_base_ptr())
|
||||||
|
staging_total_size_str = str(_alloc.get_total_size()).encode("ascii")
|
||||||
|
else:
|
||||||
|
packed_staging_base_ptr = b""
|
||||||
|
staging_total_size_str = b""
|
||||||
|
|
||||||
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
|
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
|
||||||
with lock:
|
with lock:
|
||||||
sock.send_multipart(
|
sock.send_multipart(
|
||||||
@@ -1288,6 +1743,8 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
|||||||
dst_kv_item_len,
|
dst_kv_item_len,
|
||||||
packed_state_item_lens,
|
packed_state_item_lens,
|
||||||
packed_state_dim_per_tensor,
|
packed_state_dim_per_tensor,
|
||||||
|
packed_staging_base_ptr,
|
||||||
|
staging_total_size_str,
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1311,6 +1768,15 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
|||||||
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
|
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
if (
|
||||||
|
self.kv_mgr.enable_staging
|
||||||
|
and self.kv_mgr._staging_ctx.allocator is not None
|
||||||
|
):
|
||||||
|
self.chunk_staging_infos = []
|
||||||
|
self.kv_mgr.register_staging_room_bootstrap(
|
||||||
|
self.bootstrap_room, self.bootstrap_infos, self
|
||||||
|
)
|
||||||
|
|
||||||
for bootstrap_info in self.bootstrap_infos:
|
for bootstrap_info in self.bootstrap_infos:
|
||||||
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
|
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
|
||||||
is_dummy = bootstrap_info["is_dummy"]
|
is_dummy = bootstrap_info["is_dummy"]
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
poll_and_all_reduce_attn_cp_tp_group,
|
poll_and_all_reduce_attn_cp_tp_group,
|
||||||
prepare_abort,
|
prepare_abort,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.schedule_batch import (
|
from sglang.srt.managers.schedule_batch import (
|
||||||
FINISH_ABORT,
|
FINISH_ABORT,
|
||||||
FINISH_LENGTH,
|
FINISH_LENGTH,
|
||||||
@@ -201,6 +202,21 @@ class PrefillBootstrapQueue:
|
|||||||
self.scheduler.server_args,
|
self.scheduler.server_args,
|
||||||
self.is_mla_backend,
|
self.is_mla_backend,
|
||||||
)
|
)
|
||||||
|
# Pass KV pool tensor refs to the manager for GPU gather (staging mode)
|
||||||
|
if (
|
||||||
|
envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
|
and hasattr(kv_manager, "set_kv_buffer_tensors")
|
||||||
|
and not self.is_mla_backend
|
||||||
|
):
|
||||||
|
kv_pool = self.token_to_kv_pool
|
||||||
|
if hasattr(kv_pool, "full_kv_pool"):
|
||||||
|
kv_pool = kv_pool.full_kv_pool
|
||||||
|
if hasattr(kv_pool, "k_buffer") and hasattr(kv_pool, "v_buffer"):
|
||||||
|
kv_manager.set_kv_buffer_tensors(
|
||||||
|
kv_pool.k_buffer,
|
||||||
|
kv_pool.v_buffer,
|
||||||
|
kv_pool.page_size,
|
||||||
|
)
|
||||||
return kv_manager
|
return kv_manager
|
||||||
|
|
||||||
def add(self, req: Req, num_kv_heads: int) -> None:
|
def add(self, req: Req, num_kv_heads: int) -> None:
|
||||||
@@ -336,6 +352,17 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
Mixin for Scheduler to handle disaggregation prefill
|
Mixin for Scheduler to handle disaggregation prefill
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
def maybe_prefetch_staging_for_batch(self: Scheduler, batch: ScheduleBatch) -> None:
|
||||||
|
"""Pre-send STAGING_REQ so decode allocates staging during GPU forward."""
|
||||||
|
kv_mgr = self.disagg_prefill_bootstrap_queue.kv_manager
|
||||||
|
prefetch = getattr(kv_mgr, "_prefetch_staging_reqs", None)
|
||||||
|
if prefetch is None:
|
||||||
|
return
|
||||||
|
for req in batch.reqs:
|
||||||
|
room = getattr(req, "bootstrap_room", None)
|
||||||
|
if room is not None and room in kv_mgr.transfer_infos:
|
||||||
|
prefetch(room)
|
||||||
|
|
||||||
def get_next_disagg_prefill_batch_to_run(
|
def get_next_disagg_prefill_batch_to_run(
|
||||||
self: Scheduler,
|
self: Scheduler,
|
||||||
) -> Optional[ScheduleBatch]:
|
) -> Optional[ScheduleBatch]:
|
||||||
@@ -356,6 +383,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def event_loop_normal_disagg_prefill(self: Scheduler) -> None:
|
def event_loop_normal_disagg_prefill(self: Scheduler) -> None:
|
||||||
"""A normal scheduler loop for prefill worker in disaggregation mode."""
|
"""A normal scheduler loop for prefill worker in disaggregation mode."""
|
||||||
|
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
@@ -371,6 +399,8 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
|
|
||||||
# Launch the current batch
|
# Launch the current batch
|
||||||
if batch:
|
if batch:
|
||||||
|
if self.enable_staging:
|
||||||
|
self.maybe_prefetch_staging_for_batch(batch)
|
||||||
result = self.run_batch(batch)
|
result = self.run_batch(batch)
|
||||||
self.process_batch_result(batch, result)
|
self.process_batch_result(batch, result)
|
||||||
else:
|
else:
|
||||||
@@ -384,6 +414,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def event_loop_overlap_disagg_prefill(self: Scheduler) -> None:
|
def event_loop_overlap_disagg_prefill(self: Scheduler) -> None:
|
||||||
self.result_queue = deque()
|
self.result_queue = deque()
|
||||||
|
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
@@ -399,6 +430,8 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
|
|
||||||
# Launch the current batch
|
# Launch the current batch
|
||||||
if batch:
|
if batch:
|
||||||
|
if self.enable_staging:
|
||||||
|
self.maybe_prefetch_staging_for_batch(batch)
|
||||||
batch_result = self.run_batch(batch)
|
batch_result = self.run_batch(batch)
|
||||||
self.result_queue.append((batch.copy(), batch_result))
|
self.result_queue.append((batch.copy(), batch_result))
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -80,6 +80,30 @@ def poll_and_all_reduce_attn_cp_tp_group(
|
|||||||
return tensor_to_reduce.tolist()
|
return tensor_to_reduce.tolist()
|
||||||
|
|
||||||
|
|
||||||
|
def poll_and_all_reduce_with_staging(
|
||||||
|
decode_reqs, staging_handler, gloo_group: dist.ProcessGroup
|
||||||
|
):
|
||||||
|
"""Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
|
||||||
|
from sglang.srt.disaggregation.base import KVPoll
|
||||||
|
|
||||||
|
for decode_req in decode_reqs:
|
||||||
|
if decode_req.kv_receiver.require_staging and not staging_handler.is_done(
|
||||||
|
decode_req
|
||||||
|
):
|
||||||
|
staging_handler.advance_scatter(decode_req)
|
||||||
|
|
||||||
|
raw_polls = [int(dr.kv_receiver.poll()) for dr in decode_reqs]
|
||||||
|
for i, decode_req in enumerate(decode_reqs):
|
||||||
|
if raw_polls[i] == int(KVPoll.Success):
|
||||||
|
if decode_req.kv_receiver.require_staging and not staging_handler.is_done(
|
||||||
|
decode_req
|
||||||
|
):
|
||||||
|
raw_polls[i] = int(KVPoll.Transferring)
|
||||||
|
poll_tensor = torch.tensor(raw_polls, dtype=torch.uint8, device="cpu")
|
||||||
|
dist.all_reduce(poll_tensor, op=dist.ReduceOp.MIN, group=gloo_group)
|
||||||
|
return poll_tensor.tolist()
|
||||||
|
|
||||||
|
|
||||||
#########################
|
#########################
|
||||||
# Metadata Buffers
|
# Metadata Buffers
|
||||||
#########################
|
#########################
|
||||||
|
|||||||
@@ -287,6 +287,14 @@ class Envs:
|
|||||||
# Max fraction of cache (by token count) that can be pinned; 0 = disable pinning.
|
# Max fraction of cache (by token count) that can be pinned; 0 = disable pinning.
|
||||||
SGLANG_HICACHE_MAX_PINNED_RATIO = EnvFloat(0.0)
|
SGLANG_HICACHE_MAX_PINNED_RATIO = EnvFloat(0.0)
|
||||||
|
|
||||||
|
# Staging buffer for heterogeneous TP KV transfer
|
||||||
|
SGLANG_DISAGG_STAGING_BUFFER = EnvBool(False)
|
||||||
|
SGLANG_DISAGG_STAGING_BUFFER_SIZE_MB = EnvInt(64)
|
||||||
|
SGLANG_DISAGG_STAGING_POOL_SIZE_MB = EnvInt(4096)
|
||||||
|
# TODO(yangminl): remove SGLANG_STAGING_USE_TORCH and the torch fallback in
|
||||||
|
# staging_buffer.py once Triton kernels are fully validated in production.
|
||||||
|
SGLANG_STAGING_USE_TORCH = EnvBool(False)
|
||||||
|
|
||||||
# Mooncake KV Transfer
|
# Mooncake KV Transfer
|
||||||
SGLANG_MOONCAKE_CUSTOM_MEM_POOL = EnvStr(None)
|
SGLANG_MOONCAKE_CUSTOM_MEM_POOL = EnvStr(None)
|
||||||
ENABLE_ASCEND_TRANSFER_WITH_MOONCAKE = EnvBool(False)
|
ENABLE_ASCEND_TRANSFER_WITH_MOONCAKE = EnvBool(False)
|
||||||
|
|||||||
@@ -3238,6 +3238,17 @@ class ServerArgs:
|
|||||||
"Cuda graph is disabled for prefill server when piecewise cuda graph is not enabled."
|
"Cuda graph is disabled for prefill server when piecewise cuda graph is not enabled."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if self.disaggregation_mode in ("prefill", "decode"):
|
||||||
|
if (
|
||||||
|
envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
|
and self.disaggregation_transfer_backend != "mooncake"
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"SGLANG_DISAGG_STAGING_BUFFER requires "
|
||||||
|
f"disaggregation_transfer_backend='mooncake', "
|
||||||
|
f"got '{self.disaggregation_transfer_backend}'."
|
||||||
|
)
|
||||||
|
|
||||||
def _handle_encoder_disaggregation(self):
|
def _handle_encoder_disaggregation(self):
|
||||||
if self.enable_prefix_mm_cache and not self.encoder_only:
|
if self.enable_prefix_mm_cache and not self.encoder_only:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
Reference in New Issue
Block a user