[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:
Aurick Qiao
2026-09-17 10:18:44 +08:00
committed by GitHub
co-authored by Aurick Qiao Ke Bao
parent 0daa040e48
commit 4793f56835
2 changed files with 109 additions and 5 deletions
@@ -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)