[HiCache] Fix the compatibility between PP and HiCache (L2). (#27285)

Co-authored-by: ybyang <ybyang7@iflytek.com>
Co-authored-by: hzh0425 <hzh0425@apache.org>
Co-authored-by: shangmingc <csmthu@gmail.com>
Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
Chao Shi
2026-06-06 16:57:23 +08:00
committed by GitHub
co-authored by ybyang hzh0425 shangmingc 晟海
parent aa5213abb1
commit 42fe025280
10 changed files with 313 additions and 76 deletions
@@ -34,6 +34,8 @@ if TYPE_CHECKING:
from sglang.srt.mem_cache.memory_pool_host import HostKVCache
from sglang.srt.distributed import (
get_pipeline_model_parallel_rank,
get_pipeline_model_parallel_world_size,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
@@ -254,19 +256,19 @@ class HiCacheController:
load_cache_event: threading.Event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
write_policy: str = "write_through_selective",
io_backend: str = "",
storage_backend: Optional[str] = None,
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
):
self.tp_group = tp_group
self.attn_cp_group = attn_cp_group
self.attn_tp_group = attn_tp_group
self.pp_group = pp_group
self.prefetch_sync_groups: List[torch.distributed.ProcessGroup] = []
self.mem_pool_device_allocator = token_to_kv_pool_allocator
mem_pool_device = token_to_kv_pool_allocator.get_kvcache()
@@ -282,8 +284,6 @@ class HiCacheController:
self.enable_storage = False
self.storage_backend = None
self.storage_backend_type = None
self.pp_rank = pp_rank
self.pp_size = pp_size
self.enable_storage_metrics = enable_storage_metrics
# Draft KV pool support (best-effort piggyback on target L2/L3 ops).
@@ -624,6 +624,9 @@ class HiCacheController:
self.tp_size = get_tensor_model_parallel_world_size()
self.dp_rank = 0
self.pp_rank = get_pipeline_model_parallel_rank()
self.pp_size = get_pipeline_model_parallel_world_size()
# Currently, NPUMLATokenToKVPool is the subclass of MLATokenToKVPool.
# DeepSeekV4TokenToKVPool has compressed MLA-style rank-replicated cache
# data. storage only needs rank 0 to write it back.
+1
View File
@@ -443,6 +443,7 @@ class Scheduler(
),
ps=self.ps,
tp_group=self.tp_group,
pp_group=self.pp_group,
enable_hierarchical_cache=self.enable_hierarchical_cache,
)
self.is_hybrid_swa = result.is_hybrid_swa
@@ -25,6 +25,7 @@ class CacheInitParams:
tp_cache_group: Optional[torch.distributed.ProcessGroup] = None
attn_cp_cache_group: Optional[torch.distributed.ProcessGroup] = None
attn_tp_cache_group: Optional[torch.distributed.ProcessGroup] = None
pp_cache_group: Optional[torch.distributed.ProcessGroup] = None
eviction_policy: str = "lru"
disable_finished_insert: bool = False
+94 -19
View File
@@ -104,6 +104,7 @@ class HiRadixCache(RadixCache):
self.tp_group = params.tp_cache_group
self.attn_cp_group = params.attn_cp_cache_group
self.attn_tp_group = params.attn_tp_cache_group
self.pp_group = params.pp_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
@@ -145,14 +146,13 @@ class HiRadixCache(RadixCache):
load_cache_event=self.load_cache_event,
attn_cp_group=self.attn_cp_group,
attn_tp_group=self.attn_tp_group,
pp_group=self.pp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=server_args.hicache_storage_backend,
prefetch_threshold=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,
enable_storage_metrics=self.enable_storage_metrics,
)
self._apply_storage_runtime_config(
@@ -175,6 +175,7 @@ class HiRadixCache(RadixCache):
# track per-request tokens loaded from storage (L3 hits)
# key: request_id, value: number of tokens actually loaded from storage
self.prefetch_loaded_tokens_by_reqid: dict[str, int] = {}
self.work_list: List[torch.distributed.Work] = []
# todo: dynamically adjust the threshold
self.write_through_threshold = (
1 if server_args.hicache_write_policy == "write_through" else 2
@@ -206,6 +207,71 @@ class HiRadixCache(RadixCache):
if not waited and self.tp_world_size > 1:
torch.distributed.barrier(group=self.tp_group)
def _reap_completed_async_work(self):
"""
Poll outstanding async work and reap completed ones.
Must be called in the scheduler thread.
"""
count = 0
while count < len(self.work_list) and self.work_list[count].is_completed():
count += 1
if count > 0:
logger.debug(f"Reap {count} completed async work")
self.work_list = self.work_list[count:]
def _all_reduce(self, data: torch.Tensor, tp_reduce_op: torch.distributed.ReduceOp):
"""
Synchronize data across all TP and PP ranks.
In particular, "tp_reduce_op" is performed on all TP ranks of the first PP rank,
and then the result is propagated to all following PP ranks.
Must be called in the scheduler thread.
"""
if self.pp_rank == 0:
self._all_reduce_attn_groups(data, tp_reduce_op)
self._pp_sync(data)
def _pp_sync(self, data: torch.Tensor) -> None:
"""
Synchronize data across the PP pipeline, where PPn (n>0) will receive PP0's data.
The following diagram illustrates the behavior of _pp_sync.
time | pp0 | pp1 | pp2
------|-------------------------|-------------------------|-----------------------------
0 | _pp_sync(data=1) starts | _pp_sync(data=?) starts | _pp_sync(data=?) starts
1 | _pp_sync(data=1) ends | |
2 | | _pp_sync(data=1) ends |
3 | | | _pp_sync(data=1) ends
_pp_sync requires no synchronization point among ranks. The following case may also happen.
time | pp0 | pp1 | pp2
------|-------------------------|-------------------------|-----------------------------
0 | _pp_sync(data=1) starts | |
1 | _pp_sync(data=1) ends | |
2 | | _pp_sync(data=?) starts |
3 | | _pp_sync(data=1) ends |
4 | | | _pp_sync(data=?) starts
5 | | | _pp_sync(data=1) ends
"""
if self.pp_size <= 1 or self.pp_group is None:
return
if self.pp_rank > 0:
torch.distributed.recv(
data, group_src=self.pp_rank - 1, group=self.pp_group, tag=2
)
if self.pp_rank + 1 < self.pp_size:
# Make a copy of data, so that the caller is safe to modify `data` after this call.
# This is cheap, as _pp_sync is not to be used for transmitting large data.
copy_of_data = data.clone()
send_work = torch.distributed.isend(
copy_of_data, group_dst=self.pp_rank + 1, group=self.pp_group, tag=2
)
self.work_list.append(send_work)
def shutdown(self):
"""Best-effort auto-detach of storage backend on process shutdown.
@@ -855,15 +921,17 @@ class HiRadixCache(RadixCache):
return
finish_count = 0
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
if not finish_event.query():
break
finish_count += 1
queue_size = torch.tensor(finish_count, dtype=torch.int, device="cpu")
# Keep cache state transitions identical across CPxTP participants.
self._all_reduce_attn_groups(queue_size, torch.distributed.ReduceOp.MIN)
if self.pp_rank == 0:
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
if not finish_event.query():
break
finish_count += 1
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
self._all_reduce(finish_count_tensor, torch.distributed.ReduceOp.MIN)
finish_count = finish_count_tensor.item()
finish_count = int(queue_size.item())
if finish_count > 0:
logger.debug(f"Process {finish_count} write back operations")
while finish_count > 0:
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0)
finish_event.synchronize()
@@ -873,18 +941,24 @@ class HiRadixCache(RadixCache):
def loading_check(self):
finish_count = 0
for _, finish_event, ack_list in self.cache_controller.ack_load_queue:
if not finish_event.query():
# the KV cache loading is still ongoing
break
finish_count += 1
# no need to sync across TP workers as batch forwarding is synced
if self.pp_rank == 0:
for _, finish_event, ack_list in self.cache_controller.ack_load_queue:
if not finish_event.query():
break
finish_count += 1
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
self._all_reduce(finish_count_tensor, torch.distributed.ReduceOp.MIN)
finish_count = finish_count_tensor.item()
if finish_count > 0:
logger.debug(f"Process {finish_count} load operations")
while finish_count > 0:
_, finish_event, ack_list = self.cache_controller.ack_load_queue.pop(0)
finish_event.synchronize()
for ack_id in ack_list:
end_node = self.ongoing_load_back.pop(ack_id)
self.dec_lock_ref(end_node)
# ACK until all events are processed
del self.cache_controller.ack_load_queue[:finish_count]
finish_count -= 1
def is_load_back_event_done(self, consumer_index: int) -> bool:
"""Return True after the local load-back event is complete."""
@@ -1209,6 +1283,7 @@ class HiRadixCache(RadixCache):
self.loading_check()
if self.enable_storage:
self.drain_storage_control_queues()
self._reap_completed_async_work()
if self.enable_storage_metrics:
self.storage_metrics_collector.log_storage_metrics(
self.cache_controller.storage_backend.get_stats()
@@ -163,14 +163,13 @@ class HybridCacheController(BaseHiCacheController):
load_cache_event: threading.Event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
write_policy: str = "write_through_selective",
io_backend: str = "",
storage_backend: Optional[str] = None,
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
transfer_layer_num: Optional[int] = None,
enable_storage_metrics: bool = False,
):
@@ -184,14 +183,13 @@ class HybridCacheController(BaseHiCacheController):
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
write_policy=write_policy,
io_backend=io_backend,
storage_backend=None,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=pp_rank,
pp_size=pp_size,
enable_storage_metrics=enable_storage_metrics,
)
# Override layer_num: hybrid models transfer all layers (For example, Linear Model (KV + Mamba)),
@@ -109,14 +109,13 @@ def build_kv_only_stack(
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
use_mla: bool,
override_kv_cache_dim: Optional[int] = None,
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
transfer_layer_num = len(full_layer_mapping)
@@ -146,14 +145,13 @@ def build_kv_only_stack(
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=pp_rank,
pp_size=pp_size,
transfer_layer_num=transfer_layer_num,
enable_storage_metrics=enable_storage_metrics,
)
@@ -173,6 +171,7 @@ def build_hybrid_swa_stack(
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
use_mla: bool,
host_swa_evict_fn: Optional[Callable[[int], Any]] = None,
@@ -180,8 +179,6 @@ def build_hybrid_swa_stack(
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
transfer_layer_num = len(full_layer_mapping | swa_layer_mapping)
@@ -230,14 +227,13 @@ def build_hybrid_swa_stack(
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=pp_rank,
pp_size=pp_size,
transfer_layer_num=transfer_layer_num,
enable_storage_metrics=enable_storage_metrics,
)
@@ -279,14 +275,13 @@ def build_deepseek_v4_hicache_stack(
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
host_swa_evict_fn: Optional[Callable[[int], Any]] = None,
device_swa_evict_fn: Optional[Callable[[int], Any]] = None,
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
# TODO(hzh0425): Support PP for deepseek v4 with hicache
@@ -483,14 +478,13 @@ def build_deepseek_v4_hicache_stack(
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=pp_rank,
pp_size=pp_size,
transfer_layer_num=transfer_layer_num,
enable_storage_metrics=enable_storage_metrics,
)
@@ -510,6 +504,7 @@ def build_hybrid_mamba_stack(
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
use_mla: bool,
host_mamba_evict_fn: Optional[Callable[[int], Any]] = None,
@@ -517,8 +512,6 @@ def build_hybrid_mamba_stack(
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping)
@@ -563,14 +556,13 @@ def build_hybrid_mamba_stack(
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=pp_rank,
pp_size=pp_size,
transfer_layer_num=transfer_layer_num,
enable_storage_metrics=enable_storage_metrics,
)
@@ -589,6 +581,7 @@ def build_anchor_sidecar_stack(
load_cache_event,
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
pp_group: Optional[torch.distributed.ProcessGroup] = None,
storage_backend: Optional[str],
use_mla: bool,
override_kv_cache_dim: Optional[int] = None,
@@ -596,8 +589,6 @@ def build_anchor_sidecar_stack(
prefetch_threshold: int = 256,
model_name: Optional[str] = None,
storage_backend_extra_config: Optional[dict] = None,
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
transfer_layer_num = len(full_layer_mapping)
@@ -635,14 +626,13 @@ def build_anchor_sidecar_stack(
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=pp_group,
write_policy=server_args.hicache_write_policy,
io_backend=server_args.hicache_io_backend,
storage_backend=storage_backend,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=pp_rank,
pp_size=pp_size,
transfer_layer_num=transfer_layer_num,
enable_storage_metrics=enable_storage_metrics,
)
@@ -730,14 +720,13 @@ class _DeepSeekV4Strategy(StackStrategy):
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=params.pp_cache_group,
storage_backend=storage_backend,
host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA),
device_swa_evict_fn=lambda n: cache.evict(EvictParams(swa_num_tokens=n)),
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
sidecars = [
@@ -814,6 +803,7 @@ class _MambaStrategy(StackStrategy):
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=params.pp_cache_group,
storage_backend=storage_backend,
use_mla=kvcache.use_mla,
host_mamba_evict_fn=lambda n: cache.evict_host(n, ComponentType.MAMBA),
@@ -821,8 +811,6 @@ class _MambaStrategy(StackStrategy):
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
return StackBuildResult(
@@ -890,6 +878,7 @@ class _SwaStrategy(StackStrategy):
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=params.pp_cache_group,
storage_backend=storage_backend,
use_mla=False,
host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA),
@@ -897,8 +886,6 @@ class _SwaStrategy(StackStrategy):
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
return StackBuildResult(
@@ -965,8 +952,6 @@ class _DsaStrategy(StackStrategy):
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
return StackBuildResult(
@@ -1035,13 +1020,12 @@ class _PlainKvStrategy(StackStrategy):
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=params.pp_cache_group,
storage_backend=storage_backend,
use_mla=use_mla,
prefetch_threshold=prefetch_threshold,
model_name=model_name,
storage_backend_extra_config=storage_backend_extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
return StackBuildResult(
@@ -1178,6 +1162,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=radix_cache.pp_group,
storage_backend=server_args.hicache_storage_backend,
use_mla=True,
override_kv_cache_dim=kv.kv_cache_dim,
@@ -1190,8 +1175,6 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
),
model_name=server_args.served_model_name,
storage_backend_extra_config=extra_config,
pp_rank=radix_cache.pp_rank,
pp_size=radix_cache.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
radix_cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV)
@@ -1240,6 +1223,7 @@ def attach_hybrid_pool_to_mamba_cache(
load_cache_event=load_cache_event,
attn_cp_group=attn_cp_group,
attn_tp_group=attn_tp_group,
pp_group=params.pp_cache_group,
storage_backend=server_args.hicache_storage_backend,
use_mla=hybrid_kv.use_mla,
host_mamba_evict_fn=mamba_cache.evict_mamba_host,
@@ -1247,8 +1231,6 @@ def attach_hybrid_pool_to_mamba_cache(
prefetch_threshold=prefetch_threshold,
model_name=server_args.served_model_name,
storage_backend_extra_config=extra_config,
pp_rank=params.pp_rank,
pp_size=params.pp_size,
enable_storage_metrics=enable_storage_metrics,
)
mamba_cache.full_kv_pool_host = host_pool_group.get_pool(PoolName.KV)
@@ -142,6 +142,7 @@ def build_kv_cache(
enable_kv_cache_events: bool,
ps: "ParallelState",
tp_group: "GroupCoordinator",
pp_group: "GroupCoordinator",
enable_hierarchical_cache: bool,
) -> "KVCacheBuildResult":
sliding_window_size: Optional[int] = None
@@ -214,6 +215,7 @@ def build_kv_cache(
),
attn_cp_cache_group=attn_cp_cpu_group,
attn_tp_cache_group=attn_tp_cpu_group,
pp_cache_group=pp_group.cpu_group,
eviction_policy=server_args.radix_eviction_policy,
enable_metrics=enable_metrics,
enable_kv_cache_events=enable_kv_cache_events,
@@ -307,11 +307,15 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
self.tp_group = params.tp_cache_group
self.attn_cp_group = params.attn_cp_cache_group
self.attn_tp_group = params.attn_tp_cache_group
self.pp_group = params.pp_cache_group
self.tp_world_size = (
1
if self.tp_group is None
else torch.distributed.get_world_size(group=self.tp_group)
)
self.pp_rank = params.pp_rank
self.pp_size = params.pp_size
self.work_list: list[torch.distributed.Work] = []
# HiCache D↔H defaults (overridden by init_hicache)
self.cache_controller: Optional[HybridCacheController] = None
@@ -343,6 +347,49 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
if not waited and self.tp_world_size > 1:
torch.distributed.barrier(group=self.tp_group)
def _reap_completed_async_work(self):
"""
Poll outstanding async work and reap completed ones.
Must be called in the scheduler thread.
"""
count = 0
while count < len(self.work_list) and self.work_list[count].is_completed():
count += 1
if count > 0:
logger.debug(f"Reap {count} completed async work")
self.work_list = self.work_list[count:]
def _all_reduce(self, data: torch.Tensor, tp_reduce_op: torch.distributed.ReduceOp):
"""
Synchronize data across all TP and PP ranks.
In particular, "tp_reduce_op" is performed on all TP ranks of the first PP rank,
and then the result is propagated to all following PP ranks.
Must be called in the scheduler thread.
"""
if self.pp_rank == 0:
self._all_reduce_attn_groups(data, tp_reduce_op)
self._pp_sync(data)
def _pp_sync(self, data: torch.Tensor) -> None:
"""
Synchronize data across the PP pipeline, where PPn (n>0) will receive PP0's data.
"""
if self.pp_size <= 1 or self.pp_group is None:
return
if self.pp_rank > 0:
torch.distributed.recv(
data, group_src=self.pp_rank - 1, group=self.pp_group, tag=2
)
if self.pp_rank + 1 < self.pp_size:
copy_of_data = data.clone()
send_work = torch.distributed.isend(
copy_of_data, group_dst=self.pp_rank + 1, group=self.pp_group, tag=2
)
self.work_list.append(send_work)
def reset(self) -> None:
self._reset_full()
@@ -463,7 +510,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
self.write_through_threshold = (
1 if server_args.hicache_write_policy == "write_through" else 2
)
self.load_back_threshold = 256
self.load_back_threshold = 10
self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy
if storage_backend is not None:
@@ -2215,15 +2262,15 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
return
finish_count = 0
for _, finish_event, ack_list in cc.ack_write_queue:
if not finish_event.query():
break
finish_count += 1
if self.pp_rank == 0:
for _, finish_event, ack_list in cc.ack_write_queue:
if not finish_event.query():
break
finish_count += 1
# Keep cache state transitions identical across CPxTP participants.
queue_size = torch.tensor(finish_count, dtype=torch.int, device="cpu")
self._all_reduce_attn_groups(queue_size, torch.distributed.ReduceOp.MIN)
finish_count = int(queue_size.item())
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
self._all_reduce(finish_count_tensor, torch.distributed.ReduceOp.MIN)
finish_count = finish_count_tensor.item()
# Process completed acks
while finish_count > 0:
@@ -2239,14 +2286,22 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
if cc is None or not self.ongoing_load_back:
return
finish_count = 0
for _, finish_event, ack_list in cc.ack_load_queue:
if not finish_event.query():
break
finish_count += 1
if self.pp_rank == 0:
for _, finish_event, ack_list in cc.ack_load_queue:
if not finish_event.query():
break
finish_count += 1
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
self._all_reduce(finish_count_tensor, torch.distributed.ReduceOp.MIN)
finish_count = finish_count_tensor.item()
while finish_count > 0:
_, finish_event, ack_list = cc.ack_load_queue.pop(0)
finish_event.synchronize()
for ack_id in ack_list:
node, lock_params = self.ongoing_load_back.pop(ack_id)
self.dec_lock_ref(node, lock_params)
del cc.ack_load_queue[:finish_count]
finish_count -= 1
# ---- HiCache: Scheduler Entry Points ----
@@ -2302,6 +2357,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
self.loading_check()
if self.enable_storage:
self.drain_storage_control_queues()
self._reap_completed_async_work()
if self.enable_storage_metrics and self.storage_metrics_collector is not None:
self.storage_metrics_collector.log_storage_metrics(
self.cache_controller.storage_backend.get_stats()