[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
@@ -76,6 +76,15 @@ def _make_layer_mapper(
|
||||
return mapper
|
||||
|
||||
|
||||
def _stage_local_layer_mapping(
|
||||
layer_mapping: dict[int, int], start_layer: int
|
||||
) -> dict[int, int]:
|
||||
return {
|
||||
global_layer - start_layer: pool_layer
|
||||
for global_layer, pool_layer in layer_mapping.items()
|
||||
}
|
||||
|
||||
|
||||
def _with_mtp_layer_mapping(
|
||||
layer_mapping: dict[int, int],
|
||||
*,
|
||||
@@ -1478,8 +1487,12 @@ class _MambaStrategy(StackStrategy):
|
||||
model_name=None,
|
||||
enable_storage_metrics=False,
|
||||
):
|
||||
full_layer_mapping = dict(kvcache.full_attention_layer_id_mapping)
|
||||
mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map)
|
||||
full_layer_mapping = _stage_local_layer_mapping(
|
||||
kvcache.full_attention_layer_id_mapping, kvcache.start_layer
|
||||
)
|
||||
mamba_layer_mapping = _stage_local_layer_mapping(
|
||||
params.req_to_token_pool.mamba_map, kvcache.start_layer
|
||||
)
|
||||
host_pool_group, cache_controller = build_hybrid_mamba_stack(
|
||||
params=params,
|
||||
kv_pool=kvcache.full_kv_pool,
|
||||
@@ -1511,9 +1524,15 @@ class _MambaStrategy(StackStrategy):
|
||||
|
||||
def _swa_layer_mappings(kvcache) -> tuple[dict[int, int], dict[int, int]]:
|
||||
full = {
|
||||
gid: lid for gid, (lid, is_swa) in kvcache.layers_mapping.items() if not is_swa
|
||||
gid - kvcache.start_layer: lid
|
||||
for gid, (lid, is_swa) in kvcache.layers_mapping.items()
|
||||
if not is_swa
|
||||
}
|
||||
swa = {
|
||||
gid - kvcache.start_layer: lid
|
||||
for gid, (lid, is_swa) in kvcache.layers_mapping.items()
|
||||
if is_swa
|
||||
}
|
||||
swa = {gid: lid for gid, (lid, is_swa) in kvcache.layers_mapping.items() if is_swa}
|
||||
return full, swa
|
||||
|
||||
|
||||
@@ -1600,7 +1619,9 @@ class _MambaSwaStrategy(StackStrategy):
|
||||
enable_storage_metrics=False,
|
||||
):
|
||||
full_layer_mapping, swa_layer_mapping = _swa_layer_mappings(kvcache)
|
||||
mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map)
|
||||
mamba_layer_mapping = _stage_local_layer_mapping(
|
||||
params.req_to_token_pool.mamba_map, kvcache.start_layer
|
||||
)
|
||||
host_pool_group, cache_controller = build_hybrid_mamba_swa_stack(
|
||||
params=params,
|
||||
full_kv_pool=kvcache.full_kv_pool,
|
||||
|
||||
@@ -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