From bdde94961967f3dfcd7c3e4000baf18953cb1229 Mon Sep 17 00:00:00 2001 From: Teng Ma Date: Sat, 3 Jan 2026 13:49:21 +0800 Subject: [PATCH] [HiCache] Add PP Support with suffix pp rank (#15175) Co-authored-by: Xuchun Shang Co-authored-by: ybyang <10629930+whybeyoung@users.noreply.github.com> Co-authored-by: Shangming Cai --- .../sglang/srt/managers/cache_controller.py | 6 +++++ python/sglang/srt/managers/scheduler.py | 2 ++ .../sglang/srt/mem_cache/cache_init_params.py | 3 +++ .../sglang/srt/mem_cache/hicache_storage.py | 2 ++ python/sglang/srt/mem_cache/hiradix_cache.py | 6 +++++ .../storage/mooncake_store/mooncake_store.py | 24 ++++++++++++++----- 6 files changed, 37 insertions(+), 6 deletions(-) diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index d90c76b94..0caa63cde 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -259,6 +259,8 @@ class HiCacheController: prefetch_threshold: int = 256, model_name: Optional[str] = None, storage_backend_extra_config: Optional[dict] = None, + pp_rank: int = 0, + pp_size: int = 1, ): self.mem_pool_device_allocator = token_to_kv_pool_allocator self.mem_pool_device = token_to_kv_pool_allocator.get_kvcache() @@ -267,6 +269,8 @@ class HiCacheController: self.page_size = page_size self.io_backend = io_backend self.enable_storage = False + self.pp_rank = pp_rank + self.pp_size = pp_size if storage_backend is not None: self.storage_backend_type = storage_backend @@ -394,6 +398,8 @@ class HiCacheController: return HiCacheStorageConfig( tp_rank=self.tp_rank, tp_size=self.tp_size, + pp_rank=self.pp_rank, + pp_size=self.pp_size, is_mla_model=is_mla_backend, is_page_first_layout=self.mem_pool_host.layout == "page_first", model_name=model_name, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 1db8a3df8..b45360e9f 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -642,6 +642,8 @@ class Scheduler( enable_metrics=self.enable_metrics, enable_kv_cache_events=self.enable_kv_cache_events, enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(), + pp_rank=self.pp_rank, + pp_size=self.pp_size, ) if ( diff --git a/python/sglang/srt/mem_cache/cache_init_params.py b/python/sglang/srt/mem_cache/cache_init_params.py index d410dda0e..50430a5eb 100644 --- a/python/sglang/srt/mem_cache/cache_init_params.py +++ b/python/sglang/srt/mem_cache/cache_init_params.py @@ -30,3 +30,6 @@ class CacheInitParams: # For SWAChunkCache sliding_window_size: Optional[int] = None attention_chunk_size: Optional[int] = None + + pp_rank: int = 0 + pp_size: int = 1 diff --git a/python/sglang/srt/mem_cache/hicache_storage.py b/python/sglang/srt/mem_cache/hicache_storage.py index acaaab32b..38df15262 100644 --- a/python/sglang/srt/mem_cache/hicache_storage.py +++ b/python/sglang/srt/mem_cache/hicache_storage.py @@ -47,6 +47,8 @@ def hash_str_to_int64(hash_str: str) -> int: class HiCacheStorageConfig: tp_rank: int tp_size: int + pp_rank: int + pp_size: int is_mla_model: bool is_page_first_layout: bool model_name: Optional[str] diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index a38f1ac5f..f6cfca8b6 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -68,6 +68,8 @@ class HiRadixCache(RadixCache): self.tp_group = params.tp_cache_group 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.enable_storage = server_args.hicache_storage_backend is not None self.enable_storage_metrics = self.enable_storage and params.enable_metrics @@ -103,6 +105,8 @@ class HiRadixCache(RadixCache): prefetch_threshold=self.prefetch_threshold, model_name=server_args.served_model_name, storage_backend_extra_config=extra_config, + pp_rank=self.pp_rank, + pp_size=self.pp_size, ) if self.enable_storage_metrics: # TODO: support pp @@ -110,6 +114,8 @@ class HiRadixCache(RadixCache): "storage_backend": server_args.hicache_storage_backend, "tp_rank": self.cache_controller.tp_rank, "dp_rank": self.cache_controller.dp_rank, + "pp_rank": self.cache_controller.pp_rank, + "pp_size": self.cache_controller.pp_size, } self.storage_metrics_collector = StorageMetricsCollector(labels=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 b015db46c..0388e9a0c 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 @@ -330,9 +330,21 @@ class MooncakeStore(HiCacheStorage): if storage_config is not None: self.is_mla_backend = storage_config.is_mla_model self.local_rank = storage_config.tp_rank + self.pp_rank = storage_config.pp_rank + self.pp_size = storage_config.pp_size else: self.is_mla_backend = False self.local_rank = 0 + self.pp_rank = 0 + self.pp_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}" + else: + self.mha_suffix = f"{self.local_rank}" + self.mla_suffix = "" except ValueError as e: logger.error("Configuration loading failed: %s", e) @@ -406,8 +418,8 @@ class MooncakeStore(HiCacheStorage): ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices) key_list = [] for key_ in keys: - key_list.append(f"{key_}_{self.local_rank}_k") - key_list.append(f"{key_}_{self.local_rank}_v") + key_list.append(f"{key_}_{self.mha_suffix}_k") + key_list.append(f"{key_}_{self.mha_suffix}_v") assert len(key_list) == len(ptr_list) return key_list, ptr_list, element_size_list @@ -415,7 +427,7 @@ class MooncakeStore(HiCacheStorage): ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices) key_list = [] for key_ in keys: - key_list.append(f"{key_}_k") + key_list.append(f"{key_}_{self.mla_suffix}_k") assert len(key_list) == len(ptr_list) return key_list, ptr_list, element_size_list @@ -610,13 +622,13 @@ class MooncakeStore(HiCacheStorage): self, keys, extra_info: Optional[HiCacheStorageExtraInfo] = None ) -> int: if self.is_mla_backend: - query_keys = [f"{key}_k" for key in keys] + query_keys = [f"{key}_{self.mla_suffix}_k" for key in keys] key_multiplier = 1 else: query_keys = [] for key in keys: - query_keys.append(f"{key}_{self.local_rank}_k") - query_keys.append(f"{key}_{self.local_rank}_v") + query_keys.append(f"{key}_{self.mha_suffix}_k") + query_keys.append(f"{key}_{self.mha_suffix}_v") key_multiplier = 2 exist_result = self._batch_exist(query_keys)