[HiCache] Prevent KV cache data loss when radix tree node is split b… (#16946)
Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
This commit is contained in:
co-authored by
Zhiqiang Xie
parent
4151a04d1a
commit
dff45411da
@@ -714,34 +714,76 @@ class HiRadixCache(RadixCache):
|
|||||||
if host_indices is not None:
|
if host_indices is not None:
|
||||||
node.host_value = host_indices.clone()
|
node.host_value = host_indices.clone()
|
||||||
assert len(node.host_value) > 0
|
assert len(node.host_value) > 0
|
||||||
self.ongoing_write_through[node.id] = node
|
# Record backup_len for ack-time walk-and-concat after split.
|
||||||
|
self.ongoing_write_through[node.id] = (node, len(node.key))
|
||||||
if not write_back:
|
if not write_back:
|
||||||
# no need to lock nodes if write back
|
|
||||||
self.inc_lock_ref(node)
|
self.inc_lock_ref(node)
|
||||||
# Note: store(CPU) event is deferred to writing_check() after the
|
|
||||||
# async DMA transfer is confirmed complete.
|
|
||||||
else:
|
else:
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
return len(host_indices)
|
return len(host_indices)
|
||||||
|
|
||||||
def write_backup_storage(self, node: TreeNode):
|
def write_backup_storage(self, node: TreeNode, backup_len: Optional[int] = None):
|
||||||
|
# Recover pre-split data via walk-and-concat if node was split.
|
||||||
|
# prefix_keys anchored at chain top to avoid double-counting.
|
||||||
|
if backup_len is None or len(node.key) == backup_len:
|
||||||
|
top, key, hash_value, host_value = (
|
||||||
|
node,
|
||||||
|
node.key,
|
||||||
|
node.hash_value,
|
||||||
|
node.host_value,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
top, key, hash_value, host_value = self._concat_split_chain(
|
||||||
|
node, backup_len
|
||||||
|
)
|
||||||
|
|
||||||
prefix_keys = (
|
prefix_keys = (
|
||||||
node.get_prefix_hash_values(node.parent)
|
top.get_prefix_hash_values(top.parent)
|
||||||
if self.hicache_storage_pass_prefix_keys
|
if self.hicache_storage_pass_prefix_keys
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
operation_id = self.cache_controller.write_storage(
|
operation_id = self.cache_controller.write_storage(
|
||||||
node.host_value,
|
host_value, key, hash_value, prefix_keys, **self._get_extra_pools()
|
||||||
node.key,
|
|
||||||
node.hash_value,
|
|
||||||
prefix_keys,
|
|
||||||
**self._get_extra_pools(),
|
|
||||||
)
|
)
|
||||||
self.ongoing_backup[operation_id] = node
|
self.ongoing_backup[operation_id] = node
|
||||||
node.protect_host()
|
node.protect_host()
|
||||||
|
|
||||||
|
def _concat_split_chain(self, node: TreeNode, backup_len: int):
|
||||||
|
"""Recover enqueue-time key/hash/host by walking the split chain."""
|
||||||
|
chain, accumulated = [], 0
|
||||||
|
current = node
|
||||||
|
while current is not self.root_node and accumulated < backup_len:
|
||||||
|
chain.append(current)
|
||||||
|
accumulated += len(current.key)
|
||||||
|
current = current.parent
|
||||||
|
assert accumulated == backup_len, (
|
||||||
|
f"backup chain length mismatch for node {node.id}: "
|
||||||
|
f"expected {backup_len}, got {accumulated}"
|
||||||
|
)
|
||||||
|
chain.reverse() # parent-first
|
||||||
|
top = chain[0]
|
||||||
|
if top.key.is_bigram:
|
||||||
|
# Bigram segments share boundary tokens; drop overlap after first.
|
||||||
|
token_ids = list(chain[0].key.token_ids)
|
||||||
|
for n in chain[1:]:
|
||||||
|
token_ids.extend(n.key.token_ids[1:])
|
||||||
|
else:
|
||||||
|
token_ids = []
|
||||||
|
for n in chain:
|
||||||
|
token_ids.extend(n.key.token_ids)
|
||||||
|
key = RadixKey(token_ids, top.key.extra_key, top.key.is_bigram)
|
||||||
|
|
||||||
|
if all(n.hash_value is not None for n in chain):
|
||||||
|
hash_value = []
|
||||||
|
for n in chain:
|
||||||
|
hash_value.extend(n.hash_value)
|
||||||
|
else:
|
||||||
|
hash_value = None
|
||||||
|
host_value = torch.cat([n.host_value for n in chain])
|
||||||
|
return top, key, hash_value, host_value
|
||||||
|
|
||||||
def _inc_hit_count(self, node: TreeNode, chunked=False):
|
def _inc_hit_count(self, node: TreeNode, chunked=False):
|
||||||
# skip the hit count update for chunked requests
|
# skip the hit count update for chunked requests
|
||||||
if self.cache_controller.write_policy == "write_back" or chunked:
|
if self.cache_controller.write_policy == "write_back" or chunked:
|
||||||
@@ -760,13 +802,11 @@ class HiRadixCache(RadixCache):
|
|||||||
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
|
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
|
||||||
finish_event.synchronize()
|
finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack_list:
|
||||||
backuped_node = self.ongoing_write_through.pop(ack_id)
|
node, backup_len = self.ongoing_write_through.pop(ack_id)
|
||||||
# DMA confirmed -- block is now on host.
|
# DMA confirmed -- block is now on host.
|
||||||
self._record_store_event(
|
self._record_store_event(node, medium=StorageMedium.CPU)
|
||||||
backuped_node, medium=StorageMedium.CPU
|
|
||||||
)
|
|
||||||
if self.enable_storage:
|
if self.enable_storage:
|
||||||
self.write_backup_storage(backuped_node)
|
self.write_backup_storage(node, backup_len)
|
||||||
self.cache_controller.ack_write_queue.clear()
|
self.cache_controller.ack_write_queue.clear()
|
||||||
assert len(self.ongoing_write_through) == 0
|
assert len(self.ongoing_write_through) == 0
|
||||||
return
|
return
|
||||||
@@ -789,12 +829,12 @@ class HiRadixCache(RadixCache):
|
|||||||
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0)
|
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0)
|
||||||
finish_event.synchronize()
|
finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack_list:
|
||||||
backuped_node = self.ongoing_write_through.pop(ack_id)
|
node, backup_len = self.ongoing_write_through.pop(ack_id)
|
||||||
# DMA confirmed -- block is now on host.
|
# DMA confirmed -- block is now on host.
|
||||||
self._record_store_event(backuped_node, medium=StorageMedium.CPU)
|
self._record_store_event(node, medium=StorageMedium.CPU)
|
||||||
self.dec_lock_ref(backuped_node)
|
self.dec_lock_ref(node)
|
||||||
if self.enable_storage:
|
if self.enable_storage:
|
||||||
self.write_backup_storage(backuped_node)
|
self.write_backup_storage(node, backup_len)
|
||||||
finish_count -= 1
|
finish_count -= 1
|
||||||
|
|
||||||
def loading_check(self):
|
def loading_check(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user