[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:
Khoa Pham
2026-08-28 21:41:30 -07:00
committed by GitHub
co-authored by Cursor Claude Opus 5
parent 5f216fc33f
commit 3760296be8
10 changed files with 581 additions and 60 deletions
@@ -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:
+81 -37
View File
@@ -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()