Write-back policy fix for unified tree (#31845)

This commit is contained in:
Zhiqiang Xie
2026-07-24 09:58:18 -07:00
committed by GitHub
parent 448662e85e
commit 5da0b6ec39
3 changed files with 164 additions and 26 deletions
@@ -182,7 +182,9 @@ class FullComponent(TreeComponent):
# Only the last host node needs to be protected.
if lock_host:
cd = node.component_data[ct]
if cd.host_value is None:
# write_back mode: the anchor may be device-only (no host_value);
# pin it anyway so the host-pressure drop fallback cannot delete it.
if cd.host_value is None and not self.cache.is_write_back:
return result
cd.host_lock_ref += 1
self.cache._update_evictable_leaf_sets(node)
@@ -223,7 +225,10 @@ class FullComponent(TreeComponent):
ct = self.component_type
if lock_host:
cd = node.component_data[ct]
if cd.host_value is None or cd.host_lock_ref == 0:
if cd.host_lock_ref == 0:
return
# Mirror of `acquire`. write_back uses a pure counter.
if cd.host_value is None and not self.cache.is_write_back:
return
cd.host_lock_ref -= 1
self.cache._update_evictable_leaf_sets(node)
@@ -617,6 +617,13 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
result = self._insert_helper(self.root_node, key, value, params)
return result
@property
def is_write_back(self) -> bool:
return (
self.cache_controller is not None
and self.cache_controller.write_policy == "write_back"
)
def evict(self, params: EvictParams) -> EvictResult:
if self.disable:
return EvictResult()
@@ -1506,25 +1513,90 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
):
written = self.write_backup(node, write_back=True)
if written == 0:
if self._drop_subtree_no_host(node, tracker):
logger.warning(
"write_back: KV subtree dropped without backup "
"due to host memory pressure, root node %d",
node.id,
)
else:
logger.warning(
"write_back: backup failed under host memory "
"pressure but subtree drop declined (node "
"locked); root node %d stays device-resident "
"until host space frees",
node.id,
)
return
self.writing_check(write_back=True)
self._evict_to_host(node, tracker)
return
else:
# Write-through: node has no backup, delete entirely.
self._record_remove_event(node, medium=StorageMedium.GPU)
for comp in self._components_tuple:
self._evict_component_and_detach_lru(
node, comp, target=EvictLayer.ALL, tracker=tracker
)
self.evictable_device_leaves.discard(node)
parent = node.parent
self._remove_leaf_from_parent(node)
self._update_evictable_leaf_sets(parent)
self._iteratively_delete_tombstone_leaf(node, tracker)
self._delete_unbacked_device_leaf(node, tracker)
return
self._evict_to_host(node, tracker)
def _drop_subtree_no_host(
self, node: UnifiedTreeNode, tracker: dict[ComponentType, int]
) -> bool:
"""Write-back fallback when a D-leaf's D->H backup fails under host
memory pressure: drop the subtree rooted at the unbacked leaf so
device eviction keeps making progress instead of leaving its KV
unevictable until host space frees up."""
assert self._is_device_leaf(node), f"node {node.id} is not a D-leaf"
# A failed backup never issues the D->H copy, so the subtree root has
# no host state and no in-flight DMA reading its device slots.
assert not node.backuped and node.write_through_pending_id is None
if any(cd.host_lock_ref > 0 for cd in node.component_data):
return False
descendants: list[UnifiedTreeNode] = []
stack = list(node.children.values())
while stack:
cur = stack.pop()
if any(
cd.lock_ref > 0 or cd.host_lock_ref > 0 for cd in cur.component_data
):
return False
descendants.append(cur)
stack.extend(cur.children.values())
for desc in reversed(descendants):
# Host-only by construction: a device descendant would contradict
# this node being a D-leaf, and D-leaves evict before ancestors.
assert desc.evicted and desc.backuped, f"node {desc.id} not host-only"
assert desc.write_through_pending_id is None
self._release_all_component_layers(desc, StorageMedium.CPU, tracker)
self._remove_leaf_from_parent(desc)
self._delete_unbacked_device_leaf(node, tracker)
return True
def _release_all_component_layers(
self,
node: UnifiedTreeNode,
medium: StorageMedium,
tracker: dict[ComponentType, int],
) -> None:
"""Free every component layer on the node and detach it from the LRU
lists and evictable leaf sets."""
self._record_remove_event(node, medium=medium)
for comp in self._components_tuple:
self._evict_component_and_detach_lru(
node, comp, target=EvictLayer.ALL, tracker=tracker
)
self.evictable_device_leaves.discard(node)
self.evictable_host_leaves.discard(node)
def _delete_unbacked_device_leaf(
self, node: UnifiedTreeNode, tracker: dict[ComponentType, int]
) -> None:
"""Delete a device leaf that has no host backup, freeing all layers."""
self._release_all_component_layers(node, StorageMedium.GPU, tracker)
parent = node.parent
self._remove_leaf_from_parent(node)
self._update_evictable_leaf_sets(parent)
self._iteratively_delete_tombstone_leaf(node, tracker)
def _evict_host_leaf(
self, node: UnifiedTreeNode, tracker: dict[ComponentType, int]
) -> None: