[Bugfix, NSA HiCache] Fix missing override_kv_cache_dim in attach_hybrid_nsa_pool_to_hiradix_cache (#25022)

Co-authored-by: Taegeon Um <taegeon.um@samsung.com>
Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
Taegeon Um
2026-05-13 11:45:08 +08:00
committed by GitHub
co-authored by Taegeon Um Zhangheng
parent 5ed9a494d0
commit 4e35c30cbe
3 changed files with 12 additions and 24 deletions
@@ -164,8 +164,6 @@ class HybridCacheController(BaseHiCacheController):
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,
transfer_layer_num: Optional[int] = None,
enable_storage_metrics: bool = False,
):
@@ -110,8 +110,6 @@ def build_kv_only_stack(
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)
@@ -149,8 +147,6 @@ def build_kv_only_stack(
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,
)
@@ -168,6 +164,8 @@ def build_hybrid_swa_stack(
page_size: int,
tp_group,
load_cache_event,
attn_cp_group: Optional["torch.distributed.ProcessGroup"] = None,
attn_tp_group: Optional["torch.distributed.ProcessGroup"] = None,
storage_backend: Optional[str],
use_mla: bool,
host_swa_evict_fn: Optional[Callable[[int], Any]] = None,
@@ -177,8 +175,6 @@ def build_hybrid_swa_stack(
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 | swa_layer_mapping)
@@ -225,6 +221,8 @@ def build_hybrid_swa_stack(
page_size,
tp_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
@@ -233,8 +231,6 @@ def build_hybrid_swa_stack(
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,
)
@@ -263,8 +259,6 @@ def build_hybrid_mamba_stack(
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)
@@ -317,8 +311,6 @@ def build_hybrid_mamba_stack(
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,
)
@@ -346,8 +338,6 @@ def build_shared_anchor_stack(
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)
@@ -394,8 +384,6 @@ def build_shared_anchor_stack(
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,
)
@@ -505,6 +493,8 @@ def attach_hybrid_pool_to_unified_cache(
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),
@@ -538,8 +528,11 @@ def attach_hybrid_pool_to_unified_cache(
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,
shared_host_pool_factory=lambda kv_host_pool: NSAIndexerPoolHost(
full_kv_pool,
kv_host_pool,
@@ -548,8 +541,6 @@ def attach_hybrid_pool_to_unified_cache(
),
pp_rank=params.pp_rank,
pp_size=params.pp_size,
attn_cp_rank=params.attn_cp_rank,
attn_cp_size=params.attn_cp_size,
)
cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV)
cache.host_pool_group = host_pool_group
@@ -645,6 +636,7 @@ def attach_hybrid_nsa_pool_to_hiradix_cache(
attn_tp_group=attn_tp_group,
storage_backend=server_args.hicache_storage_backend,
use_mla=True,
override_kv_cache_dim=kv.kv_cache_dim,
prefetch_threshold=prefetch_threshold,
shared_host_pool_factory=lambda kv_host_pool: NSAIndexerPoolHost(
kv,
@@ -656,8 +648,6 @@ def attach_hybrid_nsa_pool_to_hiradix_cache(
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,
enable_storage_metrics=enable_storage_metrics,
)
radix_cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV)
@@ -715,8 +705,6 @@ def attach_hybrid_pool_to_mamba_cache(
storage_backend_extra_config=extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
attn_cp_rank=params.attn_cp_rank,
attn_cp_size=params.attn_cp_size,
enable_storage_metrics=enable_storage_metrics,
)
mamba_cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV)
@@ -319,6 +319,8 @@ class UnifiedRadixCache(BasePrefixCache):
params,
server_args,
load_cache_event=self.load_cache_event,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
)
# State initialization