[Unified Cache][7/N] Support MTP, EAGLE, and DSpark draft KV caches in the external linker (#37914)

Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
huangtingwei
2026-09-13 18:22:06 +08:00
committed by GitHub
co-authored by hzh0425
parent cebca698e2
commit a7cf4a6fbc
4 changed files with 183 additions and 31 deletions
@@ -1300,7 +1300,9 @@ class _DeepSeekV4Strategy(StackStrategy):
_build_deepseek_v4_device_pool_group, _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( def build(
self, self,
@@ -1572,7 +1574,9 @@ class _DsaStrategy(StackStrategy):
_build_dsa_device_pool_group, _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( def build(
self, self,
@@ -26,7 +26,7 @@ class DevicePoolEntry:
indices_from_pool: PoolName, indices_from_pool: PoolName,
device_pool: Any, device_pool: Any,
components: Sequence[Sequence[torch.Tensor]], components: Sequence[Sequence[torch.Tensor]],
layer_mapping: dict[int, int], layer_mapping: dict[int, int | Sequence[int]],
page_size: int, page_size: int,
rows_are_pages: bool, rows_are_pages: bool,
packed: bool = True, packed: bool = True,
@@ -127,14 +127,16 @@ class DevicePoolEntry:
return self._rows(indices) return self._rows(indices)
def get_prepared_layer_range_meta(self, locations: list[int], layer: int): def get_prepared_layer_range_meta(self, locations: list[int], layer: int):
buffer_index = self.layer_mapping.get(layer) mapped = self.layer_mapping.get(layer)
if buffer_index is None: if mapped is None:
return None return None
buffer_indices = [mapped] if isinstance(mapped, int) else list(mapped)
items = [] items = []
for component, offsets in zip(self.buffer_meta, self._component_offsets): for component, offsets in zip(self.buffer_meta, self._component_offsets):
base_ptr, row_stride, size = component[buffer_index] for buffer_index in buffer_indices:
items.append((base_ptr, row_stride, size, offsets[buffer_index])) base_ptr, row_stride, size = component[buffer_index]
items.append((base_ptr, row_stride, size, offsets[buffer_index]))
ptrs, sizes, offsets = [], [], [] ptrs, sizes, offsets = [], [], []
for row in locations: for row in locations:
@@ -234,8 +236,28 @@ def _deepseek_v4_state_views(state_pools: list[Any], global_layers: list[int]):
return views 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( def _build_deepseek_v4_device_pool_group(
kvcache: Any, page_size: int kvcache: Any,
page_size: int,
mtp_draft_device_pools: tuple[Any, ...] = (),
) -> DevicePoolGroup: ) -> DevicePoolGroup:
from sglang.srt.mem_cache.deepseek_v4_memory_pool import HiSparseC4DevicePool from sglang.srt.mem_cache.deepseek_v4_memory_pool import HiSparseC4DevicePool
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( 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: " "DeepSeek V4 SWA page size must match the tree page size: "
f"{kvcache.swa_page_size} != {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( entries.append(
DevicePoolEntry( DevicePoolEntry(
name=PoolName.SWA, name=PoolName.SWA,
indices_from_pool=PoolName.SWA, indices_from_pool=PoolName.SWA,
device_pool=kvcache.swa_kv_pool, device_pool=kvcache.swa_kv_pool,
components=[kvcache.swa_kv_pool.kv_buffer], components=[[*kvcache.swa_kv_pool.kv_buffer, *draft_swa_buffers]],
layer_mapping=mappings.swa, layer_mapping=swa_mapping,
page_size=page_size, page_size=page_size,
rows_are_pages=True, 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: if kvcache.page_size != page_size:
raise ValueError( raise ValueError(
"DSA KV page size must match the tree page size: " "DSA KV page size must match the tree page size: "
f"{kvcache.page_size} != {page_size}." f"{kvcache.page_size} != {page_size}."
) )
num_layers = kvcache.layer_num 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 = [ entries = [
DevicePoolEntry( DevicePoolEntry(
name=PoolName.KV, name=PoolName.KV,
indices_from_pool=PoolName.KV, indices_from_pool=PoolName.KV,
device_pool=kvcache, device_pool=kvcache,
components=[kvcache.kv_buffer], components=[[*kvcache.kv_buffer, *draft_kv_buffers]],
layer_mapping=identity, layer_mapping=layer_mapping,
page_size=page_size, page_size=page_size,
rows_are_pages=False, rows_are_pages=False,
), ),
@@ -358,8 +410,8 @@ def _build_dsa_device_pool_group(kvcache: Any, page_size: int) -> DevicePoolGrou
name=PoolName.INDEXER, name=PoolName.INDEXER,
indices_from_pool=PoolName.KV, indices_from_pool=PoolName.KV,
device_pool=kvcache, device_pool=kvcache,
components=[kvcache.index_k_with_scale_buffer], components=[[*kvcache.index_k_with_scale_buffer, *draft_indexer_buffers]],
layer_mapping=identity, layer_mapping=layer_mapping,
page_size=page_size, page_size=page_size,
rows_are_pages=True, rows_are_pages=True,
), ),
@@ -253,6 +253,7 @@ class BaseSpecWorker(ABC):
spec_algorithm = target_model_runner.spec_algorithm spec_algorithm = target_model_runner.spec_algorithm
if not ( if not (
get_memory().enable_hierarchical_cache get_memory().enable_hierarchical_cache
or get_memory().enable_unified_cache_external_linker
or get_disagg().disaggregation_decode_retraction_backup == "host_pool" or get_disagg().disaggregation_decode_retraction_backup == "host_pool"
): ):
return HiCacheDraftPlan() return HiCacheDraftPlan()
@@ -276,6 +277,11 @@ class BaseSpecWorker(ABC):
device_pools=draft_pools, 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( return HiCacheDraftPlan(
mode=HiCacheDraftMode.SIDECAR, mode=HiCacheDraftMode.SIDECAR,
# Preserve the legacy non-packed HiCache behavior: multi-layer # Preserve the legacy non-packed HiCache behavior: multi-layer
@@ -2,7 +2,7 @@
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import Mock, call from unittest.mock import Mock, call, patch
import torch import torch
@@ -207,11 +207,21 @@ class TestHybridDevicePoolAssembler(CustomTestCase):
None, None,
state_pool(), 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( group = resolve_hybrid_device_pool_group(
kvcache=kvcache, kvcache=kvcache,
page_size=2, 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}, components={ComponentType.FULL, ComponentType.SWA},
) )
@@ -240,6 +250,14 @@ class TestHybridDevicePoolAssembler(CustomTestCase):
self.assertEqual(offsets, [[5]]) self.assertEqual(offsets, [[5]])
self.assertIsNone(c4_pool.get_prepared_layer_range_meta([0], 1)) 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): def test_unified_deepseek_v4_uses_only_compressed_pools(self):
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4LayerItem 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): def test_dsa_uses_hybrid_assembler_strategy(self):
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
kvcache = DSATokenToKVPool.__new__(DSATokenToKVPool) def dsa_pool(kv_width, index_width):
kvcache.page_size = 2 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.layer_num = 2
kvcache.kv_buffer = [ kvcache.kv_buffer.append(torch.zeros((8, 5), dtype=torch.uint8))
torch.zeros((8, 3), dtype=torch.uint8), kvcache.index_key_cache.buffer.append(torch.zeros((4, 11), dtype=torch.uint8))
torch.zeros((8, 5), dtype=torch.uint8), draft_pools = (dsa_pool(13, 17), dsa_pool(19, 23))
]
kvcache.index_key_cache = SimpleNamespace(
buffer=[
torch.zeros((4, 7), dtype=torch.uint8),
torch.zeros((4, 11), dtype=torch.uint8),
]
)
group = resolve_hybrid_device_pool_group( group = resolve_hybrid_device_pool_group(
kvcache=kvcache, kvcache=kvcache,
page_size=2, page_size=2,
params=SimpleNamespace(), params=SimpleNamespace(mtp_draft_device_pools=draft_pools),
components={ComponentType.FULL}, components={ComponentType.FULL},
) )
@@ -363,6 +383,76 @@ class TestHybridDevicePoolAssembler(CustomTestCase):
PoolName.INDEXER: PoolName.KV, 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): def test_unsupported_strategy_fails_with_context(self):
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool