[HiCache] Remove redundant parameters of build_xxx_stack and others (#31308)

Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
Chao Shi
2026-07-21 17:54:06 +08:00
committed by GitHub
co-authored by Zhangheng
parent dcd9014f15
commit f369a820d4
4 changed files with 36 additions and 150 deletions
@@ -153,8 +153,6 @@ class HiMambaRadixCache(MambaRadixCache):
prefetch_threshold=prefetch_threshold,
load_cache_event=self.load_cache_event,
enable_storage_metrics=self.enable_storage_metrics,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
)
self._apply_storage_runtime_config(
storage_backend=server_args.hicache_storage_backend,
@@ -140,8 +140,6 @@ class HiRadixCache(RadixCache):
prefetch_threshold=prefetch_threshold,
enable_storage_metrics=self.enable_storage_metrics,
load_cache_event=self.load_cache_event,
attn_cp_group=self.attn_cp_group,
attn_tp_group=self.attn_tp_group,
)
elif isinstance(self.kv_cache, MiniMaxSparseKVPool):
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
@@ -156,8 +154,6 @@ class HiRadixCache(RadixCache):
prefetch_threshold=prefetch_threshold,
enable_storage_metrics=self.enable_storage_metrics,
load_cache_event=self.load_cache_event,
attn_cp_group=self.attn_cp_group,
attn_tp_group=self.attn_tp_group,
)
else:
self.cache_controller = HiCacheController(
@@ -109,12 +109,7 @@ def build_kv_only_stack(
server_args: ServerArgs,
kv_pool: Any,
full_layer_mapping: dict[int, int],
page_size: int,
tp_group,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
use_mla: bool,
override_kv_cache_dim: Optional[int] = None,
@@ -126,7 +121,7 @@ def build_kv_only_stack(
transfer_layer_num = len(full_layer_mapping)
kv_host_pool = build_kv_host_pool(
kv_pool=kv_pool,
page_size=page_size,
page_size=params.page_size,
server_args=server_args,
use_mla=use_mla,
override_kv_cache_dim=override_kv_cache_dim,
@@ -145,12 +140,12 @@ def build_kv_only_stack(
cache_controller = HybridCacheController(
params.token_to_kv_pool_allocator,
host_pool_group,
page_size,
tp_group,
params.page_size,
params.tp_cache_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
@@ -171,12 +166,7 @@ def build_hybrid_swa_stack(
swa_kv_pool: Any,
full_layer_mapping: dict[int, int],
swa_layer_mapping: dict[int, int],
page_size: int,
tp_group,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
use_mla: bool,
host_swa_evict_fn: Optional[Callable[[int], Any]] = None,
@@ -189,13 +179,13 @@ def build_hybrid_swa_stack(
transfer_layer_num = len(full_layer_mapping | swa_layer_mapping)
kv_host_pool = build_kv_host_pool(
kv_pool=full_kv_pool,
page_size=page_size,
page_size=params.page_size,
server_args=server_args,
use_mla=use_mla,
)
swa_host_pool = build_kv_host_pool(
kv_pool=swa_kv_pool,
page_size=page_size,
page_size=params.page_size,
server_args=server_args,
use_mla=use_mla,
)
@@ -227,12 +217,12 @@ def build_hybrid_swa_stack(
cache_controller = HybridCacheController(
params.token_to_kv_pool_allocator,
host_pool_group,
page_size,
tp_group,
params.page_size,
params.tp_cache_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
@@ -286,12 +276,7 @@ def build_deepseek_v4_hicache_stack(
params: CacheInitParams,
server_args: ServerArgs,
kvcache: Any,
page_size: int,
tp_group,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
host_swa_evict_fn: Optional[Callable[[int], Any]] = None,
device_swa_evict_fn: Optional[Callable[[int], Any]] = None,
@@ -300,6 +285,7 @@ def build_deepseek_v4_hicache_stack(
storage_backend_extra_config: Optional[dict] = None,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
page_size = params.page_size
transfer_layer_num = kvcache.end_layer - kvcache.start_layer
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
@@ -500,11 +486,11 @@ def build_deepseek_v4_hicache_stack(
params.token_to_kv_pool_allocator,
host_pool_group,
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,
pp_group=pp_group,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
@@ -525,12 +511,7 @@ def build_hybrid_mamba_stack(
mamba_pool: Any,
full_layer_mapping: dict[int, int],
mamba_layer_mapping: dict[int, int],
page_size: int,
tp_group,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
use_mla: bool,
host_mamba_evict_fn: Optional[Callable[[int], Any]] = None,
@@ -544,7 +525,7 @@ def build_hybrid_mamba_stack(
mamba_allocator = params.req_to_token_pool.mamba_allocator
kv_host_pool = build_kv_host_pool(
kv_pool=kv_pool,
page_size=page_size,
page_size=params.page_size,
server_args=server_args,
use_mla=use_mla,
)
@@ -580,12 +561,12 @@ def build_hybrid_mamba_stack(
cache_controller = HybridCacheController(
params.token_to_kv_pool_allocator,
host_pool_group,
page_size,
tp_group,
params.page_size,
params.tp_cache_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
@@ -709,12 +690,7 @@ def build_anchor_sidecar_stack(
kv_pool: Any,
sidecar_pool_name: PoolName,
full_layer_mapping: dict[int, int],
page_size: int,
tp_group,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
use_mla: bool,
override_kv_cache_dim: Optional[int] = None,
@@ -727,7 +703,7 @@ def build_anchor_sidecar_stack(
transfer_layer_num = len(full_layer_mapping)
kv_host_pool = build_kv_host_pool(
kv_pool=kv_pool,
page_size=page_size,
page_size=params.page_size,
server_args=server_args,
use_mla=use_mla,
override_kv_cache_dim=override_kv_cache_dim,
@@ -754,12 +730,12 @@ def build_anchor_sidecar_stack(
cache_controller = HybridCacheController(
params.token_to_kv_pool_allocator,
host_pool_group,
page_size,
tp_group,
params.page_size,
params.tp_cache_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
pp_group=params.pp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
@@ -804,8 +780,6 @@ class StackStrategy:
params: CacheInitParams,
server_args: ServerArgs,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
prefetch_threshold: int = 256,
@@ -834,8 +808,6 @@ class _DeepSeekV4Strategy(StackStrategy):
params,
server_args,
load_cache_event,
attn_cp_group=None,
attn_tp_group=None,
storage_backend=None,
storage_backend_extra_config=None,
prefetch_threshold=256,
@@ -848,12 +820,7 @@ class _DeepSeekV4Strategy(StackStrategy):
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,
pp_group=params.pp_cache_group,
storage_backend=storage_backend,
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)),
@@ -917,8 +884,6 @@ class _MambaStrategy(StackStrategy):
params,
server_args,
load_cache_event,
attn_cp_group=None,
attn_tp_group=None,
storage_backend=None,
storage_backend_extra_config=None,
prefetch_threshold=256,
@@ -936,12 +901,7 @@ class _MambaStrategy(StackStrategy):
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,
pp_group=params.pp_cache_group,
storage_backend=storage_backend,
use_mla=kvcache.use_mla,
host_mamba_evict_fn=lambda n: cache.evict_host(n, ComponentType.MAMBA),
@@ -993,8 +953,6 @@ class _SwaStrategy(StackStrategy):
params,
server_args,
load_cache_event,
attn_cp_group=None,
attn_tp_group=None,
storage_backend=None,
storage_backend_extra_config=None,
prefetch_threshold=256,
@@ -1011,12 +969,7 @@ class _SwaStrategy(StackStrategy):
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,
pp_group=params.pp_cache_group,
storage_backend=storage_backend,
use_mla=False,
host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA),
@@ -1129,8 +1082,6 @@ class _DsaStrategy(StackStrategy):
params,
server_args,
load_cache_event,
attn_cp_group=None,
attn_tp_group=None,
storage_backend=None,
storage_backend_extra_config=None,
prefetch_threshold=256,
@@ -1148,11 +1099,7 @@ class _DsaStrategy(StackStrategy):
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=storage_backend,
use_mla=use_mla,
override_kv_cache_dim=full_kv_pool.kv_cache_dim,
@@ -1200,8 +1147,6 @@ class _MiniMaxSparseStrategy(StackStrategy):
params,
server_args,
load_cache_event,
attn_cp_group=None,
attn_tp_group=None,
storage_backend=None,
storage_backend_extra_config=None,
prefetch_threshold=256,
@@ -1212,17 +1157,11 @@ class _MiniMaxSparseStrategy(StackStrategy):
params=params,
server_args=server_args,
sparse_pool=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=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
sidecars = []
@@ -1280,8 +1219,6 @@ class _PlainKvStrategy(StackStrategy):
params,
server_args,
load_cache_event,
attn_cp_group=None,
attn_tp_group=None,
storage_backend=None,
storage_backend_extra_config=None,
prefetch_threshold=256,
@@ -1298,12 +1235,7 @@ class _PlainKvStrategy(StackStrategy):
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,
pp_group=params.pp_cache_group,
storage_backend=storage_backend,
use_mla=use_mla,
prefetch_threshold=prefetch_threshold,
@@ -1386,8 +1318,6 @@ def attach_hybrid_pool_to_unified_cache(
server_args: ServerArgs,
*,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str] = None,
storage_extra_config: Optional[dict] = None,
storage_prefetch_threshold: int = 256,
@@ -1403,8 +1333,6 @@ def attach_hybrid_pool_to_unified_cache(
params=params,
server_args=server_args,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
storage_backend=storage_backend,
storage_backend_extra_config=storage_extra_config,
prefetch_threshold=storage_prefetch_threshold,
@@ -1422,23 +1350,17 @@ def build_minimax_sparse_hicache_stack(
params: CacheInitParams,
server_args: ServerArgs,
sparse_pool: Any,
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],
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
"""KV (main_pool) + INDEXER (index_k_pool) host stack for MiniMax M3 sparse."""
# Mappings are stage-local keyed (controller iterates 0..transfer_layer_num).
# PP>1 stays gated below pending end-to-end validation of the sparse host path.
if pp_size > 1:
if params.pp_size > 1:
raise NotImplementedError(
"MiniMax-M3 sparse HiCache does not support pipeline parallelism "
"(pp_size>1) yet."
@@ -1458,7 +1380,7 @@ def build_minimax_sparse_hicache_stack(
kv_host_pool = build_kv_host_pool(
kv_pool=main_pool,
page_size=page_size,
page_size=params.page_size,
server_args=server_args,
use_mla=False,
)
@@ -1498,11 +1420,11 @@ def build_minimax_sparse_hicache_stack(
cache_controller = HybridCacheController(
params.token_to_kv_pool_allocator,
host_pool_group,
page_size,
tp_group,
params.page_size,
params.tp_cache_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
attn_cp_group=params.attn_cp_cache_group,
attn_tp_group=params.attn_tp_cache_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
@@ -1525,8 +1447,6 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
prefetch_threshold: int,
enable_storage_metrics: bool,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
) -> None:
"""Attach HostPoolGroup (KV + index K) + HybridCacheController for HiRadixCache."""
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
@@ -1554,18 +1474,12 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
full_layer_mapping={
layer_id: layer_id for layer_id in range(main_pool.layer_num)
},
page_size=radix_cache.page_size,
tp_group=radix_cache.tp_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
storage_backend=server_args.hicache_storage_backend,
use_mla=False,
prefetch_threshold=prefetch_threshold,
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,
enable_storage_metrics=enable_storage_metrics,
)
pools_desc = "KV"
@@ -1574,17 +1488,11 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
params=params,
server_args=server_args,
sparse_pool=sparse_pool,
page_size=radix_cache.page_size,
tp_group=radix_cache.tp_group,
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
storage_backend=server_args.hicache_storage_backend,
prefetch_threshold=prefetch_threshold,
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,
enable_storage_metrics=enable_storage_metrics,
)
pools_desc = "KV + INDEXER(k-only)"
@@ -1614,8 +1522,6 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
prefetch_threshold: int,
enable_storage_metrics: bool,
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
) -> None:
"""Attach HostPoolGroup (KV + indexer) + HybridCacheController for HiRadixCache.
@@ -1630,12 +1536,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
kv_pool=kv,
sidecar_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,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=radix_cache.pp_group,
storage_backend=server_args.hicache_storage_backend,
use_mla=True,
override_kv_cache_dim=kv.kv_cache_dim,
@@ -1672,8 +1573,6 @@ def attach_hybrid_pool_to_mamba_cache(
prefetch_threshold: int,
load_cache_event,
enable_storage_metrics: bool = False,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
) -> None:
"""Attach HostPoolGroup (KV + Mamba) + HybridCacheController for HiMambaRadixCache.
@@ -1691,12 +1590,7 @@ def attach_hybrid_pool_to_mamba_cache(
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,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=params.pp_cache_group,
storage_backend=server_args.hicache_storage_backend,
use_mla=hybrid_kv.use_mla,
host_mamba_evict_fn=mamba_cache.evict_mamba_host,
@@ -540,8 +540,6 @@ class UnifiedRadixCache(KVCacheEventMixin, 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,
storage_backend=storage_backend,
storage_extra_config=storage_extra_config,
storage_prefetch_threshold=storage_prefetch_threshold,