[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:
@@ -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,
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user