diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 69d07fb03..7ec83b677 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -1249,6 +1249,32 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): page_indices = kv_to_page_indices(kv_indices, kv_transfer_page_size).astype( np.int32 ) + device_page_indices = None + if ( + self.scheduler.enable_hisparse + and isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool) + and not _is_fake_transfer(decode_req.req, self.scheduler.server_args) + ): + # alloc_logical_only() already allocated the shared logical pages + # used by C4 indexer and C128 KV. These device buffers do not use + # the C4 sparse physical-slot mapping; carry their logical page IDs + # alongside the independently allocated C4 host page IDs. + full_kv_indices = self.req_to_token_pool.req_to_token[ + decode_req.req.req_pool_idx, + prefix_len:origin_input_len, + ] + device_page_indices = kv_to_page_indices( + full_kv_indices, + page_size, + ).astype(np.int32) + if self.transfer_backend != TransferBackend.MOONCAKE: + raise NotImplementedError( + "DSV4 HiSparse direct PD transfer currently requires " + "the Mooncake backend" + ) + metadata_kwargs = {"decode_prefix_len": total_prefix_len} + if device_page_indices is not None: + metadata_kwargs["device_kv_indices"] = device_page_indices if ( self.transfer_queue.enable_staging and hasattr(decode_req.kv_receiver, "require_staging") @@ -1263,7 +1289,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): page_indices, decode_req.metadata_buffer_index, state_indices, - decode_prefix_len=total_prefix_len, + **metadata_kwargs, ) if decode_req.is_rebootstrap: self.kv_manager.submit_prefill_recompute( diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 87cb6b45e..782661658 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -85,6 +85,7 @@ class TransferInfo: required_dst_info_num: int is_dummy: bool decode_prefix_len: Optional[int] = None + dst_device_kv_indices: Optional[npt.NDArray[np.int32]] = None # Note: always put the optional staging field at the final (it will be set through 'STAGING_RSP' pkg when needed) staging: Optional[StagingTransferInfo] = None @@ -113,6 +114,11 @@ class TransferInfo: decode_prefix_len=( int(msg[8].decode("ascii")) if len(msg) > 8 and msg[8] != b"" else None ), + dst_device_kv_indices=( + np.frombuffer(msg[9], dtype=np.int32) + if len(msg) > 9 and msg[9] != b"" + else None + ), ) @@ -623,6 +629,8 @@ class MooncakeKVManager(CommonKVManager): force_flat: bool = False, src_layer_ids: Optional[List[int]] = None, dst_layer_ids: Optional[List[int]] = None, + dst_device_data_indices: Optional[npt.NDArray[np.int32]] = None, + dst_device_data_ptrs: Optional[set[int]] = None, ) -> int: """ Generic KV cache transfer supporting both MHA and MLA architectures. @@ -632,10 +640,18 @@ class MooncakeKVManager(CommonKVManager): even on a non-MLA backend, for K-only state buffers (e.g. MiniMax sparse index) whose per-layer list must not be half-split into K/V. """ - # Group by indices for optimization + # Host and device buffers may use different destination page spaces. + # Build both transfer plans once, then select per destination buffer. prefill_kv_blocks, dst_kv_blocks = group_concurrent_contiguous( prefill_data_indices, dst_data_indices ) + device_prefill_kv_blocks = device_dst_kv_blocks = None + if dst_device_data_indices is not None: + device_prefill_kv_blocks, device_dst_kv_blocks = ( + group_concurrent_contiguous( + prefill_data_indices, dst_device_data_indices + ) + ) layers_params = None @@ -701,7 +717,18 @@ class MooncakeKVManager(CommonKVManager): src_ptr: int, dst_ptr: int, item_len: int ) -> List[Tuple[int, int, int]]: transfer_blocks = [] - for prefill_index, decode_index in zip(prefill_kv_blocks, dst_kv_blocks): + if dst_device_data_ptrs and int(dst_ptr) in dst_device_data_ptrs: + assert ( + device_prefill_kv_blocks is not None + and device_dst_kv_blocks is not None + ) + src_blocks, dst_blocks = ( + device_prefill_kv_blocks, + device_dst_kv_blocks, + ) + else: + src_blocks, dst_blocks = prefill_kv_blocks, dst_kv_blocks + for prefill_index, decode_index in zip(src_blocks, dst_blocks): src_addr = src_ptr + int(prefill_index[0]) * item_len dst_addr = dst_ptr + int(decode_index[0]) * item_len length = item_len * len(prefill_index) @@ -750,7 +777,19 @@ class MooncakeKVManager(CommonKVManager): dst_kv_indices: npt.NDArray[np.int32], executor: concurrent.futures.ThreadPoolExecutor, dst_layer_ids: Optional[List[int]] = None, + dst_device_kv_indices: Optional[npt.NDArray[np.int32]] = None, ): + dst_device_kv_ptrs = None + if dst_device_kv_indices is not None: + compression_ratios = self.kv_args.mla_compression_ratios + assert compression_ratios is not None + if len(dst_kv_ptrs) == len(self.kv_args.kv_data_ptrs): + start = self.kv_args.prefill_start_layer + end = self.kv_args.prefill_end_layer + assert end is not None + compression_ratios = compression_ratios[start:end] + c4_layer_num = sum(ratio == 4 for ratio in compression_ratios) + dst_device_kv_ptrs = set(dst_kv_ptrs[c4_layer_num:]) return self._send_kvcache_generic( mooncake_session_id=mooncake_session_id, src_data_ptrs=self.kv_args.kv_data_ptrs, @@ -761,6 +800,8 @@ class MooncakeKVManager(CommonKVManager): executor=executor, src_layer_ids=self.kv_args.kv_layer_ids, dst_layer_ids=dst_layer_ids, + dst_device_data_indices=dst_device_kv_indices, + dst_device_data_ptrs=dst_device_kv_ptrs, ) def send_kvcache_dcp( @@ -1557,12 +1598,22 @@ class MooncakeKVManager(CommonKVManager): is_dcp_transfer = ( target_rank_registration_info.requires_dcp_relayout ) + chunked_dst_device_kv_indice = None if is_dcp_transfer: + if req.dst_device_kv_indices is not None: + raise RuntimeError( + "HiSparse destination device indices are not " + "supported by PD DCP relayout" + ) chunked_dst_kv_indice = req.dst_kv_indices else: chunked_dst_kv_indice = req.dst_kv_indices[ kv_chunk.index_slice ] + if req.dst_device_kv_indices is not None: + chunked_dst_device_kv_indice = ( + req.dst_device_kv_indices[kv_chunk.index_slice] + ) # NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices # is mismatched with the dst_kv_indices when page size > 1, this should never happen. @@ -1577,6 +1628,12 @@ class MooncakeKVManager(CommonKVManager): : len(chunked_dst_kv_indice) ] ) + if chunked_dst_device_kv_indice is not None: + chunked_dst_device_kv_indice = ( + chunked_dst_device_kv_indice[ + : len(kv_chunk.prefill_kv_indices) + ] + ) skip_kv, skip_state = self._get_dsa_cache_transfer_skip_flags( target_rank_registration_info @@ -1620,7 +1677,8 @@ class MooncakeKVManager(CommonKVManager): target_rank_registration_info.dst_kv_ptrs, chunked_dst_kv_indice, executor, - target_rank_registration_info.dst_kv_layer_ids, + dst_layer_ids=target_rank_registration_info.dst_kv_layer_ids, + dst_device_kv_indices=chunked_dst_device_kv_indice, ) elif ( self.enable_staging @@ -2244,6 +2302,7 @@ class MooncakeKVReceiver(CommonKVReceiver): aux_index: Optional[int] = None, state_indices: Optional[List] = None, decode_prefix_len: Optional[int] = None, + device_kv_indices: Optional[npt.NDArray[np.int32]] = None, ): if self.bootstrap_infos is None: self.kv_mgr.record_failure( @@ -2282,6 +2341,11 @@ class MooncakeKVReceiver(CommonKVReceiver): ), str(self.required_dst_info_num).encode("ascii"), str(decode_prefix_len or 0).encode("ascii"), + ( + np.asarray(device_kv_indices, dtype=np.int32).tobytes() + if not is_dummy and device_kv_indices is not None + else b"" + ), ] ) except zmq.ZMQError: diff --git a/test/registered/unit/mem_cache/test_hisparse_allocator.py b/test/registered/unit/mem_cache/test_hisparse_allocator.py index 592a5d528..bbf0428e1 100644 --- a/test/registered/unit/mem_cache/test_hisparse_allocator.py +++ b/test/registered/unit/mem_cache/test_hisparse_allocator.py @@ -1,7 +1,9 @@ import unittest +from concurrent.futures import ThreadPoolExecutor from types import SimpleNamespace from unittest.mock import MagicMock +import numpy as np import torch from sglang.srt.mem_cache.allocator.hisparse import ( @@ -78,8 +80,8 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase): fill_len = 512 swa_tail_len = 128 - kv_loc = torch.arange(fill_len, dtype=torch.int64) - host_indices = torch.arange(1000, 1000 + fill_len, dtype=torch.int64) + kv_loc = torch.arange(512, 512 + fill_len, dtype=torch.int64) + host_indices = torch.arange(1000, 1128, dtype=torch.int64) req = SimpleNamespace( rid="req-0", @@ -108,18 +110,17 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase): req_to_token_pool = ReqToTokenPool() allocator = SimpleNamespace( device=torch.device("cpu"), - page_size=64, + page_size=256, available_size=MagicMock(return_value=fill_len), alloc_extend_swa_tail=MagicMock(return_value=kv_loc), alloc_logical_only=MagicMock(return_value=kv_loc), ) + regular_host_alloc = MagicMock(return_value=host_indices) coordinator = SimpleNamespace( - mem_pool_host=SimpleNamespace( - alloc_paged_token_slots=MagicMock(return_value=host_indices) - ), + mem_pool_host=SimpleNamespace(alloc_paged_token_slots=regular_host_alloc), req_to_host_pool=object(), req_to_host_pool_allocated_len=object(), - host_token_len=MagicMock(side_effect=lambda length: length), + host_token_len=MagicMock(side_effect=lambda token_len: token_len // 4), ) queue = DecodePreallocQueue.__new__(DecodePreallocQueue) queue.req_to_token_pool = req_to_token_pool @@ -138,7 +139,7 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase): result = queue._pre_alloc(req) - self.assertIs(result, host_indices) + self.assertTrue(torch.equal(result, host_indices)) allocator.alloc_extend_swa_tail.assert_called_once() allocator.alloc_logical_only.assert_not_called() _, kwargs = allocator.alloc_extend_swa_tail.call_args @@ -149,7 +150,104 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase): self.assertEqual(req.kv_committed_len, fill_len) self.assertEqual(req.extend_range.length, fill_len) self.assertEqual(len(req_to_token_pool.writes), 1) - coordinator.mem_pool_host.alloc_paged_token_slots.assert_called_once() + coordinator.host_token_len.assert_called_once_with(fill_len) + regular_host_alloc.assert_called_once_with( + coordinator.req_to_host_pool, + coordinator.req_to_host_pool_allocated_len, + req.req_pool_idx, + 0, + len(host_indices), + ) + self.assertTrue(torch.equal(req_to_token_pool.writes[0][1], kv_loc)) + + # C4 indexer/C128 use the logical allocator's full-page IDs. They do not + # use either the independently allocated host pages or C4 sparse slots. + np.testing.assert_array_equal( + np.unique(kv_loc.numpy() // allocator.page_size), + np.array([2, 3]), + ) + + def test_mooncake_uses_separate_host_and_device_page_indices(self): + from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager + + manager = object.__new__(MooncakeKVManager) + manager.is_mla_backend = True + manager.is_hybrid_mla_backend = False + manager.enable_custom_mem_pool = False + manager._transfer_data = MagicMock(return_value=0) + + with ThreadPoolExecutor(max_workers=1) as executor: + ret = manager._send_kvcache_generic( + mooncake_session_id="session", + src_data_ptrs=[1000, 2000, 3000], + dst_data_ptrs=[10000, 20000, 30000], + item_lens=[100, 100, 100], + prefill_data_indices=np.array([1, 2], dtype=np.int32), + dst_data_indices=np.array([7, 8], dtype=np.int32), + executor=executor, + dst_device_data_indices=np.array([21, 22], dtype=np.int32), + dst_device_data_ptrs={20000, 30000}, + ) + + self.assertEqual(ret, 0) + manager._transfer_data.assert_called_once_with( + "session", + [ + (1100, 10700, 200), + (2100, 22100, 200), + (3100, 32100, 200), + ], + ) + + def test_mooncake_derives_device_buffers_from_local_pp_layout(self): + from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager + + manager = object.__new__(MooncakeKVManager) + manager.kv_args = SimpleNamespace( + kv_data_ptrs=[1000, 2000, 3000], + kv_item_lens=[100, 100, 100], + kv_layer_ids=[], + mla_compression_ratios=[4, 128, 4, 128], + prefill_start_layer=0, + prefill_end_layer=2, + ) + manager._send_kvcache_generic = MagicMock(return_value=0) + executor = MagicMock() + + manager.send_kvcache( + "session", + np.array([1], dtype=np.int32), + [10000, 20000, 30000], + np.array([7], dtype=np.int32), + executor, + dst_device_kv_indices=np.array([21], dtype=np.int32), + ) + + kwargs = manager._send_kvcache_generic.call_args.kwargs + self.assertEqual(kwargs["dst_device_data_ptrs"], {20000, 30000}) + + def test_mooncake_transfer_metadata_carries_device_page_indices(self): + from sglang.srt.disaggregation.mooncake.conn import TransferInfo + + host_pages = np.array([7, 8], dtype=np.int32) + device_pages = np.array([21, 22], dtype=np.int32) + info = TransferInfo.from_zmq( + [ + b"9", + b"127.0.0.1", + b"12345", + b"session", + host_pages.tobytes(), + b"0", + b"", + b"1", + b"0", + device_pages.tobytes(), + ] + ) + + np.testing.assert_array_equal(info.dst_kv_indices, host_pages) + np.testing.assert_array_equal(info.dst_device_kv_indices, device_pages) if __name__ == "__main__":