diff --git a/python/sglang/kernels/ops/kvcache/pd_dcp_gather.py b/python/sglang/kernels/ops/kvcache/pd_dcp_gather.py index e53dcafc6..a15d37f01 100644 --- a/python/sglang/kernels/ops/kvcache/pd_dcp_gather.py +++ b/python/sglang/kernels/ops/kvcache/pd_dcp_gather.py @@ -1,4 +1,4 @@ -from typing import Sequence +from typing import Optional, Sequence import torch import triton @@ -15,10 +15,11 @@ def _copy_mla_rows_into_pack_kernel( ): layer_id = tl.program_id(0) block_id = tl.program_id(1) - metadata_offset = layer_id * 3 + metadata_offset = layer_id * 4 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) + src_row_stride = tl.load(src_metadata + metadata_offset + 3) offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) layer_nbytes = num_rows * row_nbytes @@ -26,7 +27,7 @@ def _copy_mla_rows_into_pack_kernel( 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) + values = tl.load(src + src_row * src_row_stride + byte, mask=mask) tl.store(pack + pack_offset + offsets, values, mask=mask) @@ -35,11 +36,14 @@ def copy_mla_rows_into_pack( row_indices: torch.Tensor, pack: torch.Tensor, token_item_lens: Sequence[int], + src_token_item_lens: Optional[Sequence[int]] = None, ) -> None: - if len(kv_data_ptrs) != len(token_item_lens): + if src_token_item_lens is None: + src_token_item_lens = token_item_lens + if not (len(kv_data_ptrs) == len(token_item_lens) == len(src_token_item_lens)): raise ValueError( - "kv_data_ptrs and token_item_lens length mismatch: " - f"{len(kv_data_ptrs)} vs {len(token_item_lens)}" + "KV pointers, copy widths, and source strides length mismatch: " + f"{len(kv_data_ptrs)}, {len(token_item_lens)}, {len(src_token_item_lens)}" ) if not kv_data_ptrs: return @@ -47,11 +51,13 @@ def copy_mla_rows_into_pack( n = int(row_indices.numel()) metadata = [] offset = 0 - for ptr, item_len in zip(kv_data_ptrs, token_item_lens): + for ptr, item_len, src_item_len in zip( + kv_data_ptrs, token_item_lens, src_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)) + metadata.extend((int(ptr), item_len, offset, int(src_item_len))) offset += n * item_len src_metadata = torch.tensor(metadata, dtype=torch.int64, device=pack.device) diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 526edf2ac..b1a999036 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -382,7 +382,9 @@ class CommonKVManager(BaseKVManager): f"{type(self).__name__} does not support staging memory registration" ) - def _init_dcp_pack_buffers_once(self, dcp_size: int) -> None: + def _init_dcp_pack_buffers_once( + self, dcp_size: int, *, include_draft: bool = False + ) -> None: if self._dcp_pack_buffers is not None: return if not self.kv_args.kv_item_lens: @@ -397,6 +399,7 @@ class CommonKVManager(BaseKVManager): len(self.transfer_queues), dcp_size, max_tokens, + include_draft=include_draft, ) self._dcp_pack_max_tokens = max_tokens diff --git a/python/sglang/srt/disaggregation/common/dcp_pack.py b/python/sglang/srt/disaggregation/common/dcp_pack.py index 1d5ae6ab1..122dba6dd 100644 --- a/python/sglang/srt/disaggregation/common/dcp_pack.py +++ b/python/sglang/srt/disaggregation/common/dcp_pack.py @@ -41,6 +41,7 @@ def try_pack_dcp_src( kv_data_ptrs: Sequence[int], src_token_indices: npt.NDArray[np.integer], token_item_lens: Sequence[int], + src_token_item_lens: Optional[Sequence[int]] = None, pack_offset_bytes: int = 0, pack_capacity_bytes: Optional[int] = None, ) -> Optional[Tuple[List[int], npt.NDArray[np.int64]]]: @@ -75,7 +76,9 @@ def try_pack_dcp_src( 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) + copy_mla_rows_into_pack( + kv_data_ptrs, row_indices, pack, token_item_lens, src_token_item_lens + ) gather_stream.synchronize() packed_ptrs: List[int] = [] @@ -93,14 +96,16 @@ def init_dcp_pack_buffers( count: int, dcp_size: int, max_tokens: int, + *, + include_draft: bool = False, ) -> List[StagingBuffer]: from sglang.srt.disaggregation.common.staging_handler import ( _get_custom_mem_pool, ) kv_item_lens = kv_args.kv_item_lens - if kv_args.num_draft_entries > 0: - kv_item_lens = kv_item_lens[: len(kv_item_lens) - kv_args.num_draft_entries] + if not include_draft and kv_args.num_draft_entries: + kv_item_lens = kv_item_lens[: -kv_args.num_draft_entries] # 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. diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 25cbf4c80..9d1730bfe 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -1173,6 +1173,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): for entry in range(num_target) ] sliced_draft_params = [] + draft_to_pack = [] if num_draft > 0 and plan.draft_src_token_indices.size: if not dst_kv_item_lens and dst_attn_tp_size not in ( None, @@ -1224,14 +1225,45 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): dst_rank = dst_tp_rank // max(1, dst_span // src_span) src_offset = (dst_rank * dst_width) % src_width dst_offset = (src_rank * src_width) % dst_width - sliced_draft_params.append( - ( - src_kv_ptrs[entry] + src_offset, - dst_kv_ptrs[entry] + dst_offset, - src_width, - dst_width, - copy_width, - ) + params = ( + src_kv_ptrs[entry] + src_offset, + dst_kv_ptrs[entry] + dst_offset, + src_width, + dst_width, + copy_width, + ) + if pack_buffer is not None and src_width > dst_width: + draft_to_pack.append(params) + else: + sliced_draft_params.append(params) + + if draft_to_pack: + from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src + + draft_src_ptrs, draft_dst_ptrs, src_strides, _, copy_widths = zip( + *draft_to_pack + ) + target_pack_bytes = plan.target_src_token_indices.size * sum( + dcp_token_item_lens[:num_target] + ) + packed = try_pack_dcp_src( + pack_buffer=pack_buffer, + kv_data_ptrs=draft_src_ptrs, + src_token_indices=plan.draft_src_token_indices, + token_item_lens=copy_widths, + src_token_item_lens=src_strides, + pack_offset_bytes=target_pack_bytes, + ) + if packed is None: + sliced_draft_params.extend(draft_to_pack) + else: + packed_ptrs, packed_indices = packed + packed_groups = group_concurrent_contiguous( + packed_indices, plan.draft_dst_token_indices + ) + layers_params.extend( + (src, dst, width, packed_groups) + for src, dst, width in zip(packed_ptrs, draft_dst_ptrs, copy_widths) ) def process_sliced_draft(params) -> int: @@ -1284,7 +1316,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): executor.submit(process_sliced_draft, [params]) for params in sliced_draft_params ) - return self._await_transfer_futures(futures) + try: + return self._await_transfer_futures(futures) + finally: + if pack_buffer is not None: + concurrent.futures.wait(futures) transfer_blocks = [] for layer_params in layers_params: @@ -2492,7 +2528,9 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager): decode_kv_args.dst_dcp_size, ) ) - self._init_dcp_pack_buffers_once(decode_kv_args.dst_dcp_size) + self._init_dcp_pack_buffers_once( + decode_kv_args.dst_dcp_size, include_draft=True + ) 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/test/registered/kernels/ops/kvcache/test_pd_dcp_gather.py b/test/registered/kernels/ops/kvcache/test_pd_dcp_gather.py index 5f27e85b3..723916e35 100644 --- a/test/registered/kernels/ops/kvcache/test_pd_dcp_gather.py +++ b/test/registered/kernels/ops/kvcache/test_pd_dcp_gather.py @@ -1,8 +1,13 @@ +import concurrent.futures import unittest +from types import SimpleNamespace +import numpy as np 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.disaggregation.mooncake.conn import MooncakeKVManager from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase @@ -36,6 +41,135 @@ class TestPdDcpGather(CustomTestCase): torch.testing.assert_close(packed0, kv0[row_indices], rtol=0, atol=0) torch.testing.assert_close(packed1, kv1[row_indices], rtol=0, atol=0) + def test_packed_tp2_pp2_to_dcp4_preserves_kv(self): + """Packing must preserve both target rows and draft head shards across PP stages.""" + for custom_pool in (False, True): + for capacity in (256 * (2 * 64 + 2 * 512), 256 * 2 * 64 // 4): + for rank in range(4): + with self.subTest( + custom_pool=custom_pool, capacity=capacity, rank=rank + ): + self._check_packed_transfer(rank, custom_pool, capacity) + + def _check_packed_transfer(self, rank, custom_pool, capacity): + page, tokens, chunk = 64, 521, 256 + src_pages = np.array([7, 1, 9, 3, 4, 11, 2, 5, 8], dtype=np.int32) + dst_pages = np.array([4, 1, 6], dtype=np.int32) + layers, widths = [3, 11, 19, 27, 28, 28], [64] * 4 + [256] * 2 + logical = torch.arange(tokens, device="cuda") + src_rows = ( + torch.as_tensor(src_pages, device="cuda")[logical // page] * page + + logical % page + ) + values = [ + ( + (logical[:, None] + 256) * 13 + + torch.arange(width, device="cuda") * 7 + + entry * 31 + ) + .remainder(251) + .to(torch.uint8) + for entry, width in enumerate([64] * 4 + [1024] * 2) + ] + destinations = [ + torch.full((2048, w), 165, dtype=torch.uint8, device="cuda") for w in widths + ] + expected = [x.clone() for x in destinations] + owned = logical[rank::4] + target_rows = ( + torch.as_tensor(dst_pages, device="cuda")[owned // 256] * page + + owned % 256 // 4 + ) + draft_rows = ( + torch.as_tensor(dst_pages, device="cuda")[logical // 256] * 256 + + logical % 256 + ) + for entry in range(4): + expected[entry][target_rows] = values[entry][owned] + for entry in (4, 5): + expected[entry][draft_rows] = values[entry][ + :, rank * 256 : (rank + 1) * 256 + ] + + pack = StagingBuffer(capacity, "cuda:0", 0) + with concurrent.futures.ThreadPoolExecutor(max_workers=3) as executor: + for stage, entries in enumerate(([0, 1], [2, 3, 4, 5])): + sources = [] + for entry in entries: + data = values[entry] + if entry >= 4: + start = (rank // 2) * 512 + data = data[:, start : start + 512] + source = torch.full( + (1024, data.shape[1]), 165, dtype=torch.uint8, device="cuda" + ) + source[src_rows] = data + sources.append(source) + buffers = sources + destinations + [pack.buffer] + + def transfer(session, blocks): + def view(ptr, size): + for tensor in buffers: + offset = ptr - tensor.data_ptr() + if 0 <= offset and offset + size <= tensor.numel(): + return tensor.flatten()[offset : offset + size] + raise AssertionError( + f"Transfer outside registered buffers: {ptr}, {size}" + ) + + for src, dst, size in blocks: + view(dst, size).copy_(view(src, size)) + torch.cuda.synchronize() + return 0 + + manager = SimpleNamespace( + is_mla_backend=False, + kv_args=SimpleNamespace( + page_size=page, + kv_layer_ids=[layers[e] for e in entries], + kv_data_ptrs=[x.data_ptr() for x in sources], + num_draft_entries=2 if stage else 0, + engine_rank=stage * 2 + rank // 2, + ), + attn_tp_size=2, + max_transfer_batch_indices=37, + enable_custom_mem_pool=custom_pool, + enable_deferred_decode_kv_release=False, + _transfer_data=transfer, + ) + manager._await_transfer_futures = lambda futures: ( + MooncakeKVManager._await_transfer_futures(manager, futures) + ) + for start in range(0, tokens, chunk): + count = min(chunk, tokens - start) + result = MooncakeKVManager.send_kvcache_dcp( + manager, + "session", + src_pages[start // page : (start + count + page - 1) // page], + [x.data_ptr() for x in destinations], + dst_pages, + dcp_token_item_lens=[x.shape[1] for x in sources], + dst_dcp_size=4, + dst_dcp_rank=rank, + src_page_offset=start // page, + decode_prefix_len=256, + num_kv_tokens=count, + executor=executor, + dst_layer_ids=layers, + pack_buffer=pack, + dst_kv_item_lens=[ + page * w * (4 if e >= 4 else 1) + for e, w in enumerate(widths) + ], + dst_tp_rank=rank, + dst_attn_tp_size=4, + ) + self.assertEqual(result, 0) + for entry in entries: + torch.testing.assert_close( + destinations[entry], expected[entry], rtol=0, atol=0 + ) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/disaggregation/test_mooncake_transfer_batching.py b/test/registered/unit/disaggregation/test_mooncake_transfer_batching.py index c9cb9e5c6..bc2b34de2 100644 --- a/test/registered/unit/disaggregation/test_mooncake_transfer_batching.py +++ b/test/registered/unit/disaggregation/test_mooncake_transfer_batching.py @@ -1,12 +1,14 @@ import concurrent.futures import unittest +from threading import Event from types import SimpleNamespace -from unittest.mock import MagicMock, call +from unittest.mock import MagicMock, call, patch import numpy as np from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=1, suite="base-a-test-cpu") @@ -286,5 +288,67 @@ class TestDcpDraftHeadTransfer(unittest.TestCase): self.assertFalse(dst_buffers[1000000].any()) +class TestDcpPackLifetime(CustomTestCase): + def test_failed_transfer_drains_before_pack_buffer_reuse(self): + """A failed layer must not release the pack buffer while another transfer reads it.""" + manager = TestMooncakeTransferBatching._make_manager( + enable_custom_mem_pool=True + ) + manager.kv_args = SimpleNamespace( + page_size=1, kv_layer_ids=[], kv_data_ptrs=[1000, 2000], num_draft_entries=0 + ) + source = np.array([11], dtype=np.uint8) + observed = [] + running, release = Event(), Event() + + def transfer(session, blocks): + if blocks[0][0] == 1000: + self.assertTrue(running.wait(10)) + return 17 + running.set() + self.assertTrue(release.wait(10)) + observed.append(int(source[0])) + return 0 + + def send(executor): + result = MooncakeKVManager.send_kvcache_dcp( + manager, + "session", + np.array([0, 1], dtype=np.int32), + [5000, 6000], + np.array([0], dtype=np.int32), + dcp_token_item_lens=[1, 1], + dst_dcp_size=2, + dst_dcp_rank=0, + src_page_offset=0, + decode_prefix_len=0, + num_kv_tokens=2, + executor=executor, + dst_layer_ids=[], + pack_buffer=object(), + ) + source[0] = 22 + return result + + manager._transfer_data = transfer + with ( + patch( + "sglang.srt.disaggregation.common.dcp_pack.try_pack_dcp_src", + return_value=([1000, 2000], np.array([0], dtype=np.int64)), + ), + concurrent.futures.ThreadPoolExecutor(max_workers=2) as transfers, + concurrent.futures.ThreadPoolExecutor(max_workers=1) as worker, + ): + future = worker.submit(send, transfers) + try: + self.assertTrue(running.wait(10)) + with self.assertRaises(concurrent.futures.TimeoutError): + future.result(timeout=1) + finally: + release.set() + self.assertEqual(future.result(timeout=10), 17) + self.assertEqual(observed, [11]) + + if __name__ == "__main__": unittest.main()