[UnifiedTree]: Replace anonymous tuples with NamedTuples in UnifiedRadixCache (#28375)

This commit is contained in:
Zhangheng
2026-06-16 14:54:04 +08:00
committed by GitHub
parent 175336ff73
commit 6c908b3a3a
@@ -8,7 +8,7 @@ from array import array
from collections import defaultdict
from functools import partial
from queue import Empty, Queue
from typing import TYPE_CHECKING, Any, Iterator, Optional, TypeVar
from typing import TYPE_CHECKING, Any, Iterator, NamedTuple, Optional, TypeVar
import torch
@@ -275,6 +275,33 @@ COMPONENT_REGISTRY: dict[ComponentType, type[TreeComponent]] = {
logger = logging.getLogger(__name__)
class _OngoingWriteThrough(NamedTuple):
"""Tracks an in-flight D→H write-through operation."""
node: UnifiedTreeNode
lock_params: Optional[DecLockRefParams]
publish_nodes: list[UnifiedTreeNode]
class _OngoingLoadBack(NamedTuple):
"""Tracks an in-flight H→D load-back operation."""
node: UnifiedTreeNode
lock_params: DecLockRefParams
host_lock_params: DecLockRefParams
class _OngoingPrefetch(NamedTuple):
"""Tracks an in-flight storage→host prefetch operation."""
anchor_node: UnifiedTreeNode
prefetch_key: RadixKey
host_indices: torch.Tensor
operation: PrefetchOperation
anchor_lock_params: DecLockRefParams
comp_xfers: dict[ComponentType, list[PoolTransfer]]
class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
def __init__(
self,
@@ -437,31 +464,11 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
ct: UnifiedLRUList(ct, self.tree_components, use_host_ptr=True)
for ct in self.tree_components
}
self.ongoing_write_through: dict[
int,
tuple[
UnifiedTreeNode,
Optional[DecLockRefParams],
list[UnifiedTreeNode],
],
] = {}
self.ongoing_load_back: dict[
int,
tuple[UnifiedTreeNode, DecLockRefParams, DecLockRefParams],
] = {}
self.ongoing_write_through: dict[int, _OngoingWriteThrough] = {}
self.ongoing_load_back: dict[int, _OngoingLoadBack] = {}
self.enable_storage = False
self.prefetch_loaded_tokens_by_reqid: dict[str, int] = {}
self.ongoing_prefetch: dict[
str,
tuple[
UnifiedTreeNode,
RadixKey,
torch.Tensor,
PrefetchOperation,
DecLockRefParams,
dict[ComponentType, list[PoolTransfer]],
],
] = {}
self.ongoing_prefetch: dict[str, _OngoingPrefetch] = {}
self.ongoing_backup: dict[int, tuple[UnifiedTreeNode, DecLockRefParams]] = {}
if self.cache_controller is not None:
@@ -1561,7 +1568,9 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
lock_params: Optional[DecLockRefParams],
) -> None:
node.write_through_pending_id = node.id
self.ongoing_write_through[node.id] = (node, lock_params, [node])
self.ongoing_write_through[node.id] = _OngoingWriteThrough(
node, lock_params, [node]
)
def _replace_pending_write_through_node(
self, old_node: UnifiedTreeNode, new_nodes: list[UnifiedTreeNode]
@@ -1589,7 +1598,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
for node in new_nodes:
node.write_through_pending_id = ack_id
self.ongoing_write_through[ack_id] = (
self.ongoing_write_through[ack_id] = _OngoingWriteThrough(
lock_node,
lock_params,
updated_nodes,
@@ -1698,7 +1707,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
)
self._update_evictable_leaf_sets(best_match_node)
self.ongoing_load_back[best_match_node.id] = (
self.ongoing_load_back[best_match_node.id] = _OngoingLoadBack(
best_match_node,
self.inc_lock_ref(best_match_node).to_dec_params(),
host_anchor_params,
@@ -1905,7 +1914,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
prefix_keys,
extra_pools=aux_xfers or None,
)
self.ongoing_prefetch[req_id] = (
self.ongoing_prefetch[req_id] = _OngoingPrefetch(
last_host_node,
prefetch_key,
host_indices,
@@ -2039,7 +2048,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
def terminate_prefetch(self, req_id: str) -> None:
if req_id not in self.ongoing_prefetch:
return
_, _, _, operation, _, _ = self.ongoing_prefetch[req_id]
operation = self.ongoing_prefetch[req_id].operation
if operation.host_indices is None:
return
operation.mark_terminate()