|
|
|
@@ -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 = (
|
|
|
|
|