From 4793f5683543106b758fd865870db899f738f4d5 Mon Sep 17 00:00:00 2001 From: Aurick Qiao Date: Wed, 16 Sep 2026 19:18:44 -0700 Subject: [PATCH] [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 --- .../hybrid_cache/hybrid_pool_assembler.py | 31 +++++-- .../mem_cache/test_hybrid_pool_assembler.py | 83 +++++++++++++++++++ 2 files changed, 109 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 006b52a72..c10f11320 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -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, diff --git a/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py b/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py index 8d203fc7a..152aff6e7 100644 --- a/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py +++ b/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py @@ -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)