[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:
co-authored by
Xuchun Shang
ybyang
Shangming Cai
parent
b23e7ed13c
commit
bdde949619
@@ -259,6 +259,8 @@ class HiCacheController:
|
|||||||
prefetch_threshold: int = 256,
|
prefetch_threshold: int = 256,
|
||||||
model_name: Optional[str] = None,
|
model_name: Optional[str] = None,
|
||||||
storage_backend_extra_config: Optional[dict] = 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_allocator = token_to_kv_pool_allocator
|
||||||
self.mem_pool_device = token_to_kv_pool_allocator.get_kvcache()
|
self.mem_pool_device = token_to_kv_pool_allocator.get_kvcache()
|
||||||
@@ -267,6 +269,8 @@ class HiCacheController:
|
|||||||
self.page_size = page_size
|
self.page_size = page_size
|
||||||
self.io_backend = io_backend
|
self.io_backend = io_backend
|
||||||
self.enable_storage = False
|
self.enable_storage = False
|
||||||
|
self.pp_rank = pp_rank
|
||||||
|
self.pp_size = pp_size
|
||||||
|
|
||||||
if storage_backend is not None:
|
if storage_backend is not None:
|
||||||
self.storage_backend_type = storage_backend
|
self.storage_backend_type = storage_backend
|
||||||
@@ -394,6 +398,8 @@ class HiCacheController:
|
|||||||
return HiCacheStorageConfig(
|
return HiCacheStorageConfig(
|
||||||
tp_rank=self.tp_rank,
|
tp_rank=self.tp_rank,
|
||||||
tp_size=self.tp_size,
|
tp_size=self.tp_size,
|
||||||
|
pp_rank=self.pp_rank,
|
||||||
|
pp_size=self.pp_size,
|
||||||
is_mla_model=is_mla_backend,
|
is_mla_model=is_mla_backend,
|
||||||
is_page_first_layout=self.mem_pool_host.layout == "page_first",
|
is_page_first_layout=self.mem_pool_host.layout == "page_first",
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
|
|||||||
@@ -642,6 +642,8 @@ class Scheduler(
|
|||||||
enable_metrics=self.enable_metrics,
|
enable_metrics=self.enable_metrics,
|
||||||
enable_kv_cache_events=self.enable_kv_cache_events,
|
enable_kv_cache_events=self.enable_kv_cache_events,
|
||||||
enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(),
|
enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(),
|
||||||
|
pp_rank=self.pp_rank,
|
||||||
|
pp_size=self.pp_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -30,3 +30,6 @@ class CacheInitParams:
|
|||||||
# For SWAChunkCache
|
# For SWAChunkCache
|
||||||
sliding_window_size: Optional[int] = None
|
sliding_window_size: Optional[int] = None
|
||||||
attention_chunk_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:
|
class HiCacheStorageConfig:
|
||||||
tp_rank: int
|
tp_rank: int
|
||||||
tp_size: int
|
tp_size: int
|
||||||
|
pp_rank: int
|
||||||
|
pp_size: int
|
||||||
is_mla_model: bool
|
is_mla_model: bool
|
||||||
is_page_first_layout: bool
|
is_page_first_layout: bool
|
||||||
model_name: Optional[str]
|
model_name: Optional[str]
|
||||||
|
|||||||
@@ -68,6 +68,8 @@ class HiRadixCache(RadixCache):
|
|||||||
|
|
||||||
self.tp_group = params.tp_cache_group
|
self.tp_group = params.tp_cache_group
|
||||||
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_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 = server_args.hicache_storage_backend is not None
|
||||||
self.enable_storage_metrics = self.enable_storage and params.enable_metrics
|
self.enable_storage_metrics = self.enable_storage and params.enable_metrics
|
||||||
|
|
||||||
@@ -103,6 +105,8 @@ class HiRadixCache(RadixCache):
|
|||||||
prefetch_threshold=self.prefetch_threshold,
|
prefetch_threshold=self.prefetch_threshold,
|
||||||
model_name=server_args.served_model_name,
|
model_name=server_args.served_model_name,
|
||||||
storage_backend_extra_config=extra_config,
|
storage_backend_extra_config=extra_config,
|
||||||
|
pp_rank=self.pp_rank,
|
||||||
|
pp_size=self.pp_size,
|
||||||
)
|
)
|
||||||
if self.enable_storage_metrics:
|
if self.enable_storage_metrics:
|
||||||
# TODO: support pp
|
# TODO: support pp
|
||||||
@@ -110,6 +114,8 @@ class HiRadixCache(RadixCache):
|
|||||||
"storage_backend": server_args.hicache_storage_backend,
|
"storage_backend": server_args.hicache_storage_backend,
|
||||||
"tp_rank": self.cache_controller.tp_rank,
|
"tp_rank": self.cache_controller.tp_rank,
|
||||||
"dp_rank": self.cache_controller.dp_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)
|
self.storage_metrics_collector = StorageMetricsCollector(labels=labels)
|
||||||
|
|
||||||
|
|||||||
@@ -330,9 +330,21 @@ class MooncakeStore(HiCacheStorage):
|
|||||||
if storage_config is not None:
|
if storage_config is not None:
|
||||||
self.is_mla_backend = storage_config.is_mla_model
|
self.is_mla_backend = storage_config.is_mla_model
|
||||||
self.local_rank = storage_config.tp_rank
|
self.local_rank = storage_config.tp_rank
|
||||||
|
self.pp_rank = storage_config.pp_rank
|
||||||
|
self.pp_size = storage_config.pp_size
|
||||||
else:
|
else:
|
||||||
self.is_mla_backend = False
|
self.is_mla_backend = False
|
||||||
self.local_rank = 0
|
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:
|
except ValueError as e:
|
||||||
logger.error("Configuration loading failed: %s", 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)
|
ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices)
|
||||||
key_list = []
|
key_list = []
|
||||||
for key_ in keys:
|
for key_ in keys:
|
||||||
key_list.append(f"{key_}_{self.local_rank}_k")
|
key_list.append(f"{key_}_{self.mha_suffix}_k")
|
||||||
key_list.append(f"{key_}_{self.local_rank}_v")
|
key_list.append(f"{key_}_{self.mha_suffix}_v")
|
||||||
assert len(key_list) == len(ptr_list)
|
assert len(key_list) == len(ptr_list)
|
||||||
return key_list, ptr_list, element_size_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)
|
ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices)
|
||||||
key_list = []
|
key_list = []
|
||||||
for key_ in keys:
|
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)
|
assert len(key_list) == len(ptr_list)
|
||||||
return key_list, ptr_list, element_size_list
|
return key_list, ptr_list, element_size_list
|
||||||
|
|
||||||
@@ -610,13 +622,13 @@ class MooncakeStore(HiCacheStorage):
|
|||||||
self, keys, extra_info: Optional[HiCacheStorageExtraInfo] = None
|
self, keys, extra_info: Optional[HiCacheStorageExtraInfo] = None
|
||||||
) -> int:
|
) -> int:
|
||||||
if self.is_mla_backend:
|
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
|
key_multiplier = 1
|
||||||
else:
|
else:
|
||||||
query_keys = []
|
query_keys = []
|
||||||
for key in keys:
|
for key in keys:
|
||||||
query_keys.append(f"{key}_{self.local_rank}_k")
|
query_keys.append(f"{key}_{self.mha_suffix}_k")
|
||||||
query_keys.append(f"{key}_{self.local_rank}_v")
|
query_keys.append(f"{key}_{self.mha_suffix}_v")
|
||||||
key_multiplier = 2
|
key_multiplier = 2
|
||||||
|
|
||||||
exist_result = self._batch_exist(query_keys)
|
exist_result = self._batch_exist(query_keys)
|
||||||
|
|||||||
Reference in New Issue
Block a user