feat: restore CP HiCache host hits as logical locs
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user