[HiCache] Add CP support for HiCache (#20977)
Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -261,6 +261,8 @@ class HiCacheController:
|
|||||||
storage_backend_extra_config: Optional[dict] = None,
|
storage_backend_extra_config: Optional[dict] = None,
|
||||||
pp_rank: int = 0,
|
pp_rank: int = 0,
|
||||||
pp_size: int = 1,
|
pp_size: int = 1,
|
||||||
|
attn_cp_rank: int = 0,
|
||||||
|
attn_cp_size: int = 1,
|
||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
):
|
):
|
||||||
self.tp_group = tp_group
|
self.tp_group = tp_group
|
||||||
@@ -280,6 +282,8 @@ class HiCacheController:
|
|||||||
self.storage_backend_type = None
|
self.storage_backend_type = None
|
||||||
self.pp_rank = pp_rank
|
self.pp_rank = pp_rank
|
||||||
self.pp_size = pp_size
|
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
|
self.enable_storage_metrics = enable_storage_metrics
|
||||||
|
|
||||||
# Default storage page IO functions (may be overridden by attach).
|
# Default storage page IO functions (may be overridden by attach).
|
||||||
@@ -611,6 +615,8 @@ class HiCacheController:
|
|||||||
tp_size=self.tp_size,
|
tp_size=self.tp_size,
|
||||||
pp_rank=self.pp_rank,
|
pp_rank=self.pp_rank,
|
||||||
pp_size=self.pp_size,
|
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,
|
is_mla_model=is_mla_backend,
|
||||||
enable_storage_metrics=self.enable_storage_metrics,
|
enable_storage_metrics=self.enable_storage_metrics,
|
||||||
is_page_first_layout=self.mem_pool_host.layout == "page_first",
|
is_page_first_layout=self.mem_pool_host.layout == "page_first",
|
||||||
|
|||||||
@@ -794,6 +794,8 @@ class Scheduler(
|
|||||||
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_rank=self.pp_rank,
|
||||||
pp_size=self.pp_size,
|
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,
|
chunked_prefill_size=effective_chunked_prefill_size,
|
||||||
sliding_window_size=self.sliding_window_size,
|
sliding_window_size=self.sliding_window_size,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -30,6 +30,9 @@ class CacheInitParams:
|
|||||||
pp_rank: int = 0
|
pp_rank: int = 0
|
||||||
pp_size: int = 1
|
pp_size: int = 1
|
||||||
|
|
||||||
|
attn_cp_rank: int = 0
|
||||||
|
attn_cp_size: int = 1
|
||||||
|
|
||||||
chunked_prefill_size: Optional[int] = None
|
chunked_prefill_size: Optional[int] = None
|
||||||
|
|
||||||
sliding_window_size: Optional[int] = None
|
sliding_window_size: Optional[int] = None
|
||||||
|
|||||||
@@ -19,6 +19,8 @@ class HiCacheStorageConfig:
|
|||||||
tp_size: int
|
tp_size: int
|
||||||
pp_rank: int
|
pp_rank: int
|
||||||
pp_size: int
|
pp_size: int
|
||||||
|
attn_cp_rank: int
|
||||||
|
attn_cp_size: int
|
||||||
is_mla_model: bool
|
is_mla_model: bool
|
||||||
enable_storage_metrics: bool
|
enable_storage_metrics: bool
|
||||||
is_page_first_layout: bool
|
is_page_first_layout: bool
|
||||||
|
|||||||
@@ -94,6 +94,8 @@ class HiRadixCache(RadixCache):
|
|||||||
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_rank = params.pp_rank
|
||||||
self.pp_size = params.pp_size
|
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 = 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
|
||||||
self.extra_metric_labels = server_args.extra_metric_labels
|
self.extra_metric_labels = server_args.extra_metric_labels
|
||||||
@@ -126,6 +128,8 @@ class HiRadixCache(RadixCache):
|
|||||||
storage_backend_extra_config=extra_config,
|
storage_backend_extra_config=extra_config,
|
||||||
pp_rank=self.pp_rank,
|
pp_rank=self.pp_rank,
|
||||||
pp_size=self.pp_size,
|
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,
|
enable_storage_metrics=self.enable_storage_metrics,
|
||||||
)
|
)
|
||||||
self._apply_storage_runtime_config(
|
self._apply_storage_runtime_config(
|
||||||
@@ -204,6 +208,8 @@ class HiRadixCache(RadixCache):
|
|||||||
"dp_rank": self.cache_controller.dp_rank,
|
"dp_rank": self.cache_controller.dp_rank,
|
||||||
"pp_rank": self.cache_controller.pp_rank,
|
"pp_rank": self.cache_controller.pp_rank,
|
||||||
"pp_size": self.cache_controller.pp_size,
|
"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:
|
if extra_metric_labels:
|
||||||
labels.update(extra_metric_labels)
|
labels.update(extra_metric_labels)
|
||||||
|
|||||||
@@ -394,17 +394,24 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
self.local_rank = storage_config.tp_rank
|
self.local_rank = storage_config.tp_rank
|
||||||
self.pp_rank = storage_config.pp_rank
|
self.pp_rank = storage_config.pp_rank
|
||||||
self.pp_size = storage_config.pp_size
|
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
|
self.enable_storage_metrics = storage_config.enable_storage_metrics
|
||||||
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_rank = 0
|
||||||
self.pp_size = 1
|
self.pp_size = 1
|
||||||
|
self.attn_cp_rank = 0
|
||||||
|
self.attn_cp_size = 1
|
||||||
|
|
||||||
self.enable_pp = self.pp_size > 1
|
self.enable_pp = self.pp_size > 1
|
||||||
if self.enable_pp:
|
self.enable_cp = self.attn_cp_size > 1
|
||||||
self.mha_suffix = f"{self.local_rank}_{self.pp_rank}"
|
if self.enable_pp or self.enable_cp:
|
||||||
self.mla_suffix = f"{self.pp_rank}"
|
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:
|
else:
|
||||||
self.mha_suffix = f"{self.local_rank}"
|
self.mha_suffix = f"{self.local_rank}"
|
||||||
self.mla_suffix = ""
|
self.mla_suffix = ""
|
||||||
@@ -417,9 +424,10 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
)
|
)
|
||||||
base_rank = self.local_rank * self.split_factor
|
base_rank = self.local_rank * self.split_factor
|
||||||
target_ranks = [base_rank + i for i in range(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 = [
|
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:
|
else:
|
||||||
self.mha_suffix = [f"{rank}" for rank in target_ranks]
|
self.mha_suffix = [f"{rank}" for rank in target_ranks]
|
||||||
|
|||||||
Reference in New Issue
Block a user