feat: restore CP HiCache host hits as logical locs

This commit is contained in:
2026-05-08 02:34:16 +08:00
parent f10a8983c3
commit 829f3ebceb
2 changed files with 195 additions and 6 deletions

View File

@@ -1180,6 +1180,71 @@ class HiRadixCache(RadixCache):
self, node: TreeNode, mem_quota: Optional[int] = None
) -> Optional[torch.Tensor]:
if self._uses_cp_hicache:
start_time = time.perf_counter()
last_hit_node = node
nodes_to_load = []
while node.evicted:
assert self._node_backuped(
node
), "No backup available on evicted nodes, should not happen"
nodes_to_load.insert(0, node)
node = node.parent
else:
ancester_node = node
# protect the ancestor nodes from eviction
result = self.inc_lock_ref(ancester_node)
delta = result.delta
# load it all or not at all
host_hit_len = sum(self._node_host_len(n) for n in nodes_to_load)
if host_hit_len < self.load_back_threshold or (
host_hit_len > mem_quota + delta if mem_quota is not None else False
):
# skip loading back if the total size is too small or exceeding the memory quota
self.dec_lock_ref(ancester_node)
return None
device_indices = self.cache_controller.load_cp(
nodes_to_load, node_id=last_hit_node.id
)
if device_indices is None:
self.evict(EvictParams(num_tokens=host_hit_len))
device_indices = self.cache_controller.load_cp(
nodes_to_load, node_id=last_hit_node.id
)
self.dec_lock_ref(ancester_node)
if device_indices is None:
# no sufficient GPU memory to load back KV caches
logger.warning(
"load_back: FAILED to load %d tokens for node %d "
"even after eviction (evictable_size=%d)",
host_hit_len,
last_hit_node.id,
self.evictable_size_,
)
return None
self.ongoing_load_back[last_hit_node.id] = last_hit_node
offset = 0
for loaded_node in nodes_to_load:
host_len = self._node_host_len(loaded_node)
loaded_node.value = device_indices[offset : offset + host_len].clone()
offset += host_len
self.evictable_size_ += len(device_indices)
self.inc_lock_ref(last_hit_node)
if self.metrics_collector is not None:
self.metrics_collector.observe_load_back_duration(
time.perf_counter() - start_time
)
self.metrics_collector.increment_load_back_num_tokens(
len(device_indices)
)
return device_indices
start_time = time.perf_counter()
last_hit_node = node
nodes_to_load = []
@@ -1464,11 +1529,23 @@ class HiRadixCache(RadixCache):
host_hit_length = 0
last_host_node = last_node
while last_node.evicted:
host_hit_length += len(last_node.host_value)
last_node = last_node.parent
while not last_host_node.backuped:
last_host_node = last_host_node.parent
if self._uses_cp_hicache:
while last_node.evicted:
host_hit_length += self._node_host_len(last_node)
last_node = last_node.parent
while (
last_host_node != self.root_node
and not self._node_backuped(last_host_node)
):
last_host_node = last_host_node.parent
if not self._node_backuped(last_host_node):
last_host_node = self.root_node
else:
while last_node.evicted:
host_hit_length += len(last_node.host_value)
last_node = last_node.parent
while not last_host_node.backuped:
last_host_node = last_host_node.parent
return MatchResult(
device_indices=value,