From ab3ce02de9711c4f27147e8c0772d782c6092a58 Mon Sep 17 00:00:00 2001 From: Zhangheng Date: Tue, 21 Apr 2026 10:45:23 +0800 Subject: [PATCH] [Hybrid-Cache]: Refactor hybrid_pool_assembler.py (#23243) --- .../srt/mem_cache/hi_mamba_radix_cache.py | 4 +- python/sglang/srt/mem_cache/hiradix_cache.py | 6 +- .../hybrid_cache/hybrid_pool_assembler.py | 561 ++++++++++++++---- .../sglang/srt/mem_cache/memory_pool_host.py | 3 + 4 files changed, 440 insertions(+), 134 deletions(-) diff --git a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py index 42c7ea43e..de64ad8ab 100644 --- a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py @@ -27,7 +27,7 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( PrefetchOperation, ) from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( - build_mamba_hybrid_stack, + attach_hybrid_pool_to_mamba_cache, ) from sglang.srt.mem_cache.mamba_radix_cache import ( LRUList, @@ -135,7 +135,7 @@ class HiMambaRadixCache(MambaRadixCache): self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy self.load_cache_event = threading.Event() - build_mamba_hybrid_stack( + attach_hybrid_pool_to_mamba_cache( self, params, server_args, diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index e8e463c61..c8e869cc4 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -35,7 +35,7 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( HybridCacheController, ) from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( - build_nsa_hybrid_stack, + attach_hybrid_nsa_pool_to_hiradix_cache, ) from sglang.srt.mem_cache.memory_pool import ( MHATokenToKVPool, @@ -81,7 +81,7 @@ class HiRadixCache(RadixCache): allocator_type=server_args.hicache_storage_backend, ) elif isinstance(self.kv_cache, NSATokenToKVPool): - # Filled by build_nsa_hybrid_stack after storage extra_config is parsed. + # Filled by attach_hybrid_nsa_pool_to_hiradix_cache after storage extra_config is parsed. self.token_to_kv_pool_host = None elif isinstance(self.kv_cache, MLATokenToKVPool): self.token_to_kv_pool_host = MLATokenToKVPoolHost( @@ -122,7 +122,7 @@ class HiRadixCache(RadixCache): self.load_cache_event = threading.Event() if isinstance(self.kv_cache, NSATokenToKVPool): - build_nsa_hybrid_stack( + attach_hybrid_nsa_pool_to_hiradix_cache( self, params, server_args, 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 27d2b60d9..0bd0f037a 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,7 +1,7 @@ from __future__ import annotations import logging -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING, Any, Callable, Optional from sglang.srt.mem_cache.hicache_storage import PoolName from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( @@ -20,173 +20,477 @@ if TYPE_CHECKING: from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.hi_mamba_radix_cache import HiMambaRadixCache from sglang.srt.mem_cache.hiradix_cache import HiRadixCache + from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache from sglang.srt.server_args import ServerArgs logger = logging.getLogger(__name__) -def build_nsa_hybrid_stack( - radix_cache: "HiRadixCache", - params: "CacheInitParams", - server_args: "ServerArgs", +def _make_layer_mapper( + layer_mapping: dict[int, int], + transfer_layer_num: int, +) -> Callable[[int], Optional[int]]: + def mapper(layer_id: int) -> Optional[int]: + if not 0 <= layer_id < transfer_layer_num: + return None + return layer_mapping.get(layer_id) + + return mapper + + +def build_kv_host_pool( + *, + kv_pool: Any, + page_size: int, + server_args: ServerArgs, + use_mla: bool, + override_kv_cache_dim: Optional[int] = None, +): + kv_host_pool_cls = MLATokenToKVPoolHost if use_mla else MHATokenToKVPoolHost + kwargs = {} + if override_kv_cache_dim is not None: + kwargs["override_kv_cache_dim"] = override_kv_cache_dim + return kv_host_pool_cls( + kv_pool, + server_args.hicache_ratio, + server_args.hicache_size, + page_size, + server_args.hicache_mem_layout, + allocator_type=server_args.hicache_storage_backend, + **kwargs, + ) + + +def build_pool_entry( + *, + name: PoolName, + host_pool: Any, + device_pool: Any, + layer_mapping: dict[int, int], + transfer_layer_num: int, + is_anchor: bool = False, + share_indices_with_anchor: bool = False, + host_evict_fn: Optional[Callable[[int], Any]] = None, + device_evict_fn: Optional[Callable[[int], Any]] = None, +) -> PoolEntry: + return PoolEntry( + name=name, + host_pool=host_pool, + device_pool=device_pool, + layer_mapper=_make_layer_mapper(layer_mapping, transfer_layer_num), + is_primary_index_anchor=is_anchor, + share_indices_with_anchor=share_indices_with_anchor, + host_evict_fn=host_evict_fn, + device_evict_fn=device_evict_fn, + ) + + +def build_kv_only_stack( + *, + params: CacheInitParams, + server_args: ServerArgs, + kv_pool: Any, + full_layer_mapping: dict[int, int], + page_size: int, + tp_group, + load_cache_event, + storage_backend: Optional[str], + use_mla: bool, + override_kv_cache_dim: Optional[int] = None, + prefetch_threshold: int = 256, + model_name: Optional[str] = None, + storage_backend_extra_config: Optional[dict] = None, + pp_rank: int = 0, + pp_size: int = 1, + attn_cp_rank: int = 0, + attn_cp_size: int = 1, + enable_storage_metrics: bool = False, +) -> tuple[HostPoolGroup, HybridCacheController]: + transfer_layer_num = len(full_layer_mapping) + kv_host_pool = build_kv_host_pool( + kv_pool=kv_pool, + page_size=page_size, + server_args=server_args, + use_mla=use_mla, + override_kv_cache_dim=override_kv_cache_dim, + ) + entries = [ + build_pool_entry( + name=PoolName.KV, + host_pool=kv_host_pool, + device_pool=kv_pool, + layer_mapping=full_layer_mapping, + transfer_layer_num=transfer_layer_num, + is_anchor=True, + ) + ] + host_pool_group = HostPoolGroup(entries) + cache_controller = HybridCacheController( + params.token_to_kv_pool_allocator, + host_pool_group, + page_size, + tp_group, + load_cache_event=load_cache_event, + write_policy=server_args.hicache_write_policy, + io_backend=server_args.hicache_io_backend, + storage_backend=storage_backend, + prefetch_threshold=prefetch_threshold, + model_name=model_name, + storage_backend_extra_config=storage_backend_extra_config, + pp_rank=pp_rank, + pp_size=pp_size, + attn_cp_rank=attn_cp_rank, + attn_cp_size=attn_cp_size, + transfer_layer_num=transfer_layer_num, + enable_storage_metrics=enable_storage_metrics, + ) + return host_pool_group, cache_controller + + +def build_hybrid_mamba_stack( + *, + params: CacheInitParams, + server_args: ServerArgs, + kv_pool: Any, + mamba_pool: Any, + full_layer_mapping: dict[int, int], + mamba_layer_mapping: dict[int, int], + page_size: int, + tp_group, + load_cache_event, + storage_backend: Optional[str], + use_mla: bool, + host_mamba_evict_fn: Optional[Callable[[int], Any]] = None, + device_mamba_evict_fn: Optional[Callable[[int], Any]] = None, + prefetch_threshold: int = 256, + model_name: Optional[str] = None, + storage_backend_extra_config: Optional[dict] = None, + pp_rank: int = 0, + pp_size: int = 1, + attn_cp_rank: int = 0, + attn_cp_size: int = 1, + enable_storage_metrics: bool = False, +) -> tuple[HostPoolGroup, HybridCacheController]: + transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping) + kv_host_pool = build_kv_host_pool( + kv_pool=kv_pool, + page_size=page_size, + server_args=server_args, + use_mla=use_mla, + ) + mamba_host_pool = MambaPoolHost( + mamba_pool, + server_args.hicache_ratio, + server_args.hicache_size, + allocator_type=server_args.hicache_storage_backend, + layout=server_args.hicache_mem_layout, + ) + entries = [ + build_pool_entry( + name=PoolName.KV, + host_pool=kv_host_pool, + device_pool=kv_pool, + layer_mapping=full_layer_mapping, + transfer_layer_num=transfer_layer_num, + is_anchor=True, + ), + build_pool_entry( + name=PoolName.MAMBA, + host_pool=mamba_host_pool, + device_pool=mamba_pool, + layer_mapping=mamba_layer_mapping, + transfer_layer_num=transfer_layer_num, + host_evict_fn=host_mamba_evict_fn, + device_evict_fn=device_mamba_evict_fn, + ), + ] + host_pool_group = HostPoolGroup(entries) + cache_controller = HybridCacheController( + params.token_to_kv_pool_allocator, + host_pool_group, + page_size, + tp_group, + load_cache_event=load_cache_event, + write_policy=server_args.hicache_write_policy, + io_backend=server_args.hicache_io_backend, + storage_backend=storage_backend, + prefetch_threshold=prefetch_threshold, + model_name=model_name, + storage_backend_extra_config=storage_backend_extra_config, + pp_rank=pp_rank, + pp_size=pp_size, + attn_cp_rank=attn_cp_rank, + attn_cp_size=attn_cp_size, + transfer_layer_num=transfer_layer_num, + enable_storage_metrics=enable_storage_metrics, + ) + return host_pool_group, cache_controller + + +def build_shared_anchor_stack( + *, + params: CacheInitParams, + server_args: ServerArgs, + kv_pool: Any, + shared_pool_name: PoolName, + full_layer_mapping: dict[int, int], + page_size: int, + tp_group, + load_cache_event, + storage_backend: Optional[str], + use_mla: bool, + override_kv_cache_dim: Optional[int] = None, + shared_host_pool_factory: Callable[[Any], Any], + prefetch_threshold: int = 256, + model_name: Optional[str] = None, + storage_backend_extra_config: Optional[dict] = None, + pp_rank: int = 0, + pp_size: int = 1, + attn_cp_rank: int = 0, + attn_cp_size: int = 1, + enable_storage_metrics: bool = False, +) -> tuple[HostPoolGroup, HybridCacheController]: + transfer_layer_num = len(full_layer_mapping) + kv_host_pool = build_kv_host_pool( + kv_pool=kv_pool, + page_size=page_size, + server_args=server_args, + use_mla=use_mla, + override_kv_cache_dim=override_kv_cache_dim, + ) + shared_host_pool = shared_host_pool_factory(kv_host_pool) + entries = [ + build_pool_entry( + name=PoolName.KV, + host_pool=kv_host_pool, + device_pool=kv_pool, + layer_mapping=full_layer_mapping, + transfer_layer_num=transfer_layer_num, + is_anchor=True, + ), + build_pool_entry( + name=shared_pool_name, + host_pool=shared_host_pool, + device_pool=kv_pool, + layer_mapping=full_layer_mapping, + transfer_layer_num=transfer_layer_num, + share_indices_with_anchor=True, + ), + ] + host_pool_group = HostPoolGroup(entries) + cache_controller = HybridCacheController( + params.token_to_kv_pool_allocator, + host_pool_group, + page_size, + tp_group, + load_cache_event=load_cache_event, + write_policy=server_args.hicache_write_policy, + io_backend=server_args.hicache_io_backend, + storage_backend=storage_backend, + prefetch_threshold=prefetch_threshold, + model_name=model_name, + storage_backend_extra_config=storage_backend_extra_config, + pp_rank=pp_rank, + pp_size=pp_size, + attn_cp_rank=attn_cp_rank, + attn_cp_size=attn_cp_size, + transfer_layer_num=transfer_layer_num, + enable_storage_metrics=enable_storage_metrics, + ) + return host_pool_group, cache_controller + + +def attach_hybrid_pool_to_unified_cache( + cache: UnifiedRadixCache, + params: CacheInitParams, + server_args: ServerArgs, + *, + load_cache_event, +) -> None: + """Attach HostPoolGroup + HybridCacheController to UnifiedRadixCache.""" + from sglang.srt.mem_cache.base_prefix_cache import EvictParams + from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, MLATokenToKVPool + from sglang.srt.mem_cache.unified_cache_components import ComponentType + + try: + kvcache = params.token_to_kv_pool_allocator.get_kvcache() + if isinstance(kvcache, HybridLinearKVPool): + 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." + 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." + + mamba_stack = isinstance(kvcache, HybridLinearKVPool) + if 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, + 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) + 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, + 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 + ) + + logger.info( + "Attached hybrid pool stack to UnifiedRadixCache: pools=%s, transfer_layer_num=%s", + "KV + MAMBA" if mamba_stack else "KV", + transfer_layer_num, + ) + except Exception: + logger.exception("attach_hybrid_pool_to_unified_cache failed") + raise + + +def attach_hybrid_nsa_pool_to_hiradix_cache( + radix_cache: HiRadixCache, + params: CacheInitParams, + server_args: ServerArgs, *, extra_config: dict, prefetch_threshold: int, enable_storage_metrics: bool, load_cache_event, ) -> None: - """HostPoolGroup (KV + indexer) + HybridCacheController for NSA (DSA).""" + """Attach HostPoolGroup (KV + indexer) + HybridCacheController for HiRadixCache. + + This entrypoint is currently intended only for HiRadixCache's NSA path. + """ try: kv = radix_cache.kv_cache - mla_host = MLATokenToKVPoolHost( - kv, - server_args.hicache_ratio, - server_args.hicache_size, - radix_cache.page_size, - server_args.hicache_mem_layout, - allocator_type=server_args.hicache_storage_backend, - override_kv_cache_dim=kv.kv_cache_dim, - ) - indexer_host = NSAIndexerPoolHost( - kv, - mla_host, - server_args.hicache_mem_layout, - allocator_type=server_args.hicache_storage_backend, - ) - layer_num = kv.layer_num - - def layer_mapper(layer_id: int): - if 0 <= layer_id < layer_num: - return layer_id - return None - - host_pool_group = HostPoolGroup( - [ - PoolEntry( - name=PoolName.KV, - host_pool=mla_host, - device_pool=kv, - layer_mapper=layer_mapper, - is_primary_index_anchor=True, - ), - PoolEntry( - name=PoolName.INDEXER, - host_pool=indexer_host, - device_pool=kv, - layer_mapper=layer_mapper, - share_indices_with_anchor=True, - ), - ] - ) - cache_controller = HybridCacheController( - params.token_to_kv_pool_allocator, - host_pool_group, - radix_cache.page_size, - radix_cache.tp_group, + layer_mapping = {layer_id: layer_id for layer_id in range(kv.layer_num)} + host_pool_group, cache_controller = build_shared_anchor_stack( + params=params, + server_args=server_args, + kv_pool=kv, + shared_pool_name=PoolName.INDEXER, + full_layer_mapping=layer_mapping, + page_size=radix_cache.page_size, + tp_group=radix_cache.tp_group, load_cache_event=load_cache_event, - write_policy=server_args.hicache_write_policy, - io_backend=server_args.hicache_io_backend, storage_backend=server_args.hicache_storage_backend, + use_mla=True, prefetch_threshold=prefetch_threshold, + shared_host_pool_factory=lambda kv_host_pool: NSAIndexerPoolHost( + kv, + kv_host_pool, + server_args.hicache_mem_layout, + allocator_type=server_args.hicache_storage_backend, + ), model_name=server_args.served_model_name, storage_backend_extra_config=extra_config, pp_rank=radix_cache.pp_rank, pp_size=radix_cache.pp_size, attn_cp_rank=params.attn_cp_rank, attn_cp_size=params.attn_cp_size, - transfer_layer_num=layer_num, enable_storage_metrics=enable_storage_metrics, ) - radix_cache.full_kv_pool_host = mla_host + radix_cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV) radix_cache.token_to_kv_pool_host = host_pool_group radix_cache.cache_controller = cache_controller logger.info( - "Hybrid hierarchical cache: HostPoolGroup(KV + INDEXER), HybridCacheController, " + "Attached hybrid NSA pool stack to HiRadixCache: pools=KV + INDEXER, " "transfer_layer_num=%s", - layer_num, + len(layer_mapping), ) except Exception: - logger.exception("build_nsa_hybrid_stack failed") + logger.exception("attach_hybrid_nsa_pool_to_hiradix_cache failed") raise -def build_mamba_hybrid_stack( - mamba_cache: "HiMambaRadixCache", - params: "CacheInitParams", - server_args: "ServerArgs", +def attach_hybrid_pool_to_mamba_cache( + mamba_cache: HiMambaRadixCache, + params: CacheInitParams, + server_args: ServerArgs, *, extra_config: dict, prefetch_threshold: int, load_cache_event, enable_storage_metrics: bool = False, ) -> None: - """HostPoolGroup (KV + Mamba) + HybridCacheController for hybrid SSM models.""" + """Attach HostPoolGroup (KV + Mamba) + HybridCacheController for HiMambaRadixCache. + + This entrypoint is currently intended only for HiMambaRadixCache. + """ try: hybrid_kv = mamba_cache.hybrid_kv_cache kvcache = mamba_cache.kvcache - kv_host_pool_cls = ( - MLATokenToKVPoolHost if hybrid_kv.use_mla else MHATokenToKVPoolHost - ) - full_kv_pool_host = kv_host_pool_cls( - kvcache, - server_args.hicache_ratio, - server_args.hicache_size, - params.page_size, - server_args.hicache_mem_layout, - allocator_type=server_args.hicache_storage_backend, - ) - mamba_pool_host = MambaPoolHost( - params.req_to_token_pool.mamba_pool, - server_args.hicache_ratio, - server_args.hicache_size, - allocator_type=server_args.hicache_storage_backend, - layout=server_args.hicache_mem_layout, - ) - - full_layer_ids = sorted(hybrid_kv.full_attention_layer_id_mapping.keys()) - mamba_layer_ids = sorted(params.req_to_token_pool.mamba_map.keys()) - transfer_layer_num = len(set(full_layer_ids) | set(mamba_layer_ids)) full_layer_mapping = dict(hybrid_kv.full_attention_layer_id_mapping) mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map) - - def kv_layer_mapper(layer_id: int) -> Optional[int]: - if not 0 <= layer_id < transfer_layer_num: - return None - return full_layer_mapping.get(layer_id) - - def mamba_layer_mapper(layer_id: int) -> Optional[int]: - if not 0 <= layer_id < transfer_layer_num: - return None - return mamba_layer_mapping.get(layer_id) - - host_pool_group = HostPoolGroup( - [ - PoolEntry( - name=PoolName.KV, - host_pool=full_kv_pool_host, - device_pool=kvcache, - layer_mapper=kv_layer_mapper, - is_primary_index_anchor=True, - ), - PoolEntry( - name=PoolName.MAMBA, - host_pool=mamba_pool_host, - device_pool=params.req_to_token_pool.mamba_pool, - layer_mapper=mamba_layer_mapper, - host_evict_fn=mamba_cache.evict_mamba_host, - device_evict_fn=mamba_cache.evict_mamba, - ), - ] - ) - cache_controller = HybridCacheController( - params.token_to_kv_pool_allocator, - host_pool_group, - params.page_size, - params.tp_cache_group, + host_pool_group, cache_controller = build_hybrid_mamba_stack( + params=params, + server_args=server_args, + kv_pool=kvcache, + mamba_pool=params.req_to_token_pool.mamba_pool, + full_layer_mapping=full_layer_mapping, + mamba_layer_mapping=mamba_layer_mapping, + page_size=params.page_size, + tp_group=params.tp_cache_group, load_cache_event=load_cache_event, - write_policy=server_args.hicache_write_policy, - io_backend=server_args.hicache_io_backend, storage_backend=server_args.hicache_storage_backend, + use_mla=hybrid_kv.use_mla, + host_mamba_evict_fn=mamba_cache.evict_mamba_host, + device_mamba_evict_fn=mamba_cache.evict_mamba, prefetch_threshold=prefetch_threshold, model_name=server_args.served_model_name, storage_backend_extra_config=extra_config, @@ -194,12 +498,11 @@ def build_mamba_hybrid_stack( pp_size=params.pp_size, attn_cp_rank=params.attn_cp_rank, attn_cp_size=params.attn_cp_size, - transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, ) - mamba_cache.full_kv_pool_host = full_kv_pool_host - mamba_cache.mamba_pool_host = mamba_pool_host - mamba_cache.transfer_layer_num = transfer_layer_num + mamba_cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV) + mamba_cache.mamba_pool_host = host_pool_group.get_pool(PoolName.MAMBA) + mamba_cache.transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping) mamba_cache.host_pool_group = host_pool_group mamba_cache.cache_controller = cache_controller params.req_to_token_pool.register_layer_transfer_counter( @@ -207,10 +510,10 @@ def build_mamba_hybrid_stack( ) hybrid_kv.register_layer_transfer_counter(cache_controller.layer_done_counter) logger.info( - "Hybrid hierarchical cache: HostPoolGroup(KV + MAMBA), HybridCacheController, " + "Attached hybrid Mamba pool stack to HiMambaRadixCache: pools=KV + MAMBA, " "transfer_layer_num=%s", - transfer_layer_num, + mamba_cache.transfer_layer_num, ) except Exception: - logger.exception("build_mamba_hybrid_stack failed") + logger.exception("attach_hybrid_pool_to_mamba_cache failed") raise diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index da8308f9c..3132d9844 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -1715,6 +1715,9 @@ class HostPoolGroup: def get_ksize_per_token(self): return self.anchor_entry.host_pool.get_ksize_per_token() + def get_pool(self, name: PoolName): + return self.entry_map[name].host_pool + def get_page_buffer_meta(self, indices): return self.anchor_entry.host_pool.get_page_buffer_meta(indices)