From 1c76f322df5c3ee5887ef1e979d760041ffab139 Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Fri, 10 Apr 2026 17:52:51 +0800 Subject: [PATCH] [HiCache] Add CP support for HiCache (#20977) Signed-off-by: Shangming Cai --- python/sglang/srt/managers/cache_controller.py | 6 ++++++ python/sglang/srt/managers/scheduler.py | 2 ++ .../sglang/srt/mem_cache/cache_init_params.py | 3 +++ python/sglang/srt/mem_cache/hicache_storage.py | 2 ++ python/sglang/srt/mem_cache/hiradix_cache.py | 6 ++++++ .../storage/mooncake_store/mooncake_store.py | 18 +++++++++++++----- 6 files changed, 32 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 6a735a542..ddd34fd16 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -261,6 +261,8 @@ class HiCacheController: 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, ): self.tp_group = tp_group @@ -280,6 +282,8 @@ class HiCacheController: self.storage_backend_type = None self.pp_rank = pp_rank self.pp_size = pp_size + self.attn_cp_rank = attn_cp_rank + self.attn_cp_size = attn_cp_size self.enable_storage_metrics = enable_storage_metrics # Default storage page IO functions (may be overridden by attach). @@ -611,6 +615,8 @@ class HiCacheController: tp_size=self.tp_size, pp_rank=self.pp_rank, pp_size=self.pp_size, + attn_cp_rank=self.attn_cp_rank, + attn_cp_size=self.attn_cp_size, is_mla_model=is_mla_backend, enable_storage_metrics=self.enable_storage_metrics, is_page_first_layout=self.mem_pool_host.layout == "page_first", diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index fb4d42ef6..6a69c2b02 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -794,6 +794,8 @@ class Scheduler( enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(), pp_rank=self.pp_rank, pp_size=self.pp_size, + attn_cp_rank=self.attn_cp_rank, + attn_cp_size=self.attn_cp_size, chunked_prefill_size=effective_chunked_prefill_size, sliding_window_size=self.sliding_window_size, ) diff --git a/python/sglang/srt/mem_cache/cache_init_params.py b/python/sglang/srt/mem_cache/cache_init_params.py index 6de3f984f..d8731160c 100644 --- a/python/sglang/srt/mem_cache/cache_init_params.py +++ b/python/sglang/srt/mem_cache/cache_init_params.py @@ -30,6 +30,9 @@ class CacheInitParams: pp_rank: int = 0 pp_size: int = 1 + attn_cp_rank: int = 0 + attn_cp_size: int = 1 + chunked_prefill_size: Optional[int] = None sliding_window_size: Optional[int] = None diff --git a/python/sglang/srt/mem_cache/hicache_storage.py b/python/sglang/srt/mem_cache/hicache_storage.py index f1e5520f4..b2195c4c3 100644 --- a/python/sglang/srt/mem_cache/hicache_storage.py +++ b/python/sglang/srt/mem_cache/hicache_storage.py @@ -19,6 +19,8 @@ class HiCacheStorageConfig: tp_size: int pp_rank: int pp_size: int + attn_cp_rank: int + attn_cp_size: int is_mla_model: bool enable_storage_metrics: bool is_page_first_layout: bool diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index 47b91fdab..b4a5eb5b8 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -94,6 +94,8 @@ class HiRadixCache(RadixCache): self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group) self.pp_rank = params.pp_rank self.pp_size = params.pp_size + self.attn_cp_rank = params.attn_cp_rank + self.attn_cp_size = params.attn_cp_size self.enable_storage = server_args.hicache_storage_backend is not None self.enable_storage_metrics = self.enable_storage and params.enable_metrics self.extra_metric_labels = server_args.extra_metric_labels @@ -126,6 +128,8 @@ class HiRadixCache(RadixCache): storage_backend_extra_config=extra_config, pp_rank=self.pp_rank, pp_size=self.pp_size, + attn_cp_rank=self.attn_cp_rank, + attn_cp_size=self.attn_cp_size, enable_storage_metrics=self.enable_storage_metrics, ) self._apply_storage_runtime_config( @@ -204,6 +208,8 @@ class HiRadixCache(RadixCache): "dp_rank": self.cache_controller.dp_rank, "pp_rank": self.cache_controller.pp_rank, "pp_size": self.cache_controller.pp_size, + "attn_cp_rank": self.cache_controller.attn_cp_rank, + "attn_cp_size": self.cache_controller.attn_cp_size, } if extra_metric_labels: labels.update(extra_metric_labels) diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py index 2c815fd7e..1923d5baf 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py @@ -394,17 +394,24 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): self.local_rank = storage_config.tp_rank self.pp_rank = storage_config.pp_rank self.pp_size = storage_config.pp_size + self.attn_cp_rank = storage_config.attn_cp_rank + self.attn_cp_size = storage_config.attn_cp_size self.enable_storage_metrics = storage_config.enable_storage_metrics else: self.is_mla_backend = False self.local_rank = 0 self.pp_rank = 0 self.pp_size = 1 + self.attn_cp_rank = 0 + self.attn_cp_size = 1 self.enable_pp = self.pp_size > 1 - if self.enable_pp: - self.mha_suffix = f"{self.local_rank}_{self.pp_rank}" - self.mla_suffix = f"{self.pp_rank}" + self.enable_cp = self.attn_cp_size > 1 + if self.enable_pp or self.enable_cp: + self.mha_suffix = ( + f"{self.local_rank}_{self.pp_rank}_{self.attn_cp_rank}" + ) + self.mla_suffix = f"{self.pp_rank}_{self.attn_cp_rank}" else: self.mha_suffix = f"{self.local_rank}" self.mla_suffix = "" @@ -417,9 +424,10 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): ) base_rank = self.local_rank * self.split_factor target_ranks = [base_rank + i for i in range(self.split_factor)] - if self.enable_pp: + if self.enable_pp or self.enable_cp: self.mha_suffix = [ - f"{rank}_{self.pp_rank}" for rank in target_ranks + f"{rank}_{self.pp_rank}_{self.attn_cp_rank}" + for rank in target_ranks ] else: self.mha_suffix = [f"{rank}" for rank in target_ranks]