From b13d3d18c68c5da56936b1b9aee9dcf467b65b91 Mon Sep 17 00:00:00 2001 From: Ke Bao Date: Tue, 26 May 2026 00:06:17 +0800 Subject: [PATCH] Refactor HiCache stack dispatch into strategies (#26295) --- .../hybrid_cache/hybrid_pool_assembler.py | 667 +++++++++++------- .../test_unified_radix_hicache_dispatch.py | 183 +++++ 2 files changed, 600 insertions(+), 250 deletions(-) create mode 100644 test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py 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 63cb96c4f..0090932a9 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 @@ -1,6 +1,7 @@ from __future__ import annotations import logging +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Callable, Optional from sglang.srt.mem_cache.hicache_storage import PoolName, SidecarPoolSpec @@ -18,6 +19,7 @@ from sglang.srt.mem_cache.memory_pool_host import ( MLATokenToKVPoolHost, PoolEntry, ) +from sglang.srt.mem_cache.unified_cache_components import ComponentType if TYPE_CHECKING: import torch @@ -643,6 +645,410 @@ def build_anchor_sidecar_stack( return host_pool_group, cache_controller +_COMPONENT_HOST_ATTR: dict[ComponentType, tuple[str, str]] = { + ComponentType.FULL: ("full_kv_pool_host", "_full_kv_pool_host"), + ComponentType.SWA: ("swa_kv_pool_host", "_swa_kv_pool_host"), + ComponentType.MAMBA: ("mamba_pool_host", "_mamba_pool_host"), +} + + +@dataclass +class StackBuildResult: + host_pool_group: HostPoolGroup + cache_controller: HybridCacheController + component_host_pools: dict[ComponentType, Any] + sidecars: list[SidecarPoolSpec] = field(default_factory=list) + # Mamba state lives in req_to_token_pool, not in kvcache, so its + # layer_transfer_counter has to be wired separately. + register_req_to_token_counter: bool = False + transfer_layer_num: int = 0 + pools_desc: str = "" + + +class StackStrategy: + def matches(self, kvcache: Any, components: set[ComponentType]) -> bool: + raise NotImplementedError + + def build( + self, + *, + cache: UnifiedRadixCache, + kvcache: Any, + params: CacheInitParams, + server_args: ServerArgs, + load_cache_event, + attn_cp_group: Optional[torch.distributed.ProcessGroup] = None, + attn_tp_group: Optional[torch.distributed.ProcessGroup] = None, + ) -> StackBuildResult: + raise NotImplementedError + + +class _DeepSeekV4Strategy(StackStrategy): + def matches(self, kvcache, components): + from sglang.srt.mem_cache.deepseek_v4_memory_pool import ( + DeepSeekV4TokenToKVPool, + ) + + return isinstance(kvcache, DeepSeekV4TokenToKVPool) and components == { + ComponentType.FULL, + ComponentType.SWA, + } + + def build( + self, + *, + cache, + kvcache, + params, + server_args, + load_cache_event, + attn_cp_group=None, + attn_tp_group=None, + ): + from sglang.srt.mem_cache.base_prefix_cache import EvictParams + + host_pool_group, cache_controller = build_deepseek_v4_hicache_stack( + params=params, + server_args=server_args, + kvcache=kvcache, + page_size=cache.page_size, + tp_group=params.tp_cache_group, + load_cache_event=load_cache_event, + attn_cp_group=attn_cp_group, + attn_tp_group=attn_tp_group, + storage_backend=None, + host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA), + device_swa_evict_fn=lambda n: cache.evict(EvictParams(swa_num_tokens=n)), + pp_rank=params.pp_rank, + pp_size=params.pp_size, + ) + sidecars = [ + SidecarPoolSpec(pool_name=name, indices_from_pool=src) + for name, src in ( + (PoolName.DEEPSEEK_V4_C4, PoolName.KV), + (PoolName.DEEPSEEK_V4_C4_INDEXER, PoolName.KV), + (PoolName.DEEPSEEK_V4_C128, PoolName.KV), + (PoolName.DEEPSEEK_V4_C4_STATE, PoolName.SWA), + (PoolName.DEEPSEEK_V4_C4_INDEXER_STATE, PoolName.SWA), + (PoolName.DEEPSEEK_V4_C128_STATE, PoolName.SWA), + ) + if name in host_pool_group.entry_map + ] + return StackBuildResult( + host_pool_group=host_pool_group, + cache_controller=cache_controller, + component_host_pools={ + ComponentType.FULL: host_pool_group.get_pool(PoolName.KV), + ComponentType.SWA: host_pool_group.get_pool(PoolName.SWA), + }, + sidecars=sidecars, + transfer_layer_num=kvcache.end_layer - kvcache.start_layer, + pools_desc="KV + SWA + DeepSeekV4 sidecars", + ) + + +class _MambaStrategy(StackStrategy): + def matches(self, kvcache, components): + from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool + + return isinstance(kvcache, HybridLinearKVPool) and components == { + ComponentType.FULL, + ComponentType.MAMBA, + } + + def build( + self, + *, + cache, + kvcache, + params, + server_args, + load_cache_event, + attn_cp_group=None, + attn_tp_group=None, + ): + from sglang.srt.mem_cache.base_prefix_cache import EvictParams + + full_layer_mapping = dict(kvcache.full_attention_layer_id_mapping) + mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map) + host_pool_group, cache_controller = build_hybrid_mamba_stack( + params=params, + server_args=server_args, + kv_pool=kvcache.full_kv_pool, + mamba_pool=params.req_to_token_pool.mamba_pool, + full_layer_mapping=full_layer_mapping, + mamba_layer_mapping=mamba_layer_mapping, + page_size=cache.page_size, + tp_group=params.tp_cache_group, + load_cache_event=load_cache_event, + attn_cp_group=attn_cp_group, + attn_tp_group=attn_tp_group, + storage_backend=None, + use_mla=kvcache.use_mla, + host_mamba_evict_fn=lambda n: cache.evict_host(n, ComponentType.MAMBA), + device_mamba_evict_fn=lambda n: cache.evict(EvictParams(mamba_num=n)), + pp_rank=params.pp_rank, + pp_size=params.pp_size, + ) + return StackBuildResult( + host_pool_group=host_pool_group, + cache_controller=cache_controller, + component_host_pools={ + ComponentType.FULL: host_pool_group.get_pool(PoolName.KV), + ComponentType.MAMBA: host_pool_group.get_pool(PoolName.MAMBA), + }, + register_req_to_token_counter=True, + transfer_layer_num=len(full_layer_mapping | mamba_layer_mapping), + pools_desc="KV + MAMBA", + ) + + +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 + } + swa = {gid: lid for gid, (lid, is_swa) in kvcache.layers_mapping.items() if is_swa} + return full, swa + + +class _SwaStrategy(StackStrategy): + def matches(self, kvcache, components): + from sglang.srt.mem_cache.deepseek_v4_memory_pool import ( + DeepSeekV4TokenToKVPool, + ) + from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool + + return ( + isinstance(kvcache, SWAKVPool) + and not isinstance(kvcache, DeepSeekV4TokenToKVPool) + and components == {ComponentType.FULL, ComponentType.SWA} + ) + + def build( + self, + *, + cache, + kvcache, + params, + server_args, + load_cache_event, + attn_cp_group=None, + attn_tp_group=None, + ): + from sglang.srt.mem_cache.base_prefix_cache import EvictParams + + full_layer_mapping, swa_layer_mapping = _swa_layer_mappings(kvcache) + host_pool_group, cache_controller = build_hybrid_swa_stack( + params=params, + server_args=server_args, + full_kv_pool=kvcache.full_kv_pool, + swa_kv_pool=kvcache.swa_kv_pool, + full_layer_mapping=full_layer_mapping, + swa_layer_mapping=swa_layer_mapping, + page_size=cache.page_size, + tp_group=params.tp_cache_group, + load_cache_event=load_cache_event, + attn_cp_group=attn_cp_group, + attn_tp_group=attn_tp_group, + storage_backend=None, + use_mla=False, + host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA), + device_swa_evict_fn=lambda n: cache.evict(EvictParams(swa_num_tokens=n)), + pp_rank=params.pp_rank, + pp_size=params.pp_size, + ) + return StackBuildResult( + host_pool_group=host_pool_group, + cache_controller=cache_controller, + component_host_pools={ + ComponentType.FULL: host_pool_group.get_pool(PoolName.KV), + ComponentType.SWA: host_pool_group.get_pool(PoolName.SWA), + }, + transfer_layer_num=len(full_layer_mapping | swa_layer_mapping), + pools_desc="KV + SWA", + ) + + +class _DsaStrategy(StackStrategy): + def matches(self, kvcache, components): + from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool + + return isinstance(kvcache, DSATokenToKVPool) and components == { + ComponentType.FULL + } + + def build( + self, + *, + cache, + kvcache, + params, + server_args, + load_cache_event, + attn_cp_group=None, + attn_tp_group=None, + ): + from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool + + full_kv_pool = kvcache + use_mla = isinstance(kvcache, MLATokenToKVPool) + full_layer_mapping = {i: i for i in range(full_kv_pool.layer_num)} + host_pool_group, cache_controller = build_anchor_sidecar_stack( + params=params, + server_args=server_args, + kv_pool=full_kv_pool, + sidecar_pool_name=PoolName.INDEXER, + full_layer_mapping=full_layer_mapping, + page_size=cache.page_size, + tp_group=params.tp_cache_group, + load_cache_event=load_cache_event, + attn_cp_group=attn_cp_group, + attn_tp_group=attn_tp_group, + storage_backend=None, + use_mla=use_mla, + override_kv_cache_dim=full_kv_pool.kv_cache_dim, + sidecar_host_pool_factory=lambda kv_host_pool: DSAIndexerPoolHost( + full_kv_pool, + kv_host_pool, + server_args.hicache_mem_layout, + allocator_type=server_args.hicache_storage_backend, + ), + pp_rank=params.pp_rank, + pp_size=params.pp_size, + ) + return StackBuildResult( + host_pool_group=host_pool_group, + cache_controller=cache_controller, + component_host_pools={ + ComponentType.FULL: host_pool_group.get_pool(PoolName.KV), + }, + sidecars=[ + SidecarPoolSpec( + pool_name=PoolName.INDEXER, + indices_from_pool=PoolName.KV, + ), + ], + transfer_layer_num=len(full_layer_mapping), + pools_desc="KV + INDEXER", + ) + + +class _PlainKvStrategy(StackStrategy): + def matches(self, kvcache, components): + from sglang.srt.mem_cache.deepseek_v4_memory_pool import ( + DeepSeekV4TokenToKVPool, + ) + from sglang.srt.mem_cache.memory_pool import ( + DSATokenToKVPool, + HybridLinearKVPool, + ) + from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool + + if isinstance( + kvcache, + (SWAKVPool, HybridLinearKVPool, DSATokenToKVPool, DeepSeekV4TokenToKVPool), + ): + return False + return components == {ComponentType.FULL} + + def build( + self, + *, + cache, + kvcache, + params, + server_args, + load_cache_event, + attn_cp_group=None, + attn_tp_group=None, + ): + from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool + + full_kv_pool = kvcache + use_mla = isinstance(kvcache, MLATokenToKVPool) + full_layer_mapping = {i: i for i in range(full_kv_pool.layer_num)} + host_pool_group, cache_controller = build_kv_only_stack( + params=params, + server_args=server_args, + kv_pool=full_kv_pool, + full_layer_mapping=full_layer_mapping, + page_size=cache.page_size, + tp_group=params.tp_cache_group, + load_cache_event=load_cache_event, + attn_cp_group=attn_cp_group, + attn_tp_group=attn_tp_group, + storage_backend=None, + use_mla=use_mla, + pp_rank=params.pp_rank, + pp_size=params.pp_size, + ) + return StackBuildResult( + host_pool_group=host_pool_group, + cache_controller=cache_controller, + component_host_pools={ + ComponentType.FULL: host_pool_group.get_pool(PoolName.KV), + }, + transfer_layer_num=len(full_layer_mapping), + pools_desc="KV", + ) + + +# Resolved first-to-last; _PlainKvStrategy is the catch-all fallback. +_STRATEGIES: list[StackStrategy] = [ + _DeepSeekV4Strategy(), + _MambaStrategy(), + _SwaStrategy(), + _DsaStrategy(), + _PlainKvStrategy(), +] + + +def register_stack_strategy(strategy: StackStrategy) -> None: + """Prepend a strategy so downstream forks can plug in (kvcache, components) + combinations not in the built-in list.""" + _STRATEGIES.insert(0, strategy) + + +def _select_strategy(kvcache: Any, components: set[ComponentType]) -> StackStrategy: + for strategy in _STRATEGIES: + if strategy.matches(kvcache, components): + return strategy + raise AssertionError( + f"No matching HiCache strategy for kvcache={type(kvcache).__name__}, " + f"components={sorted(c.name for c in components)}" + ) + + +def _apply_stack_result( + cache: UnifiedRadixCache, + kvcache: Any, + params: CacheInitParams, + result: StackBuildResult, +) -> None: + cache.host_pool_group = result.host_pool_group + cache.cache_controller = result.cache_controller + + for ct, host_pool in result.component_host_pools.items(): + cache_attr, component_attr = _COMPONENT_HOST_ATTR[ct] + setattr(cache, cache_attr, host_pool) + setattr(cache.components[ct], component_attr, host_pool) + + for sidecar in result.sidecars: + cache.register_sidecar_pool(sidecar) + + kvcache.register_layer_transfer_counter(result.cache_controller.layer_done_counter) + if result.register_req_to_token_counter: + params.req_to_token_pool.register_layer_transfer_counter( + result.cache_controller.layer_done_counter + ) + + logger.info( + "Attached hybrid pool stack to UnifiedRadixCache: pools=%s, transfer_layer_num=%s", + result.pools_desc, + result.transfer_layer_num, + ) + + def attach_hybrid_pool_to_unified_cache( cache: UnifiedRadixCache, params: CacheInitParams, @@ -653,259 +1059,20 @@ def attach_hybrid_pool_to_unified_cache( attn_tp_group: Optional[torch.distributed.ProcessGroup] = None, ) -> None: """Attach HostPoolGroup + HybridCacheController to UnifiedRadixCache.""" - from sglang.srt.mem_cache.base_prefix_cache import EvictParams - from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool - from sglang.srt.mem_cache.memory_pool import ( - DSATokenToKVPool, - HybridLinearKVPool, - MLATokenToKVPool, - ) - from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool - from sglang.srt.mem_cache.unified_cache_components import ComponentType - try: kvcache = params.token_to_kv_pool_allocator.get_kvcache() - swa_stack = isinstance(kvcache, SWAKVPool) - mamba_stack = isinstance(kvcache, HybridLinearKVPool) - dsa_stack = isinstance(kvcache, DSATokenToKVPool) - deepseek_v4_stack = isinstance(kvcache, DeepSeekV4TokenToKVPool) - - if deepseek_v4_stack: - use_mla = False - assert set(cache.components.keys()) == { - ComponentType.FULL, - ComponentType.SWA, - }, "DeepSeekV4TokenToKVPool requires FULL + SWA in UnifiedRadixCache." - elif mamba_stack: - full_kv_pool = kvcache.full_kv_pool - use_mla = kvcache.use_mla - assert set(cache.components.keys()) == { - ComponentType.FULL, - ComponentType.MAMBA, - }, "HybridLinearKVPool currently only supports FULL + MAMBA in UnifiedRadixCache." - elif swa_stack: - full_kv_pool = kvcache.full_kv_pool - use_mla = False - assert set(cache.components.keys()) == { - ComponentType.FULL, - ComponentType.SWA, - }, "SWAKVPool currently only supports FULL + SWA in UnifiedRadixCache." - else: - full_kv_pool = kvcache - use_mla = isinstance(kvcache, MLATokenToKVPool) - assert set(cache.components.keys()) == { - ComponentType.FULL - }, "Non-hybrid KV pool currently only supports FULL-only UnifiedRadixCache." - - if deepseek_v4_stack: - host_pool_group, cache_controller = build_deepseek_v4_hicache_stack( - params=params, - server_args=server_args, - kvcache=kvcache, - page_size=cache.page_size, - tp_group=params.tp_cache_group, - load_cache_event=load_cache_event, - attn_cp_group=attn_cp_group, - attn_tp_group=attn_tp_group, - storage_backend=None, - host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA), - device_swa_evict_fn=lambda n: cache.evict( - EvictParams(swa_num_tokens=n) - ), - pp_rank=params.pp_rank, - pp_size=params.pp_size, - ) - cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV) - cache.host_pool_group = host_pool_group - cache.cache_controller = cache_controller - cache.components[ComponentType.FULL]._full_kv_pool_host = ( - cache.full_kv_pool_host - ) - cache.swa_kv_pool_host = host_pool_group.get_pool(PoolName.SWA) - cache.components[ComponentType.SWA]._swa_kv_pool_host = ( - cache.swa_kv_pool_host - ) - for pool_name, indices_from_pool in ( - (PoolName.DEEPSEEK_V4_C4, PoolName.KV), - (PoolName.DEEPSEEK_V4_C4_INDEXER, PoolName.KV), - (PoolName.DEEPSEEK_V4_C128, PoolName.KV), - (PoolName.DEEPSEEK_V4_C4_STATE, PoolName.SWA), - (PoolName.DEEPSEEK_V4_C4_INDEXER_STATE, PoolName.SWA), - (PoolName.DEEPSEEK_V4_C128_STATE, PoolName.SWA), - ): - if pool_name in host_pool_group.entry_map: - cache.register_sidecar_pool( - SidecarPoolSpec( - pool_name=pool_name, - indices_from_pool=indices_from_pool, - ) - ) - transfer_layer_num = kvcache.end_layer - kvcache.start_layer - elif mamba_stack: - full_layer_mapping = dict(kvcache.full_attention_layer_id_mapping) - mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map) - host_pool_group, cache_controller = build_hybrid_mamba_stack( - params=params, - server_args=server_args, - kv_pool=full_kv_pool, - mamba_pool=params.req_to_token_pool.mamba_pool, - full_layer_mapping=full_layer_mapping, - mamba_layer_mapping=mamba_layer_mapping, - page_size=cache.page_size, - tp_group=params.tp_cache_group, - load_cache_event=load_cache_event, - attn_cp_group=attn_cp_group, - attn_tp_group=attn_tp_group, - storage_backend=None, - use_mla=use_mla, - host_mamba_evict_fn=lambda n: cache.evict_host(n, ComponentType.MAMBA), - device_mamba_evict_fn=lambda n: cache.evict(EvictParams(mamba_num=n)), - pp_rank=params.pp_rank, - pp_size=params.pp_size, - ) - cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV) - cache.host_pool_group = host_pool_group - cache.cache_controller = cache_controller - cache.components[ComponentType.FULL]._full_kv_pool_host = ( - cache.full_kv_pool_host - ) - cache.mamba_pool_host = host_pool_group.get_pool(PoolName.MAMBA) - cache.components[ComponentType.MAMBA]._mamba_pool_host = ( - cache.mamba_pool_host - ) - params.req_to_token_pool.register_layer_transfer_counter( - cache_controller.layer_done_counter - ) - transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping) - elif swa_stack: - full_layer_mapping = { - global_id: local_id - for global_id, (local_id, is_swa) in kvcache.layers_mapping.items() - if not is_swa - } - swa_layer_mapping = { - global_id: local_id - for global_id, (local_id, is_swa) in kvcache.layers_mapping.items() - if is_swa - } - host_pool_group, cache_controller = build_hybrid_swa_stack( - params=params, - server_args=server_args, - full_kv_pool=full_kv_pool, - swa_kv_pool=kvcache.swa_kv_pool, - full_layer_mapping=full_layer_mapping, - swa_layer_mapping=swa_layer_mapping, - page_size=cache.page_size, - tp_group=params.tp_cache_group, - load_cache_event=load_cache_event, - attn_cp_group=attn_cp_group, - attn_tp_group=attn_tp_group, - storage_backend=None, - use_mla=False, - host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA), - device_swa_evict_fn=lambda n: cache.evict( - EvictParams(swa_num_tokens=n) - ), - pp_rank=params.pp_rank, - pp_size=params.pp_size, - ) - cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV) - cache.host_pool_group = host_pool_group - cache.cache_controller = cache_controller - cache.components[ComponentType.FULL]._full_kv_pool_host = ( - cache.full_kv_pool_host - ) - cache.swa_kv_pool_host = host_pool_group.get_pool(PoolName.SWA) - cache.components[ComponentType.SWA]._swa_kv_pool_host = ( - cache.swa_kv_pool_host - ) - transfer_layer_num = len(full_layer_mapping | swa_layer_mapping) - elif dsa_stack: - full_layer_mapping = { - layer_id: layer_id for layer_id in range(full_kv_pool.layer_num) - } - host_pool_group, cache_controller = build_anchor_sidecar_stack( - params=params, - server_args=server_args, - kv_pool=full_kv_pool, - sidecar_pool_name=PoolName.INDEXER, - full_layer_mapping=full_layer_mapping, - page_size=cache.page_size, - tp_group=params.tp_cache_group, - load_cache_event=load_cache_event, - attn_cp_group=attn_cp_group, - attn_tp_group=attn_tp_group, - storage_backend=None, - use_mla=use_mla, - override_kv_cache_dim=full_kv_pool.kv_cache_dim, - sidecar_host_pool_factory=lambda kv_host_pool: DSAIndexerPoolHost( - full_kv_pool, - kv_host_pool, - server_args.hicache_mem_layout, - allocator_type=server_args.hicache_storage_backend, - ), - pp_rank=params.pp_rank, - pp_size=params.pp_size, - ) - cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV) - cache.host_pool_group = host_pool_group - cache.cache_controller = cache_controller - cache.register_sidecar_pool( - SidecarPoolSpec( - pool_name=PoolName.INDEXER, - indices_from_pool=PoolName.KV, - ) - ) - cache.components[ComponentType.FULL]._full_kv_pool_host = ( - cache.full_kv_pool_host - ) - transfer_layer_num = len(full_layer_mapping) - else: - full_layer_mapping = { - layer_id: layer_id for layer_id in range(full_kv_pool.layer_num) - } - host_pool_group, cache_controller = build_kv_only_stack( - params=params, - server_args=server_args, - kv_pool=full_kv_pool, - full_layer_mapping=full_layer_mapping, - page_size=cache.page_size, - tp_group=params.tp_cache_group, - load_cache_event=load_cache_event, - attn_cp_group=attn_cp_group, - attn_tp_group=attn_tp_group, - storage_backend=None, - use_mla=use_mla, - pp_rank=params.pp_rank, - pp_size=params.pp_size, - ) - cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV) - cache.host_pool_group = host_pool_group - cache.cache_controller = cache_controller - cache.components[ComponentType.FULL]._full_kv_pool_host = ( - cache.full_kv_pool_host - ) - transfer_layer_num = len(full_layer_mapping) - - kvcache.register_layer_transfer_counter( - cache.cache_controller.layer_done_counter - ) - - if deepseek_v4_stack: - pools_desc = "KV + SWA + DeepSeekV4 sidecars" - elif mamba_stack: - pools_desc = "KV + MAMBA" - elif swa_stack: - pools_desc = "KV + SWA" - elif dsa_stack: - pools_desc = "KV + INDEXER" - else: - pools_desc = "KV" - logger.info( - "Attached hybrid pool stack to UnifiedRadixCache: pools=%s, transfer_layer_num=%s", - pools_desc, - transfer_layer_num, + components = set(cache.components.keys()) + strategy = _select_strategy(kvcache, components) + result = strategy.build( + cache=cache, + kvcache=kvcache, + params=params, + server_args=server_args, + load_cache_event=load_cache_event, + attn_cp_group=attn_cp_group, + attn_tp_group=attn_tp_group, ) + _apply_stack_result(cache, kvcache, params, result) except Exception: logger.exception("attach_hybrid_pool_to_unified_cache failed") raise diff --git a/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py b/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py new file mode 100644 index 000000000..3ce1ed789 --- /dev/null +++ b/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py @@ -0,0 +1,183 @@ +import unittest +from unittest.mock import MagicMock + +from sglang.srt.mem_cache.hicache_storage import PoolName, SidecarPoolSpec +from sglang.srt.mem_cache.hybrid_cache import hybrid_pool_assembler +from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( + _STRATEGIES, + StackBuildResult, + StackStrategy, + _apply_stack_result, + _DeepSeekV4Strategy, + _DsaStrategy, + _MambaStrategy, + _PlainKvStrategy, + _select_strategy, + _SwaStrategy, + register_stack_strategy, +) +from sglang.srt.mem_cache.unified_cache_components import ComponentType +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +def _mock_kvcache(cls): + return MagicMock(spec=cls) + + +FULL = ComponentType.FULL +SWA = ComponentType.SWA +MAMBA = ComponentType.MAMBA + + +class TestUnifiedRadixHiCacheDispatch(unittest.TestCase): + def test_strategy_registry_ordering(self): + order = [type(s) for s in _STRATEGIES] + # DeepSeekV4 inherits from SWAKVPool, so it must resolve before _SwaStrategy. + self.assertLess(order.index(_DeepSeekV4Strategy), order.index(_SwaStrategy)) + self.assertEqual(order[-1], _PlainKvStrategy) + + def test_deepseek_v4_full_swa(self): + from sglang.srt.mem_cache.deepseek_v4_memory_pool import ( + DeepSeekV4TokenToKVPool, + ) + + kvcache = _mock_kvcache(DeepSeekV4TokenToKVPool) + strategy = _select_strategy(kvcache, {FULL, SWA}) + self.assertIsInstance(strategy, _DeepSeekV4Strategy) + + def test_mamba(self): + from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool + + kvcache = _mock_kvcache(HybridLinearKVPool) + strategy = _select_strategy(kvcache, {FULL, MAMBA}) + self.assertIsInstance(strategy, _MambaStrategy) + + def test_swa(self): + from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool + + kvcache = _mock_kvcache(SWAKVPool) + strategy = _select_strategy(kvcache, {FULL, SWA}) + self.assertIsInstance(strategy, _SwaStrategy) + + def test_dsa(self): + from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool + + kvcache = _mock_kvcache(DSATokenToKVPool) + strategy = _select_strategy(kvcache, {FULL}) + self.assertIsInstance(strategy, _DsaStrategy) + + def test_plain_kv_fallback(self): + from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool + + kvcache = _mock_kvcache(MHATokenToKVPool) + strategy = _select_strategy(kvcache, {FULL}) + self.assertIsInstance(strategy, _PlainKvStrategy) + + def test_mla_routes_to_plain(self): + from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool + + kvcache = _mock_kvcache(MLATokenToKVPool) + strategy = _select_strategy(kvcache, {FULL}) + self.assertIsInstance(strategy, _PlainKvStrategy) + + def test_unknown_combo_raises(self): + from sglang.srt.mem_cache.deepseek_v4_memory_pool import ( + DeepSeekV4TokenToKVPool, + ) + from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool + + for cls in (SWAKVPool, DeepSeekV4TokenToKVPool): + kvcache = _mock_kvcache(cls) + with self.assertRaises(AssertionError) as cm: + _select_strategy(kvcache, {FULL}) + self.assertIn("No matching HiCache strategy", str(cm.exception)) + + def test_register_custom_strategy_takes_precedence(self): + class _CustomStrategy(StackStrategy): + def matches(self, kvcache, components): + return components == {FULL} + + def build(self, **_): + raise NotImplementedError + + custom = _CustomStrategy() + original = list(hybrid_pool_assembler._STRATEGIES) + try: + register_stack_strategy(custom) + from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool + + kvcache = _mock_kvcache(MHATokenToKVPool) + self.assertIs(_select_strategy(kvcache, {FULL}), custom) + finally: + hybrid_pool_assembler._STRATEGIES[:] = original + + +class TestApplyStackResult(unittest.TestCase): + @staticmethod + def _fake_cache(component_types): + cache = MagicMock() + cache.components = {ct: MagicMock() for ct in component_types} + return cache + + def test_wires_components_sidecars_and_counters(self): + full_host, swa_host, mamba_host = MagicMock(), MagicMock(), MagicMock() + cache = self._fake_cache([FULL, SWA, MAMBA]) + kvcache = MagicMock() + params = MagicMock() + controller = MagicMock() + sidecar = SidecarPoolSpec( + pool_name=PoolName.INDEXER, indices_from_pool=PoolName.KV + ) + result = StackBuildResult( + host_pool_group=MagicMock(), + cache_controller=controller, + component_host_pools={FULL: full_host, SWA: swa_host, MAMBA: mamba_host}, + sidecars=[sidecar], + register_req_to_token_counter=True, + transfer_layer_num=8, + pools_desc="KV + SWA + MAMBA", + ) + + _apply_stack_result(cache, kvcache, params, result) + + self.assertIs(cache.host_pool_group, result.host_pool_group) + self.assertIs(cache.cache_controller, controller) + self.assertIs(cache.full_kv_pool_host, full_host) + self.assertIs(cache.swa_kv_pool_host, swa_host) + self.assertIs(cache.mamba_pool_host, mamba_host) + self.assertIs(cache.components[FULL]._full_kv_pool_host, full_host) + self.assertIs(cache.components[SWA]._swa_kv_pool_host, swa_host) + self.assertIs(cache.components[MAMBA]._mamba_pool_host, mamba_host) + cache.register_sidecar_pool.assert_called_once_with(sidecar) + kvcache.register_layer_transfer_counter.assert_called_once_with( + controller.layer_done_counter + ) + params.req_to_token_pool.register_layer_transfer_counter.assert_called_once_with( + controller.layer_done_counter + ) + + def test_skips_req_to_token_counter_when_flag_false(self): + cache = self._fake_cache([FULL]) + kvcache = MagicMock() + params = MagicMock() + result = StackBuildResult( + host_pool_group=MagicMock(), + cache_controller=MagicMock(), + component_host_pools={FULL: MagicMock()}, + sidecars=[], + register_req_to_token_counter=False, + transfer_layer_num=1, + pools_desc="KV", + ) + + _apply_stack_result(cache, kvcache, params, result) + + kvcache.register_layer_transfer_counter.assert_called_once() + params.req_to_token_pool.register_layer_transfer_counter.assert_not_called() + cache.register_sidecar_pool.assert_not_called() + + +if __name__ == "__main__": + unittest.main()