[Radix Cache] Add test-only TreeCore inspector for shared backend tests (#35791)
This commit is contained in:
@@ -1537,47 +1537,23 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
):
|
):
|
||||||
"""Cascade eviction from trigger to lower-or-equal priority components."""
|
"""Cascade eviction from trigger to lower-or-equal priority components."""
|
||||||
|
|
||||||
is_leaf = False
|
is_leaf = self._is_cascade_evict_leaf(node, target)
|
||||||
if target == EvictLayer.DEVICE:
|
|
||||||
is_leaf = node in self.evictable_device_leaves
|
|
||||||
elif target == EvictLayer.HOST:
|
|
||||||
is_leaf = node in self.evictable_host_leaves
|
|
||||||
|
|
||||||
trigger_priority = trigger.eviction_priority(is_leaf)
|
|
||||||
base_evicted = False
|
base_evicted = False
|
||||||
|
|
||||||
for comp in self.components:
|
for comp in self.components:
|
||||||
if comp.eviction_priority(is_leaf) <= trigger_priority:
|
if self._should_cascade_evict_component(
|
||||||
if comp is not trigger and comp.node_has_component_data(node, target):
|
node, trigger, comp, target, is_leaf
|
||||||
cd = node.component_data[comp.component_type]
|
):
|
||||||
# A comp whose TRUE internal priority outranks the trigger
|
self._evict_component_and_detach_lru(
|
||||||
# is only in this loop because leaf-collapse flattened
|
node,
|
||||||
# priorities; a lock on it is a legit pin and must be
|
comp,
|
||||||
# spared. A lock on a strictly-lower-priority tier is a
|
target=target,
|
||||||
# real strand — fall through to the assert below.
|
tracker=tracker,
|
||||||
if comp.eviction_priority(
|
device_frees=device_frees,
|
||||||
is_leaf=False
|
host_frees=host_frees,
|
||||||
) >= trigger.eviction_priority(is_leaf=False):
|
)
|
||||||
if EvictLayer.DEVICE in target and cd.lock_ref != 0:
|
if comp.component_type == BASE_COMPONENT_TYPE:
|
||||||
continue
|
base_evicted = True
|
||||||
if EvictLayer.HOST in target and cd.host_lock_ref != 0:
|
|
||||||
continue
|
|
||||||
if cd.session_ref > 0 and trigger.session_ref(node) == 0:
|
|
||||||
continue
|
|
||||||
if EvictLayer.DEVICE in target:
|
|
||||||
assert cd.lock_ref == 0
|
|
||||||
if EvictLayer.HOST in target:
|
|
||||||
assert cd.host_lock_ref == 0
|
|
||||||
self._evict_component_and_detach_lru(
|
|
||||||
node,
|
|
||||||
comp,
|
|
||||||
target=target,
|
|
||||||
tracker=tracker,
|
|
||||||
device_frees=device_frees,
|
|
||||||
host_frees=host_frees,
|
|
||||||
)
|
|
||||||
if comp.component_type == BASE_COMPONENT_TYPE:
|
|
||||||
base_evicted = True
|
|
||||||
|
|
||||||
# Now that all components (including SWA which depends on Full.value)
|
# Now that all components (including SWA which depends on Full.value)
|
||||||
# have been freed, we can safely tombstone Full.value.
|
# have been freed, we can safely tombstone Full.value.
|
||||||
@@ -1593,6 +1569,48 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
|
|
||||||
self._update_evictable_leaf_sets(node)
|
self._update_evictable_leaf_sets(node)
|
||||||
|
|
||||||
|
def _is_cascade_evict_leaf(self, node: UnifiedTreeNode, target: EvictLayer) -> bool:
|
||||||
|
if target == EvictLayer.DEVICE:
|
||||||
|
return node in self.evictable_device_leaves
|
||||||
|
if target == EvictLayer.HOST:
|
||||||
|
return node in self.evictable_host_leaves
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _should_cascade_evict_component(
|
||||||
|
node: UnifiedTreeNode,
|
||||||
|
trigger: TreeComponent,
|
||||||
|
comp: TreeComponent,
|
||||||
|
target: EvictLayer,
|
||||||
|
is_leaf: bool,
|
||||||
|
) -> bool:
|
||||||
|
"""Return whether a component is an unlocked cascade-eviction target."""
|
||||||
|
trigger_priority = trigger.eviction_priority(is_leaf)
|
||||||
|
if comp.eviction_priority(is_leaf) > trigger_priority:
|
||||||
|
return False
|
||||||
|
if comp is trigger or not comp.node_has_component_data(node, target):
|
||||||
|
return False
|
||||||
|
|
||||||
|
cd = node.component_data[comp.component_type]
|
||||||
|
# A comp whose TRUE internal priority outranks the trigger is only a
|
||||||
|
# candidate because leaf-collapse flattened priorities; a lock on it is
|
||||||
|
# a legitimate pin and must be spared. A lock on a strictly-lower-
|
||||||
|
# priority tier is a real strand and must trip the assertions below.
|
||||||
|
if comp.eviction_priority(is_leaf=False) >= trigger.eviction_priority(
|
||||||
|
is_leaf=False
|
||||||
|
):
|
||||||
|
if EvictLayer.DEVICE in target and cd.lock_ref != 0:
|
||||||
|
return False
|
||||||
|
if EvictLayer.HOST in target and cd.host_lock_ref != 0:
|
||||||
|
return False
|
||||||
|
if cd.session_ref > 0 and trigger.session_ref(node) == 0:
|
||||||
|
return False
|
||||||
|
if EvictLayer.DEVICE in target:
|
||||||
|
assert cd.lock_ref == 0
|
||||||
|
if EvictLayer.HOST in target:
|
||||||
|
assert cd.host_lock_ref == 0
|
||||||
|
return True
|
||||||
|
|
||||||
def _remove_leaf_from_parent(self, node: UnifiedTreeNode):
|
def _remove_leaf_from_parent(self, node: UnifiedTreeNode):
|
||||||
for component in self.components:
|
for component in self.components:
|
||||||
component.discard_deleted_session_leaf(node)
|
component.discard_deleted_session_leaf(node)
|
||||||
@@ -1648,8 +1666,19 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
- Full host present → keep as H-leaf
|
- Full host present → keep as H-leaf
|
||||||
- neither → evict all remaining data, delete, continue up
|
- neither → evict all remaining data, delete, continue up
|
||||||
"""
|
"""
|
||||||
|
self._iteratively_delete_tombstone_ancestors(
|
||||||
|
deleted_node.parent, tracker, device_frees, host_frees
|
||||||
|
)
|
||||||
|
|
||||||
|
def _iteratively_delete_tombstone_ancestors(
|
||||||
|
self,
|
||||||
|
cur: UnifiedTreeNode,
|
||||||
|
tracker: dict[ComponentType, int],
|
||||||
|
device_frees: dict[ComponentType, list[torch.Tensor]],
|
||||||
|
host_frees: dict[ComponentType, list[torch.Tensor]],
|
||||||
|
) -> None:
|
||||||
|
"""Delete childless tombstone ancestors until a live or locked node is reached."""
|
||||||
ct = BASE_COMPONENT_TYPE
|
ct = BASE_COMPONENT_TYPE
|
||||||
cur = deleted_node.parent
|
|
||||||
while cur != self.root_node and len(cur.children) == 0:
|
while cur != self.root_node and len(cur.children) == 0:
|
||||||
if any(
|
if any(
|
||||||
cd.lock_ref > 0 or cd.host_lock_ref > 0 for cd in cur.component_data
|
cd.lock_ref > 0 or cd.host_lock_ref > 0 for cd in cur.component_data
|
||||||
|
|||||||
@@ -3,6 +3,8 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from unittest import mock
|
from unittest import mock
|
||||||
|
|
||||||
|
from unified_tree_core_inspection_interface import UnifiedTreeCoreInspectionInterface
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
@@ -98,6 +100,7 @@ class TreeCoreRegistryTest(CustomTestCase):
|
|||||||
components={ComponentType.FULL: component},
|
components={ComponentType.FULL: component},
|
||||||
)
|
)
|
||||||
self.assertIsInstance(core, UnifiedTreeCore)
|
self.assertIsInstance(core, UnifiedTreeCore)
|
||||||
|
self.assertNotIsInstance(core, UnifiedTreeCoreInspectionInterface)
|
||||||
self.assertIs(component.tree_core, core)
|
self.assertIs(component.tree_core, core)
|
||||||
|
|
||||||
def test_unknown_backend_raises_naming_the_known_backends(self):
|
def test_unknown_backend_raises_naming_the_known_backends(self):
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,250 @@
|
|||||||
|
"""Test-only TreeCore surface used by backend-neutral white-box tests.
|
||||||
|
|
||||||
|
Production cache and controller code depends only on ``UnifiedTreeCoreInterface``.
|
||||||
|
A TreeCore backend implements this extended interface only when it opts into the
|
||||||
|
shared unified radix-cache conformance suite. Some methods intentionally mutate
|
||||||
|
internal state to construct edge cases; they must not be used by runtime code.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from abc import abstractmethod
|
||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.unified_cache.unified_tree_core_interface import (
|
||||||
|
BaseEvictionResult,
|
||||||
|
NodeId,
|
||||||
|
UnifiedTreeCoreInterface,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(
|
||||||
|
est_time=0, suite="base-a-test-cpu", disabled="TreeCore inspection test helper"
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import MatchPrefixParams, MatchResult
|
||||||
|
from sglang.srt.mem_cache.unified_cache.components import ComponentType, EvictLayer
|
||||||
|
|
||||||
|
|
||||||
|
class UnifiedTreeCoreInspectionInterface(UnifiedTreeCoreInterface):
|
||||||
|
"""Test-only inspection and control contract for shared backend tests."""
|
||||||
|
|
||||||
|
# ==== Read-only inspection ====
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def contains_node(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node id is live in the tree."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_parent_node_id(self, node_id: NodeId) -> Optional[NodeId]:
|
||||||
|
"""The parent node id, or None for the root."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_child_node_ids(self, node_id: NodeId) -> list[NodeId]:
|
||||||
|
"""The node's child ids."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_node_key_length(self, node_id: NodeId) -> int:
|
||||||
|
"""The node's logical radix-key length."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_node_token_ids(self, node_id: NodeId) -> list[int]:
|
||||||
|
"""The raw token ids spanned by the node's radix key."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def is_node_key_bigram(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node's radix key uses bigram atoms."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_component_host_value(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> Optional[torch.Tensor]:
|
||||||
|
"""The component's host value on the node, or None if absent."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_component_device_lock_ref(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> int:
|
||||||
|
"""The component's device lock count on the node."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_node_hit_count(self, node_id: NodeId) -> int:
|
||||||
|
"""The node's accumulated match count."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_write_through_pending_id(self, node_id: NodeId) -> Optional[int]:
|
||||||
|
"""The node's pending write-through id, if any."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def is_node_in_device_lru(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> bool:
|
||||||
|
"""Whether the node belongs to the component's device LRU."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def is_node_in_host_lru(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> bool:
|
||||||
|
"""Whether the node belongs to the component's host LRU."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_component_device_lru_node_ids(
|
||||||
|
self, component_type: ComponentType
|
||||||
|
) -> list[NodeId]:
|
||||||
|
"""The component's device LRU members from most to least recent."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def is_device_evictable_leaf(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node belongs to the device-evictable leaf set."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def is_host_evictable_leaf(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node belongs to the host-evictable leaf set."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def is_device_leaf(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node has no device-resident descendants."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_all_node_ids(self) -> list[NodeId]:
|
||||||
|
"""All live tree node ids."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def component_protected_size(self, component_type: ComponentType) -> int:
|
||||||
|
"""Protected token count for one component (0 if the component is absent)."""
|
||||||
|
...
|
||||||
|
|
||||||
|
# ==== White-box state controls ====
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def set_node_hash_values(
|
||||||
|
self, node_id: NodeId, hash_values: Optional[list[str]]
|
||||||
|
) -> None:
|
||||||
|
"""Replace the node's page-hash field."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def set_component_device_value_raw(
|
||||||
|
self,
|
||||||
|
node_id: NodeId,
|
||||||
|
component_type: ComponentType,
|
||||||
|
value: Optional[torch.Tensor],
|
||||||
|
) -> None:
|
||||||
|
"""Replace the device-value field without updating tree bookkeeping."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def set_component_host_value_raw(
|
||||||
|
self,
|
||||||
|
node_id: NodeId,
|
||||||
|
component_type: ComponentType,
|
||||||
|
value: Optional[torch.Tensor],
|
||||||
|
) -> None:
|
||||||
|
"""Replace the host-value field without updating tree bookkeeping."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def set_component_device_lock_ref(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType, lock_ref: int
|
||||||
|
) -> None:
|
||||||
|
"""Replace the component's device lock count."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def remove_node_from_device_lru(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> None:
|
||||||
|
"""Remove the node from the component's device LRU."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def insert_node_into_host_lru(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> None:
|
||||||
|
"""Insert the node as the component's most-recent host entry."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def set_component_evictable_size(
|
||||||
|
self, component_type: ComponentType, value: int
|
||||||
|
) -> None:
|
||||||
|
"""Replace the component's evictable device-token count."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def set_component_protected_size(
|
||||||
|
self, component_type: ComponentType, value: int
|
||||||
|
) -> None:
|
||||||
|
"""Replace the component's protected device-token count."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def update_duplicate_tracking(self, node_id: NodeId) -> None:
|
||||||
|
"""Refresh duplicate-host tracking for the node."""
|
||||||
|
...
|
||||||
|
|
||||||
|
# ==== Targeted white-box operations ====
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def evict_component(
|
||||||
|
self,
|
||||||
|
node_id: NodeId,
|
||||||
|
component_type: ComponentType,
|
||||||
|
target: EvictLayer,
|
||||||
|
) -> BaseEvictionResult:
|
||||||
|
"""Evict one component layer from a node and detach its LRU entry."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def validate_cascade_evict(
|
||||||
|
self,
|
||||||
|
node_id: NodeId,
|
||||||
|
component_type: ComponentType,
|
||||||
|
target: EvictLayer,
|
||||||
|
) -> None:
|
||||||
|
"""Validate the locks for a component-triggered cascade eviction."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def cleanup_tombstone_ancestors(self, node_id: NodeId) -> BaseEvictionResult:
|
||||||
|
"""Delete childless tombstone ancestors until a live or locked node is reached."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def finalize_component_match_result(
|
||||||
|
self,
|
||||||
|
component_type: ComponentType,
|
||||||
|
result: MatchResult,
|
||||||
|
params: MatchPrefixParams,
|
||||||
|
value_chunks: list[torch.Tensor],
|
||||||
|
best_value_len: int,
|
||||||
|
) -> MatchResult:
|
||||||
|
"""Run one component's match finalizer with NodeId boundaries."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def build_backup_node_ids(
|
||||||
|
self, node_id: NodeId, write_back: bool = False
|
||||||
|
) -> list[NodeId]:
|
||||||
|
"""Build the ordered node list for a device-to-host backup."""
|
||||||
|
...
|
||||||
@@ -0,0 +1,274 @@
|
|||||||
|
"""Test-only Python TreeCore implementation with white-box capabilities."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from unified_tree_core_inspection_interface import (
|
||||||
|
UnifiedTreeCoreInspectionInterface,
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import MatchPrefixParams, MatchResult
|
||||||
|
from sglang.srt.mem_cache.unified_cache.components import ComponentType, EvictLayer
|
||||||
|
from sglang.srt.mem_cache.unified_cache.unified_tree_core import (
|
||||||
|
UnifiedLRUList,
|
||||||
|
UnifiedTreeCore,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.unified_cache.unified_tree_core_interface import (
|
||||||
|
BaseEvictionResult,
|
||||||
|
NodeId,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(
|
||||||
|
est_time=0, suite="base-a-test-cpu", disabled="Python TreeCore test inspector"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class UnifiedTreeCoreInspector(UnifiedTreeCore, UnifiedTreeCoreInspectionInterface):
|
||||||
|
"""Python TreeCore variant used by the shared backend-conformance tests."""
|
||||||
|
|
||||||
|
def contains_node(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node id is live in the tree."""
|
||||||
|
return node_id in self._node_arena
|
||||||
|
|
||||||
|
def get_parent_node_id(self, node_id: NodeId) -> Optional[NodeId]:
|
||||||
|
"""The parent node id, or None for the root."""
|
||||||
|
parent = self.node_by_id(node_id).parent
|
||||||
|
return None if parent is None else parent.id
|
||||||
|
|
||||||
|
def get_child_node_ids(self, node_id: NodeId) -> list[NodeId]:
|
||||||
|
"""The node's child ids."""
|
||||||
|
return [child.id for child in self.node_by_id(node_id).children.values()]
|
||||||
|
|
||||||
|
def get_node_key_length(self, node_id: NodeId) -> int:
|
||||||
|
"""The node's logical radix-key length."""
|
||||||
|
key = self.node_by_id(node_id).key
|
||||||
|
assert key is not None
|
||||||
|
return len(key)
|
||||||
|
|
||||||
|
def get_node_token_ids(self, node_id: NodeId) -> list[int]:
|
||||||
|
"""The raw token ids spanned by the node's radix key."""
|
||||||
|
key = self.node_by_id(node_id).key
|
||||||
|
assert key is not None
|
||||||
|
return list(key.raw_token_ids())
|
||||||
|
|
||||||
|
def is_node_key_bigram(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node's radix key uses bigram atoms."""
|
||||||
|
key = self.node_by_id(node_id).key
|
||||||
|
assert key is not None
|
||||||
|
return key.is_bigram
|
||||||
|
|
||||||
|
def get_component_host_value(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> Optional[torch.Tensor]:
|
||||||
|
"""The component's host value on the node, or None if absent."""
|
||||||
|
return self.node_by_id(node_id).component_data[component_type].host_value
|
||||||
|
|
||||||
|
def get_component_device_lock_ref(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> int:
|
||||||
|
"""The component's device lock count on the node."""
|
||||||
|
return self.node_by_id(node_id).component_data[component_type].lock_ref
|
||||||
|
|
||||||
|
def get_node_hit_count(self, node_id: NodeId) -> int:
|
||||||
|
"""The node's accumulated match count."""
|
||||||
|
return self.node_by_id(node_id).hit_count
|
||||||
|
|
||||||
|
def get_write_through_pending_id(self, node_id: NodeId) -> Optional[int]:
|
||||||
|
"""The node's pending write-through id, if any."""
|
||||||
|
return self.node_by_id(node_id).write_through_pending_id
|
||||||
|
|
||||||
|
def is_node_in_device_lru(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> bool:
|
||||||
|
"""Whether the node belongs to the component's device LRU."""
|
||||||
|
lru = self.lru_lists.get(component_type)
|
||||||
|
return lru is not None and lru.in_list(self.node_by_id(node_id))
|
||||||
|
|
||||||
|
def is_node_in_host_lru(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> bool:
|
||||||
|
"""Whether the node belongs to the component's host LRU."""
|
||||||
|
lru = self.host_lru_lists.get(component_type)
|
||||||
|
return lru is not None and lru.in_list(self.node_by_id(node_id))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _lru_node_ids(lru: UnifiedLRUList) -> list[NodeId]:
|
||||||
|
"""Return real LRU members from most to least recent."""
|
||||||
|
node_ids = []
|
||||||
|
node = lru.head.lru_next[lru._pt]
|
||||||
|
while node is not lru.tail:
|
||||||
|
if node.id in lru.cache:
|
||||||
|
node_ids.append(node.id)
|
||||||
|
node = node.lru_next[lru._pt]
|
||||||
|
return node_ids
|
||||||
|
|
||||||
|
def get_component_device_lru_node_ids(
|
||||||
|
self, component_type: ComponentType
|
||||||
|
) -> list[NodeId]:
|
||||||
|
"""The component's device LRU members from most to least recent."""
|
||||||
|
lru = self.lru_lists.get(component_type)
|
||||||
|
return [] if lru is None else self._lru_node_ids(lru)
|
||||||
|
|
||||||
|
def is_device_evictable_leaf(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node belongs to the device-evictable leaf set."""
|
||||||
|
node = self._node_arena.get(node_id)
|
||||||
|
return node is not None and node in self.evictable_device_leaves
|
||||||
|
|
||||||
|
def is_host_evictable_leaf(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node belongs to the host-evictable leaf set."""
|
||||||
|
node = self._node_arena.get(node_id)
|
||||||
|
return node is not None and node in self.evictable_host_leaves
|
||||||
|
|
||||||
|
def is_device_leaf(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node has no device-resident descendants."""
|
||||||
|
return self._is_device_leaf(self.node_by_id(node_id))
|
||||||
|
|
||||||
|
def get_all_node_ids(self) -> list[NodeId]:
|
||||||
|
"""All live tree node ids."""
|
||||||
|
return [node.id for node in self._collect_all_nodes()]
|
||||||
|
|
||||||
|
def component_protected_size(self, component_type: ComponentType) -> int:
|
||||||
|
"""Protected token count for one component (0 if the component is absent)."""
|
||||||
|
return self.component_protected_size_.get(component_type, 0)
|
||||||
|
|
||||||
|
def set_node_hash_values(
|
||||||
|
self, node_id: NodeId, hash_values: Optional[list[str]]
|
||||||
|
) -> None:
|
||||||
|
"""Replace the node's page-hash field."""
|
||||||
|
self.node_by_id(node_id).hash_value = hash_values
|
||||||
|
|
||||||
|
def set_component_device_value_raw(
|
||||||
|
self,
|
||||||
|
node_id: NodeId,
|
||||||
|
component_type: ComponentType,
|
||||||
|
value: Optional[torch.Tensor],
|
||||||
|
) -> None:
|
||||||
|
"""Replace the device-value field without updating tree bookkeeping."""
|
||||||
|
self.node_by_id(node_id).component_data[component_type].value = value
|
||||||
|
|
||||||
|
def set_component_host_value_raw(
|
||||||
|
self,
|
||||||
|
node_id: NodeId,
|
||||||
|
component_type: ComponentType,
|
||||||
|
value: Optional[torch.Tensor],
|
||||||
|
) -> None:
|
||||||
|
"""Replace the host-value field without updating tree bookkeeping."""
|
||||||
|
self.node_by_id(node_id).component_data[component_type].host_value = value
|
||||||
|
|
||||||
|
def set_component_device_lock_ref(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType, lock_ref: int
|
||||||
|
) -> None:
|
||||||
|
"""Replace the component's device lock count."""
|
||||||
|
assert lock_ref >= 0
|
||||||
|
self.node_by_id(node_id).component_data[component_type].lock_ref = lock_ref
|
||||||
|
|
||||||
|
def remove_node_from_device_lru(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> None:
|
||||||
|
"""Remove the node from the component's device LRU."""
|
||||||
|
self.lru_lists[component_type].remove_node(self.node_by_id(node_id))
|
||||||
|
|
||||||
|
def insert_node_into_host_lru(
|
||||||
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
) -> None:
|
||||||
|
"""Insert the node as the component's most-recent host entry."""
|
||||||
|
self.host_lru_lists[component_type].insert_mru(self.node_by_id(node_id))
|
||||||
|
|
||||||
|
def set_component_evictable_size(
|
||||||
|
self, component_type: ComponentType, value: int
|
||||||
|
) -> None:
|
||||||
|
"""Replace the component's evictable device-token count."""
|
||||||
|
assert value >= 0
|
||||||
|
self.component_evictable_size_[component_type] = value
|
||||||
|
|
||||||
|
def set_component_protected_size(
|
||||||
|
self, component_type: ComponentType, value: int
|
||||||
|
) -> None:
|
||||||
|
"""Replace the component's protected device-token count."""
|
||||||
|
assert value >= 0
|
||||||
|
self.component_protected_size_[component_type] = value
|
||||||
|
|
||||||
|
def update_duplicate_tracking(self, node_id: NodeId) -> None:
|
||||||
|
"""Refresh duplicate-host tracking for the node."""
|
||||||
|
self._update_duplicate_tracking(self.node_by_id(node_id))
|
||||||
|
|
||||||
|
def evict_component(
|
||||||
|
self,
|
||||||
|
node_id: NodeId,
|
||||||
|
component_type: ComponentType,
|
||||||
|
target: EvictLayer,
|
||||||
|
) -> BaseEvictionResult:
|
||||||
|
"""Evict one component layer from a node and detach its LRU entry."""
|
||||||
|
result = BaseEvictionResult()
|
||||||
|
self._evict_component_and_detach_lru(
|
||||||
|
self.node_by_id(node_id),
|
||||||
|
self.components_by_type[component_type],
|
||||||
|
result.device_frees,
|
||||||
|
result.host_frees,
|
||||||
|
target,
|
||||||
|
result.tracker,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
def validate_cascade_evict(
|
||||||
|
self,
|
||||||
|
node_id: NodeId,
|
||||||
|
component_type: ComponentType,
|
||||||
|
target: EvictLayer,
|
||||||
|
) -> None:
|
||||||
|
"""Validate the locks for a component-triggered cascade eviction."""
|
||||||
|
node = self.node_by_id(node_id)
|
||||||
|
trigger = self.components_by_type[component_type]
|
||||||
|
is_leaf = self._is_cascade_evict_leaf(node, target)
|
||||||
|
for comp in self.components:
|
||||||
|
self._should_cascade_evict_component(node, trigger, comp, target, is_leaf)
|
||||||
|
|
||||||
|
def cleanup_tombstone_ancestors(self, node_id: NodeId) -> BaseEvictionResult:
|
||||||
|
"""Delete childless tombstone ancestors until a live or locked node is reached."""
|
||||||
|
result = BaseEvictionResult()
|
||||||
|
self._iteratively_delete_tombstone_ancestors(
|
||||||
|
self.node_by_id(node_id),
|
||||||
|
result.tracker,
|
||||||
|
result.device_frees,
|
||||||
|
result.host_frees,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
def finalize_component_match_result(
|
||||||
|
self,
|
||||||
|
component_type: ComponentType,
|
||||||
|
result: MatchResult,
|
||||||
|
params: MatchPrefixParams,
|
||||||
|
value_chunks: list[torch.Tensor],
|
||||||
|
best_value_len: int,
|
||||||
|
) -> MatchResult:
|
||||||
|
"""Run one component's match finalizer with NodeId boundaries."""
|
||||||
|
node_result = result._replace(
|
||||||
|
last_device_node=self.node_by_id(result.last_device_node),
|
||||||
|
last_host_node=self.node_by_id(result.last_host_node),
|
||||||
|
best_match_node=self.node_by_id(result.best_match_node),
|
||||||
|
)
|
||||||
|
finalized = self.components_by_type[
|
||||||
|
component_type
|
||||||
|
].finalize_match_result_in_tree_core(
|
||||||
|
result=node_result,
|
||||||
|
params=params,
|
||||||
|
value_chunks=value_chunks,
|
||||||
|
best_value_len=best_value_len,
|
||||||
|
)
|
||||||
|
return finalized._replace(
|
||||||
|
last_device_node=finalized.last_device_node.id,
|
||||||
|
last_host_node=finalized.last_host_node.id,
|
||||||
|
best_match_node=finalized.best_match_node.id,
|
||||||
|
)
|
||||||
|
|
||||||
|
def build_backup_node_ids(
|
||||||
|
self, node_id: NodeId, write_back: bool = False
|
||||||
|
) -> list[NodeId]:
|
||||||
|
"""Build the ordered node list for a device-to-host backup."""
|
||||||
|
return self._build_backup_kv_action(
|
||||||
|
self.node_by_id(node_id), write_back
|
||||||
|
).node_ids
|
||||||
Reference in New Issue
Block a user