diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 401e1f547..3cb29e016 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -1300,7 +1300,9 @@ class _DeepSeekV4Strategy(StackStrategy): _build_deepseek_v4_device_pool_group, ) - return _build_deepseek_v4_device_pool_group(kvcache, page_size) + return _build_deepseek_v4_device_pool_group( + kvcache, page_size, params.mtp_draft_device_pools + ) def build( self, @@ -1572,7 +1574,9 @@ class _DsaStrategy(StackStrategy): _build_dsa_device_pool_group, ) - return _build_dsa_device_pool_group(kvcache, page_size) + return _build_dsa_device_pool_group( + kvcache, page_size, params.mtp_draft_device_pools + ) def build( self, diff --git a/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py index 5ccb585b7..f7973f0d6 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py @@ -26,7 +26,7 @@ class DevicePoolEntry: indices_from_pool: PoolName, device_pool: Any, components: Sequence[Sequence[torch.Tensor]], - layer_mapping: dict[int, int], + layer_mapping: dict[int, int | Sequence[int]], page_size: int, rows_are_pages: bool, packed: bool = True, @@ -127,14 +127,16 @@ class DevicePoolEntry: return self._rows(indices) def get_prepared_layer_range_meta(self, locations: list[int], layer: int): - buffer_index = self.layer_mapping.get(layer) - if buffer_index is None: + mapped = self.layer_mapping.get(layer) + if mapped is None: return None + buffer_indices = [mapped] if isinstance(mapped, int) else list(mapped) items = [] for component, offsets in zip(self.buffer_meta, self._component_offsets): - base_ptr, row_stride, size = component[buffer_index] - items.append((base_ptr, row_stride, size, offsets[buffer_index])) + for buffer_index in buffer_indices: + base_ptr, row_stride, size = component[buffer_index] + items.append((base_ptr, row_stride, size, offsets[buffer_index])) ptrs, sizes, offsets = [], [], [] for row in locations: @@ -234,8 +236,28 @@ def _deepseek_v4_state_views(state_pools: list[Any], global_layers: list[int]): return views +def _with_packed_draft_mapping( + layer_mapping: dict[int, int], + *, + target_device_layer_num: int, + draft_layer_num: int, +) -> dict[int, int | tuple[int, ...]]: + """Attach draft depth N to the same transfer layer as target layer N.""" + if draft_layer_num > len(layer_mapping): + raise ValueError( + "Packed draft layers exceed the target transfer layer count: " + f"{draft_layer_num} > {len(layer_mapping)}." + ) + result: dict[int, int | tuple[int, ...]] = dict(layer_mapping) + for depth in range(draft_layer_num): + result[depth] = (layer_mapping[depth], target_device_layer_num + depth) + return result + + def _build_deepseek_v4_device_pool_group( - kvcache: Any, page_size: int + kvcache: Any, + page_size: int, + mtp_draft_device_pools: tuple[Any, ...] = (), ) -> DevicePoolGroup: from sglang.srt.mem_cache.deepseek_v4_memory_pool import HiSparseC4DevicePool from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( @@ -256,13 +278,23 @@ def _build_deepseek_v4_device_pool_group( "DeepSeek V4 SWA page size must match the tree page size: " f"{kvcache.swa_page_size} != {page_size}." ) + draft_swa_buffers = [ + buffer + for pool in mtp_draft_device_pools + for buffer in pool.swa_kv_pool.kv_buffer + ] + swa_mapping = _with_packed_draft_mapping( + mappings.swa, + target_device_layer_num=len(kvcache.swa_kv_pool.kv_buffer), + draft_layer_num=len(draft_swa_buffers), + ) entries.append( DevicePoolEntry( name=PoolName.SWA, indices_from_pool=PoolName.SWA, device_pool=kvcache.swa_kv_pool, - components=[kvcache.swa_kv_pool.kv_buffer], - layer_mapping=mappings.swa, + components=[[*kvcache.swa_kv_pool.kv_buffer, *draft_swa_buffers]], + layer_mapping=swa_mapping, page_size=page_size, rows_are_pages=True, ) @@ -336,21 +368,41 @@ def _build_deepseek_v4_device_pool_group( ) -def _build_dsa_device_pool_group(kvcache: Any, page_size: int) -> DevicePoolGroup: +def _build_dsa_device_pool_group( + kvcache: Any, + page_size: int, + mtp_draft_device_pools: tuple[Any, ...] = (), +) -> DevicePoolGroup: if kvcache.page_size != page_size: raise ValueError( "DSA KV page size must match the tree page size: " f"{kvcache.page_size} != {page_size}." ) num_layers = kvcache.layer_num - identity = {layer: layer for layer in range(num_layers)} + if any(pool.page_size != page_size for pool in mtp_draft_device_pools): + raise ValueError("DSA MTP page size must match the tree page size.") + draft_kv_buffers = [ + buffer for pool in mtp_draft_device_pools for buffer in pool.kv_buffer + ] + draft_indexer_buffers = [ + buffer + for pool in mtp_draft_device_pools + for buffer in pool.index_k_with_scale_buffer + ] + if len(draft_kv_buffers) != len(draft_indexer_buffers): + raise ValueError("DSA MTP KV and indexer draft layer counts must match.") + layer_mapping = _with_packed_draft_mapping( + {layer: layer for layer in range(num_layers)}, + target_device_layer_num=num_layers, + draft_layer_num=len(draft_kv_buffers), + ) entries = [ DevicePoolEntry( name=PoolName.KV, indices_from_pool=PoolName.KV, device_pool=kvcache, - components=[kvcache.kv_buffer], - layer_mapping=identity, + components=[[*kvcache.kv_buffer, *draft_kv_buffers]], + layer_mapping=layer_mapping, page_size=page_size, rows_are_pages=False, ), @@ -358,8 +410,8 @@ def _build_dsa_device_pool_group(kvcache: Any, page_size: int) -> DevicePoolGrou name=PoolName.INDEXER, indices_from_pool=PoolName.KV, device_pool=kvcache, - components=[kvcache.index_k_with_scale_buffer], - layer_mapping=identity, + components=[[*kvcache.index_k_with_scale_buffer, *draft_indexer_buffers]], + layer_mapping=layer_mapping, page_size=page_size, rows_are_pages=True, ), diff --git a/python/sglang/srt/speculative/base_spec_worker.py b/python/sglang/srt/speculative/base_spec_worker.py index c5f348d60..2e9a9fafc 100644 --- a/python/sglang/srt/speculative/base_spec_worker.py +++ b/python/sglang/srt/speculative/base_spec_worker.py @@ -253,6 +253,7 @@ class BaseSpecWorker(ABC): spec_algorithm = target_model_runner.spec_algorithm if not ( get_memory().enable_hierarchical_cache + or get_memory().enable_unified_cache_external_linker or get_disagg().disaggregation_decode_retraction_backup == "host_pool" ): return HiCacheDraftPlan() @@ -276,6 +277,11 @@ class BaseSpecWorker(ABC): device_pools=draft_pools, ) + if get_memory().enable_unified_cache_external_linker: + raise NotImplementedError( + "The external linker only supports packed draft KV caches." + ) + return HiCacheDraftPlan( mode=HiCacheDraftMode.SIDECAR, # Preserve the legacy non-packed HiCache behavior: multi-layer diff --git a/test/registered/unit/mem_cache/test_linker_pool_assembler.py b/test/registered/unit/mem_cache/test_linker_pool_assembler.py index 00547a140..069cad2d6 100644 --- a/test/registered/unit/mem_cache/test_linker_pool_assembler.py +++ b/test/registered/unit/mem_cache/test_linker_pool_assembler.py @@ -2,7 +2,7 @@ import unittest from types import SimpleNamespace -from unittest.mock import Mock, call +from unittest.mock import Mock, call, patch import torch @@ -207,11 +207,21 @@ class TestHybridDevicePoolAssembler(CustomTestCase): None, state_pool(), ] + draft_swa_buffers = [ + torch.zeros((8, 13), dtype=torch.uint8), + torch.zeros((8, 17), dtype=torch.uint8), + ] group = resolve_hybrid_device_pool_group( kvcache=kvcache, page_size=2, - params=SimpleNamespace(), + params=SimpleNamespace( + mtp_draft_device_pools=( + SimpleNamespace( + swa_kv_pool=SimpleNamespace(kv_buffer=draft_swa_buffers) + ), + ) + ), components={ComponentType.FULL, ComponentType.SWA}, ) @@ -240,6 +250,14 @@ class TestHybridDevicePoolAssembler(CustomTestCase): self.assertEqual(offsets, [[5]]) self.assertIsNone(c4_pool.get_prepared_layer_range_meta([0], 1)) + swa_pool = group.entry_map[PoolName.SWA] + _, sizes, offsets = swa_pool.get_prepared_layer_range_meta([0], 0) + self.assertEqual(sizes, [[3, 13]]) + self.assertEqual(offsets, [[0, 9]]) + _, sizes, offsets = swa_pool.get_prepared_layer_range_meta([0], 1) + self.assertEqual(sizes, [[3, 17]]) + self.assertEqual(offsets, [[3, 22]]) + def test_unified_deepseek_v4_uses_only_compressed_pools(self): from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4LayerItem @@ -332,24 +350,26 @@ class TestHybridDevicePoolAssembler(CustomTestCase): def test_dsa_uses_hybrid_assembler_strategy(self): from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool - kvcache = DSATokenToKVPool.__new__(DSATokenToKVPool) - kvcache.page_size = 2 + def dsa_pool(kv_width, index_width): + pool = DSATokenToKVPool.__new__(DSATokenToKVPool) + pool.page_size = 2 + pool.layer_num = 1 + pool.kv_buffer = [torch.zeros((8, kv_width), dtype=torch.uint8)] + pool.index_key_cache = SimpleNamespace( + buffer=[torch.zeros((4, index_width), dtype=torch.uint8)] + ) + return pool + + kvcache = dsa_pool(3, 7) kvcache.layer_num = 2 - kvcache.kv_buffer = [ - torch.zeros((8, 3), dtype=torch.uint8), - torch.zeros((8, 5), dtype=torch.uint8), - ] - kvcache.index_key_cache = SimpleNamespace( - buffer=[ - torch.zeros((4, 7), dtype=torch.uint8), - torch.zeros((4, 11), dtype=torch.uint8), - ] - ) + kvcache.kv_buffer.append(torch.zeros((8, 5), dtype=torch.uint8)) + kvcache.index_key_cache.buffer.append(torch.zeros((4, 11), dtype=torch.uint8)) + draft_pools = (dsa_pool(13, 17), dsa_pool(19, 23)) group = resolve_hybrid_device_pool_group( kvcache=kvcache, page_size=2, - params=SimpleNamespace(), + params=SimpleNamespace(mtp_draft_device_pools=draft_pools), components={ComponentType.FULL}, ) @@ -363,6 +383,76 @@ class TestHybridDevicePoolAssembler(CustomTestCase): PoolName.INDEXER: PoolName.KV, }, ) + _, sizes, offsets = group.entry_map[PoolName.KV].get_prepared_layer_range_meta( + [0], 0 + ) + self.assertEqual(sizes, [[6, 26]]) + self.assertEqual(offsets, [[0, 16]]) + _, sizes, offsets = group.entry_map[ + PoolName.INDEXER + ].get_prepared_layer_range_meta([0], 0) + self.assertEqual(sizes, [[7, 17]]) + self.assertEqual(offsets, [[0, 18]]) + _, sizes, offsets = group.entry_map[ + PoolName.INDEXER + ].get_prepared_layer_range_meta([0], 1) + self.assertEqual(sizes, [[11, 23]]) + self.assertEqual(offsets, [[7, 35]]) + + def test_linker_requires_packed_draft(self): + """Do not accept draft state that the linker would omit from storage.""" + from sglang.srt.speculative import base_spec_worker as spec + from sglang.srt.speculative.spec_info import SpeculativeAlgorithm + + draft = SimpleNamespace( + token_to_kv_pool=object(), + model_config=SimpleNamespace( + num_nextn_predict_layers=0, + hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]), + ), + ) + target = SimpleNamespace(spec_algorithm=SpeculativeAlgorithm.EAGLE) + worker = SimpleNamespace( + target_worker=SimpleNamespace(model_runner=target), + _draft_model_runners=lambda: (draft,), + ) + for linker_enabled, nextn_layers in ( + (False, 0), + (True, 0), + (False, 1), + (True, 1), + ): + draft.model_config.num_nextn_predict_layers = nextn_layers + with ( + self.subTest(linker=linker_enabled, nextn=nextn_layers), + patch.object( + spec, + "get_memory", + return_value=SimpleNamespace( + enable_hierarchical_cache=not linker_enabled, + enable_unified_cache_external_linker=linker_enabled, + ), + ), + ): + if linker_enabled and not nextn_layers: + with self.assertRaisesRegex( + NotImplementedError, "only supports packed" + ): + spec.BaseSpecWorker._build_hicache_draft_plan(worker) + self.assertEqual(target.mtp_draft_device_pools, ()) + else: + plan = spec.BaseSpecWorker._build_hicache_draft_plan(worker) + self.assertEqual( + plan.mode, + spec.HiCacheDraftMode.PACKED + if nextn_layers + else spec.HiCacheDraftMode.SIDECAR, + ) + self.assertEqual(plan.device_pools, (draft.token_to_kv_pool,)) + self.assertEqual( + target.mtp_draft_device_pools, + plan.device_pools if nextn_layers else (), + ) def test_unsupported_strategy_fails_with_context(self): from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool