[PD] Pack DCP1→DCP-N PD KV transfers into dest-contiguous RDMA blocks (#35762)
Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Cursor
Claude Opus 5
parent
5f216fc33f
commit
3760296be8
@@ -61,6 +61,7 @@ __all__ = ["reshape_and_cache_flash"]
|
||||
_TRITON_KERNELS = [
|
||||
("cache_ops", "concat_and_cast_mha_k_triton"),
|
||||
("cache_ops", "launch_reshape_and_cache_flash"),
|
||||
("pd_dcp_gather", "copy_mla_rows_into_pack"),
|
||||
("kv_indices", "create_flashinfer_kv_indices_triton"),
|
||||
("kv_indices", "create_flashmla_kv_indices_triton"),
|
||||
("kv_indices", "create_chunked_prefix_cache_kv_indices"),
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
from typing import Sequence
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _copy_mla_rows_into_pack_kernel(
|
||||
src_metadata,
|
||||
row_indices,
|
||||
pack,
|
||||
num_rows,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
layer_id = tl.program_id(0)
|
||||
block_id = tl.program_id(1)
|
||||
metadata_offset = layer_id * 3
|
||||
src = tl.load(src_metadata + metadata_offset).to(pack.dtype)
|
||||
row_nbytes = tl.load(src_metadata + metadata_offset + 1)
|
||||
pack_offset = tl.load(src_metadata + metadata_offset + 2)
|
||||
|
||||
offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
|
||||
layer_nbytes = num_rows * row_nbytes
|
||||
mask = offsets < layer_nbytes
|
||||
row = offsets // row_nbytes
|
||||
byte = offsets % row_nbytes
|
||||
src_row = tl.load(row_indices + row, mask=mask, other=0)
|
||||
values = tl.load(src + src_row * row_nbytes + byte, mask=mask)
|
||||
tl.store(pack + pack_offset + offsets, values, mask=mask)
|
||||
|
||||
|
||||
def copy_mla_rows_into_pack(
|
||||
kv_data_ptrs: Sequence[int],
|
||||
row_indices: torch.Tensor,
|
||||
pack: torch.Tensor,
|
||||
token_item_lens: Sequence[int],
|
||||
) -> None:
|
||||
if len(kv_data_ptrs) != len(token_item_lens):
|
||||
raise ValueError(
|
||||
"kv_data_ptrs and token_item_lens length mismatch: "
|
||||
f"{len(kv_data_ptrs)} vs {len(token_item_lens)}"
|
||||
)
|
||||
if not kv_data_ptrs:
|
||||
return
|
||||
|
||||
n = int(row_indices.numel())
|
||||
metadata = []
|
||||
offset = 0
|
||||
for ptr, item_len in zip(kv_data_ptrs, token_item_lens):
|
||||
item_len = int(item_len)
|
||||
if item_len <= 0:
|
||||
raise ValueError(f"MLA token item length must be positive, got {item_len}")
|
||||
metadata.extend((int(ptr), item_len, offset))
|
||||
offset += n * item_len
|
||||
|
||||
src_metadata = torch.tensor(metadata, dtype=torch.int64, device=pack.device)
|
||||
max_item_len = max(int(item_len) for item_len in token_item_lens)
|
||||
grid = (len(kv_data_ptrs), triton.cdiv(n * max_item_len, 1024))
|
||||
_copy_mla_rows_into_pack_kernel[grid](
|
||||
src_metadata,
|
||||
row_indices,
|
||||
pack,
|
||||
n,
|
||||
BLOCK_SIZE=1024,
|
||||
)
|
||||
@@ -167,6 +167,7 @@ class CommonKVManager(BaseKVManager):
|
||||
self.enable_deferred_decode_kv_release = (
|
||||
envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get()
|
||||
)
|
||||
self._dcp_pack_buffers = None
|
||||
# for p/d multi node infer
|
||||
self.bootstrap_host = get_serving().host
|
||||
self.bootstrap_port = get_disagg().disaggregation_bootstrap_port
|
||||
@@ -339,6 +340,25 @@ class CommonKVManager(BaseKVManager):
|
||||
)
|
||||
return src_token_lens
|
||||
|
||||
def _register_staging_memory(self, ptr: int, size: int) -> None:
|
||||
raise NotImplementedError(
|
||||
f"{type(self).__name__} does not support staging memory registration"
|
||||
)
|
||||
|
||||
def _init_dcp_pack_buffers_once(self, dcp_size: int) -> None:
|
||||
if self._dcp_pack_buffers is not None:
|
||||
return
|
||||
if not self.kv_args.kv_item_lens:
|
||||
return
|
||||
from sglang.srt.disaggregation.common.dcp_pack import init_dcp_pack_buffers
|
||||
|
||||
self._dcp_pack_buffers = init_dcp_pack_buffers(
|
||||
self._register_staging_memory,
|
||||
self.kv_args,
|
||||
len(self.transfer_queues),
|
||||
dcp_size,
|
||||
)
|
||||
|
||||
def check_status(self, bootstrap_room: int) -> KVPoll:
|
||||
return self.request_status[bootstrap_room]
|
||||
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import List, Optional, Sequence, Tuple
|
||||
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.kvcache.pd_dcp_gather import copy_mla_rows_into_pack
|
||||
from sglang.srt.disaggregation.common.staging_buffer import StagingBuffer
|
||||
from sglang.srt.runtime_context import get_schedule, max_prefill_buffer_tokens
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def dcp_pack_buffer_bytes(
|
||||
kv_item_lens: Sequence[int], page_size: int, max_tokens: int, dcp_size: int = 1
|
||||
) -> int:
|
||||
if page_size <= 0:
|
||||
raise ValueError(f"page_size must be positive, got {page_size}")
|
||||
if dcp_size <= 0:
|
||||
raise ValueError(f"dcp_size must be positive, got {dcp_size}")
|
||||
if any(item_len < page_size for item_len in kv_item_lens):
|
||||
raise ValueError(
|
||||
"PD DCP pack requires each kv_item_len to span at least one page, "
|
||||
f"got {list(kv_item_lens)} with page_size={page_size}"
|
||||
)
|
||||
if any(item_len % page_size != 0 for item_len in kv_item_lens):
|
||||
raise ValueError(
|
||||
"PD DCP pack requires page-aligned kv_item_lens, "
|
||||
f"got {list(kv_item_lens)} with page_size={page_size}"
|
||||
)
|
||||
token_item_lens = [item_len // page_size for item_len in kv_item_lens]
|
||||
rank_tokens = (max_tokens + dcp_size - 1) // dcp_size
|
||||
return dcp_size * rank_tokens * sum(token_item_lens)
|
||||
|
||||
|
||||
def try_pack_dcp_src(
|
||||
*,
|
||||
pack_buffer: StagingBuffer,
|
||||
kv_data_ptrs: Sequence[int],
|
||||
src_token_indices: npt.NDArray[np.integer],
|
||||
token_item_lens: Sequence[int],
|
||||
pack_offset_bytes: int = 0,
|
||||
) -> Optional[Tuple[List[int], npt.NDArray[np.int64]]]:
|
||||
if pack_offset_bytes < 0:
|
||||
raise ValueError(
|
||||
f"pack_offset_bytes must be non-negative, got {pack_offset_bytes}"
|
||||
)
|
||||
n = int(src_token_indices.size)
|
||||
if n == 0:
|
||||
empty = np.empty((0,), dtype=np.int64)
|
||||
return [], empty
|
||||
required = n * sum(int(item_len) for item_len in token_item_lens)
|
||||
required_end = pack_offset_bytes + required
|
||||
if not pack_buffer.fits(required_end):
|
||||
logger.warning(
|
||||
"PD DCP pack buffer too small for byte range [%s, %s) (have %s); "
|
||||
"falling back to per-token RDMA",
|
||||
pack_offset_bytes,
|
||||
required_end,
|
||||
pack_buffer.get_size(),
|
||||
)
|
||||
return None
|
||||
|
||||
pack = pack_buffer.buffer.narrow(0, pack_offset_bytes, required)
|
||||
row_indices = torch.as_tensor(
|
||||
src_token_indices, device=pack.device, dtype=torch.int64
|
||||
)
|
||||
gather_stream = pack_buffer.get_gather_stream()
|
||||
gather_stream.wait_stream(torch.cuda.default_stream(pack.device))
|
||||
with torch.cuda.stream(gather_stream):
|
||||
copy_mla_rows_into_pack(kv_data_ptrs, row_indices, pack, token_item_lens)
|
||||
gather_stream.synchronize()
|
||||
|
||||
packed_ptrs: List[int] = []
|
||||
offset = 0
|
||||
base = pack_buffer.get_ptr() + pack_offset_bytes
|
||||
for item_len in token_item_lens:
|
||||
packed_ptrs.append(base + offset)
|
||||
offset += n * int(item_len)
|
||||
return packed_ptrs, np.arange(n, dtype=np.int64)
|
||||
|
||||
|
||||
def init_dcp_pack_buffers(
|
||||
register_fn,
|
||||
kv_args,
|
||||
count: int,
|
||||
dcp_size: int,
|
||||
) -> List[StagingBuffer]:
|
||||
from sglang.srt.disaggregation.common.staging_handler import (
|
||||
_get_custom_mem_pool,
|
||||
)
|
||||
|
||||
max_tokens = max_prefill_buffer_tokens()
|
||||
if max_tokens <= 0:
|
||||
max_tokens = get_schedule().max_prefill_tokens
|
||||
# Note(kpham-sgl): size = dcp_size x ceil(max_tokens / dcp_size)
|
||||
# x sum(per-layer token bytes). At 32,768 tokens and 61 MLA layers
|
||||
# x 576 bf16 dims x 2 B: 2.14 GiB/buffer, 8.58 GiB for 4 queues.
|
||||
size_bytes = dcp_pack_buffer_bytes(
|
||||
kv_args.kv_item_lens, kv_args.page_size, max_tokens, dcp_size
|
||||
)
|
||||
gpu_id = kv_args.gpu_id
|
||||
device = f"cuda:{gpu_id}"
|
||||
custom_mem_pool, _ = _get_custom_mem_pool(device)
|
||||
|
||||
buffers = []
|
||||
for _ in range(count):
|
||||
buf = StagingBuffer(size_bytes, device, gpu_id, custom_mem_pool=custom_mem_pool)
|
||||
register_fn(buf.get_ptr(), buf.get_size())
|
||||
buffers.append(buf)
|
||||
logger.info(
|
||||
"PD DCP pack buffers allocated: %d x %.1f MB (max_tokens=%d)",
|
||||
count,
|
||||
size_bytes / (1024 * 1024),
|
||||
max_tokens,
|
||||
)
|
||||
return buffers
|
||||
@@ -133,6 +133,7 @@ class StagingBuffer:
|
||||
self.size_bytes = size_bytes
|
||||
self.device = device
|
||||
self.gpu_id = gpu_id
|
||||
self._gather_stream: Optional[torch.cuda.Stream] = None
|
||||
|
||||
torch.cuda.set_device(gpu_id)
|
||||
if custom_mem_pool is not None:
|
||||
@@ -158,6 +159,11 @@ class StagingBuffer:
|
||||
def fits(self, required_bytes: int) -> bool:
|
||||
return required_bytes <= self.size_bytes
|
||||
|
||||
def get_gather_stream(self) -> torch.cuda.Stream:
|
||||
if self._gather_stream is None:
|
||||
self._gather_stream = torch.cuda.Stream(device=self.device)
|
||||
return self._gather_stream
|
||||
|
||||
|
||||
class StagingAllocator:
|
||||
"""Decode-side dynamic staging ring buffer allocator with overcommit.
|
||||
@@ -350,16 +356,12 @@ def _gather_all_layers_torch(
|
||||
|
||||
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))
|
||||
)
|
||||
gather_stream = staging_buffer.get_gather_stream()
|
||||
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):
|
||||
with torch.cuda.stream(gather_stream):
|
||||
for layer_id in range(num_layers):
|
||||
dst = (
|
||||
staging_view[offset : offset + per_layer_bytes]
|
||||
@@ -389,7 +391,7 @@ def _gather_all_layers_torch(
|
||||
)
|
||||
offset += per_layer_bytes
|
||||
|
||||
staging_buffer._gather_stream.synchronize()
|
||||
gather_stream.synchronize()
|
||||
return offset
|
||||
|
||||
|
||||
@@ -430,17 +432,13 @@ def _gather_all_layers_triton(
|
||||
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))
|
||||
)
|
||||
gather_stream = staging_buffer.get_gather_stream()
|
||||
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):
|
||||
with torch.cuda.stream(gather_stream):
|
||||
_fused_gather_to_staging_kernel[grid](
|
||||
layer_ptrs,
|
||||
page_idx_tensor,
|
||||
@@ -454,7 +452,7 @@ def _gather_all_layers_triton(
|
||||
BLOCK_SIZE,
|
||||
)
|
||||
|
||||
staging_buffer._gather_stream.synchronize()
|
||||
gather_stream.synchronize()
|
||||
return total_bytes
|
||||
|
||||
|
||||
|
||||
@@ -353,13 +353,16 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
"slot_layer_ids": list(slot_layer_ids or []),
|
||||
}
|
||||
|
||||
def _register_staging_memory(self, ptr: int, size: int) -> None:
|
||||
self.engine.batch_register([ptr], [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(
|
||||
lambda ptr, size: self.engine.batch_register([ptr], [size]),
|
||||
self._register_staging_memory,
|
||||
self.kv_args,
|
||||
count,
|
||||
get_schedule().chunked_prefill_size,
|
||||
@@ -372,7 +375,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
)
|
||||
|
||||
self._staging_ctx.allocator = init_staging_allocator(
|
||||
lambda ptr, size: self.engine.batch_register([ptr], [size]),
|
||||
self._register_staging_memory,
|
||||
self.kv_args,
|
||||
)
|
||||
self.kv_buffer_tensors = None
|
||||
@@ -910,6 +913,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
num_kv_tokens: int,
|
||||
executor: concurrent.futures.ThreadPoolExecutor,
|
||||
dst_layer_ids: List[int],
|
||||
pack_buffer=None,
|
||||
) -> int:
|
||||
if num_kv_tokens is None:
|
||||
raise ValueError("PD DCP transfer requires num_kv_tokens")
|
||||
@@ -942,10 +946,24 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
self.kv_args.kv_data_ptrs,
|
||||
dst_kv_ptrs,
|
||||
)
|
||||
src_token_indices = plan.src_token_indices
|
||||
dst_token_indices = plan.dst_token_indices
|
||||
if pack_buffer is not None:
|
||||
from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src
|
||||
|
||||
packed = try_pack_dcp_src(
|
||||
pack_buffer=pack_buffer,
|
||||
kv_data_ptrs=src_kv_ptrs,
|
||||
src_token_indices=src_token_indices,
|
||||
token_item_lens=dcp_token_item_lens[: len(src_kv_ptrs)],
|
||||
)
|
||||
if packed is not None:
|
||||
src_kv_ptrs, src_token_indices = packed
|
||||
|
||||
layers_current_pp_stage = len(src_kv_ptrs)
|
||||
src_groups, dst_groups = group_concurrent_contiguous(
|
||||
plan.src_token_indices,
|
||||
plan.dst_token_indices,
|
||||
src_token_indices,
|
||||
dst_token_indices,
|
||||
)
|
||||
|
||||
layers_params = [
|
||||
@@ -1771,6 +1789,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
target_rank_registration_info.dcp_token_item_lens
|
||||
)
|
||||
assert dcp_token_item_lens is not None
|
||||
pack_buffer = (
|
||||
self._dcp_pack_buffers[worker_index]
|
||||
if self._dcp_pack_buffers
|
||||
else None
|
||||
)
|
||||
ret = self.send_kvcache_dcp(
|
||||
req.mooncake_session_id,
|
||||
kv_chunk.prefill_kv_indices,
|
||||
@@ -1786,6 +1809,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
dst_layer_ids=(
|
||||
target_rank_registration_info.dst_kv_layer_ids
|
||||
),
|
||||
pack_buffer=pack_buffer,
|
||||
)
|
||||
elif (
|
||||
self.is_mla_backend
|
||||
@@ -2073,6 +2097,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
* len(self.kv_args.kv_item_lens)
|
||||
)
|
||||
)
|
||||
self._init_dcp_pack_buffers_once(decode_kv_args.dst_dcp_size)
|
||||
self.decode_kv_args_table[mooncake_session_id] = decode_kv_args
|
||||
with self.session_lock:
|
||||
if mooncake_session_id in self.failed_sessions:
|
||||
|
||||
@@ -506,7 +506,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
)
|
||||
threading.Thread(
|
||||
target=self.transfer_worker,
|
||||
args=(queue, staging_buffer),
|
||||
args=(queue, staging_buffer, i),
|
||||
daemon=True,
|
||||
).start()
|
||||
self._start_bootstrap_thread()
|
||||
@@ -545,9 +545,8 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
init_staging_buffers,
|
||||
)
|
||||
|
||||
gpu_id = self.kv_args.gpu_id
|
||||
self._staging_ctx.buffers = init_staging_buffers(
|
||||
lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id),
|
||||
self._register_staging_memory,
|
||||
self.kv_args,
|
||||
count,
|
||||
get_schedule().chunked_prefill_size,
|
||||
@@ -558,15 +557,14 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
init_staging_allocator,
|
||||
)
|
||||
|
||||
gpu_id = self.kv_args.gpu_id
|
||||
self._staging_ctx.allocator = init_staging_allocator(
|
||||
lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id),
|
||||
self._register_staging_memory,
|
||||
self.kv_args,
|
||||
)
|
||||
|
||||
def _register_staging_memory(self, ptr: int, size: int, gpu_id: int):
|
||||
def _register_staging_memory(self, ptr: int, size: int):
|
||||
"""Register a staging buffer with the NIXL agent."""
|
||||
addrs = [(ptr, size, gpu_id, "")]
|
||||
addrs = [(ptr, size, self.kv_args.gpu_id, "")]
|
||||
descs = self.agent.register_memory(addrs, "VRAM")
|
||||
if not descs:
|
||||
raise RuntimeError(
|
||||
@@ -1080,7 +1078,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
dst_mem_kind=dst_mem_kind,
|
||||
)
|
||||
|
||||
def transfer_worker(self, queue: FastQueue, staging_buffer=None):
|
||||
def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0):
|
||||
# Per-worker staging strategy: lazy-created on first chunk so we
|
||||
# see kv_buffer_tensors (set by ModelRunner after engine init).
|
||||
# Never cache on self -- multiple workers would race the ring.
|
||||
@@ -1120,6 +1118,10 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
self.update_status(room, KVPoll.Transferring)
|
||||
|
||||
reqs_to_be_processed = list(self.transfer_infos[room].values())
|
||||
# Note(kpham-sgl): Pack each DCP rank once into its fixed region.
|
||||
# NIXL reads regions asynchronously; the chunk barrier prevents
|
||||
# reuse until every transfer completes.
|
||||
packed_source_by_dcp_rank = {}
|
||||
|
||||
# Set when staging allocation/watermark is not yet ready and
|
||||
# the chunk has been re-enqueued. We then break out of the
|
||||
@@ -1209,15 +1211,37 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
|
||||
if kv_xfer_handle is None:
|
||||
if is_dcp_transfer:
|
||||
kv_xfer_handle = self.send_kvcache_dcp(
|
||||
req.agent_name,
|
||||
pack_buffer = (
|
||||
self._dcp_pack_buffers[worker_index]
|
||||
if self._dcp_pack_buffers
|
||||
else None
|
||||
)
|
||||
if kv_chunk.num_kv_tokens is None:
|
||||
raise ValueError(
|
||||
"PD DCP transfer requires num_kv_tokens"
|
||||
)
|
||||
plan = build_dcp_token_transfer_plan(
|
||||
src_prefill_kv_indices,
|
||||
dst_info,
|
||||
chunked_dst_kv_indice,
|
||||
src_page_offset=kv_chunk.index_slice.start or 0,
|
||||
physical_page_size=self.kv_args.page_size,
|
||||
dcp_size=dst_info.dst_dcp_size,
|
||||
dcp_rank=dst_info.dst_dcp_rank,
|
||||
src_page_offset=(kv_chunk.index_slice.start or 0),
|
||||
decode_prefix_len=req.decode_prefix_len or 0,
|
||||
num_kv_tokens=kv_chunk.num_kv_tokens,
|
||||
notif=notif,
|
||||
)
|
||||
packed_src = self._pack_dcp_rank_once(
|
||||
pack_buffer,
|
||||
dst_info,
|
||||
plan.src_token_indices,
|
||||
packed_source_by_dcp_rank,
|
||||
)
|
||||
kv_xfer_handle = self.send_kvcache_dcp(
|
||||
req.agent_name,
|
||||
dst_info,
|
||||
plan,
|
||||
notif,
|
||||
packed_src,
|
||||
)
|
||||
elif (
|
||||
self.is_mla_backend
|
||||
@@ -1434,6 +1458,8 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
decode_kv_args.requires_dcp_relayout = self.requires_dcp_relayout(
|
||||
decode_kv_args.dst_dcp_size, decode_kv_args.dst_dcp_rank
|
||||
)
|
||||
if decode_kv_args.requires_dcp_relayout:
|
||||
self._init_dcp_pack_buffers_once(decode_kv_args.dst_dcp_size)
|
||||
self.decode_kv_args_table[agent_name] = decode_kv_args
|
||||
self.agent.add_remote_agent(decode_kv_args.agent_metadata)
|
||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
@@ -1632,36 +1658,51 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
dst_mem_kind=dst_mem_kind,
|
||||
)
|
||||
|
||||
def _pack_dcp_rank_once(
|
||||
self,
|
||||
pack_buffer,
|
||||
dst_info: KVArgsRegisterInfo,
|
||||
src_token_indices,
|
||||
packed_source_by_dcp_rank,
|
||||
):
|
||||
"""Pack one source region per DCP rank for the current chunk.
|
||||
|
||||
Note(kpham-sgl): TP ranks may share a DCP rank (`tp_rank % dcp_size`),
|
||||
so they reuse one packed source while sending to distinct destination GPUs.
|
||||
"""
|
||||
rank = dst_info.dst_dcp_rank
|
||||
if rank in packed_source_by_dcp_rank:
|
||||
return packed_source_by_dcp_rank[rank]
|
||||
if pack_buffer is None or src_token_indices.size == 0:
|
||||
packed_source_by_dcp_rank[rank] = None
|
||||
return None
|
||||
|
||||
from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src
|
||||
|
||||
token_item_lens = dst_info.dcp_token_item_lens
|
||||
assert token_item_lens is not None
|
||||
rank_stride = pack_buffer.get_size() // dst_info.dst_dcp_size
|
||||
packed_source_by_dcp_rank[rank] = try_pack_dcp_src(
|
||||
pack_buffer=pack_buffer,
|
||||
kv_data_ptrs=self.kv_args.kv_data_ptrs,
|
||||
src_token_indices=src_token_indices,
|
||||
token_item_lens=token_item_lens[: len(self.kv_args.kv_data_ptrs)],
|
||||
pack_offset_bytes=rank * rank_stride,
|
||||
)
|
||||
return packed_source_by_dcp_rank[rank]
|
||||
|
||||
def send_kvcache_dcp(
|
||||
self,
|
||||
peer_name: str,
|
||||
prefill_kv_indices: npt.NDArray[np.int32],
|
||||
dst_info: KVArgsRegisterInfo,
|
||||
dst_kv_indices: npt.NDArray[np.int32],
|
||||
*,
|
||||
src_page_offset: int,
|
||||
decode_prefix_len: int,
|
||||
num_kv_tokens: int,
|
||||
plan,
|
||||
notif: str,
|
||||
packed_src,
|
||||
):
|
||||
if self.src_mem_kind is None:
|
||||
raise RuntimeError("Missing NIXL source KV memory kind")
|
||||
if dst_info.dst_homogeneous_mem_kind is None:
|
||||
raise RuntimeError("Missing NIXL destination KV memory kind")
|
||||
if num_kv_tokens is None:
|
||||
raise ValueError("PD DCP transfer requires num_kv_tokens")
|
||||
|
||||
physical_page_size = self.kv_args.page_size
|
||||
plan = build_dcp_token_transfer_plan(
|
||||
prefill_kv_indices,
|
||||
dst_kv_indices,
|
||||
physical_page_size=physical_page_size,
|
||||
dcp_size=dst_info.dst_dcp_size,
|
||||
dcp_rank=dst_info.dst_dcp_rank,
|
||||
src_page_offset=src_page_offset,
|
||||
decode_prefix_len=decode_prefix_len,
|
||||
num_kv_tokens=num_kv_tokens,
|
||||
)
|
||||
if plan.src_token_indices.size == 0:
|
||||
self.agent.send_notif(peer_name, notif.encode("ascii"))
|
||||
return None
|
||||
@@ -1671,15 +1712,18 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
||||
dst_kv_ptrs = [
|
||||
dst_info.dst_kv_ptrs[dst_idx] for dst_idx in dst_info.dcp_dst_region_indices
|
||||
]
|
||||
src_kv_ptrs = self.kv_args.kv_data_ptrs
|
||||
src_token_indices = plan.src_token_indices
|
||||
if packed_src is not None:
|
||||
src_kv_ptrs, src_token_indices = packed_src
|
||||
token_item_lens = token_item_lens[: len(src_kv_ptrs)]
|
||||
|
||||
# Prepared handles encode page-level offsets, while DCP relayout needs
|
||||
# flat descriptors for the selected token rows.
|
||||
return self._send_kvcache_generic(
|
||||
peer_name=peer_name,
|
||||
src_data_ptrs=self.kv_args.kv_data_ptrs,
|
||||
src_data_ptrs=src_kv_ptrs,
|
||||
dst_data_ptrs=dst_kv_ptrs,
|
||||
item_lens=token_item_lens,
|
||||
prefill_data_indices=plan.src_token_indices,
|
||||
prefill_data_indices=src_token_indices,
|
||||
dst_data_indices=plan.dst_token_indices,
|
||||
dst_gpu_id=dst_info.gpu_id,
|
||||
notif=notif,
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.kvcache.pd_dcp_gather import copy_mla_rows_into_pack
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
|
||||
class TestPdDcpGather(CustomTestCase):
|
||||
def test_gathers_strided_rows_layer_major(self):
|
||||
dim = 8
|
||||
kv0 = torch.arange(32 * dim, dtype=torch.float32, device="cuda").view(
|
||||
32, 1, dim
|
||||
)
|
||||
kv1 = torch.arange(32 * 5, dtype=torch.float16, device="cuda").view(32, 1, 5)
|
||||
row_indices = torch.tensor([0, 4, 9, 12], dtype=torch.int64, device="cuda")
|
||||
item_lens = [int(kv0[0].nbytes), int(kv1[0].nbytes)]
|
||||
pack = torch.zeros(
|
||||
row_indices.numel() * sum(item_lens), dtype=torch.uint8, device="cuda"
|
||||
)
|
||||
|
||||
copy_mla_rows_into_pack(
|
||||
[kv0.data_ptr(), kv1.data_ptr()],
|
||||
row_indices,
|
||||
pack,
|
||||
item_lens,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
split = row_indices.numel() * item_lens[0]
|
||||
packed0 = pack[:split].view(torch.float32).view(4, 1, dim)
|
||||
packed1 = pack[split:].view(torch.float16).view(4, 1, 5)
|
||||
torch.testing.assert_close(packed0, kv0[row_indices], rtol=0, atol=0)
|
||||
torch.testing.assert_close(packed1, kv1[row_indices], rtol=0, atol=0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,120 @@
|
||||
import unittest
|
||||
from contextlib import nullcontext
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from sglang.srt.disaggregation.common.dcp_pack import (
|
||||
dcp_pack_buffer_bytes,
|
||||
try_pack_dcp_src,
|
||||
)
|
||||
from sglang.srt.disaggregation.common.utils import (
|
||||
build_dcp_token_transfer_plan,
|
||||
group_concurrent_contiguous,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class TestPackedDcpGrouping(CustomTestCase):
|
||||
def test_packed_groups_collapse_cyclic_src(self):
|
||||
page_size = 64
|
||||
dcp_size = 4
|
||||
src_pages = np.arange(4, dtype=np.int32)
|
||||
dst_pages = np.array([7], dtype=np.int32)
|
||||
plan = build_dcp_token_transfer_plan(
|
||||
src_pages,
|
||||
dst_pages,
|
||||
physical_page_size=page_size,
|
||||
dcp_size=dcp_size,
|
||||
dcp_rank=0,
|
||||
num_kv_tokens=256,
|
||||
)
|
||||
raw_src, _ = group_concurrent_contiguous(
|
||||
plan.src_token_indices, plan.dst_token_indices
|
||||
)
|
||||
self.assertEqual(len(raw_src), 64)
|
||||
self.assertTrue(all(len(group) == 1 for group in raw_src))
|
||||
|
||||
packed_src = np.arange(plan.dst_token_indices.size, dtype=np.int64)
|
||||
packed_groups, _ = group_concurrent_contiguous(
|
||||
packed_src, plan.dst_token_indices
|
||||
)
|
||||
self.assertEqual(len(packed_groups), 1)
|
||||
self.assertEqual(len(packed_groups[0]), 64)
|
||||
|
||||
|
||||
class TestDcpPackBufferBytes(CustomTestCase):
|
||||
def test_sizes_fixed_regions_for_each_dcp_rank(self):
|
||||
self.assertEqual(
|
||||
dcp_pack_buffer_bytes(
|
||||
[64 * 16, 64 * 16],
|
||||
page_size=64,
|
||||
max_tokens=10,
|
||||
dcp_size=4,
|
||||
),
|
||||
4 * 3 * (16 + 16),
|
||||
)
|
||||
|
||||
def test_rejects_invalid_item_lens(self):
|
||||
with self.assertRaisesRegex(ValueError, "at least one page"):
|
||||
dcp_pack_buffer_bytes([0], page_size=64, max_tokens=8)
|
||||
with self.assertRaisesRegex(ValueError, "page-aligned"):
|
||||
dcp_pack_buffer_bytes([100], page_size=64, max_tokens=8)
|
||||
|
||||
|
||||
class TestTryDcpPack(CustomTestCase):
|
||||
def test_try_pack_uses_requested_region_and_dense_indices(self):
|
||||
dim = 4
|
||||
kv = torch.arange(16 * dim, dtype=torch.float32).view(16, 1, dim)
|
||||
item_len = int(kv[0].nbytes)
|
||||
pack = torch.zeros(8 * item_len, dtype=torch.uint8)
|
||||
gather_stream = Mock()
|
||||
buf = type(
|
||||
"Buf",
|
||||
(),
|
||||
{
|
||||
"buffer": pack,
|
||||
"fits": lambda self, n: n <= pack.numel(),
|
||||
"get_ptr": lambda self: 0x1000,
|
||||
"get_size": lambda self: pack.numel(),
|
||||
"get_gather_stream": lambda self: gather_stream,
|
||||
},
|
||||
)()
|
||||
src = np.array([1, 5, 9, 13], dtype=np.int64)
|
||||
pack_offset = 2 * item_len
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.disaggregation.common.dcp_pack.torch.cuda.default_stream"
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.disaggregation.common.dcp_pack.torch.cuda.stream",
|
||||
return_value=nullcontext(),
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.disaggregation.common.dcp_pack.copy_mla_rows_into_pack"
|
||||
) as copy_mock,
|
||||
):
|
||||
packed = try_pack_dcp_src(
|
||||
pack_buffer=buf,
|
||||
kv_data_ptrs=[kv.data_ptr()],
|
||||
src_token_indices=src,
|
||||
token_item_lens=[item_len],
|
||||
pack_offset_bytes=pack_offset,
|
||||
)
|
||||
|
||||
gather_stream.synchronize.assert_called_once_with()
|
||||
self.assertIsNotNone(packed)
|
||||
ptrs, indices = packed
|
||||
self.assertEqual(ptrs, [0x1000 + pack_offset])
|
||||
np.testing.assert_array_equal(indices, np.arange(4))
|
||||
pack_view = copy_mock.call_args.args[2]
|
||||
self.assertEqual(pack_view.storage_offset(), pack_offset)
|
||||
self.assertEqual(pack_view.numel(), src.size * item_len)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -535,6 +535,92 @@ class TestNixlTransferWorker(CustomTestCase):
|
||||
self.assertIn(room, mgr.req_to_decode_prefix_len)
|
||||
mgr.send_kvcache.assert_called_once()
|
||||
|
||||
def test_dcp_destinations_use_disjoint_pack_regions_before_chunk_barrier(self):
|
||||
room = 23
|
||||
mgr = self._make_manager(room)
|
||||
agents = ("agent0a", "agent0b", "agent1")
|
||||
dcp_ranks = (0, 0, 1)
|
||||
mgr.transfer_infos[room] = {
|
||||
agent: TransferInfo(
|
||||
room=room,
|
||||
endpoint="127.0.0.1",
|
||||
dst_port=5555 + i,
|
||||
agent_name=agent,
|
||||
dst_kv_indices=np.array([2 + i], dtype=np.int32),
|
||||
dst_aux_index=0,
|
||||
required_dst_info_num=len(agents),
|
||||
dst_state_indices=[],
|
||||
)
|
||||
for i, agent in enumerate(agents)
|
||||
}
|
||||
mgr.decode_kv_args_table = {
|
||||
agent: SimpleNamespace(
|
||||
decode_tp_size=len(agents),
|
||||
dst_kv_ptrs=[0x3000 + i * 0x100],
|
||||
dst_aux_ptrs=[0],
|
||||
gpu_id=0,
|
||||
staging_base_ptr=0,
|
||||
staging_total_size=0,
|
||||
kv_xfer_segments=None,
|
||||
dst_homogeneous_mem_kind="VRAM",
|
||||
requires_dcp_relayout=True,
|
||||
dst_dcp_size=2,
|
||||
dst_dcp_rank=dcp_rank,
|
||||
dcp_dst_region_indices=[0],
|
||||
dcp_token_item_lens=[4],
|
||||
)
|
||||
for i, (agent, dcp_rank) in enumerate(zip(agents, dcp_ranks))
|
||||
}
|
||||
mgr.kv_args = SimpleNamespace(
|
||||
engine_rank=0,
|
||||
kv_data_ptrs=[0x1000],
|
||||
page_size=4,
|
||||
)
|
||||
mgr._dcp_pack_buffers = [SimpleNamespace(get_size=lambda: 16)]
|
||||
|
||||
packed_rank0 = ([0x9000], np.arange(2, dtype=np.int64))
|
||||
packed_rank1 = ([0x9008], np.arange(2, dtype=np.int64))
|
||||
try_pack = MagicMock(side_effect=[packed_rank0, packed_rank1])
|
||||
dcp_pack_module = types.ModuleType("sglang.srt.disaggregation.common.dcp_pack")
|
||||
dcp_pack_module.try_pack_dcp_src = try_pack
|
||||
submitted = []
|
||||
|
||||
def send_kvcache_dcp(*args, **kwargs):
|
||||
submitted.append((args[0], args[-1]))
|
||||
return f"handle-{args[0]}"
|
||||
|
||||
mgr.send_kvcache_dcp = MagicMock(side_effect=send_kvcache_dcp)
|
||||
submitted_counts_at_poll = []
|
||||
|
||||
def check_xfer_state(_handle):
|
||||
submitted_counts_at_poll.append(len(submitted))
|
||||
return "DONE"
|
||||
|
||||
mgr.agent = SimpleNamespace(check_xfer_state=check_xfer_state)
|
||||
chunk = self._make_chunk(room, [1], is_last_chunk=False)
|
||||
chunk.num_kv_tokens = 4
|
||||
|
||||
with patch.dict(
|
||||
sys.modules,
|
||||
{"sglang.srt.disaggregation.common.dcp_pack": dcp_pack_module},
|
||||
):
|
||||
self._run_worker_once(mgr, chunk)
|
||||
|
||||
self.assertEqual(try_pack.call_count, 2)
|
||||
self.assertEqual(
|
||||
[call.kwargs["pack_offset_bytes"] for call in try_pack.call_args_list],
|
||||
[0, 8],
|
||||
)
|
||||
self.assertEqual(
|
||||
submitted,
|
||||
[
|
||||
("agent0a", packed_rank0),
|
||||
("agent0b", packed_rank0),
|
||||
("agent1", packed_rank1),
|
||||
],
|
||||
)
|
||||
self.assertEqual(submitted_counts_at_poll, [3, 3, 3])
|
||||
|
||||
|
||||
class TestNixlNotifications(CustomTestCase):
|
||||
def _make_manager(self, messages, required=None):
|
||||
@@ -788,16 +874,16 @@ class TestNixlStaging(CustomTestCase):
|
||||
agent = StagingFakeAgent(register_result=["staging"])
|
||||
mgr = self._make_manager(agent)
|
||||
|
||||
mgr._register_staging_memory(0x1000, 4096, 3)
|
||||
mgr._register_staging_memory(0x1000, 4096)
|
||||
|
||||
self.assertEqual(
|
||||
agent.register_memory_calls,
|
||||
[([(0x1000, 4096, 3, "")], "VRAM")],
|
||||
[([(0x1000, 4096, 1, "")], "VRAM")],
|
||||
)
|
||||
|
||||
mgr = self._make_manager(StagingFakeAgent(register_result=[]))
|
||||
with self.assertRaisesRegex(RuntimeError, "staging buffer"):
|
||||
mgr._register_staging_memory(0x1000, 4096, 3)
|
||||
mgr._register_staging_memory(0x1000, 4096)
|
||||
|
||||
def test_prefetch_staging_reqs_noops_when_disabled_or_missing_kv_buffers(self):
|
||||
mgr = self._make_manager()
|
||||
|
||||
Reference in New Issue
Block a user