[HiCache] Keep hybrid transfer layer maps stage-local under PP (#39699)
Co-authored-by: Aurick Qiao <6137920+aurickq@users.noreply.github.com> Co-authored-by: Ke Bao <ispobaoke@gmail.com>
This commit is contained in:
co-authored by
Aurick Qiao
Ke Bao
parent
0daa040e48
commit
4793f56835
@@ -5,10 +5,14 @@ from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||
from sglang.srt.mem_cache.hybrid_cache import hybrid_pool_assembler
|
||||
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
||||
_evict_mamba_for_device_alloc,
|
||||
_evict_swa_for_device_alloc,
|
||||
_MambaStrategy,
|
||||
_MambaSwaStrategy,
|
||||
_split_hicache_size,
|
||||
_SwaStrategy,
|
||||
build_full_draft_pools,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
|
||||
@@ -76,6 +80,85 @@ class TestSplitHicacheSize(CustomTestCase):
|
||||
self.assertEqual(sum(shares), 100) # total budget preserved, not doubled
|
||||
|
||||
|
||||
class TestHybridStageLayerMappings(CustomTestCase):
|
||||
def test_strategies_pass_stage_local_maps_without_changing_device_maps(self):
|
||||
"""Later pipeline stages must transfer every layer before signaling completion."""
|
||||
cases = [
|
||||
(
|
||||
_MambaStrategy,
|
||||
"build_hybrid_mamba_stack",
|
||||
{"full": {3: 0}, "mamba": {0: 2, 1: 0, 2: 1}},
|
||||
),
|
||||
(
|
||||
_SwaStrategy,
|
||||
"build_hybrid_swa_stack",
|
||||
{"full": {3: 0}, "swa": {0: 2, 1: 0, 2: 1}},
|
||||
),
|
||||
(
|
||||
_MambaSwaStrategy,
|
||||
"build_hybrid_mamba_swa_stack",
|
||||
{"full": {3: 0}, "swa": {0: 2, 1: 0, 2: 1}, "mamba": {0: 1, 3: 0}},
|
||||
),
|
||||
]
|
||||
for strategy_cls, builder_name, local_maps in cases:
|
||||
for start_layer in (0, 4):
|
||||
with self.subTest(strategy=strategy_cls.__name__, start=start_layer):
|
||||
global_maps = {
|
||||
name: {
|
||||
layer + start_layer: index
|
||||
for layer, index in mapping.items()
|
||||
}
|
||||
for name, mapping in local_maps.items()
|
||||
}
|
||||
layers_mapping = {
|
||||
layer: (index, name == "swa")
|
||||
for name in ("full", "swa")
|
||||
for layer, index in global_maps.get(name, {}).items()
|
||||
}
|
||||
kvcache = SimpleNamespace(
|
||||
start_layer=start_layer,
|
||||
full_attention_layer_id_mapping=global_maps["full"].copy(),
|
||||
layers_mapping=layers_mapping.copy(),
|
||||
full_kv_pool=object(),
|
||||
swa_kv_pool=object(),
|
||||
use_mla=False,
|
||||
)
|
||||
req_pool = SimpleNamespace(
|
||||
mamba_map=global_maps.get("mamba", {}).copy(),
|
||||
mamba_pool=object(),
|
||||
)
|
||||
params = SimpleNamespace(
|
||||
req_to_token_pool=req_pool,
|
||||
tp_cache_group=None,
|
||||
pp_cache_group=None,
|
||||
)
|
||||
with patch.object(
|
||||
hybrid_pool_assembler,
|
||||
builder_name,
|
||||
return_value=(MagicMock(), object()),
|
||||
) as build_stack:
|
||||
result = strategy_cls().build(
|
||||
cache=SimpleNamespace(page_size=1),
|
||||
kvcache=kvcache,
|
||||
params=params,
|
||||
server_args=None,
|
||||
load_cache_event=None,
|
||||
)
|
||||
|
||||
build_stack.assert_called_once()
|
||||
for name, mapping in local_maps.items():
|
||||
self.assertEqual(
|
||||
build_stack.call_args.kwargs[f"{name}_layer_mapping"],
|
||||
mapping,
|
||||
)
|
||||
self.assertEqual(result.transfer_layer_num, 4)
|
||||
self.assertEqual(
|
||||
kvcache.full_attention_layer_id_mapping, global_maps["full"]
|
||||
)
|
||||
self.assertEqual(kvcache.layers_mapping, layers_mapping)
|
||||
self.assertEqual(req_pool.mamba_map, global_maps.get("mamba", {}))
|
||||
|
||||
|
||||
class TestDraftSidecarPoolDispatch(CustomTestCase):
|
||||
def test_full_builder_unwraps_empty_hybrid_linear_pool(self):
|
||||
draft_kv_pool = object.__new__(HybridLinearKVPool)
|
||||
|
||||
Reference in New Issue
Block a user