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 249c970cc..b9b7a7074 100644 --- a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py @@ -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, diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index 392802451..7455ca09e 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -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( 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 2d5791ec1..a7af5e51a 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 @@ -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, diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index f513c299e..71042a21e 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -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,