[HiCache] Add PP Support with suffix pp rank (#15175)

Co-authored-by: Xuchun Shang <xuchun.shang@gmail.com>
Co-authored-by: ybyang <10629930+whybeyoung@users.noreply.github.com>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
Teng Ma
2026-01-03 13:49:21 +08:00
committed by GitHub
parent b23e7ed13c
commit bdde949619
6 changed files with 37 additions and 6 deletions

View File

@@ -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,

View File

@@ -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 (

View File

@@ -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

View File

@@ -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]

View File

@@ -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)

View File

@@ -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)