[UnifiedTree]: Replace anonymous tuples with NamedTuples in UnifiedRadixCache (#28375)
This commit is contained in:
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user