diff --git a/python/sglang/kernels/ops/kvcache/__init__.py b/python/sglang/kernels/ops/kvcache/__init__.py index c04e8212c..40cbb6f98 100644 --- a/python/sglang/kernels/ops/kvcache/__init__.py +++ b/python/sglang/kernels/ops/kvcache/__init__.py @@ -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"), diff --git a/python/sglang/kernels/ops/kvcache/pd_dcp_gather.py b/python/sglang/kernels/ops/kvcache/pd_dcp_gather.py new file mode 100644 index 000000000..e53dcafc6 --- /dev/null +++ b/python/sglang/kernels/ops/kvcache/pd_dcp_gather.py @@ -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, + ) diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 7d1ce18ef..36955e72a 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -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] diff --git a/python/sglang/srt/disaggregation/common/dcp_pack.py b/python/sglang/srt/disaggregation/common/dcp_pack.py new file mode 100644 index 000000000..db0a6c534 --- /dev/null +++ b/python/sglang/srt/disaggregation/common/dcp_pack.py @@ -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 diff --git a/python/sglang/srt/disaggregation/common/staging_buffer.py b/python/sglang/srt/disaggregation/common/staging_buffer.py index 0824af25a..824a34bbb 100644 --- a/python/sglang/srt/disaggregation/common/staging_buffer.py +++ b/python/sglang/srt/disaggregation/common/staging_buffer.py @@ -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 diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 7eac514df..202d014e2 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -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: diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index fb45a2220..157c88356 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -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, diff --git a/test/registered/kernels/ops/kvcache/test_pd_dcp_gather.py b/test/registered/kernels/ops/kvcache/test_pd_dcp_gather.py new file mode 100644 index 000000000..5f27e85b3 --- /dev/null +++ b/test/registered/kernels/ops/kvcache/test_pd_dcp_gather.py @@ -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() diff --git a/test/registered/unit/disaggregation/test_dcp_pack.py b/test/registered/unit/disaggregation/test_dcp_pack.py new file mode 100644 index 000000000..9bd7d9b6e --- /dev/null +++ b/test/registered/unit/disaggregation/test_dcp_pack.py @@ -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() diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py index 6e6ab2238..b5f8c8e5c 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -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()