[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 collections import defaultdict
from functools import partial from functools import partial
from queue import Empty, Queue 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 import torch
@@ -275,6 +275,33 @@ COMPONENT_REGISTRY: dict[ComponentType, type[TreeComponent]] = {
logger = logging.getLogger(__name__) 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): class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
def __init__( def __init__(
self, self,
@@ -437,31 +464,11 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
ct: UnifiedLRUList(ct, self.tree_components, use_host_ptr=True) ct: UnifiedLRUList(ct, self.tree_components, use_host_ptr=True)
for ct in self.tree_components for ct in self.tree_components
} }
self.ongoing_write_through: dict[ self.ongoing_write_through: dict[int, _OngoingWriteThrough] = {}
int, self.ongoing_load_back: dict[int, _OngoingLoadBack] = {}
tuple[
UnifiedTreeNode,
Optional[DecLockRefParams],
list[UnifiedTreeNode],
],
] = {}
self.ongoing_load_back: dict[
int,
tuple[UnifiedTreeNode, DecLockRefParams, DecLockRefParams],
] = {}
self.enable_storage = False self.enable_storage = False
self.prefetch_loaded_tokens_by_reqid: dict[str, int] = {} self.prefetch_loaded_tokens_by_reqid: dict[str, int] = {}
self.ongoing_prefetch: dict[ self.ongoing_prefetch: dict[str, _OngoingPrefetch] = {}
str,
tuple[
UnifiedTreeNode,
RadixKey,
torch.Tensor,
PrefetchOperation,
DecLockRefParams,
dict[ComponentType, list[PoolTransfer]],
],
] = {}
self.ongoing_backup: dict[int, tuple[UnifiedTreeNode, DecLockRefParams]] = {} self.ongoing_backup: dict[int, tuple[UnifiedTreeNode, DecLockRefParams]] = {}
if self.cache_controller is not None: if self.cache_controller is not None:
@@ -1561,7 +1568,9 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
lock_params: Optional[DecLockRefParams], lock_params: Optional[DecLockRefParams],
) -> None: ) -> None:
node.write_through_pending_id = node.id 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( def _replace_pending_write_through_node(
self, old_node: UnifiedTreeNode, new_nodes: list[UnifiedTreeNode] self, old_node: UnifiedTreeNode, new_nodes: list[UnifiedTreeNode]
@@ -1589,7 +1598,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
for node in new_nodes: for node in new_nodes:
node.write_through_pending_id = ack_id node.write_through_pending_id = ack_id
self.ongoing_write_through[ack_id] = ( self.ongoing_write_through[ack_id] = _OngoingWriteThrough(
lock_node, lock_node,
lock_params, lock_params,
updated_nodes, updated_nodes,
@@ -1698,7 +1707,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
) )
self._update_evictable_leaf_sets(best_match_node) 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, best_match_node,
self.inc_lock_ref(best_match_node).to_dec_params(), self.inc_lock_ref(best_match_node).to_dec_params(),
host_anchor_params, host_anchor_params,
@@ -1905,7 +1914,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
prefix_keys, prefix_keys,
extra_pools=aux_xfers or None, extra_pools=aux_xfers or None,
) )
self.ongoing_prefetch[req_id] = ( self.ongoing_prefetch[req_id] = _OngoingPrefetch(
last_host_node, last_host_node,
prefetch_key, prefetch_key,
host_indices, host_indices,
@@ -2039,7 +2048,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
def terminate_prefetch(self, req_id: str) -> None: def terminate_prefetch(self, req_id: str) -> None:
if req_id not in self.ongoing_prefetch: if req_id not in self.ongoing_prefetch:
return return
_, _, _, operation, _, _ = self.ongoing_prefetch[req_id] operation = self.ongoing_prefetch[req_id].operation
if operation.host_indices is None: if operation.host_indices is None:
return return
operation.mark_terminate() operation.mark_terminate()