[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:
@@ -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,
|
||||
|
||||
@@ -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,12 +127,14 @@ 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):
|
||||
for buffer_index in buffer_indices:
|
||||
base_ptr, row_stride, size = component[buffer_index]
|
||||
items.append((base_ptr, row_stride, size, offsets[buffer_index]))
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
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),
|
||||
]
|
||||
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.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
|
||||
|
||||
Reference in New Issue
Block a user