[HiCache] write_back: reclaim duplicated host copy first under host pressure (#33777)
This commit is contained in:
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user