[HiCache] write_back: reclaim duplicated host copy first under host pressure (#33777)

This commit is contained in:
Zhiqiang Xie
2026-08-07 15:59:59 -07:00
committed by GitHub
parent eb3cc879e0
commit 3dc91366ac
4 changed files with 209 additions and 11 deletions
@@ -84,6 +84,11 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
# 42 bits: digest * 1000003 (< 2^20) stays under 2^62, so the update never
# overflows int64 with plain (non-wrapping) arithmetic in the Rust port, and
# the TP consistency check can still all_reduce [digest, -digest] in int64.
_RECLAIM_DIGEST_MASK = (1 << 42) - 1
class StorageBackupSpec(NamedTuple):
"""A node's device->storage backup spec, gathered tree-side."""
@@ -123,6 +128,9 @@ class UnifiedTreeNode:
self.id = UnifiedTreeNode.counter
UnifiedTreeNode.counter += 1
self.write_through_pending_id: Optional[int] = None
# Anchor NodeId of an in-flight H->D load-back reading this node's
# host slots; such host copies must not be reclaimed until the ack.
self.load_back_pending_id: Optional[int] = None
def component(self, component_type: ComponentType) -> ComponentData:
return self.component_data[component_type]
@@ -457,6 +465,11 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
)
for ct in self.component_types
}
# Full KV on both tiers -> redundant host copy, reclaimed first by
# write_back; insertion-ordered dict keeps victims TP-deterministic.
self.full_host_duplicates: dict[NodeId, UnifiedTreeNode] = {}
# Rolling digest of reclaim victim ids, cross-checked across TP ranks.
self.write_back_duplicate_reclaim_digest: int = 0
self._empty_match_result = MatchResult(
device_indices=torch.empty(
@@ -1021,6 +1034,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
new_node.key = child.key[:split_len]
new_node.hit_count = child.hit_count
new_node.creation_time = child.creation_time
# Split fragments stay on the anchor's root path for the ack's walk.
new_node.load_back_pending_id = child.load_back_pending_id
self._for_each_component_lru(child, UnifiedLRUList.remove_node)
@@ -1056,6 +1071,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
self._update_evictable_leaf_sets(new_node)
self._update_evictable_leaf_sets(child)
# Only the new fragment needs qualifying; the child keeps its id.
self._update_duplicate_tracking(new_node)
return new_node, action
def _add_new_node(
@@ -1091,6 +1108,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
cd.value = fresh_value.clone()
self.component_evictable_size_[ct] += n
self._update_evictable_leaf_sets(node)
# A backuped node restored from fresh KV is a duplicate right away.
self._update_duplicate_tracking(node)
if node.parent is not None:
self._update_evictable_leaf_sets(node.parent)
self._record_store_event(node, medium=StorageMedium.GPU)
@@ -1107,6 +1126,26 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
else:
self.evictable_host_leaves.discard(node)
def _update_duplicate_tracking(self, node: UnifiedTreeNode) -> None:
"""Register where duplicates are born (acks, split, unevict);
deregistration is lazy, so entries may be stale and re-checked live."""
if self._is_settled_full_host_duplicate(node):
self.full_host_duplicates.setdefault(node.id, node)
else:
self.full_host_duplicates.pop(node.id, None)
def _is_settled_full_host_duplicate(self, node: UnifiedTreeNode) -> bool:
"""Full KV present on both tiers with no in-flight DMA on the node's
host slots; mid-transfer nodes join the tracking at their ack."""
cd = node.component_data[BASE_COMPONENT_TYPE]
return (
node is not self.root_node
and cd.value is not None
and cd.host_value is not None
and node.write_through_pending_id is None
and node.load_back_pending_id is None
)
def _for_each_component_lru(
self,
node: UnifiedTreeNode,
@@ -1272,10 +1311,18 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
def drive_host_eviction(
self, component_type: ComponentType, num_tokens: int
) -> DriveHostEvictionResult:
"""Evict a component's host-side resources; no-op if the component is absent."""
"""Evict a component's host-side resources; no-op if absent. Under
write_back, FULL pressure reclaims redundant Full host copies first."""
result = DriveHostEvictionResult()
comp = self.components_by_type.get(component_type)
if comp is not None:
if self.is_write_back and component_type == BASE_COMPONENT_TYPE:
self._reclaim_full_host_duplicates(
num_tokens,
result.tracker,
result.device_frees,
result.host_frees,
)
comp.drive_host_eviction(
num_tokens,
result.tracker,
@@ -1294,6 +1341,74 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
self.node_by_id(tail_node_id), device_frees, host_frees
)
def _reclaim_full_host_duplicates(
self,
num_tokens: int,
tracker: dict[ComponentType, int],
device_frees: dict[ComponentType, list[torch.Tensor]],
host_frees: dict[ComponentType, list[torch.Tensor]],
) -> None:
"""Reclaim Full host duplicates until num_tokens are freed; pass 1
spares evictable D-leaves (imminent free demotes), pass 2 takes them."""
swept_ids: list[NodeId] = []
for spare_imminent_demotes in (True, False):
if tracker[BASE_COMPONENT_TYPE] >= num_tokens:
break
for node in self.full_host_duplicates.values():
if tracker[BASE_COMPONENT_TYPE] >= num_tokens:
break
cd = node.component_data[BASE_COMPONENT_TYPE]
if cd.value is None or cd.host_value is None:
swept_ids.append(node.id) # stale entry
continue
if spare_imminent_demotes and node in self.evictable_device_leaves:
continue
if not self._can_reclaim_full_host_duplicate(node):
continue
self._release_full_host_duplicate(
node, tracker, device_frees, host_frees
)
swept_ids.append(node.id) # released -> no longer a duplicate
# Sweep after the walk: the dict must not be mutated mid-iteration.
for nid in swept_ids:
self.full_host_duplicates.pop(nid, None)
def _can_reclaim_full_host_duplicate(self, node: UnifiedTreeNode) -> bool:
"""Full on both tiers, no in-flight DMA, no Full host lock; checked
live because tracking may be stale."""
cd = node.component_data[BASE_COMPONENT_TYPE]
if node is self.root_node or cd.value is None or cd.host_value is None:
return False
if (
node.write_through_pending_id is not None
or node.load_back_pending_id is not None
):
return False
return cd.host_lock_ref == 0
def _release_full_host_duplicate(
self,
node: UnifiedTreeNode,
tracker: dict[ComponentType, int],
device_frees: dict[ComponentType, list[torch.Tensor]],
host_frees: dict[ComponentType, list[torch.Tensor]],
) -> None:
"""Free only the Full host layer; aux host slices stay under their own
pools' LRU (a host-only aux slice may be a sole copy)."""
assert self._can_reclaim_full_host_duplicate(node)
self._record_remove_event(node, medium=StorageMedium.CPU)
self._evict_component_and_detach_lru(
node,
self.components_by_type[BASE_COMPONENT_TYPE],
target=EvictLayer.HOST,
tracker=tracker,
device_frees=device_frees,
host_frees=host_frees,
)
self.write_back_duplicate_reclaim_digest = (
self.write_back_duplicate_reclaim_digest * 1000003 + node.id + 1
) & _RECLAIM_DIGEST_MASK
def _evict_host_leaf(
self,
node: UnifiedTreeNode,
@@ -1435,6 +1550,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
key = node.key.child_key(self.page_size)
v = node.parent.children.pop(key, None)
assert v == node
# Deleted nodes must not linger in duplicate tracking as ghosts.
self.full_host_duplicates.pop(node.id, None)
self._unregister_node(node)
def _evict_component_and_detach_lru(
@@ -1552,7 +1669,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
"""H-leaf: evicted, Full host value present, no children, unlocked, not root.
Only the Full (base) component host_value is required; auxiliary
components are not mandatory for H-leaf membership."""
components are not mandatory for H-leaf membership. In-flight DMA
marks need no check: marked nodes are never ``evicted``."""
if node is self.root_node or not node.evicted:
return False
if not node.backuped:
@@ -1803,6 +1921,19 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
rebuild is deferred to the orchestration layer."""
node = self.node_by_id(node_id)
cache_actions: list[CacheAction | ComponentAction] = []
# Pin every node whose host slots the in-flight DMA reads (including
# aux-only nodes) against reclaim until the ack.
for xfers in ([kv_xfer], *comp_xfers.values()):
for xfer in xfers:
for nid in xfer.nodes_to_load or ():
pinned = self.node_by_id(nid)
# One live load-back per node; only the same anchor may
# re-pin (a node can sit in Full and aux transfer lists).
assert pinned.load_back_pending_id in (None, node_id), (
f"node {nid} pinned by load-back "
f"{pinned.load_back_pending_id}, new anchor {node_id}"
)
pinned.load_back_pending_id = node_id
kv_xfer.device_indices = device_indices
self.components_by_type[BASE_COMPONENT_TYPE].commit_hicache_transfer(
node,
@@ -1811,8 +1942,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
cache_actions=cache_actions,
)
for nid in kv_xfer.nodes_to_load or ():
loaded = self.node_by_id(nid)
self._record_store_event(loaded, medium=StorageMedium.GPU)
self._record_store_event(self.node_by_id(nid), medium=StorageMedium.GPU)
for ct, xfers in comp_xfers.items():
self.components_by_type[ct].commit_hicache_transfer(
node,
@@ -1823,6 +1953,17 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
self._update_evictable_leaf_sets(node)
return cache_actions
def finish_load_back(self, anchor_node_id: NodeId) -> None:
"""Clear the in-flight H->D marks along the anchor's root path at ack
time; split fragments stay on the path, so the walk covers them."""
node = self.node_by_id(anchor_node_id)
while node is not None and node is not self.root_node:
if node.load_back_pending_id == anchor_node_id:
node.load_back_pending_id = None
# The loaded copies become tracked duplicates only now.
self._update_duplicate_tracking(node)
node = node.parent
def mark_write_through_pending(self, node_id: NodeId) -> None:
"""Mark a node as having an in-flight write-through backup."""
node = self.node_by_id(node_id)
@@ -1835,6 +1976,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
node = self.node_by_id(node_id)
if node.write_through_pending_id == ack_id:
node.write_through_pending_id = None
# The backed-up copy becomes a tracked duplicate only now.
self._update_duplicate_tracking(node)
self._record_store_event(node, medium=StorageMedium.CPU)
def set_component_device_value(
@@ -1902,6 +2045,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
# ── PART 2: Per-node state machine and leaf qualification ──
expected_dev_leaves: set[UnifiedTreeNode] = set()
expected_hst_leaves: set[UnifiedTreeNode] = set()
expected_duplicates: set[UnifiedTreeNode] = set()
for node in all_nodes:
if node is self.root_node:
@@ -1918,7 +2062,10 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
if cd.value is not None and not full_dev:
E(f"node {nid} {ct} device present but Full.value=None")
if cd.host_value is not None and not full_hst:
E(f"node {nid} {ct} host present but Full.host_value=None")
# write_back reclaim takes only the Full host layer; an
# aux host slice may outlive it while Full device is live.
if not (self.is_write_back and full_dev):
E(f"node {nid} {ct} host present but Full.host_value=None")
# Every node must keep Full data on at least one layer.
if not full_dev and not full_hst:
@@ -1951,6 +2098,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
expected_dev_leaves.add(node)
if self._is_host_leaf(node):
expected_hst_leaves.add(node)
if self._is_settled_full_host_duplicate(node):
expected_duplicates.add(node)
# ── PART 3: Tracking structures ──
@@ -1972,6 +2121,16 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
if missing:
E(f"H-leaf missing: {[n.id for n in list(missing)[:5]]}")
# Lazy deregistration: stale extras are legal; settled duplicates must
# be tracked and entries must not outlive their node.
expected_ids = {n.id for n in expected_duplicates}
dup_ids = set(self.full_host_duplicates.keys())
if expected_ids - dup_ids:
E(f"Duplicate missing: {list(expected_ids - dup_ids)[:5]}")
ghost_ids = dup_ids - {n.id for n in all_nodes}
if ghost_ids:
E(f"Duplicate ghosts: {list(ghost_ids)[:5]}")
# D-leaf ∩ H-leaf = ∅
overlap = self.evictable_device_leaves & self.evictable_host_leaves
if overlap:
@@ -2083,6 +2242,18 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
E(
f"[Ongoing] load_back node {nid} lock_ref={n.component_data[FCT].lock_ref}"
)
# Every in-flight H->D mark must belong to a live load-back; a leaked
# mark would pin the node's host copy against reclaim forever.
ongoing_load_ids = {node_id for _, node_id in ongoing_load_back}
for node in all_nodes:
if (
node.load_back_pending_id is not None
and node.load_back_pending_id not in ongoing_load_ids
):
E(
f"[Ongoing] node {node.id} load_back_pending_id="
f"{node.load_back_pending_id} has no live load-back"
)
if errors:
msg = (
@@ -460,6 +460,15 @@ class UnifiedTreeCoreInterface(KVCacheEventMixin, ABC):
"""Commit a successful H->D load-back onto the node; returns any cache actions."""
...
@abstractmethod
def finish_load_back(self, anchor_node_id: NodeId) -> None:
"""Clear the in-flight H->D marks on the anchor's root path at ack time."""
...
# Order-sensitive digest of write_back duplicate-reclaim victim ids,
# cross-checked across TP ranks; cores that never reclaim keep 0.
write_back_duplicate_reclaim_digest: int = 0
@abstractmethod
def mark_write_through_pending(self, node_id: NodeId) -> None:
"""Mark a node as having an in-flight write-through backup."""
@@ -1839,22 +1839,30 @@ class UnifiedRadixCache(BasePrefixCache):
else ()
)
# Piggybacked TP check: [digest, -digest] MIN-reduces to [min, -max],
# equal iff reclaim victim order matched on every rank.
digest = self.tree_core.write_back_duplicate_reclaim_digest
ready_counts = torch.tensor(
[
write_acks,
load_acks,
*storage_queue_sizes,
digest,
-digest,
],
dtype=torch.int,
dtype=torch.int64,
device="cpu",
)
self._all_reduce(ready_counts, torch.distributed.ReduceOp.MIN)
count_values = list(map(int, ready_counts.tolist()))
assert (
count_values[-2] == -count_values[-1]
), "write_back duplicate-reclaim victims diverged across TP ranks"
return (
count_values[0],
count_values[1],
tuple(count_values[2:]),
tuple(count_values[2:-2]),
extra_pool_names,
)
@@ -1926,11 +1934,17 @@ class UnifiedRadixCache(BasePrefixCache):
finish_count = 0
if self.pp_rank == 0:
finish_count = self._count_ready_acks(cc.ack_load_queue)
finish_count_tensor = torch.tensor(
finish_count, dtype=torch.int, device="cpu"
# Piggybacked TP check: [digest, -digest] MIN-reduces to [min, -max],
# equal iff reclaim victim order matched on every rank.
digest = self.tree_core.write_back_duplicate_reclaim_digest
sync_tensor = torch.tensor(
[finish_count, digest, -digest], dtype=torch.int64, device="cpu"
)
self._all_reduce(finish_count_tensor, torch.distributed.ReduceOp.MIN)
finish_count = finish_count_tensor.item()
self._all_reduce(sync_tensor, torch.distributed.ReduceOp.MIN)
finish_count = int(sync_tensor[0].item())
assert (
sync_tensor[1].item() == -sync_tensor[2].item()
), "write_back duplicate-reclaim victims diverged across TP ranks"
while finish_count > 0:
ack = cc.ack_load_queue.pop(0)
@@ -1939,6 +1953,8 @@ class UnifiedRadixCache(BasePrefixCache):
node, lock_params, host_lock_params = self.ongoing_load_back.pop(ack_id)
self.dec_lock_ref(node, lock_params)
self.dec_host_lock_ref(node, host_lock_params)
# Unpin the loaded nodes; host copies stay as reclaimable duplicates.
self.tree_core.finish_load_back(node)
if self.metrics_collector is not None:
for pool, num_tokens in (ack.num_tokens_by_pool or {}).items():