[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,
|
_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,12 +127,14 @@ 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):
|
||||||
|
for buffer_index in buffer_indices:
|
||||||
base_ptr, row_stride, size = component[buffer_index]
|
base_ptr, row_stride, size = component[buffer_index]
|
||||||
items.append((base_ptr, row_stride, size, offsets[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
|
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)
|
||||||
kvcache.layer_num = 2
|
pool.page_size = 2
|
||||||
kvcache.kv_buffer = [
|
pool.layer_num = 1
|
||||||
torch.zeros((8, 3), dtype=torch.uint8),
|
pool.kv_buffer = [torch.zeros((8, kv_width), dtype=torch.uint8)]
|
||||||
torch.zeros((8, 5), dtype=torch.uint8),
|
pool.index_key_cache = SimpleNamespace(
|
||||||
]
|
buffer=[torch.zeros((4, index_width), dtype=torch.uint8)]
|
||||||
kvcache.index_key_cache = SimpleNamespace(
|
|
||||||
buffer=[
|
|
||||||
torch.zeros((4, 7), dtype=torch.uint8),
|
|
||||||
torch.zeros((4, 11), 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(
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user