[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__)
|
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):
|
class StorageBackupSpec(NamedTuple):
|
||||||
"""A node's device->storage backup spec, gathered tree-side."""
|
"""A node's device->storage backup spec, gathered tree-side."""
|
||||||
@@ -123,6 +128,9 @@ class UnifiedTreeNode:
|
|||||||
self.id = UnifiedTreeNode.counter
|
self.id = UnifiedTreeNode.counter
|
||||||
UnifiedTreeNode.counter += 1
|
UnifiedTreeNode.counter += 1
|
||||||
self.write_through_pending_id: Optional[int] = None
|
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:
|
def component(self, component_type: ComponentType) -> ComponentData:
|
||||||
return self.component_data[component_type]
|
return self.component_data[component_type]
|
||||||
@@ -457,6 +465,11 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
)
|
)
|
||||||
for ct in self.component_types
|
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(
|
self._empty_match_result = MatchResult(
|
||||||
device_indices=torch.empty(
|
device_indices=torch.empty(
|
||||||
@@ -1021,6 +1034,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
new_node.key = child.key[:split_len]
|
new_node.key = child.key[:split_len]
|
||||||
new_node.hit_count = child.hit_count
|
new_node.hit_count = child.hit_count
|
||||||
new_node.creation_time = child.creation_time
|
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)
|
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(new_node)
|
||||||
self._update_evictable_leaf_sets(child)
|
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
|
return new_node, action
|
||||||
|
|
||||||
def _add_new_node(
|
def _add_new_node(
|
||||||
@@ -1091,6 +1108,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
cd.value = fresh_value.clone()
|
cd.value = fresh_value.clone()
|
||||||
self.component_evictable_size_[ct] += n
|
self.component_evictable_size_[ct] += n
|
||||||
self._update_evictable_leaf_sets(node)
|
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:
|
if node.parent is not None:
|
||||||
self._update_evictable_leaf_sets(node.parent)
|
self._update_evictable_leaf_sets(node.parent)
|
||||||
self._record_store_event(node, medium=StorageMedium.GPU)
|
self._record_store_event(node, medium=StorageMedium.GPU)
|
||||||
@@ -1107,6 +1126,26 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
else:
|
else:
|
||||||
self.evictable_host_leaves.discard(node)
|
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(
|
def _for_each_component_lru(
|
||||||
self,
|
self,
|
||||||
node: UnifiedTreeNode,
|
node: UnifiedTreeNode,
|
||||||
@@ -1272,10 +1311,18 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
def drive_host_eviction(
|
def drive_host_eviction(
|
||||||
self, component_type: ComponentType, num_tokens: int
|
self, component_type: ComponentType, num_tokens: int
|
||||||
) -> DriveHostEvictionResult:
|
) -> 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()
|
result = DriveHostEvictionResult()
|
||||||
comp = self.components_by_type.get(component_type)
|
comp = self.components_by_type.get(component_type)
|
||||||
if comp is not None:
|
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(
|
comp.drive_host_eviction(
|
||||||
num_tokens,
|
num_tokens,
|
||||||
result.tracker,
|
result.tracker,
|
||||||
@@ -1294,6 +1341,74 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
self.node_by_id(tail_node_id), device_frees, host_frees
|
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(
|
def _evict_host_leaf(
|
||||||
self,
|
self,
|
||||||
node: UnifiedTreeNode,
|
node: UnifiedTreeNode,
|
||||||
@@ -1435,6 +1550,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
key = node.key.child_key(self.page_size)
|
key = node.key.child_key(self.page_size)
|
||||||
v = node.parent.children.pop(key, None)
|
v = node.parent.children.pop(key, None)
|
||||||
assert v == node
|
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)
|
self._unregister_node(node)
|
||||||
|
|
||||||
def _evict_component_and_detach_lru(
|
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.
|
"""H-leaf: evicted, Full host value present, no children, unlocked, not root.
|
||||||
|
|
||||||
Only the Full (base) component host_value is required; auxiliary
|
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:
|
if node is self.root_node or not node.evicted:
|
||||||
return False
|
return False
|
||||||
if not node.backuped:
|
if not node.backuped:
|
||||||
@@ -1803,6 +1921,19 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
rebuild is deferred to the orchestration layer."""
|
rebuild is deferred to the orchestration layer."""
|
||||||
node = self.node_by_id(node_id)
|
node = self.node_by_id(node_id)
|
||||||
cache_actions: list[CacheAction | ComponentAction] = []
|
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
|
kv_xfer.device_indices = device_indices
|
||||||
self.components_by_type[BASE_COMPONENT_TYPE].commit_hicache_transfer(
|
self.components_by_type[BASE_COMPONENT_TYPE].commit_hicache_transfer(
|
||||||
node,
|
node,
|
||||||
@@ -1811,8 +1942,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
cache_actions=cache_actions,
|
cache_actions=cache_actions,
|
||||||
)
|
)
|
||||||
for nid in kv_xfer.nodes_to_load or ():
|
for nid in kv_xfer.nodes_to_load or ():
|
||||||
loaded = self.node_by_id(nid)
|
self._record_store_event(self.node_by_id(nid), medium=StorageMedium.GPU)
|
||||||
self._record_store_event(loaded, medium=StorageMedium.GPU)
|
|
||||||
for ct, xfers in comp_xfers.items():
|
for ct, xfers in comp_xfers.items():
|
||||||
self.components_by_type[ct].commit_hicache_transfer(
|
self.components_by_type[ct].commit_hicache_transfer(
|
||||||
node,
|
node,
|
||||||
@@ -1823,6 +1953,17 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
self._update_evictable_leaf_sets(node)
|
self._update_evictable_leaf_sets(node)
|
||||||
return cache_actions
|
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:
|
def mark_write_through_pending(self, node_id: NodeId) -> None:
|
||||||
"""Mark a node as having an in-flight write-through backup."""
|
"""Mark a node as having an in-flight write-through backup."""
|
||||||
node = self.node_by_id(node_id)
|
node = self.node_by_id(node_id)
|
||||||
@@ -1835,6 +1976,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
node = self.node_by_id(node_id)
|
node = self.node_by_id(node_id)
|
||||||
if node.write_through_pending_id == ack_id:
|
if node.write_through_pending_id == ack_id:
|
||||||
node.write_through_pending_id = None
|
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)
|
self._record_store_event(node, medium=StorageMedium.CPU)
|
||||||
|
|
||||||
def set_component_device_value(
|
def set_component_device_value(
|
||||||
@@ -1902,6 +2045,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
# ── PART 2: Per-node state machine and leaf qualification ──
|
# ── PART 2: Per-node state machine and leaf qualification ──
|
||||||
expected_dev_leaves: set[UnifiedTreeNode] = set()
|
expected_dev_leaves: set[UnifiedTreeNode] = set()
|
||||||
expected_hst_leaves: set[UnifiedTreeNode] = set()
|
expected_hst_leaves: set[UnifiedTreeNode] = set()
|
||||||
|
expected_duplicates: set[UnifiedTreeNode] = set()
|
||||||
|
|
||||||
for node in all_nodes:
|
for node in all_nodes:
|
||||||
if node is self.root_node:
|
if node is self.root_node:
|
||||||
@@ -1918,7 +2062,10 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
if cd.value is not None and not full_dev:
|
if cd.value is not None and not full_dev:
|
||||||
E(f"node {nid} {ct} device present but Full.value=None")
|
E(f"node {nid} {ct} device present but Full.value=None")
|
||||||
if cd.host_value is not None and not full_hst:
|
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.
|
# Every node must keep Full data on at least one layer.
|
||||||
if not full_dev and not full_hst:
|
if not full_dev and not full_hst:
|
||||||
@@ -1951,6 +2098,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
expected_dev_leaves.add(node)
|
expected_dev_leaves.add(node)
|
||||||
if self._is_host_leaf(node):
|
if self._is_host_leaf(node):
|
||||||
expected_hst_leaves.add(node)
|
expected_hst_leaves.add(node)
|
||||||
|
if self._is_settled_full_host_duplicate(node):
|
||||||
|
expected_duplicates.add(node)
|
||||||
|
|
||||||
# ── PART 3: Tracking structures ──
|
# ── PART 3: Tracking structures ──
|
||||||
|
|
||||||
@@ -1972,6 +2121,16 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
if missing:
|
if missing:
|
||||||
E(f"H-leaf missing: {[n.id for n in list(missing)[:5]]}")
|
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 = ∅
|
# D-leaf ∩ H-leaf = ∅
|
||||||
overlap = self.evictable_device_leaves & self.evictable_host_leaves
|
overlap = self.evictable_device_leaves & self.evictable_host_leaves
|
||||||
if overlap:
|
if overlap:
|
||||||
@@ -2083,6 +2242,18 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
E(
|
E(
|
||||||
f"[Ongoing] load_back node {nid} lock_ref={n.component_data[FCT].lock_ref}"
|
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:
|
if errors:
|
||||||
msg = (
|
msg = (
|
||||||
|
|||||||
@@ -460,6 +460,15 @@ class UnifiedTreeCoreInterface(KVCacheEventMixin, ABC):
|
|||||||
"""Commit a successful H->D load-back onto the node; returns any cache actions."""
|
"""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
|
@abstractmethod
|
||||||
def mark_write_through_pending(self, node_id: NodeId) -> None:
|
def mark_write_through_pending(self, node_id: NodeId) -> None:
|
||||||
"""Mark a node as having an in-flight write-through backup."""
|
"""Mark a node as having an in-flight write-through backup."""
|
||||||
|
|||||||
@@ -1839,22 +1839,30 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
else ()
|
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(
|
ready_counts = torch.tensor(
|
||||||
[
|
[
|
||||||
write_acks,
|
write_acks,
|
||||||
load_acks,
|
load_acks,
|
||||||
*storage_queue_sizes,
|
*storage_queue_sizes,
|
||||||
|
digest,
|
||||||
|
-digest,
|
||||||
],
|
],
|
||||||
dtype=torch.int,
|
dtype=torch.int64,
|
||||||
device="cpu",
|
device="cpu",
|
||||||
)
|
)
|
||||||
self._all_reduce(ready_counts, torch.distributed.ReduceOp.MIN)
|
self._all_reduce(ready_counts, torch.distributed.ReduceOp.MIN)
|
||||||
|
|
||||||
count_values = list(map(int, ready_counts.tolist()))
|
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 (
|
return (
|
||||||
count_values[0],
|
count_values[0],
|
||||||
count_values[1],
|
count_values[1],
|
||||||
tuple(count_values[2:]),
|
tuple(count_values[2:-2]),
|
||||||
extra_pool_names,
|
extra_pool_names,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1926,11 +1934,17 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
finish_count = 0
|
finish_count = 0
|
||||||
if self.pp_rank == 0:
|
if self.pp_rank == 0:
|
||||||
finish_count = self._count_ready_acks(cc.ack_load_queue)
|
finish_count = self._count_ready_acks(cc.ack_load_queue)
|
||||||
finish_count_tensor = torch.tensor(
|
# Piggybacked TP check: [digest, -digest] MIN-reduces to [min, -max],
|
||||||
finish_count, dtype=torch.int, device="cpu"
|
# 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)
|
self._all_reduce(sync_tensor, torch.distributed.ReduceOp.MIN)
|
||||||
finish_count = finish_count_tensor.item()
|
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:
|
while finish_count > 0:
|
||||||
ack = cc.ack_load_queue.pop(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)
|
node, lock_params, host_lock_params = self.ongoing_load_back.pop(ack_id)
|
||||||
self.dec_lock_ref(node, lock_params)
|
self.dec_lock_ref(node, lock_params)
|
||||||
self.dec_host_lock_ref(node, host_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:
|
if self.metrics_collector is not None:
|
||||||
for pool, num_tokens in (ack.num_tokens_by_pool or {}).items():
|
for pool, num_tokens in (ack.num_tokens_by_pool or {}).items():
|
||||||
|
|||||||
@@ -2786,6 +2786,8 @@ class UnifiedRadixCacheSuite:
|
|||||||
cd = ancestor.component_data[ct]
|
cd = ancestor.component_data[ct]
|
||||||
if cd.value is not None and cd.host_value is None:
|
if cd.value is not None and cd.host_value is None:
|
||||||
cd.host_value = cd.value.clone()
|
cd.host_value = cd.value.clone()
|
||||||
|
# A real backup registers duplicate tracking at its ack.
|
||||||
|
cache.tree_core._update_duplicate_tracking(ancestor)
|
||||||
|
|
||||||
def _simulate_backup_tree(self, cache):
|
def _simulate_backup_tree(self, cache):
|
||||||
"""Backup all non-root nodes (simulates write-through)."""
|
"""Backup all non-root nodes (simulates write-through)."""
|
||||||
|
|||||||
Reference in New Issue
Block a user