[Rust TreeCore] Support external cache linker (#37306)
This commit is contained in:
@@ -679,15 +679,11 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def enable_external_cache_linker(self) -> bool:
|
def enable_external_cache_linker(self) -> bool:
|
||||||
return False
|
return self._binding.enable_external_cache_linker()
|
||||||
|
|
||||||
@enable_external_cache_linker.setter
|
@enable_external_cache_linker.setter
|
||||||
def enable_external_cache_linker(self, value: bool) -> None:
|
def enable_external_cache_linker(self, value: bool) -> None:
|
||||||
# TODO(Jialin): Port external cache linker support from #37091 and #37151.
|
self._binding.set_enable_external_cache_linker(value)
|
||||||
if value:
|
|
||||||
raise ValueError(
|
|
||||||
"External cache linker is not supported by the Rust TreeCore"
|
|
||||||
)
|
|
||||||
|
|
||||||
def insert_host(
|
def insert_host(
|
||||||
self,
|
self,
|
||||||
@@ -913,6 +909,27 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
def finish_load_back(self, anchor_node_id: NodeId) -> None:
|
def finish_load_back(self, anchor_node_id: NodeId) -> None:
|
||||||
self._binding.finish_load_back(anchor_node_id)
|
self._binding.finish_load_back(anchor_node_id)
|
||||||
|
|
||||||
|
def build_external_linker_offload_transfers(
|
||||||
|
self, node_id: NodeId
|
||||||
|
) -> Optional[list[PoolTransfer]]:
|
||||||
|
transfers = self._binding.build_external_linker_offload_transfers(node_id)
|
||||||
|
if transfers is None:
|
||||||
|
return None
|
||||||
|
return [_transfer_from_binding(transfer) for transfer in transfers]
|
||||||
|
|
||||||
|
def mark_external_cache_stored_path(
|
||||||
|
self, from_node_id: NodeId, until_node_id: NodeId
|
||||||
|
) -> None:
|
||||||
|
self._binding.mark_external_cache_stored_path(from_node_id, until_node_id)
|
||||||
|
|
||||||
|
def mark_external_linker_offload_pending(self, node_id: NodeId) -> None:
|
||||||
|
self._binding.mark_external_linker_offload_pending(node_id)
|
||||||
|
|
||||||
|
def finish_external_linker_offload(
|
||||||
|
self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool
|
||||||
|
) -> None:
|
||||||
|
self._binding.finish_external_linker_offload(list(node_ids), ack_id, success)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def write_back_duplicate_reclaim_digest(self) -> int:
|
def write_back_duplicate_reclaim_digest(self) -> int:
|
||||||
return self._binding.write_back_duplicate_reclaim_digest()
|
return self._binding.write_back_duplicate_reclaim_digest()
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixKey
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
from sglang.srt.mem_cache.unified_cache.components import (
|
from sglang.srt.mem_cache.unified_cache.components import (
|
||||||
|
ComponentType,
|
||||||
ExternalLinkerLoadPhase,
|
ExternalLinkerLoadPhase,
|
||||||
LinkerTransferPhase,
|
LinkerTransferPhase,
|
||||||
TreeComponent,
|
TreeComponent,
|
||||||
@@ -48,6 +49,14 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||||
|
|
||||||
|
|
||||||
|
_EXTERNAL_LINKER_SUPPORTED_COMPONENTS = frozenset(
|
||||||
|
{
|
||||||
|
ComponentType.FULL,
|
||||||
|
ComponentType.SWA,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class UnifiedCacheLinker(ABC):
|
class UnifiedCacheLinker(ABC):
|
||||||
"""External KV store reached directly from the device pools."""
|
"""External KV store reached directly from the device pools."""
|
||||||
|
|
||||||
@@ -139,6 +148,16 @@ class UnifiedCacheLinkerWrapper:
|
|||||||
cache: UnifiedRadixCache,
|
cache: UnifiedRadixCache,
|
||||||
cache_linker: UnifiedCacheLinker,
|
cache_linker: UnifiedCacheLinker,
|
||||||
):
|
):
|
||||||
|
unsupported = set(cache.tree_components) - _EXTERNAL_LINKER_SUPPORTED_COMPONENTS
|
||||||
|
if unsupported:
|
||||||
|
names = ", ".join(
|
||||||
|
component.name for component in sorted(unsupported, key=int)
|
||||||
|
)
|
||||||
|
raise ValueError(
|
||||||
|
"External cache linker supports only Full and SWA tree "
|
||||||
|
f"components; unsupported: {names}"
|
||||||
|
)
|
||||||
|
|
||||||
self.cache = cache
|
self.cache = cache
|
||||||
self.cache_linker = cache_linker
|
self.cache_linker = cache_linker
|
||||||
# rid -> what match found, consumed by the next init_load_back.
|
# rid -> what match found, consumed by the next init_load_back.
|
||||||
@@ -162,6 +181,7 @@ class UnifiedCacheLinkerWrapper:
|
|||||||
|
|
||||||
def match(self, key: RadixKey, req: Req, result: MatchResult) -> MatchResult:
|
def match(self, key: RadixKey, req: Req, result: MatchResult) -> MatchResult:
|
||||||
cache = self.cache
|
cache = self.cache
|
||||||
|
key, _ = key.maybe_to_bigram_view(cache.tree_core.is_eagle)
|
||||||
page = cache.page_size
|
page = cache.page_size
|
||||||
device_hit_len = int(result.device_indices.numel())
|
device_hit_len = int(result.device_indices.numel())
|
||||||
if device_hit_len >= len(key):
|
if device_hit_len >= len(key):
|
||||||
@@ -343,10 +363,9 @@ class UnifiedCacheLinkerWrapper:
|
|||||||
|
|
||||||
self._queue_load(req.rid, insert_result.last_device_node, load_transfers)
|
self._queue_load(req.rid, insert_result.last_device_node, load_transfers)
|
||||||
|
|
||||||
node = cache.resolve_node_handle(insert_result.last_device_node)
|
cache.tree_core.mark_external_cache_stored_path(
|
||||||
while node.id != req.last_node:
|
insert_result.last_device_node, req.last_node
|
||||||
node.external_cache_stored = True
|
)
|
||||||
node = node.parent
|
|
||||||
return canonical_tail, insert_result.last_device_node
|
return canonical_tail, insert_result.last_device_node
|
||||||
|
|
||||||
def _queue_load(
|
def _queue_load(
|
||||||
@@ -461,20 +480,14 @@ class UnifiedCacheLinkerWrapper:
|
|||||||
def offload_nodes(self, node_ids: Sequence[NodeId]) -> None:
|
def offload_nodes(self, node_ids: Sequence[NodeId]) -> None:
|
||||||
"""Persist a write-through chain, skipping nodes already in the store."""
|
"""Persist a write-through chain, skipping nodes already in the store."""
|
||||||
for node_id in node_ids:
|
for node_id in node_ids:
|
||||||
if not self.cache.resolve_node_handle(node_id).external_cache_stored:
|
transfers = self.cache.tree_core.build_external_linker_offload_transfers(
|
||||||
self._offload_node(node_id)
|
node_id
|
||||||
|
|
||||||
def _offload_node(self, node_id: NodeId) -> None:
|
|
||||||
cache = self.cache
|
|
||||||
node = cache.resolve_node_handle(node_id)
|
|
||||||
transfers = []
|
|
||||||
for component in cache._components_tuple:
|
|
||||||
transfer = component.build_external_linker_transfer(
|
|
||||||
LinkerTransferPhase.OFFLOAD, node, None
|
|
||||||
)
|
)
|
||||||
if transfer is not None:
|
if transfers is not None:
|
||||||
transfers.append(transfer)
|
self._offload_node(node_id, transfers)
|
||||||
|
|
||||||
|
def _offload_node(self, node_id: NodeId, transfers: list[PoolTransfer]) -> None:
|
||||||
|
cache = self.cache
|
||||||
lock_params = cache.inc_lock_ref(node_id).to_dec_params()
|
lock_params = cache.inc_lock_ref(node_id).to_dec_params()
|
||||||
try:
|
try:
|
||||||
queued = self.cache_linker.offload(transfers)
|
queued = self.cache_linker.offload(transfers)
|
||||||
@@ -485,8 +498,7 @@ class UnifiedCacheLinkerWrapper:
|
|||||||
cache.dec_lock_ref(node_id, lock_params)
|
cache.dec_lock_ref(node_id, lock_params)
|
||||||
return
|
return
|
||||||
|
|
||||||
cache.tree_core.mark_write_through_pending([node_id], ack_id=node_id)
|
cache.tree_core.mark_external_linker_offload_pending(node_id)
|
||||||
node.external_cache_stored = True
|
|
||||||
self.pending_offloads.append(_PendingOffload(node_id, lock_params, [node_id]))
|
self.pending_offloads.append(_PendingOffload(node_id, lock_params, [node_id]))
|
||||||
|
|
||||||
def replace_pending_offload_node(
|
def replace_pending_offload_node(
|
||||||
@@ -528,11 +540,9 @@ class UnifiedCacheLinkerWrapper:
|
|||||||
assert len(successes) <= len(self.pending_offloads)
|
assert len(successes) <= len(self.pending_offloads)
|
||||||
for success in successes:
|
for success in successes:
|
||||||
pending = self.pending_offloads.pop(0)
|
pending = self.pending_offloads.pop(0)
|
||||||
for node_id in pending.publish_node_ids:
|
self.cache.tree_core.finish_external_linker_offload(
|
||||||
node = self.cache.resolve_node_handle(node_id)
|
pending.publish_node_ids, pending.lock_node_id, success
|
||||||
if node.write_through_pending_id == pending.lock_node_id:
|
)
|
||||||
node.write_through_pending_id = None
|
|
||||||
node.external_cache_stored = success
|
|
||||||
self.cache.dec_lock_ref(pending.lock_node_id, pending.lock_params)
|
self.cache.dec_lock_ref(pending.lock_node_id, pending.lock_params)
|
||||||
|
|
||||||
def start_layer_wise_loading(self) -> int:
|
def start_layer_wise_loading(self) -> int:
|
||||||
@@ -550,11 +560,9 @@ class UnifiedCacheLinkerWrapper:
|
|||||||
self.cache.dec_lock_ref(node_id, lock_params)
|
self.cache.dec_lock_ref(node_id, lock_params)
|
||||||
self.pending_loads.clear()
|
self.pending_loads.clear()
|
||||||
for pending in self.pending_offloads:
|
for pending in self.pending_offloads:
|
||||||
for node_id in pending.publish_node_ids:
|
self.cache.tree_core.finish_external_linker_offload(
|
||||||
node = self.cache.resolve_node_handle(node_id)
|
pending.publish_node_ids, pending.lock_node_id, False
|
||||||
if node.write_through_pending_id == pending.lock_node_id:
|
)
|
||||||
node.write_through_pending_id = None
|
|
||||||
node.external_cache_stored = False
|
|
||||||
self.cache.dec_lock_ref(pending.lock_node_id, pending.lock_params)
|
self.cache.dec_lock_ref(pending.lock_node_id, pending.lock_params)
|
||||||
self.pending_offloads.clear()
|
self.pending_offloads.clear()
|
||||||
|
|
||||||
|
|||||||
@@ -60,6 +60,7 @@ from sglang.srt.mem_cache.unified_cache.components import (
|
|||||||
ComponentData,
|
ComponentData,
|
||||||
ComponentType,
|
ComponentType,
|
||||||
EvictLayer,
|
EvictLayer,
|
||||||
|
LinkerTransferPhase,
|
||||||
LRURefreshPhase,
|
LRURefreshPhase,
|
||||||
TreeComponent,
|
TreeComponent,
|
||||||
get_and_increase_time_counter,
|
get_and_increase_time_counter,
|
||||||
@@ -954,7 +955,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
|
|
||||||
if self.enable_external_cache_linker:
|
if self.enable_external_cache_linker:
|
||||||
return (
|
return (
|
||||||
not node.external_cache_stored
|
self._needs_external_linker_offload(node)
|
||||||
and node.hit_count >= self.write_through_threshold
|
and node.hit_count >= self.write_through_threshold
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -964,6 +965,11 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
and node.hit_count >= self.write_through_threshold
|
and node.hit_count >= self.write_through_threshold
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _needs_external_linker_offload(node: UnifiedTreeNode) -> bool:
|
||||||
|
"""Whether neither a confirmed nor an in-flight external copy exists."""
|
||||||
|
return not node.external_cache_stored and node.write_through_pending_id is None
|
||||||
|
|
||||||
def begin_insert(self, params: InsertParams) -> InsertStepResult:
|
def begin_insert(self, params: InsertParams) -> InsertStepResult:
|
||||||
"""Start the insert, running to its first barrier or completion."""
|
"""Start the insert, running to its first barrier or completion."""
|
||||||
# Insert walks are single-flight; a live walk means re-entrancy.
|
# Insert walks are single-flight; a live walk means re-entrancy.
|
||||||
@@ -2140,7 +2146,12 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
while (
|
while (
|
||||||
ancestor is not None
|
ancestor is not None
|
||||||
and ancestor is not self.root_node
|
and ancestor is not self.root_node
|
||||||
and not (ancestor.backuped or ancestor.external_cache_stored)
|
and not ancestor.backuped
|
||||||
|
and not ancestor.external_cache_stored
|
||||||
|
and (
|
||||||
|
not self.enable_external_cache_linker
|
||||||
|
or ancestor.write_through_pending_id is None
|
||||||
|
)
|
||||||
):
|
):
|
||||||
chain.append(ancestor)
|
chain.append(ancestor)
|
||||||
ancestor = ancestor.parent
|
ancestor = ancestor.parent
|
||||||
@@ -2279,6 +2290,69 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
|||||||
node = node.parent
|
node = node.parent
|
||||||
return depth
|
return depth
|
||||||
|
|
||||||
|
def build_external_linker_offload_transfers(
|
||||||
|
self, node_id: NodeId
|
||||||
|
) -> Optional[list[PoolTransfer]]:
|
||||||
|
"""Build transfers for a node with no stored or pending external copy."""
|
||||||
|
node = self.node_by_id(node_id)
|
||||||
|
if not self._needs_external_linker_offload(node):
|
||||||
|
return None
|
||||||
|
|
||||||
|
transfers = []
|
||||||
|
for component in self.components:
|
||||||
|
transfer = component.build_external_linker_transfer(
|
||||||
|
LinkerTransferPhase.OFFLOAD, node, None
|
||||||
|
)
|
||||||
|
if transfer is not None:
|
||||||
|
transfers.append(transfer)
|
||||||
|
return transfers
|
||||||
|
|
||||||
|
def mark_external_cache_stored_path(
|
||||||
|
self, from_node_id: NodeId, until_node_id: NodeId
|
||||||
|
) -> None:
|
||||||
|
"""Mark an externally restored path, excluding its existing anchor."""
|
||||||
|
until_node = self.node_by_id(until_node_id)
|
||||||
|
node = self.node_by_id(from_node_id)
|
||||||
|
path = []
|
||||||
|
while node is not until_node:
|
||||||
|
if node.parent is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"node {until_node_id} is not an ancestor of node {from_node_id}"
|
||||||
|
)
|
||||||
|
path.append(node)
|
||||||
|
node = node.parent
|
||||||
|
|
||||||
|
for node in path:
|
||||||
|
node.external_cache_stored = True
|
||||||
|
|
||||||
|
def mark_external_linker_offload_pending(self, node_id: NodeId) -> None:
|
||||||
|
"""Publish an accepted external offload as pending."""
|
||||||
|
node = self.node_by_id(node_id)
|
||||||
|
if not self._needs_external_linker_offload(node):
|
||||||
|
raise AssertionError(
|
||||||
|
f"invalid external offload state for node {node_id}: "
|
||||||
|
f"stored={node.external_cache_stored}, "
|
||||||
|
f"pending={node.write_through_pending_id}"
|
||||||
|
)
|
||||||
|
node.write_through_pending_id = node_id
|
||||||
|
|
||||||
|
def finish_external_linker_offload(
|
||||||
|
self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool
|
||||||
|
) -> None:
|
||||||
|
"""Finalize external-store state for an offload and its split fragments."""
|
||||||
|
nodes = [self.node_by_id(node_id) for node_id in node_ids]
|
||||||
|
for node_id, node in zip(node_ids, nodes):
|
||||||
|
if node.write_through_pending_id != ack_id:
|
||||||
|
raise AssertionError(
|
||||||
|
f"invalid external offload state for node {node_id}: "
|
||||||
|
f"expected pending={ack_id}; got "
|
||||||
|
f"stored={node.external_cache_stored}, "
|
||||||
|
f"pending={node.write_through_pending_id}"
|
||||||
|
)
|
||||||
|
for node in nodes:
|
||||||
|
node.write_through_pending_id = None
|
||||||
|
node.external_cache_stored |= success
|
||||||
|
|
||||||
def finish_write_through(self, node_ids: list[NodeId], ack_id: int) -> None:
|
def finish_write_through(self, node_ids: list[NodeId], ack_id: int) -> None:
|
||||||
"""Clear the write-through-pending mark (when it matches ack_id) and record the
|
"""Clear the write-through-pending mark (when it matches ack_id) and record the
|
||||||
host store event for each acked node."""
|
host store event for each acked node."""
|
||||||
|
|||||||
@@ -148,6 +148,7 @@ class UnifiedTreeCoreInterface(ABC):
|
|||||||
device: torch.device
|
device: torch.device
|
||||||
enable_hicache: bool
|
enable_hicache: bool
|
||||||
enable_storage: bool
|
enable_storage: bool
|
||||||
|
enable_external_cache_linker: bool
|
||||||
write_through_threshold: int
|
write_through_threshold: int
|
||||||
is_write_back: bool
|
is_write_back: bool
|
||||||
has_swa_host_pool: bool
|
has_swa_host_pool: bool
|
||||||
@@ -533,6 +534,41 @@ class UnifiedTreeCoreInterface(ABC):
|
|||||||
"""Clear the in-flight H->D marks on the anchor's root path at ack time."""
|
"""Clear the in-flight H->D marks on the anchor's root path at ack time."""
|
||||||
...
|
...
|
||||||
|
|
||||||
|
# ==== External Cache Linker ====
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def build_external_linker_offload_transfers(
|
||||||
|
self, node_id: NodeId
|
||||||
|
) -> Optional[list[PoolTransfer]]:
|
||||||
|
"""Build direct device-to-external-store transfers for an eligible node.
|
||||||
|
|
||||||
|
Return None when the node is stored externally or has an offload pending.
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def mark_external_cache_stored_path(
|
||||||
|
self, from_node_id: NodeId, until_node_id: NodeId
|
||||||
|
) -> None:
|
||||||
|
"""Mark the path from ``from_node_id`` to, but excluding, ``until_node_id``."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def mark_external_linker_offload_pending(self, node_id: NodeId) -> None:
|
||||||
|
"""Publish an accepted external offload as pending."""
|
||||||
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def finish_external_linker_offload(
|
||||||
|
self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool
|
||||||
|
) -> None:
|
||||||
|
"""Finish one external offload for every current fragment of its node.
|
||||||
|
|
||||||
|
A successful write confirms external storage. A failed redundant write
|
||||||
|
preserves storage already confirmed independently by a concurrent load.
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
# Order-sensitive digest of write_back duplicate-reclaim victim ids,
|
# Order-sensitive digest of write_back duplicate-reclaim victim ids,
|
||||||
# cross-checked across TP ranks; cores that never reclaim keep 0.
|
# cross-checked across TP ranks; cores that never reclaim keep 0.
|
||||||
write_back_duplicate_reclaim_digest: int = 0
|
write_back_duplicate_reclaim_digest: int = 0
|
||||||
|
|||||||
@@ -427,6 +427,25 @@ impl<K: ChildKeyType> TreeComponent<K> for FullComponent {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn build_external_linker_offload_transfer(
|
||||||
|
&self,
|
||||||
|
tree_core: &UnifiedTreeCore<K>,
|
||||||
|
node_id: NodeIdx_,
|
||||||
|
) -> Option<PoolTransfer> {
|
||||||
|
let node = tree_core.arena.node(node_id);
|
||||||
|
let keys = node
|
||||||
|
.hash_value
|
||||||
|
.as_ref()
|
||||||
|
.filter(|hashes| !hashes.is_empty())?;
|
||||||
|
let device_indices = node.try_device_value(FULL)?;
|
||||||
|
Some(PoolTransfer {
|
||||||
|
name: PoolName::Kv,
|
||||||
|
device_indices: Some(device_indices.to_kind(Kind::Int64)),
|
||||||
|
keys: Some(keys.clone()),
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
fn commit_hicache_transfer(
|
fn commit_hicache_transfer(
|
||||||
&self,
|
&self,
|
||||||
tree_core: &mut UnifiedTreeCore<K>,
|
tree_core: &mut UnifiedTreeCore<K>,
|
||||||
|
|||||||
@@ -392,6 +392,15 @@ pub trait TreeComponent<K: ChildKeyType> {
|
|||||||
unimplemented!("TreeComponent.build_hicache_transfers")
|
unimplemented!("TreeComponent.build_hicache_transfers")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Build this component's direct device-to-external-store transfer for a node.
|
||||||
|
fn build_external_linker_offload_transfer(
|
||||||
|
&self,
|
||||||
|
_tree_core: &UnifiedTreeCore<K>,
|
||||||
|
_node_id: NodeIdx_,
|
||||||
|
) -> Option<PoolTransfer> {
|
||||||
|
None
|
||||||
|
}
|
||||||
|
|
||||||
/// Post-transfer bookkeeping: store host indices, update LRU, etc.
|
/// Post-transfer bookkeeping: store host indices, update LRU, etc.
|
||||||
fn commit_hicache_transfer(
|
fn commit_hicache_transfer(
|
||||||
&self,
|
&self,
|
||||||
|
|||||||
@@ -951,6 +951,35 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn build_external_linker_offload_transfer(
|
||||||
|
&self,
|
||||||
|
tree_core: &UnifiedTreeCore<K>,
|
||||||
|
node_id: NodeIdx_,
|
||||||
|
) -> Option<PoolTransfer> {
|
||||||
|
let node = tree_core.arena.node(node_id);
|
||||||
|
let hashes = node
|
||||||
|
.hash_value
|
||||||
|
.as_ref()
|
||||||
|
.filter(|hashes| !hashes.is_empty())?;
|
||||||
|
let value = node.try_device_value(SWA)?;
|
||||||
|
let num_pages = value.size()[0] as usize / tree_core.page_size;
|
||||||
|
if num_pages == 0 {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
let num_tokens = num_pages * tree_core.page_size;
|
||||||
|
Some(PoolTransfer {
|
||||||
|
name: PoolName::Swa,
|
||||||
|
device_indices: Some(
|
||||||
|
value
|
||||||
|
.narrow(0, value.size()[0] - num_tokens as i64, num_tokens as i64)
|
||||||
|
.to_kind(Kind::Int64),
|
||||||
|
),
|
||||||
|
keys: Some(hashes[hashes.len().saturating_sub(num_pages)..].to_vec()),
|
||||||
|
hit_policy: PoolHitPolicy::TrailingPages,
|
||||||
|
..Default::default()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
fn commit_hicache_transfer(
|
fn commit_hicache_transfer(
|
||||||
&self,
|
&self,
|
||||||
tree_core: &mut UnifiedTreeCore<K>,
|
tree_core: &mut UnifiedTreeCore<K>,
|
||||||
|
|||||||
@@ -175,6 +175,8 @@ pub struct Node<K: ChildKeyType> {
|
|||||||
/// Per-page hash chain; None when the node was never hashed.
|
/// Per-page hash chain; None when the node was never hashed.
|
||||||
/// TODO: Store raw digests and hex-encode only at the Python or storage boundary.
|
/// TODO: Store raw digests and hex-encode only at the Python or storage boundary.
|
||||||
pub hash_value: Option<Vec<String>>,
|
pub hash_value: Option<Vec<String>>,
|
||||||
|
/// Whether this node is available through the direct external-cache linker.
|
||||||
|
pub external_cache_stored: bool,
|
||||||
/// The in-flight write-through backup's ack id.
|
/// The in-flight write-through backup's ack id.
|
||||||
pub write_through_pending_id: Option<usize>,
|
pub write_through_pending_id: Option<usize>,
|
||||||
/// Load-back anchor currently reading this node's host slots.
|
/// Load-back anchor currently reading this node's host slots.
|
||||||
@@ -396,6 +398,7 @@ impl<K: ChildKeyType> Node<K> {
|
|||||||
swa_uuid: None,
|
swa_uuid: None,
|
||||||
swa_host_uuid: None,
|
swa_host_uuid: None,
|
||||||
hash_value: Some(Vec::new()),
|
hash_value: Some(Vec::new()),
|
||||||
|
external_cache_stored: false,
|
||||||
write_through_pending_id: None,
|
write_through_pending_id: None,
|
||||||
load_back_pending_id: None,
|
load_back_pending_id: None,
|
||||||
last_access_counter: 0,
|
last_access_counter: 0,
|
||||||
@@ -418,6 +421,7 @@ impl<K: ChildKeyType> Node<K> {
|
|||||||
swa_uuid: None,
|
swa_uuid: None,
|
||||||
swa_host_uuid: None,
|
swa_host_uuid: None,
|
||||||
hash_value: None,
|
hash_value: None,
|
||||||
|
external_cache_stored: false,
|
||||||
write_through_pending_id: None,
|
write_through_pending_id: None,
|
||||||
load_back_pending_id: None,
|
load_back_pending_id: None,
|
||||||
last_access_counter: 0,
|
last_access_counter: 0,
|
||||||
@@ -753,6 +757,24 @@ pub enum TreeCoreRuntimeError {
|
|||||||
#[cfg(any(test, feature = "inspection"))]
|
#[cfg(any(test, feature = "inspection"))]
|
||||||
#[error("{0}")]
|
#[error("{0}")]
|
||||||
InspectionAssertion(String),
|
InspectionAssertion(String),
|
||||||
|
/// Direct external-cache linking does not support this tree component.
|
||||||
|
#[error("external cache linker does not support component {component_type:?}")]
|
||||||
|
ExternalCacheLinkerUnsupportedComponent { component_type: ComponentType },
|
||||||
|
/// The existing device anchor must be on the restored endpoint's root path.
|
||||||
|
#[error("node {until_node_id} is not an ancestor of node {from_node_id}")]
|
||||||
|
ExternalCachePathNotAncestor {
|
||||||
|
from_node_id: NodeId,
|
||||||
|
until_node_id: NodeId,
|
||||||
|
},
|
||||||
|
/// External offload lifecycle calls must observe valid state transitions.
|
||||||
|
#[error(
|
||||||
|
"invalid external offload state for node {node_id}: stored={stored}, pending={pending_id:?}"
|
||||||
|
)]
|
||||||
|
InvalidExternalCacheOffloadState {
|
||||||
|
node_id: NodeId,
|
||||||
|
stored: bool,
|
||||||
|
pending_id: Option<NodeId>,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
// Unigram and bigram child keys.
|
// Unigram and bigram child keys.
|
||||||
|
|||||||
@@ -1775,6 +1775,17 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
|
|||||||
py.allow_threads(|| self.core().enable_storage)
|
py.allow_threads(|| self.core().enable_storage)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Enable or disable the direct external-cache linker.
|
||||||
|
fn set_enable_external_cache_linker(&self, py: Python<'_>, value: bool) -> PyResult<()> {
|
||||||
|
py.allow_threads(|| self.core().set_enable_external_cache_linker(value))
|
||||||
|
.map_err(tree_core_assertion_error)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Whether the direct external-cache linker is wired.
|
||||||
|
fn enable_external_cache_linker(&self, py: Python<'_>) -> bool {
|
||||||
|
py.allow_threads(|| self.core().enable_external_cache_linker)
|
||||||
|
}
|
||||||
|
|
||||||
/// Queue the all-cleared placement event.
|
/// Queue the all-cleared placement event.
|
||||||
fn record_all_cleared_event(&self, py: Python<'_>) {
|
fn record_all_cleared_event(&self, py: Python<'_>) {
|
||||||
py.allow_threads(|| self.core().record_all_cleared_event());
|
py.allow_threads(|| self.core().record_all_cleared_event());
|
||||||
@@ -1872,6 +1883,64 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
|
|||||||
.map_err(node_access_error)
|
.map_err(node_access_error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Build transfers for a node with no stored or pending external copy.
|
||||||
|
fn build_external_linker_offload_transfers(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
node_id: NodeId,
|
||||||
|
) -> PyResult<Option<Vec<Py<PyAny>>>> {
|
||||||
|
let transfers = py
|
||||||
|
.allow_threads(|| self.core().build_external_linker_offload_transfers(node_id))
|
||||||
|
.map_err(node_access_error)?;
|
||||||
|
transfers
|
||||||
|
.map(|transfers| {
|
||||||
|
transfers
|
||||||
|
.into_iter()
|
||||||
|
.map(|transfer| transfer_to_py(py, transfer))
|
||||||
|
.collect::<PyResult<Vec<_>>>()
|
||||||
|
})
|
||||||
|
.transpose()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Mark an externally restored path, excluding its existing anchor.
|
||||||
|
fn mark_external_cache_stored_path(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
from_node_id: NodeId,
|
||||||
|
until_node_id: NodeId,
|
||||||
|
) -> PyResult<()> {
|
||||||
|
py.allow_threads(|| {
|
||||||
|
self.core()
|
||||||
|
.mark_external_cache_stored_path(from_node_id, until_node_id)
|
||||||
|
})
|
||||||
|
.map_err(tree_core_runtime_error)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Publish an accepted external offload as pending.
|
||||||
|
fn mark_external_linker_offload_pending(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
node_id: NodeId,
|
||||||
|
) -> PyResult<()> {
|
||||||
|
py.allow_threads(|| self.core().mark_external_linker_offload_pending(node_id))
|
||||||
|
.map_err(tree_core_assertion_error)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Finalize external-store state for an offload and its split fragments.
|
||||||
|
fn finish_external_linker_offload(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
node_ids: Vec<NodeId>,
|
||||||
|
ack_id: NodeId,
|
||||||
|
success: bool,
|
||||||
|
) -> PyResult<()> {
|
||||||
|
py.allow_threads(|| {
|
||||||
|
self.core()
|
||||||
|
.finish_external_linker_offload(&node_ids, ack_id, success)
|
||||||
|
})
|
||||||
|
.map_err(tree_core_assertion_error)
|
||||||
|
}
|
||||||
|
|
||||||
/// Order-sensitive digest of reclaimed coexisting host values.
|
/// Order-sensitive digest of reclaimed coexisting host values.
|
||||||
fn write_back_coexist_reclaim_digest(&self, py: Python<'_>) -> i64 {
|
fn write_back_coexist_reclaim_digest(&self, py: Python<'_>) -> i64 {
|
||||||
py.allow_threads(|| self.core().write_back_coexist_reclaim_digest)
|
py.allow_threads(|| self.core().write_back_coexist_reclaim_digest)
|
||||||
@@ -1968,6 +2037,11 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
|
|||||||
.map_err(node_access_error)
|
.map_err(node_access_error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn inspect_is_external_cache_stored(&self, py: Python<'_>, node_id: NodeId) -> PyResult<bool> {
|
||||||
|
py.allow_threads(|| self.core().inspect_is_external_cache_stored(node_id))
|
||||||
|
.map_err(node_access_error)
|
||||||
|
}
|
||||||
|
|
||||||
fn inspect_is_node_in_device_lru(
|
fn inspect_is_node_in_device_lru(
|
||||||
&self,
|
&self,
|
||||||
py: Python<'_>,
|
py: Python<'_>,
|
||||||
@@ -2839,6 +2913,20 @@ macro_rules! tree_core_binding {
|
|||||||
self.inner.enable_storage(py)
|
self.inner.enable_storage(py)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Enable or disable the direct external-cache linker.
|
||||||
|
fn set_enable_external_cache_linker(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
value: bool,
|
||||||
|
) -> PyResult<()> {
|
||||||
|
self.inner.set_enable_external_cache_linker(py, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Whether the direct external-cache linker is wired.
|
||||||
|
fn enable_external_cache_linker(&self, py: Python<'_>) -> bool {
|
||||||
|
self.inner.enable_external_cache_linker(py)
|
||||||
|
}
|
||||||
|
|
||||||
/// Queue the all-cleared placement event.
|
/// Queue the all-cleared placement event.
|
||||||
fn record_all_cleared_event(&self, py: Python<'_>) {
|
fn record_all_cleared_event(&self, py: Python<'_>) {
|
||||||
self.inner.record_all_cleared_event(py)
|
self.inner.record_all_cleared_event(py)
|
||||||
@@ -2884,6 +2972,53 @@ macro_rules! tree_core_binding {
|
|||||||
self.inner.finish_load_back(py, anchor_node_id)
|
self.inner.finish_load_back(py, anchor_node_id)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Build transfers for a node with no stored or pending external copy.
|
||||||
|
fn build_external_linker_offload_transfers(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
node_id: NodeId,
|
||||||
|
) -> PyResult<Option<Vec<Py<PyAny>>>> {
|
||||||
|
self.inner
|
||||||
|
.build_external_linker_offload_transfers(py, node_id)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Mark an externally restored path, excluding its existing anchor.
|
||||||
|
fn mark_external_cache_stored_path(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
from_node_id: NodeId,
|
||||||
|
until_node_id: NodeId,
|
||||||
|
) -> PyResult<()> {
|
||||||
|
self.inner.mark_external_cache_stored_path(
|
||||||
|
py,
|
||||||
|
from_node_id,
|
||||||
|
until_node_id,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Publish an accepted external offload as pending.
|
||||||
|
fn mark_external_linker_offload_pending(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
node_id: NodeId,
|
||||||
|
) -> PyResult<()> {
|
||||||
|
self.inner
|
||||||
|
.mark_external_linker_offload_pending(py, node_id)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Finalize external-store state for an offload and its split fragments.
|
||||||
|
fn finish_external_linker_offload(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
node_ids: Vec<NodeId>,
|
||||||
|
ack_id: NodeId,
|
||||||
|
success: bool,
|
||||||
|
) -> PyResult<()> {
|
||||||
|
self.inner.finish_external_linker_offload(
|
||||||
|
py, node_ids, ack_id, success,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
/// Order-sensitive digest of reclaimed coexisting host values.
|
/// Order-sensitive digest of reclaimed coexisting host values.
|
||||||
#[pyo3(name = "write_back_duplicate_reclaim_digest")]
|
#[pyo3(name = "write_back_duplicate_reclaim_digest")]
|
||||||
fn write_back_coexist_reclaim_digest(&self, py: Python<'_>) -> i64 {
|
fn write_back_coexist_reclaim_digest(&self, py: Python<'_>) -> i64 {
|
||||||
@@ -2994,6 +3129,15 @@ macro_rules! tree_core_binding {
|
|||||||
.inspect_get_write_through_pending_id(py, node_id)
|
.inspect_get_write_through_pending_id(py, node_id)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(feature = "inspection")]
|
||||||
|
fn inspect_is_external_cache_stored(
|
||||||
|
&self,
|
||||||
|
py: Python<'_>,
|
||||||
|
node_id: NodeId,
|
||||||
|
) -> PyResult<bool> {
|
||||||
|
self.inner.inspect_is_external_cache_stored(py, node_id)
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(feature = "inspection")]
|
#[cfg(feature = "inspection")]
|
||||||
fn inspect_is_node_in_device_lru(
|
fn inspect_is_node_in_device_lru(
|
||||||
&self,
|
&self,
|
||||||
|
|||||||
@@ -2201,6 +2201,248 @@ fn insert_threshold_crossing_emits_the_backup_kv_action() {
|
|||||||
assert_eq!(backups, vec![vec![leaf]]);
|
assert_eq!(backups, vec![vec![leaf]]);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn external_linker_hashes_new_nodes_and_triggers_offload_action() {
|
||||||
|
let params = CacheInitParams {
|
||||||
|
page_size: 2,
|
||||||
|
write_through_threshold: 1,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let mut tc: UnifiedTreeCore<Vec<i64>> = UnifiedTreeCore::new(params, vec![FULL]);
|
||||||
|
tc.set_enable_external_cache_linker(true).unwrap();
|
||||||
|
|
||||||
|
let result = tc.insert(&insert_params(&vec![1, 2], &[10, 11]));
|
||||||
|
let leaf = result.last_device_node_id.unwrap();
|
||||||
|
assert!(result.cache_actions.iter().any(|action| {
|
||||||
|
matches!(action, CacheAction::BackupKV(backup) if backup.node_ids == vec![leaf])
|
||||||
|
}));
|
||||||
|
|
||||||
|
let transfers = tc
|
||||||
|
.build_external_linker_offload_transfers(leaf)
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(transfers.len(), 1);
|
||||||
|
assert_eq!(transfers[0].name, PoolName::Kv);
|
||||||
|
assert_eq!(transfers[0].hit_policy, PoolHitPolicy::AllPages);
|
||||||
|
assert!(
|
||||||
|
transfers[0]
|
||||||
|
.device_indices
|
||||||
|
.as_ref()
|
||||||
|
.unwrap()
|
||||||
|
.equal(&Tensor::from_slice(&[10i64, 11]))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
transfers[0].keys,
|
||||||
|
tc.arena
|
||||||
|
.node(tc.arena.resolve(leaf).expect("live test node"))
|
||||||
|
.hash_value
|
||||||
|
.clone()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn external_linker_swa_offload_uses_complete_trailing_pages() {
|
||||||
|
let params = CacheInitParams {
|
||||||
|
page_size: 2,
|
||||||
|
swa_sliding_window_size: Some(4),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let mut tc: UnifiedTreeCore<Vec<i64>> = UnifiedTreeCore::new(params, vec![FULL, SWA]);
|
||||||
|
tc.set_enable_external_cache_linker(true).unwrap();
|
||||||
|
let root = tc.arena.root();
|
||||||
|
let node = tc.add_new_node_(
|
||||||
|
root,
|
||||||
|
vec![1, 2, 3, 4, 5, 6],
|
||||||
|
&Tensor::from_slice(&[10i64, 11, 12, 13, 14, 15]),
|
||||||
|
0,
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
tc.arena.node_mut(node).values[SWA.idx()].value =
|
||||||
|
Some(Tensor::from_slice(&[20i64, 21, 22, 23, 24]));
|
||||||
|
|
||||||
|
let transfers = tc
|
||||||
|
.build_external_linker_offload_transfers(tc.arena.node(node).id)
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
assert_eq!(transfers.len(), 2);
|
||||||
|
assert_eq!(transfers[0].name, PoolName::Kv);
|
||||||
|
assert_eq!(transfers[1].name, PoolName::Swa);
|
||||||
|
assert_eq!(transfers[1].hit_policy, PoolHitPolicy::TrailingPages);
|
||||||
|
assert!(
|
||||||
|
transfers[1]
|
||||||
|
.device_indices
|
||||||
|
.as_ref()
|
||||||
|
.unwrap()
|
||||||
|
.equal(&Tensor::from_slice(&[21i64, 22, 23, 24]))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
transfers[1].keys.as_ref().unwrap(),
|
||||||
|
&tc.arena.node(node).hash_value.as_ref().unwrap()[1..]
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn external_linker_rejects_mamba_trees() {
|
||||||
|
let params = CacheInitParams {
|
||||||
|
mamba_cache_chunk_size: Some(1),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let mut tc: UnifiedTreeCore<Vec<i64>> = UnifiedTreeCore::new(params, vec![FULL, MAMBA]);
|
||||||
|
let error = tc.set_enable_external_cache_linker(true).unwrap_err();
|
||||||
|
assert!(matches!(
|
||||||
|
&error,
|
||||||
|
TreeCoreRuntimeError::ExternalCacheLinkerUnsupportedComponent {
|
||||||
|
component_type
|
||||||
|
} if *component_type == MAMBA
|
||||||
|
));
|
||||||
|
assert!(error.to_string().contains("Mamba"));
|
||||||
|
assert!(!tc.enable_external_cache_linker);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn external_linker_state_follows_load_offload_and_split_lifecycle() {
|
||||||
|
let mut tc = core();
|
||||||
|
tc.set_enable_external_cache_linker(true).unwrap();
|
||||||
|
tc.insert(&insert_params(&vec![1, 2], &[10, 11]));
|
||||||
|
tc.insert(&insert_params(&vec![1, 2, 3, 4], &[10, 11, 12, 13]));
|
||||||
|
let anchor = tc
|
||||||
|
.match_prefix(&match_params(&vec![1, 2]))
|
||||||
|
.best_match_node_id;
|
||||||
|
let leaf = tc
|
||||||
|
.match_prefix(&match_params(&vec![1, 2, 3, 4]))
|
||||||
|
.best_match_node_id;
|
||||||
|
|
||||||
|
tc.mark_external_cache_stored_path(leaf, anchor).unwrap();
|
||||||
|
assert!(
|
||||||
|
tc.arena
|
||||||
|
.node(tc.arena.resolve(leaf).expect("live test node"))
|
||||||
|
.external_cache_stored
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
!tc.arena
|
||||||
|
.node(tc.arena.resolve(anchor).expect("live test node"))
|
||||||
|
.external_cache_stored
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
tc.build_external_linker_offload_transfers(leaf)
|
||||||
|
.unwrap()
|
||||||
|
.is_none()
|
||||||
|
);
|
||||||
|
assert!(matches!(
|
||||||
|
tc.mark_external_linker_offload_pending(leaf),
|
||||||
|
Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState {
|
||||||
|
node_id,
|
||||||
|
stored: true,
|
||||||
|
pending_id: None,
|
||||||
|
}) if node_id == leaf
|
||||||
|
));
|
||||||
|
|
||||||
|
let leaf_idx = tc.arena.resolve(leaf).expect("live test node");
|
||||||
|
tc.arena.node_mut(leaf_idx).external_cache_stored = false;
|
||||||
|
tc.mark_external_linker_offload_pending(leaf).unwrap();
|
||||||
|
assert!(
|
||||||
|
tc.build_external_linker_offload_transfers(leaf)
|
||||||
|
.unwrap()
|
||||||
|
.is_none()
|
||||||
|
);
|
||||||
|
assert!(matches!(
|
||||||
|
tc.mark_external_linker_offload_pending(leaf),
|
||||||
|
Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState {
|
||||||
|
node_id,
|
||||||
|
stored: false,
|
||||||
|
pending_id: Some(pending_id),
|
||||||
|
}) if node_id == leaf && pending_id == leaf
|
||||||
|
));
|
||||||
|
let (new_parent, action) = tc.split_node_(leaf_idx, 1);
|
||||||
|
let new_parent_handle = tc.arena.node(new_parent).id;
|
||||||
|
assert!(action.is_some());
|
||||||
|
assert!(!tc.arena.node(new_parent).external_cache_stored);
|
||||||
|
assert!(!tc.arena.node(leaf_idx).external_cache_stored);
|
||||||
|
|
||||||
|
tc.insert(&insert_params(&vec![9], &[19]));
|
||||||
|
let independent = tc.match_prefix(&match_params(&vec![9])).best_match_node_id;
|
||||||
|
tc.mark_external_linker_offload_pending(independent)
|
||||||
|
.unwrap();
|
||||||
|
assert!(matches!(
|
||||||
|
tc.finish_external_linker_offload(&[independent, new_parent_handle], independent, false),
|
||||||
|
Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState {
|
||||||
|
node_id,
|
||||||
|
stored: false,
|
||||||
|
pending_id: Some(pending_id),
|
||||||
|
}) if node_id == new_parent_handle && pending_id == leaf
|
||||||
|
));
|
||||||
|
assert_eq!(
|
||||||
|
tc.arena
|
||||||
|
.node(tc.arena.resolve(independent).expect("live test node"))
|
||||||
|
.write_through_pending_id,
|
||||||
|
Some(independent)
|
||||||
|
);
|
||||||
|
for node_id in [new_parent_handle, leaf] {
|
||||||
|
let node = tc
|
||||||
|
.arena
|
||||||
|
.node(tc.arena.resolve(node_id).expect("live test node"));
|
||||||
|
assert_eq!(node.write_through_pending_id, Some(leaf));
|
||||||
|
assert!(!node.external_cache_stored);
|
||||||
|
}
|
||||||
|
|
||||||
|
tc.finish_external_linker_offload(&[independent], independent, false)
|
||||||
|
.unwrap();
|
||||||
|
tc.finish_external_linker_offload(&[new_parent_handle, leaf], leaf, false)
|
||||||
|
.unwrap();
|
||||||
|
for node_id in [new_parent_handle, leaf] {
|
||||||
|
let node = tc
|
||||||
|
.arena
|
||||||
|
.node(tc.arena.resolve(node_id).expect("live test node"));
|
||||||
|
assert_eq!(node.write_through_pending_id, None);
|
||||||
|
assert!(!node.external_cache_stored);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn failed_external_offload_preserves_independently_confirmed_state() {
|
||||||
|
let mut tc = core();
|
||||||
|
tc.set_enable_external_cache_linker(true).unwrap();
|
||||||
|
tc.insert(&insert_params(&vec![1], &[10]));
|
||||||
|
tc.insert(&insert_params(&vec![1, 2], &[10, 11]));
|
||||||
|
let anchor = tc.match_prefix(&match_params(&vec![1])).best_match_node_id;
|
||||||
|
let leaf = tc
|
||||||
|
.match_prefix(&match_params(&vec![1, 2]))
|
||||||
|
.best_match_node_id;
|
||||||
|
|
||||||
|
tc.mark_external_linker_offload_pending(leaf).unwrap();
|
||||||
|
tc.mark_external_cache_stored_path(leaf, anchor).unwrap();
|
||||||
|
tc.finish_external_linker_offload(&[leaf], leaf, false)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let leaf = tc
|
||||||
|
.arena
|
||||||
|
.node(tc.arena.resolve(leaf).expect("live test node"));
|
||||||
|
assert_eq!(leaf.write_through_pending_id, None);
|
||||||
|
assert!(leaf.external_cache_stored);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn external_linker_path_validation_is_atomic() {
|
||||||
|
let mut tc = core();
|
||||||
|
tc.insert(&insert_params(&vec![1], &[10]));
|
||||||
|
tc.insert(&insert_params(&vec![2], &[20]));
|
||||||
|
let left = tc.match_prefix(&match_params(&vec![1])).best_match_node_id;
|
||||||
|
let right = tc.match_prefix(&match_params(&vec![2])).best_match_node_id;
|
||||||
|
|
||||||
|
assert!(matches!(
|
||||||
|
tc.mark_external_cache_stored_path(left, right),
|
||||||
|
Err(TreeCoreRuntimeError::ExternalCachePathNotAncestor {
|
||||||
|
from_node_id,
|
||||||
|
until_node_id,
|
||||||
|
}) if from_node_id == left && until_node_id == right
|
||||||
|
));
|
||||||
|
assert!(
|
||||||
|
!tc.arena
|
||||||
|
.node(tc.arena.resolve(left).expect("live test node"))
|
||||||
|
.external_cache_stored
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn mark_write_through_pending_stamps_the_supplied_ack() {
|
fn mark_write_through_pending_stamps_the_supplied_ack() {
|
||||||
let mut tc = core();
|
let mut tc = core();
|
||||||
@@ -2325,6 +2567,38 @@ fn backup_kv_action_chains_unbacked_ancestors_first() {
|
|||||||
assert_eq!(action.node_ids, vec![c]);
|
assert_eq!(action.node_ids, vec![c]);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn backup_kv_action_stops_at_an_externally_stored_or_pending_ancestor() {
|
||||||
|
let mut tc = core();
|
||||||
|
tc.set_enable_external_cache_linker(true).unwrap();
|
||||||
|
tc.insert(&insert_params(&vec![1], &[10]));
|
||||||
|
tc.insert(&insert_params(&vec![1, 2], &[10, 11]));
|
||||||
|
tc.insert(&insert_params(&vec![1, 2, 3], &[10, 11, 12]));
|
||||||
|
let a = tc.match_prefix(&match_params(&vec![1])).best_match_node_id;
|
||||||
|
let b = tc
|
||||||
|
.match_prefix(&match_params(&vec![1, 2]))
|
||||||
|
.best_match_node_id;
|
||||||
|
let c = tc
|
||||||
|
.match_prefix(&match_params(&vec![1, 2, 3]))
|
||||||
|
.best_match_node_id;
|
||||||
|
let a_idx = tc.arena.resolve(a).expect("live test node");
|
||||||
|
tc.arena.node_mut(a_idx).external_cache_stored = true;
|
||||||
|
|
||||||
|
let action = tc.build_backup_kv_action_(
|
||||||
|
tc.arena.node(tc.arena.resolve(c).expect("live test node")),
|
||||||
|
/* write_back = */ false,
|
||||||
|
);
|
||||||
|
assert_eq!(action.node_ids, vec![b, c]);
|
||||||
|
|
||||||
|
tc.arena.node_mut(a_idx).external_cache_stored = false;
|
||||||
|
tc.mark_external_linker_offload_pending(a).unwrap();
|
||||||
|
let action = tc.build_backup_kv_action_(
|
||||||
|
tc.arena.node(tc.arena.resolve(c).expect("live test node")),
|
||||||
|
/* write_back = */ false,
|
||||||
|
);
|
||||||
|
assert_eq!(action.node_ids, vec![b, c]);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn split_of_a_pending_node_transfers_the_ack_and_emits_the_replace_action() {
|
fn split_of_a_pending_node_transfers_the_ack_and_emits_the_replace_action() {
|
||||||
let mut tc = core();
|
let mut tc = core();
|
||||||
@@ -3958,6 +4232,35 @@ fn fallible_node_boundaries_reject_stale_handles() {
|
|||||||
Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id }))
|
Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id }))
|
||||||
if node_id == stale_root
|
if node_id == stale_root
|
||||||
));
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
tc.build_external_linker_offload_transfers(stale_root),
|
||||||
|
Err(NodeAccessError { node_id }) if node_id == stale_root
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
tc.mark_external_cache_stored_path(stale_root, live_root),
|
||||||
|
Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id }))
|
||||||
|
if node_id == stale_root
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
tc.mark_external_cache_stored_path(live_root, stale_root),
|
||||||
|
Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id }))
|
||||||
|
if node_id == stale_root
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
tc.mark_external_linker_offload_pending(stale_root),
|
||||||
|
Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id }))
|
||||||
|
if node_id == stale_root
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
tc.finish_external_linker_offload(&[live_root, stale_root], live_root, true),
|
||||||
|
Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id }))
|
||||||
|
if node_id == stale_root
|
||||||
|
));
|
||||||
|
assert!(
|
||||||
|
!tc.arena
|
||||||
|
.node(tc.arena.resolve(live_root).expect("live root"))
|
||||||
|
.external_cache_stored
|
||||||
|
);
|
||||||
assert!(matches!(
|
assert!(matches!(
|
||||||
tc.get_hash_values(stale_root),
|
tc.get_hash_values(stale_root),
|
||||||
Err(NodeAccessError { node_id }) if node_id == stale_root
|
Err(NodeAccessError { node_id }) if node_id == stale_root
|
||||||
@@ -7964,6 +8267,10 @@ fn inspection_rejects_stale_handles_without_panicking() {
|
|||||||
assert_eq!(tc.inspect_get_parent_node_id(stale_root), Err(expected));
|
assert_eq!(tc.inspect_get_parent_node_id(stale_root), Err(expected));
|
||||||
assert_eq!(tc.inspect_get_child_node_ids(stale_root), Err(expected));
|
assert_eq!(tc.inspect_get_child_node_ids(stale_root), Err(expected));
|
||||||
assert_eq!(tc.inspect_get_node_key_length(stale_root), Err(expected));
|
assert_eq!(tc.inspect_get_node_key_length(stale_root), Err(expected));
|
||||||
|
assert_eq!(
|
||||||
|
tc.inspect_is_external_cache_stored(stale_root),
|
||||||
|
Err(expected)
|
||||||
|
);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
tc.inspect_set_node_hash_values(stale_root, None),
|
tc.inspect_set_node_hash_values(stale_root, None),
|
||||||
Err(expected)
|
Err(expected)
|
||||||
@@ -7993,6 +8300,7 @@ fn inspection_rejects_stale_handles_without_panicking() {
|
|||||||
assert!(!tc.inspect_is_host_evictable_leaf(stale_root));
|
assert!(!tc.inspect_is_host_evictable_leaf(stale_root));
|
||||||
|
|
||||||
assert_eq!(tc.inspect_get_parent_node_id(live_root), Ok(None));
|
assert_eq!(tc.inspect_get_parent_node_id(live_root), Ok(None));
|
||||||
|
assert_eq!(tc.inspect_is_external_cache_stored(live_root), Ok(false));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -534,6 +534,8 @@ pub struct UnifiedTreeCore<K: ChildKeyType> {
|
|||||||
pub(crate) enable_hicache: bool,
|
pub(crate) enable_hicache: bool,
|
||||||
/// Whether the storage tier (L3) is wired; gates page-hash computation.
|
/// Whether the storage tier (L3) is wired; gates page-hash computation.
|
||||||
pub(crate) enable_storage: bool,
|
pub(crate) enable_storage: bool,
|
||||||
|
/// Whether a direct device-to-external-cache linker is wired.
|
||||||
|
pub(crate) enable_external_cache_linker: bool,
|
||||||
/// Whether the cache wired a host SWA pool (HiCache).
|
/// Whether the cache wired a host SWA pool (HiCache).
|
||||||
pub(crate) has_swa_host_pool: bool,
|
pub(crate) has_swa_host_pool: bool,
|
||||||
/// Whether tree mutations emit BlockStored/BlockRemoved events.
|
/// Whether tree mutations emit BlockStored/BlockRemoved events.
|
||||||
@@ -717,6 +719,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
is_write_back: params.is_write_back,
|
is_write_back: params.is_write_back,
|
||||||
enable_hicache: params.enable_hicache,
|
enable_hicache: params.enable_hicache,
|
||||||
enable_storage: false,
|
enable_storage: false,
|
||||||
|
enable_external_cache_linker: false,
|
||||||
has_swa_host_pool: params.has_swa_host_pool,
|
has_swa_host_pool: params.has_swa_host_pool,
|
||||||
enable_kv_cache_events: params.enable_kv_cache_events,
|
enable_kv_cache_events: params.enable_kv_cache_events,
|
||||||
kv_event_queue: Vec::new(),
|
kv_event_queue: Vec::new(),
|
||||||
@@ -1319,6 +1322,12 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
node.hit_count += 1;
|
node.hit_count += 1;
|
||||||
|
|
||||||
|
if self.enable_external_cache_linker {
|
||||||
|
return Self::needs_external_linker_offload_(node)
|
||||||
|
&& node.hit_count >= self.write_through_threshold;
|
||||||
|
}
|
||||||
|
|
||||||
self.enable_hicache && !node.backuped() && node.hit_count >= self.write_through_threshold
|
self.enable_hicache && !node.backuped() && node.hit_count >= self.write_through_threshold
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1763,6 +1772,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
let child = self.arena.node(child_id);
|
let child = self.arena.node(child_id);
|
||||||
let parent_id = child.parent();
|
let parent_id = child.parent();
|
||||||
let child_namespace = child.namespace.clone();
|
let child_namespace = child.namespace.clone();
|
||||||
|
let child_external_cache_stored = child.external_cache_stored;
|
||||||
let (key_head, key_tail) = child.key.split_at(split_len);
|
let (key_head, key_tail) = child.key.split_at(split_len);
|
||||||
// key_head keeps the original key's first page, which keys the parent's child map.
|
// key_head keeps the original key's first page, which keys the parent's child map.
|
||||||
let parent_map_key = key_head.child_key(page_size);
|
let parent_map_key = key_head.child_key(page_size);
|
||||||
@@ -1778,6 +1788,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
(child_namespace.clone(), key_tail.child_key(page_size)),
|
(child_namespace.clone(), key_tail.child_key(page_size)),
|
||||||
child_id,
|
child_id,
|
||||||
);
|
);
|
||||||
|
self.arena.node_mut(new_node_id).external_cache_stored = child_external_cache_stored;
|
||||||
|
|
||||||
// The child's aux LRU cells detach while it is re-linked.
|
// The child's aux LRU cells detach while it is re-linked.
|
||||||
self.for_each_component_lru_(
|
self.for_each_component_lru_(
|
||||||
@@ -1897,7 +1908,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
"add_new_node_: parent {parent_id} already has a child on the new node's page"
|
"add_new_node_: parent {parent_id} already has a child on the new node's page"
|
||||||
);
|
);
|
||||||
self.inc_evictable_size(FULL, value.size()[0] as usize);
|
self.inc_evictable_size(FULL, value.size()[0] as usize);
|
||||||
if self.enable_storage {
|
if self.enable_storage || self.enable_external_cache_linker {
|
||||||
let hash_values = self.arena.compute_node_hash_values(new_node_id, page_size);
|
let hash_values = self.arena.compute_node_hash_values(new_node_id, page_size);
|
||||||
self.arena.node_mut(new_node_id).hash_value = Some(hash_values);
|
self.arena.node_mut(new_node_id).hash_value = Some(hash_values);
|
||||||
}
|
}
|
||||||
@@ -2708,6 +2719,22 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
self.enable_storage = value;
|
self.enable_storage = value;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Enable or disable the direct external-cache linker.
|
||||||
|
pub fn set_enable_external_cache_linker(
|
||||||
|
&mut self,
|
||||||
|
value: bool,
|
||||||
|
) -> Result<(), TreeCoreRuntimeError> {
|
||||||
|
if value && self.components_by_type[MAMBA.idx()].is_some() {
|
||||||
|
return Err(
|
||||||
|
TreeCoreRuntimeError::ExternalCacheLinkerUnsupportedComponent {
|
||||||
|
component_type: MAMBA,
|
||||||
|
},
|
||||||
|
);
|
||||||
|
}
|
||||||
|
self.enable_external_cache_linker = value;
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
// ==== KV cache placement events ====
|
// ==== KV cache placement events ====
|
||||||
|
|
||||||
/// Append an event, coalescing it with a compatible queue tail.
|
/// Append an event, coalescing it with a compatible queue tail.
|
||||||
@@ -3463,14 +3490,19 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
Ok(order)
|
Ok(order)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Build the backup action for a node and its unbacked ancestors.
|
/// Build the backup action for a node and its not-yet-persisted ancestors.
|
||||||
pub fn build_backup_kv_action_(&self, node: &Node<K>, write_back: bool) -> BackupKV {
|
pub fn build_backup_kv_action_(&self, node: &Node<K>, write_back: bool) -> BackupKV {
|
||||||
let mut chain = vec![node.id];
|
let mut chain = vec![node.id];
|
||||||
if !write_back {
|
if !write_back {
|
||||||
let mut ancestor = node.try_parent();
|
let mut ancestor = node.try_parent();
|
||||||
while let Some(ancestor_idx) = ancestor {
|
while let Some(ancestor_idx) = ancestor {
|
||||||
let ancestor_node = self.arena.node(ancestor_idx);
|
let ancestor_node = self.arena.node(ancestor_idx);
|
||||||
if ancestor_node.is_root() || ancestor_node.backuped() {
|
if ancestor_node.is_root()
|
||||||
|
|| ancestor_node.backuped()
|
||||||
|
|| ancestor_node.external_cache_stored
|
||||||
|
|| (self.enable_external_cache_linker
|
||||||
|
&& ancestor_node.write_through_pending_id.is_some())
|
||||||
|
{
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
chain.push(ancestor_node.id);
|
chain.push(ancestor_node.id);
|
||||||
@@ -3696,6 +3728,103 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
depth
|
depth
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Build transfers for a node with no stored or pending external copy.
|
||||||
|
pub fn build_external_linker_offload_transfers(
|
||||||
|
&self,
|
||||||
|
node_id: NodeId,
|
||||||
|
) -> Result<Option<Vec<PoolTransfer>>, NodeAccessError> {
|
||||||
|
let node_id = self.arena.resolve(node_id)?;
|
||||||
|
if !Self::needs_external_linker_offload_(self.arena.node(node_id)) {
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
|
|
||||||
|
let transfers = self
|
||||||
|
.components
|
||||||
|
.iter()
|
||||||
|
.filter_map(|component| component.build_external_linker_offload_transfer(self, node_id))
|
||||||
|
.collect();
|
||||||
|
Ok(Some(transfers))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn needs_external_linker_offload_(node: &Node<K>) -> bool {
|
||||||
|
!node.external_cache_stored && node.write_through_pending_id.is_none()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Mark the path from `from_node_id` to, but excluding, `until_node_id` as
|
||||||
|
/// available in the external cache.
|
||||||
|
pub fn mark_external_cache_stored_path(
|
||||||
|
&mut self,
|
||||||
|
from_node_id: NodeId,
|
||||||
|
until_node_id: NodeId,
|
||||||
|
) -> Result<(), TreeCoreRuntimeError> {
|
||||||
|
let from = self.arena.resolve(from_node_id)?;
|
||||||
|
let until = self.arena.resolve(until_node_id)?;
|
||||||
|
let mut path = Vec::new();
|
||||||
|
let mut current = from;
|
||||||
|
while current != until {
|
||||||
|
let node = self.arena.node(current);
|
||||||
|
let Some(parent) = node.try_parent() else {
|
||||||
|
return Err(TreeCoreRuntimeError::ExternalCachePathNotAncestor {
|
||||||
|
from_node_id,
|
||||||
|
until_node_id,
|
||||||
|
});
|
||||||
|
};
|
||||||
|
path.push(current);
|
||||||
|
current = parent;
|
||||||
|
}
|
||||||
|
for node_id in path {
|
||||||
|
self.arena.node_mut(node_id).external_cache_stored = true;
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Publish an accepted external offload as pending.
|
||||||
|
pub fn mark_external_linker_offload_pending(
|
||||||
|
&mut self,
|
||||||
|
node_id: NodeId,
|
||||||
|
) -> Result<(), TreeCoreRuntimeError> {
|
||||||
|
let node_idx = self.arena.resolve(node_id)?;
|
||||||
|
let node = self.arena.node(node_idx);
|
||||||
|
if !Self::needs_external_linker_offload_(node) {
|
||||||
|
return Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState {
|
||||||
|
node_id,
|
||||||
|
stored: node.external_cache_stored,
|
||||||
|
pending_id: node.write_through_pending_id,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
self.arena.node_mut(node_idx).write_through_pending_id = Some(node_id);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Finalize external-store state for an offload and its split fragments.
|
||||||
|
pub fn finish_external_linker_offload(
|
||||||
|
&mut self,
|
||||||
|
node_ids: &[NodeId],
|
||||||
|
ack_id: NodeId,
|
||||||
|
success: bool,
|
||||||
|
) -> Result<(), TreeCoreRuntimeError> {
|
||||||
|
let node_indices = node_ids
|
||||||
|
.iter()
|
||||||
|
.map(|&node_id| self.arena.resolve(node_id))
|
||||||
|
.collect::<Result<Vec<_>, _>>()?;
|
||||||
|
for (&node_id, &node_idx) in node_ids.iter().zip(&node_indices) {
|
||||||
|
let node = self.arena.node(node_idx);
|
||||||
|
if node.write_through_pending_id != Some(ack_id) {
|
||||||
|
return Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState {
|
||||||
|
node_id,
|
||||||
|
stored: node.external_cache_stored,
|
||||||
|
pending_id: node.write_through_pending_id,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for node_id in node_indices {
|
||||||
|
let node = self.arena.node_mut(node_id);
|
||||||
|
node.write_through_pending_id = None;
|
||||||
|
node.external_cache_stored |= success;
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
/// Clear the write-through-pending mark (when it matches ack_id) and record the
|
/// Clear the write-through-pending mark (when it matches ack_id) and record the
|
||||||
/// host store event for each acked node.
|
/// host store event for each acked node.
|
||||||
pub fn finish_write_through(
|
pub fn finish_write_through(
|
||||||
@@ -4357,6 +4486,15 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
Ok(self.arena.node(node_id).write_through_pending_id)
|
Ok(self.arena.node(node_id).write_through_pending_id)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Whether a node is known to be stored in the external cache.
|
||||||
|
pub fn inspect_is_external_cache_stored(
|
||||||
|
&self,
|
||||||
|
node_id: NodeId,
|
||||||
|
) -> Result<bool, NodeAccessError> {
|
||||||
|
let node_id = self.arena.resolve(node_id)?;
|
||||||
|
Ok(self.arena.node(node_id).external_cache_stored)
|
||||||
|
}
|
||||||
|
|
||||||
/// Whether a node is in a component's device LRU.
|
/// Whether a node is in a component's device LRU.
|
||||||
pub fn inspect_is_node_in_device_lru(
|
pub fn inspect_is_node_in_device_lru(
|
||||||
&self,
|
&self,
|
||||||
|
|||||||
+5
-3
@@ -13,6 +13,7 @@ from sglang.test.test_utils import (
|
|||||||
find_available_port,
|
find_available_port,
|
||||||
popen_launch_server,
|
popen_launch_server,
|
||||||
terminate_and_kill_process_tree,
|
terminate_and_kill_process_tree,
|
||||||
|
unified_radix_tree_server_env,
|
||||||
)
|
)
|
||||||
|
|
||||||
GLM52_MODEL = os.environ.get("SGLANG_LINKER_GLM52_MODEL", "zai-org/GLM-5.2-FP8")
|
GLM52_MODEL = os.environ.get("SGLANG_LINKER_GLM52_MODEL", "zai-org/GLM-5.2-FP8")
|
||||||
@@ -22,6 +23,7 @@ register_cuda_ci(est_time=383, stage="extra-b", runner_config="8-gpu-h200")
|
|||||||
|
|
||||||
|
|
||||||
class TestGLM52UnifiedCacheLinkerKL(UnifiedRadixTreeTestMixin, CustomTestCase):
|
class TestGLM52UnifiedCacheLinkerKL(UnifiedRadixTreeTestMixin, CustomTestCase):
|
||||||
|
tree_core_backend = "rust"
|
||||||
page_size = 64
|
page_size = 64
|
||||||
kl_threshold = 0.03
|
kl_threshold = 0.03
|
||||||
sampling_temperature = 0
|
sampling_temperature = 0
|
||||||
@@ -61,10 +63,10 @@ class TestGLM52UnifiedCacheLinkerKL(UnifiedRadixTreeTestMixin, CustomTestCase):
|
|||||||
"--hicache-storage-backend-extra-config",
|
"--hicache-storage-backend-extra-config",
|
||||||
json.dumps({"enable_group_semantics": True}),
|
json.dumps({"enable_group_semantics": True}),
|
||||||
],
|
],
|
||||||
env={
|
env=unified_radix_tree_server_env(
|
||||||
|
cls.tree_core_backend,
|
||||||
**cls.mooncake.server_env(),
|
**cls.mooncake.server_env(),
|
||||||
"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1",
|
),
|
||||||
},
|
|
||||||
)
|
)
|
||||||
cls.input_ids = get_input_ids(cls.model, num_samples=18)
|
cls.input_ids = get_input_ids(cls.model, num_samples=18)
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|||||||
@@ -82,6 +82,9 @@ class RustUnifiedTreeCoreInspector(
|
|||||||
def get_write_through_pending_id(self, node_id: NodeId) -> Optional[int]:
|
def get_write_through_pending_id(self, node_id: NodeId) -> Optional[int]:
|
||||||
return self._binding.inspect_get_write_through_pending_id(node_id)
|
return self._binding.inspect_get_write_through_pending_id(node_id)
|
||||||
|
|
||||||
|
def is_external_cache_stored(self, node_id: NodeId) -> bool:
|
||||||
|
return self._binding.inspect_is_external_cache_stored(node_id)
|
||||||
|
|
||||||
def is_node_in_device_lru(
|
def is_node_in_device_lru(
|
||||||
self, node_id: NodeId, component_type: ComponentType
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
) -> bool:
|
) -> bool:
|
||||||
|
|||||||
@@ -266,6 +266,21 @@ def test_stale_handle_operations_raise_key_error_without_poisoning_the_core():
|
|||||||
stale_root, empty, PoolTransfer(name=PoolName.KV), {}
|
stale_root, empty, PoolTransfer(name=PoolName.KV), {}
|
||||||
),
|
),
|
||||||
"build_load_back_spec": lambda: core.build_load_back_spec(stale_root),
|
"build_load_back_spec": lambda: core.build_load_back_spec(stale_root),
|
||||||
|
"build_external_linker_offload_transfers": lambda: (
|
||||||
|
core.build_external_linker_offload_transfers(stale_root)
|
||||||
|
),
|
||||||
|
"mark_external_cache_stored_path/from": lambda: (
|
||||||
|
core.mark_external_cache_stored_path(stale_root, live_root)
|
||||||
|
),
|
||||||
|
"mark_external_cache_stored_path/until": lambda: (
|
||||||
|
core.mark_external_cache_stored_path(live_root, stale_root)
|
||||||
|
),
|
||||||
|
"mark_external_linker_offload_pending": lambda: (
|
||||||
|
core.mark_external_linker_offload_pending(stale_root)
|
||||||
|
),
|
||||||
|
"finish_external_linker_offload": lambda: core.finish_external_linker_offload(
|
||||||
|
[live_root, stale_root], live_root, True
|
||||||
|
),
|
||||||
"evict_excess_path_states": lambda: core.evict_excess_path_states(
|
"evict_excess_path_states": lambda: core.evict_excess_path_states(
|
||||||
stale_root, {}, {}
|
stale_root, {}, {}
|
||||||
),
|
),
|
||||||
@@ -499,13 +514,19 @@ def test_configuration_reads_the_locked_rust_state():
|
|||||||
assert swa_core.has_swa_host_pool is True
|
assert swa_core.has_swa_host_pool is True
|
||||||
|
|
||||||
|
|
||||||
def test_external_cache_linker_is_rejected():
|
def test_external_cache_linker_enablement_and_component_guard():
|
||||||
core = _tree_core()
|
core = _tree_core()
|
||||||
assert core.enable_external_cache_linker is False
|
assert core.enable_external_cache_linker is False
|
||||||
with pytest.raises(ValueError, match="External cache linker"):
|
core.enable_external_cache_linker = True
|
||||||
core.enable_external_cache_linker = True
|
assert core.enable_external_cache_linker is True
|
||||||
|
core.enable_external_cache_linker = False
|
||||||
assert core.enable_external_cache_linker is False
|
assert core.enable_external_cache_linker is False
|
||||||
|
|
||||||
|
mamba_core = _mamba_tree_core()
|
||||||
|
with pytest.raises(AssertionError, match="(?i)mamba"):
|
||||||
|
mamba_core.enable_external_cache_linker = True
|
||||||
|
assert mamba_core.enable_external_cache_linker is False
|
||||||
|
|
||||||
|
|
||||||
def test_sanity_check_passes_after_the_full_flow():
|
def test_sanity_check_passes_after_the_full_flow():
|
||||||
core = _tree_core()
|
core = _tree_core()
|
||||||
@@ -2179,6 +2200,7 @@ def test_stale_inspection_handles_raise_key_error_or_report_absence():
|
|||||||
"get_write_through_pending_id": lambda: core.get_write_through_pending_id(
|
"get_write_through_pending_id": lambda: core.get_write_through_pending_id(
|
||||||
stale_root
|
stale_root
|
||||||
),
|
),
|
||||||
|
"is_external_cache_stored": lambda: core.is_external_cache_stored(stale_root),
|
||||||
"is_node_in_device_lru": lambda: core.is_node_in_device_lru(
|
"is_node_in_device_lru": lambda: core.is_node_in_device_lru(
|
||||||
stale_root, ComponentType.FULL
|
stale_root, ComponentType.FULL
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -1,10 +1,33 @@
|
|||||||
|
import sys
|
||||||
|
import unittest
|
||||||
|
from array import array
|
||||||
|
from collections import defaultdict
|
||||||
|
from dataclasses import replace
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import test_unified_radix_cache_unittest as shared_cache_suite
|
||||||
import torch
|
import torch
|
||||||
|
from test_unified_radix_cache_unittest import (
|
||||||
|
CacheConfig,
|
||||||
|
_device_lock_ref,
|
||||||
|
_device_value,
|
||||||
|
_InsertWalkSuite,
|
||||||
|
_node_children,
|
||||||
|
build_fixture,
|
||||||
|
)
|
||||||
|
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import InsertResult
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer
|
InitLoadBackParams,
|
||||||
|
InsertResult,
|
||||||
|
MatchPrefixParams,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.hicache_storage import (
|
||||||
|
PoolHitPolicy,
|
||||||
|
PoolName,
|
||||||
|
PoolTransfer,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
from sglang.srt.mem_cache.unified_cache.cache_action import (
|
from sglang.srt.mem_cache.unified_cache.cache_action import (
|
||||||
ReplaceWriteThroughOnNodeSplit,
|
ReplaceWriteThroughOnNodeSplit,
|
||||||
)
|
)
|
||||||
@@ -13,16 +36,16 @@ from sglang.srt.mem_cache.unified_cache.components.full_component import FullCom
|
|||||||
from sglang.srt.mem_cache.unified_cache.components.swa_component import SWAComponent
|
from sglang.srt.mem_cache.unified_cache.components.swa_component import SWAComponent
|
||||||
from sglang.srt.mem_cache.unified_cache.components.tree_component import (
|
from sglang.srt.mem_cache.unified_cache.components.tree_component import (
|
||||||
ExternalLinkerLoadPhase,
|
ExternalLinkerLoadPhase,
|
||||||
LinkerTransferPhase,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.unified_cache.unified_cache_linker import (
|
from sglang.srt.mem_cache.unified_cache.unified_cache_linker import (
|
||||||
UnifiedCacheLinker,
|
UnifiedCacheLinker,
|
||||||
UnifiedCacheLinkerWrapper,
|
UnifiedCacheLinkerWrapper,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
register_cuda_ci(est_time=100, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
|
||||||
class _FakeLinker(UnifiedCacheLinker):
|
class _FakeLinker(UnifiedCacheLinker):
|
||||||
@@ -83,9 +106,41 @@ class _MappingRecorder:
|
|||||||
self.mapping.append((full.clone(), swa.clone()))
|
self.mapping.append((full.clone(), swa.clone()))
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeExternalTreeCore:
|
||||||
|
def __init__(self, nodes=None, offload_transfers=None):
|
||||||
|
self.enable_external_cache_linker = False
|
||||||
|
self.nodes = nodes or {}
|
||||||
|
self.offload_transfers = offload_transfers or [
|
||||||
|
PoolTransfer(name=PoolName.KV, keys=["page"])
|
||||||
|
]
|
||||||
|
|
||||||
|
def build_external_linker_offload_transfers(self, node_id):
|
||||||
|
node = self.nodes[node_id]
|
||||||
|
if node.external_cache_stored or node.write_through_pending_id is not None:
|
||||||
|
return None
|
||||||
|
return list(self.offload_transfers)
|
||||||
|
|
||||||
|
def mark_external_linker_offload_pending(self, node_id):
|
||||||
|
node = self.nodes[node_id]
|
||||||
|
assert (
|
||||||
|
not node.external_cache_stored and node.write_through_pending_id is None
|
||||||
|
), "invalid external offload state"
|
||||||
|
node.write_through_pending_id = node_id
|
||||||
|
|
||||||
|
def finish_external_linker_offload(self, node_ids, ack_id, success):
|
||||||
|
nodes = [self.nodes[node_id] for node_id in node_ids]
|
||||||
|
assert all(node.write_through_pending_id == ack_id for node in nodes), (
|
||||||
|
"invalid external offload state"
|
||||||
|
)
|
||||||
|
for node in nodes:
|
||||||
|
node.write_through_pending_id = None
|
||||||
|
node.external_cache_stored |= success
|
||||||
|
|
||||||
|
|
||||||
def _cache_for_wrapper(**kwargs):
|
def _cache_for_wrapper(**kwargs):
|
||||||
defaults = {
|
defaults = {
|
||||||
"tree_core": SimpleNamespace(enable_external_cache_linker=False),
|
"tree_core": SimpleNamespace(enable_external_cache_linker=False),
|
||||||
|
"tree_components": (ComponentType.FULL,),
|
||||||
"write_through_threshold": 256,
|
"write_through_threshold": 256,
|
||||||
"pp_size": 1,
|
"pp_size": 1,
|
||||||
"pp_group": None,
|
"pp_group": None,
|
||||||
@@ -100,6 +155,7 @@ def test_cache_linker_attachment_is_backend_independent():
|
|||||||
enable_external_cache_linker=False,
|
enable_external_cache_linker=False,
|
||||||
write_through_threshold=256,
|
write_through_threshold=256,
|
||||||
)
|
)
|
||||||
|
cache.tree_components = (ComponentType.FULL,)
|
||||||
cache.linker = None
|
cache.linker = None
|
||||||
linker = _FakeLinker()
|
linker = _FakeLinker()
|
||||||
|
|
||||||
@@ -111,6 +167,585 @@ def test_cache_linker_attachment_is_backend_independent():
|
|||||||
assert cache.linker.layer_done_counter is linker.layer_done_counter
|
assert cache.linker.layer_done_counter is linker.layer_done_counter
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("component_type", [ComponentType.MAMBA, ComponentType.C128])
|
||||||
|
def test_cache_linker_rejects_unsupported_tree_components(component_type):
|
||||||
|
cache = _cache_for_wrapper(tree_components=(ComponentType.FULL, component_type))
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match=component_type.name):
|
||||||
|
UnifiedCacheLinkerWrapper(cache, _FakeLinker())
|
||||||
|
|
||||||
|
assert not cache.tree_core.enable_external_cache_linker
|
||||||
|
|
||||||
|
|
||||||
|
class _InMemoryUnifiedCacheLinker(UnifiedCacheLinker):
|
||||||
|
"""Controllable transport for shared Python/Rust TreeCore tests."""
|
||||||
|
|
||||||
|
def __init__(self, stored_keys=None):
|
||||||
|
self.layer_done_counter = object()
|
||||||
|
self.stored_keys = defaultdict(set) if stored_keys is None else stored_keys
|
||||||
|
self.lookup_calls = []
|
||||||
|
self.offload_calls = []
|
||||||
|
self.pending_offloads = []
|
||||||
|
self.queued_loads = {}
|
||||||
|
self.started_loads = []
|
||||||
|
self.completed_loads = []
|
||||||
|
self.completed_offloads = []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _clone_transfer(transfer):
|
||||||
|
return replace(
|
||||||
|
transfer,
|
||||||
|
host_indices=(
|
||||||
|
None if transfer.host_indices is None else transfer.host_indices.clone()
|
||||||
|
),
|
||||||
|
device_indices=(
|
||||||
|
None
|
||||||
|
if transfer.device_indices is None
|
||||||
|
else transfer.device_indices.clone()
|
||||||
|
),
|
||||||
|
keys=None if transfer.keys is None else list(transfer.keys),
|
||||||
|
nodes_to_load=(
|
||||||
|
None if transfer.nodes_to_load is None else list(transfer.nodes_to_load)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _clone_transfers(cls, transfers):
|
||||||
|
return [cls._clone_transfer(transfer) for transfer in transfers]
|
||||||
|
|
||||||
|
def lookup(self, rid, transfers):
|
||||||
|
transfers = self._clone_transfers(transfers)
|
||||||
|
self.lookup_calls.append((rid, transfers))
|
||||||
|
by_pool = {transfer.name: transfer for transfer in transfers}
|
||||||
|
kv = by_pool.get(PoolName.KV)
|
||||||
|
if kv is None or not kv.keys:
|
||||||
|
return []
|
||||||
|
|
||||||
|
restorable = []
|
||||||
|
for prefix_pages in range(1, len(kv.keys) + 1):
|
||||||
|
if not set(kv.keys[:prefix_pages]) <= self.stored_keys[PoolName.KV]:
|
||||||
|
continue
|
||||||
|
for transfer in transfers:
|
||||||
|
if transfer.name == PoolName.KV:
|
||||||
|
continue
|
||||||
|
if transfer.hit_policy == PoolHitPolicy.TRAILING_PAGES:
|
||||||
|
window_pages = max(1, len(transfer.keys or ()))
|
||||||
|
required = kv.keys[
|
||||||
|
max(0, prefix_pages - window_pages) : prefix_pages
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
required = kv.keys[:prefix_pages]
|
||||||
|
if not set(required) <= self.stored_keys[transfer.name]:
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
restorable.append(prefix_pages)
|
||||||
|
return restorable
|
||||||
|
|
||||||
|
def load(self, rid, transfers):
|
||||||
|
self.queued_loads[rid] = self._clone_transfers(transfers)
|
||||||
|
return True
|
||||||
|
|
||||||
|
def start_layer_wise_loading(self):
|
||||||
|
if not self.queued_loads:
|
||||||
|
return -1
|
||||||
|
rids = list(self.queued_loads)
|
||||||
|
self.started_loads.append(rids)
|
||||||
|
return len(self.started_loads) - 1
|
||||||
|
|
||||||
|
def cancel_queued_load(self, rid):
|
||||||
|
return self.queued_loads.pop(rid, None) is not None
|
||||||
|
|
||||||
|
def num_completed_loads(self):
|
||||||
|
return len(self.completed_loads)
|
||||||
|
|
||||||
|
def pop_completed_load(self):
|
||||||
|
rids = self.completed_loads.pop(0)
|
||||||
|
for rid in rids:
|
||||||
|
self.queued_loads.pop(rid, None)
|
||||||
|
return rids
|
||||||
|
|
||||||
|
def complete_started_loads(self):
|
||||||
|
self.completed_loads.append(self.started_loads[-1])
|
||||||
|
|
||||||
|
def offload(self, transfers):
|
||||||
|
transfers = self._clone_transfers(transfers)
|
||||||
|
self.offload_calls.append(transfers)
|
||||||
|
self.pending_offloads.append(transfers)
|
||||||
|
return True
|
||||||
|
|
||||||
|
def num_completed_offloads(self):
|
||||||
|
return len(self.completed_offloads)
|
||||||
|
|
||||||
|
def pop_completed_offload(self):
|
||||||
|
return self.completed_offloads.pop(0)
|
||||||
|
|
||||||
|
def complete_next_offload(self, success):
|
||||||
|
transfers = self.pending_offloads.pop(0)
|
||||||
|
if success:
|
||||||
|
for transfer in transfers:
|
||||||
|
self.stored_keys[transfer.name].update(transfer.keys or ())
|
||||||
|
self.completed_offloads.append(success)
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
self.queued_loads.clear()
|
||||||
|
self.pending_offloads.clear()
|
||||||
|
self.completed_loads.clear()
|
||||||
|
self.completed_offloads.clear()
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
self.reset()
|
||||||
|
|
||||||
|
|
||||||
|
class _TreeCoreBackendTestMixin:
|
||||||
|
tree_core_backend = "python"
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
previous = shared_cache_suite._TREE_CORE_TEST_BACKEND
|
||||||
|
self.addCleanup(
|
||||||
|
setattr,
|
||||||
|
shared_cache_suite,
|
||||||
|
"_TREE_CORE_TEST_BACKEND",
|
||||||
|
previous,
|
||||||
|
)
|
||||||
|
shared_cache_suite._TREE_CORE_TEST_BACKEND = self.tree_core_backend
|
||||||
|
super().setUp()
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA")
|
||||||
|
class TestUnifiedCacheLinkerPythonBackend(_TreeCoreBackendTestMixin, _InsertWalkSuite):
|
||||||
|
def test_full_offload_load_round_trip_and_dedup(self):
|
||||||
|
cfg = CacheConfig(page_size=2, kv_size=64, max_context_len=64)
|
||||||
|
self.cfg = cfg
|
||||||
|
stored_keys = defaultdict(set)
|
||||||
|
|
||||||
|
producer, producer_allocator, producer_req_pool = build_fixture(cfg)
|
||||||
|
producer_linker = _InMemoryUnifiedCacheLinker(stored_keys)
|
||||||
|
producer.init_cache_linker(producer_linker)
|
||||||
|
tokens = list(range(1, 9))
|
||||||
|
|
||||||
|
inserted = self._insert(producer, producer_allocator, producer_req_pool, tokens)
|
||||||
|
self.assertEqual(len(producer_linker.offload_calls), 1)
|
||||||
|
(kv_offload,) = producer_linker.offload_calls[0]
|
||||||
|
self.assertEqual(kv_offload.name, PoolName.KV)
|
||||||
|
self.assertEqual(
|
||||||
|
kv_offload.keys,
|
||||||
|
producer.tree_core.get_hash_values(inserted.last_device_node),
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(
|
||||||
|
kv_offload.device_indices,
|
||||||
|
_device_value(producer, inserted.last_device_node, ComponentType.FULL),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
_device_lock_ref(producer, inserted.last_device_node, ComponentType.FULL),
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
self.assertFalse(
|
||||||
|
producer.tree_core.is_external_cache_stored(inserted.last_device_node)
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
producer.tree_core.get_write_through_pending_id(inserted.last_device_node),
|
||||||
|
inserted.last_device_node,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._insert(producer, producer_allocator, producer_req_pool, tokens)
|
||||||
|
self.assertEqual(len(producer_linker.offload_calls), 1)
|
||||||
|
|
||||||
|
producer_linker.complete_next_offload(True)
|
||||||
|
producer.check_hicache_events()
|
||||||
|
self.assertEqual(
|
||||||
|
_device_lock_ref(producer, inserted.last_device_node, ComponentType.FULL),
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
producer.tree_core.is_external_cache_stored(inserted.last_device_node)
|
||||||
|
)
|
||||||
|
self._insert(producer, producer_allocator, producer_req_pool, tokens)
|
||||||
|
self.assertEqual(len(producer_linker.offload_calls), 1)
|
||||||
|
|
||||||
|
consumer, _, consumer_req_pool = build_fixture(cfg)
|
||||||
|
consumer_linker = _InMemoryUnifiedCacheLinker(stored_keys)
|
||||||
|
consumer.init_cache_linker(consumer_linker)
|
||||||
|
req = self._make_req(consumer_req_pool)
|
||||||
|
match = consumer.match_prefix(
|
||||||
|
MatchPrefixParams(key=RadixKey(array("q", tokens)), req=req)
|
||||||
|
)
|
||||||
|
self.assertEqual(match.device_indices.numel(), 0)
|
||||||
|
self.assertEqual(match.host_hit_length, len(tokens))
|
||||||
|
self._apply_match_to_req(req, match)
|
||||||
|
|
||||||
|
loaded, loaded_node = consumer.init_load_back(
|
||||||
|
InitLoadBackParams(
|
||||||
|
best_match_node=match.best_match_node,
|
||||||
|
host_hit_length=match.host_hit_length,
|
||||||
|
req=req,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(loaded.numel(), len(tokens))
|
||||||
|
self.assertNotEqual(loaded_node, consumer.root_node_handle())
|
||||||
|
(kv_load,) = consumer_linker.queued_loads[req.rid]
|
||||||
|
self.assertEqual(kv_load.name, PoolName.KV)
|
||||||
|
self.assertEqual(kv_load.keys, kv_offload.keys)
|
||||||
|
self.assertEqual(_device_lock_ref(consumer, loaded_node, ComponentType.FULL), 1)
|
||||||
|
|
||||||
|
self.assertGreaterEqual(consumer.ready_to_load_host_cache(), 0)
|
||||||
|
consumer_linker.complete_started_loads()
|
||||||
|
consumer.check_hicache_events()
|
||||||
|
self.assertEqual(_device_lock_ref(consumer, loaded_node, ComponentType.FULL), 0)
|
||||||
|
final_match = consumer.match_prefix(
|
||||||
|
MatchPrefixParams(key=RadixKey(array("q", tokens)))
|
||||||
|
)
|
||||||
|
self.assertEqual(final_match.device_indices.numel(), len(tokens))
|
||||||
|
self.assertEqual(consumer_linker.offload_calls, [])
|
||||||
|
consumer.sanity_check()
|
||||||
|
|
||||||
|
def test_eagle_lookup_uses_bigram_tail_hashes(self):
|
||||||
|
cfg = CacheConfig(
|
||||||
|
page_size=2,
|
||||||
|
is_eagle=True,
|
||||||
|
kv_size=64,
|
||||||
|
max_context_len=64,
|
||||||
|
)
|
||||||
|
self.cfg = cfg
|
||||||
|
stored_keys = defaultdict(set)
|
||||||
|
tokens = list(range(1, 10))
|
||||||
|
|
||||||
|
producer, producer_allocator, producer_req_pool = build_fixture(cfg)
|
||||||
|
producer_linker = _InMemoryUnifiedCacheLinker(stored_keys)
|
||||||
|
producer.init_cache_linker(producer_linker)
|
||||||
|
self._insert(producer, producer_allocator, producer_req_pool, tokens)
|
||||||
|
self.assertEqual(len(producer_linker.offload_calls), 1)
|
||||||
|
(kv_offload,) = producer_linker.offload_calls[0]
|
||||||
|
self.assertEqual(len(kv_offload.keys), 4)
|
||||||
|
producer_linker.complete_next_offload(True)
|
||||||
|
producer.check_hicache_events()
|
||||||
|
|
||||||
|
consumer, consumer_allocator, consumer_req_pool = build_fixture(cfg)
|
||||||
|
consumer_linker = _InMemoryUnifiedCacheLinker(stored_keys)
|
||||||
|
consumer.init_cache_linker(consumer_linker)
|
||||||
|
consumer.write_through_threshold = sys.maxsize
|
||||||
|
self._insert(
|
||||||
|
consumer,
|
||||||
|
consumer_allocator,
|
||||||
|
consumer_req_pool,
|
||||||
|
tokens[:5],
|
||||||
|
)
|
||||||
|
req = self._make_req(consumer_req_pool)
|
||||||
|
lookup_key = RadixKey(array("q", tokens))
|
||||||
|
|
||||||
|
match = consumer.match_prefix(MatchPrefixParams(key=lookup_key, req=req))
|
||||||
|
|
||||||
|
self.assertTrue(lookup_key.is_bigram)
|
||||||
|
self.assertEqual(match.device_indices.numel(), 4)
|
||||||
|
self.assertEqual(match.host_hit_length, 4)
|
||||||
|
self.assertEqual(len(consumer_linker.lookup_calls), 1)
|
||||||
|
_, transfers = consumer_linker.lookup_calls[0]
|
||||||
|
(kv_lookup,) = transfers
|
||||||
|
self.assertEqual(kv_lookup.keys, kv_offload.keys[2:])
|
||||||
|
|
||||||
|
def test_failed_split_offload_retries_and_reset_releases_locks(self):
|
||||||
|
cfg = CacheConfig(page_size=1, kv_size=64, max_context_len=64)
|
||||||
|
self.cfg = cfg
|
||||||
|
cache, allocator, req_to_token_pool = build_fixture(cfg)
|
||||||
|
linker = _InMemoryUnifiedCacheLinker()
|
||||||
|
cache.init_cache_linker(linker)
|
||||||
|
tokens = [1, 2, 3, 4]
|
||||||
|
|
||||||
|
inserted = self._insert(cache, allocator, req_to_token_pool, tokens)
|
||||||
|
original_node = inserted.last_device_node
|
||||||
|
self.assertEqual(len(linker.offload_calls), 1)
|
||||||
|
original_keys = linker.offload_calls[0][0].keys
|
||||||
|
|
||||||
|
self._insert(cache, allocator, req_to_token_pool, tokens[:2])
|
||||||
|
(parent,) = _node_children(cache, cache.root_node_handle())
|
||||||
|
(child,) = _node_children(cache, parent)
|
||||||
|
self.assertEqual(child, original_node)
|
||||||
|
self.assertEqual(
|
||||||
|
cache.linker.pending_offloads[0].publish_node_ids, [parent, child]
|
||||||
|
)
|
||||||
|
self.assertEqual(cache.tree_core.get_write_through_pending_id(parent), child)
|
||||||
|
self.assertEqual(cache.tree_core.get_write_through_pending_id(child), child)
|
||||||
|
self.assertFalse(cache.tree_core.is_external_cache_stored(parent))
|
||||||
|
self.assertFalse(cache.tree_core.is_external_cache_stored(child))
|
||||||
|
|
||||||
|
linker.complete_next_offload(False)
|
||||||
|
cache.check_hicache_events()
|
||||||
|
for node_id in (parent, child):
|
||||||
|
self.assertIsNone(cache.tree_core.get_write_through_pending_id(node_id))
|
||||||
|
self.assertFalse(cache.tree_core.is_external_cache_stored(node_id))
|
||||||
|
self.assertEqual(_device_lock_ref(cache, node_id, ComponentType.FULL), 0)
|
||||||
|
|
||||||
|
self._insert(cache, allocator, req_to_token_pool, tokens)
|
||||||
|
retry_calls = linker.offload_calls[1:]
|
||||||
|
self.assertEqual(len(retry_calls), 2)
|
||||||
|
self.assertEqual(
|
||||||
|
[key for transfers in retry_calls for key in transfers[0].keys],
|
||||||
|
original_keys,
|
||||||
|
)
|
||||||
|
for _ in retry_calls:
|
||||||
|
linker.complete_next_offload(True)
|
||||||
|
cache.check_hicache_events()
|
||||||
|
for node_id in (parent, child):
|
||||||
|
self.assertIsNone(cache.tree_core.get_write_through_pending_id(node_id))
|
||||||
|
self.assertTrue(cache.tree_core.is_external_cache_stored(node_id))
|
||||||
|
self.assertEqual(_device_lock_ref(cache, node_id, ComponentType.FULL), 0)
|
||||||
|
|
||||||
|
self._insert(cache, allocator, req_to_token_pool, tokens)
|
||||||
|
self.assertEqual(len(linker.offload_calls), 3)
|
||||||
|
|
||||||
|
extended = self._insert(cache, allocator, req_to_token_pool, tokens + [5, 6])
|
||||||
|
pending_node = extended.last_device_node
|
||||||
|
self.assertEqual(len(cache.linker.pending_offloads), 1)
|
||||||
|
self.assertEqual(
|
||||||
|
cache.tree_core.get_write_through_pending_id(pending_node), pending_node
|
||||||
|
)
|
||||||
|
self.assertEqual(_device_lock_ref(cache, pending_node, ComponentType.FULL), 1)
|
||||||
|
|
||||||
|
cache.linker.reset()
|
||||||
|
self.assertEqual(cache.linker.pending_offloads, [])
|
||||||
|
self.assertIsNone(cache.tree_core.get_write_through_pending_id(pending_node))
|
||||||
|
self.assertEqual(_device_lock_ref(cache, pending_node, ComponentType.FULL), 0)
|
||||||
|
cache.sanity_check()
|
||||||
|
cache.reset()
|
||||||
|
cache.sanity_check()
|
||||||
|
|
||||||
|
def test_swa_partial_hit_loads_only_pages_not_adopted_locally(self):
|
||||||
|
cfg = CacheConfig(
|
||||||
|
page_size=1,
|
||||||
|
components=(ComponentType.FULL, ComponentType.SWA),
|
||||||
|
sliding_window_size=2,
|
||||||
|
kv_size=64,
|
||||||
|
max_context_len=64,
|
||||||
|
)
|
||||||
|
self.cfg = cfg
|
||||||
|
stored_keys = defaultdict(set)
|
||||||
|
tokens = list(range(1, 7))
|
||||||
|
|
||||||
|
producer, producer_allocator, producer_req_pool = build_fixture(cfg)
|
||||||
|
producer_linker = _InMemoryUnifiedCacheLinker(stored_keys)
|
||||||
|
producer.init_cache_linker(producer_linker)
|
||||||
|
self._insert(producer, producer_allocator, producer_req_pool, tokens)
|
||||||
|
self.assertGreaterEqual(len(producer_linker.offload_calls), 1)
|
||||||
|
all_keys = []
|
||||||
|
for transfers in producer_linker.offload_calls:
|
||||||
|
offloads = {transfer.name: transfer for transfer in transfers}
|
||||||
|
self.assertEqual(set(offloads), {PoolName.KV, PoolName.SWA})
|
||||||
|
self.assertEqual(offloads[PoolName.KV].keys, offloads[PoolName.SWA].keys)
|
||||||
|
all_keys.extend(offloads[PoolName.KV].keys)
|
||||||
|
producer_linker.complete_next_offload(True)
|
||||||
|
producer.check_hicache_events()
|
||||||
|
|
||||||
|
stored_keys[PoolName.KV].difference_update(all_keys[-2:])
|
||||||
|
|
||||||
|
consumer, consumer_allocator, consumer_req_pool = build_fixture(cfg)
|
||||||
|
consumer_linker = _InMemoryUnifiedCacheLinker(stored_keys)
|
||||||
|
consumer.init_cache_linker(consumer_linker)
|
||||||
|
consumer.write_through_threshold = sys.maxsize
|
||||||
|
self._insert(consumer, consumer_allocator, consumer_req_pool, tokens[:2])
|
||||||
|
|
||||||
|
req = self._make_req(consumer_req_pool)
|
||||||
|
match = consumer.match_prefix(
|
||||||
|
MatchPrefixParams(key=RadixKey(array("q", tokens)), req=req)
|
||||||
|
)
|
||||||
|
self.assertEqual(match.device_indices.numel(), 2)
|
||||||
|
self.assertEqual(match.host_hit_length, 2)
|
||||||
|
self.assertEqual(match.swa_host_hit_length, 2)
|
||||||
|
self._apply_match_to_req(req, match)
|
||||||
|
|
||||||
|
self._insert(consumer, consumer_allocator, consumer_req_pool, tokens[:3])
|
||||||
|
raced_match = consumer.match_prefix(
|
||||||
|
MatchPrefixParams(key=RadixKey(array("q", tokens[:3])))
|
||||||
|
)
|
||||||
|
raced_full = raced_match.device_indices[-1:].clone()
|
||||||
|
raced_swa = consumer_allocator.translate_loc_from_full_to_swa(raced_full)
|
||||||
|
|
||||||
|
loaded, loaded_node = consumer.init_load_back(
|
||||||
|
InitLoadBackParams(
|
||||||
|
best_match_node=match.best_match_node,
|
||||||
|
host_hit_length=match.host_hit_length,
|
||||||
|
req=req,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(loaded.numel(), 2)
|
||||||
|
self.assertTrue(torch.equal(loaded[:1], raced_full))
|
||||||
|
load_by_pool = {
|
||||||
|
transfer.name: transfer
|
||||||
|
for transfer in consumer_linker.queued_loads[req.rid]
|
||||||
|
}
|
||||||
|
self.assertEqual(set(load_by_pool), {PoolName.KV, PoolName.SWA})
|
||||||
|
expected_key = all_keys[3]
|
||||||
|
for transfer in load_by_pool.values():
|
||||||
|
self.assertEqual(transfer.keys, [expected_key])
|
||||||
|
self.assertEqual(transfer.device_indices.numel(), 1)
|
||||||
|
|
||||||
|
translated = consumer_allocator.translate_loc_from_full_to_swa(loaded)
|
||||||
|
self.assertTrue(torch.equal(translated[:1], raced_swa))
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(translated[-1:], load_by_pool[PoolName.SWA].device_indices)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertGreaterEqual(consumer.ready_to_load_host_cache(), 0)
|
||||||
|
consumer_linker.complete_started_loads()
|
||||||
|
consumer.check_hicache_events()
|
||||||
|
final_match = consumer.match_prefix(
|
||||||
|
MatchPrefixParams(key=RadixKey(array("q", tokens[:4])))
|
||||||
|
)
|
||||||
|
self.assertEqual(final_match.device_indices.numel(), 4)
|
||||||
|
self.assertEqual(_device_lock_ref(consumer, loaded_node, ComponentType.FULL), 0)
|
||||||
|
consumer.sanity_check()
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA")
|
||||||
|
class TestUnifiedCacheLinkerTreeCorePythonBackend(
|
||||||
|
_TreeCoreBackendTestMixin, shared_cache_suite._InsertWalkSuite
|
||||||
|
):
|
||||||
|
"""TreeCore linker contracts shared by the Python and Rust inspectors."""
|
||||||
|
|
||||||
|
def test_builds_opaque_external_offload_transfers(self):
|
||||||
|
cfg = shared_cache_suite.CacheConfig(
|
||||||
|
page_size=2, kv_size=64, max_context_len=64
|
||||||
|
)
|
||||||
|
self.cfg = cfg
|
||||||
|
cache, allocator, req_to_token_pool = shared_cache_suite.build_fixture(cfg)
|
||||||
|
core = cache.tree_core
|
||||||
|
core.enable_external_cache_linker = True
|
||||||
|
|
||||||
|
inserted = self._insert(cache, allocator, req_to_token_pool, list(range(1, 9)))
|
||||||
|
node_id = inserted.last_device_node
|
||||||
|
transfers = core.build_external_linker_offload_transfers(node_id)
|
||||||
|
|
||||||
|
self.assertIsNotNone(transfers)
|
||||||
|
(transfer,) = transfers
|
||||||
|
self.assertEqual(transfer.name, PoolName.KV)
|
||||||
|
self.assertIsNone(transfer.host_indices)
|
||||||
|
self.assertIsNotNone(transfer.device_indices)
|
||||||
|
self.assertEqual(transfer.keys, core.get_hash_values(node_id))
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(
|
||||||
|
transfer.device_indices,
|
||||||
|
shared_cache_suite._device_value(cache, node_id, ComponentType.FULL),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
core.mark_external_linker_offload_pending(node_id)
|
||||||
|
self.assertFalse(core.is_external_cache_stored(node_id))
|
||||||
|
self.assertIsNone(core.build_external_linker_offload_transfers(node_id))
|
||||||
|
with self.assertRaisesRegex(AssertionError, "invalid external offload state"):
|
||||||
|
core.mark_external_linker_offload_pending(node_id)
|
||||||
|
self.assertEqual(core.get_write_through_pending_id(node_id), node_id)
|
||||||
|
self.assertFalse(core.is_external_cache_stored(node_id))
|
||||||
|
|
||||||
|
core.finish_external_linker_offload([node_id], node_id, success=True)
|
||||||
|
self.assertIsNone(core.get_write_through_pending_id(node_id))
|
||||||
|
self.assertTrue(core.is_external_cache_stored(node_id))
|
||||||
|
|
||||||
|
def test_external_state_updates_are_atomic_and_path_scoped(self):
|
||||||
|
cfg = shared_cache_suite.CacheConfig(
|
||||||
|
page_size=1, kv_size=64, max_context_len=64
|
||||||
|
)
|
||||||
|
self.cfg = cfg
|
||||||
|
cache, allocator, req_to_token_pool = shared_cache_suite.build_fixture(cfg)
|
||||||
|
core = cache.tree_core
|
||||||
|
core.enable_external_cache_linker = True
|
||||||
|
|
||||||
|
anchor = self._insert(cache, allocator, req_to_token_pool, [1]).last_device_node
|
||||||
|
middle = self._insert(
|
||||||
|
cache, allocator, req_to_token_pool, [1, 2]
|
||||||
|
).last_device_node
|
||||||
|
tail = self._insert(
|
||||||
|
cache, allocator, req_to_token_pool, [1, 2, 3, 4]
|
||||||
|
).last_device_node
|
||||||
|
unrelated = self._insert(
|
||||||
|
cache, allocator, req_to_token_pool, [9]
|
||||||
|
).last_device_node
|
||||||
|
self.assertEqual(core.get_parent_node_id(middle), anchor)
|
||||||
|
self.assertEqual(core.get_parent_node_id(tail), middle)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "not an ancestor"):
|
||||||
|
core.mark_external_cache_stored_path(tail, unrelated)
|
||||||
|
self.assertFalse(core.is_external_cache_stored(middle))
|
||||||
|
self.assertFalse(core.is_external_cache_stored(tail))
|
||||||
|
|
||||||
|
core.mark_external_cache_stored_path(tail, anchor)
|
||||||
|
self.assertTrue(core.is_external_cache_stored(tail))
|
||||||
|
self.assertTrue(core.is_external_cache_stored(middle))
|
||||||
|
self.assertFalse(core.is_external_cache_stored(anchor))
|
||||||
|
self.assertFalse(core.is_external_cache_stored(unrelated))
|
||||||
|
with self.assertRaisesRegex(AssertionError, "invalid external offload state"):
|
||||||
|
core.mark_external_linker_offload_pending(tail)
|
||||||
|
|
||||||
|
split_tail = self._insert(
|
||||||
|
cache, allocator, req_to_token_pool, [9, 10, 11]
|
||||||
|
).last_device_node
|
||||||
|
self.assertEqual(core.get_parent_node_id(split_tail), unrelated)
|
||||||
|
core.mark_external_linker_offload_pending(split_tail)
|
||||||
|
self._insert(cache, allocator, req_to_token_pool, [9, 10])
|
||||||
|
split_parent = core.get_parent_node_id(split_tail)
|
||||||
|
self.assertIsNotNone(split_parent)
|
||||||
|
self.assertNotEqual(split_parent, unrelated)
|
||||||
|
self.assertEqual(core.get_parent_node_id(split_parent), unrelated)
|
||||||
|
for node_id in (split_parent, split_tail):
|
||||||
|
self.assertEqual(core.get_write_through_pending_id(node_id), split_tail)
|
||||||
|
self.assertFalse(core.is_external_cache_stored(node_id))
|
||||||
|
|
||||||
|
independent = self._insert(
|
||||||
|
cache, allocator, req_to_token_pool, [20]
|
||||||
|
).last_device_node
|
||||||
|
core.mark_external_linker_offload_pending(independent)
|
||||||
|
with self.assertRaisesRegex(AssertionError, "invalid external offload state"):
|
||||||
|
core.finish_external_linker_offload(
|
||||||
|
[independent, split_parent], independent, success=False
|
||||||
|
)
|
||||||
|
self.assertEqual(core.get_write_through_pending_id(independent), independent)
|
||||||
|
for node_id in (split_parent, split_tail):
|
||||||
|
self.assertEqual(core.get_write_through_pending_id(node_id), split_tail)
|
||||||
|
self.assertFalse(core.is_external_cache_stored(node_id))
|
||||||
|
|
||||||
|
core.finish_external_linker_offload([independent], independent, success=False)
|
||||||
|
core.finish_external_linker_offload(
|
||||||
|
[split_parent, split_tail], split_tail, success=False
|
||||||
|
)
|
||||||
|
for node_id in (split_parent, split_tail):
|
||||||
|
self.assertIsNone(core.get_write_through_pending_id(node_id))
|
||||||
|
self.assertFalse(core.is_external_cache_stored(node_id))
|
||||||
|
self.assertTrue(core.is_external_cache_stored(middle))
|
||||||
|
self.assertTrue(core.is_external_cache_stored(tail))
|
||||||
|
|
||||||
|
def test_failed_offload_preserves_independently_confirmed_state(self):
|
||||||
|
cfg = shared_cache_suite.CacheConfig(
|
||||||
|
page_size=1, kv_size=64, max_context_len=64
|
||||||
|
)
|
||||||
|
self.cfg = cfg
|
||||||
|
cache, allocator, req_to_token_pool = shared_cache_suite.build_fixture(cfg)
|
||||||
|
core = cache.tree_core
|
||||||
|
core.enable_external_cache_linker = True
|
||||||
|
|
||||||
|
anchor = self._insert(cache, allocator, req_to_token_pool, [1]).last_device_node
|
||||||
|
node_id = self._insert(
|
||||||
|
cache, allocator, req_to_token_pool, [1, 2]
|
||||||
|
).last_device_node
|
||||||
|
core.mark_external_linker_offload_pending(node_id)
|
||||||
|
|
||||||
|
core.mark_external_cache_stored_path(node_id, anchor)
|
||||||
|
self.assertEqual(core.get_write_through_pending_id(node_id), node_id)
|
||||||
|
self.assertTrue(core.is_external_cache_stored(node_id))
|
||||||
|
|
||||||
|
core.finish_external_linker_offload([node_id], node_id, success=False)
|
||||||
|
self.assertIsNone(core.get_write_through_pending_id(node_id))
|
||||||
|
self.assertTrue(core.is_external_cache_stored(node_id))
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedCacheLinkerRustBackend(TestUnifiedCacheLinkerPythonBackend):
|
||||||
|
tree_core_backend = "rust"
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedCacheLinkerTreeCoreRustBackend(
|
||||||
|
TestUnifiedCacheLinkerTreeCorePythonBackend
|
||||||
|
):
|
||||||
|
tree_core_backend = "rust"
|
||||||
|
|
||||||
|
|
||||||
def test_restorable_prefix_intersects_sparse_rank_results():
|
def test_restorable_prefix_intersects_sparse_rank_results():
|
||||||
remote_mask = torch.tensor([0, 0, 1, 0, 0], dtype=torch.int)
|
remote_mask = torch.tensor([0, 0, 1, 0, 0], dtype=torch.int)
|
||||||
|
|
||||||
@@ -127,11 +762,6 @@ def test_restorable_prefix_intersects_sparse_rank_results():
|
|||||||
|
|
||||||
|
|
||||||
def test_async_offload_pins_node_until_completion():
|
def test_async_offload_pins_node_until_completion():
|
||||||
class _Component:
|
|
||||||
def build_external_linker_transfer(self, phase, node, keys):
|
|
||||||
assert phase == LinkerTransferPhase.OFFLOAD
|
|
||||||
return PoolTransfer(name=PoolName.KV, keys=["page"])
|
|
||||||
|
|
||||||
linker = _FakeLinker()
|
linker = _FakeLinker()
|
||||||
lock_params = object()
|
lock_params = object()
|
||||||
locks = []
|
locks = []
|
||||||
@@ -148,23 +778,17 @@ def test_async_offload_pins_node_until_completion():
|
|||||||
write_through_pending_id=None,
|
write_through_pending_id=None,
|
||||||
)
|
)
|
||||||
cache = _cache_for_wrapper(
|
cache = _cache_for_wrapper(
|
||||||
tree_core=SimpleNamespace(
|
tree_core=_FakeExternalTreeCore({node_id: node}),
|
||||||
enable_external_cache_linker=False,
|
|
||||||
mark_write_through_pending=lambda node_ids, ack_id: (
|
|
||||||
setattr(node, "write_through_pending_id", ack_id) or list(node_ids)
|
|
||||||
),
|
|
||||||
),
|
|
||||||
_components_tuple=(_Component(),),
|
|
||||||
inc_lock_ref=inc_lock_ref,
|
inc_lock_ref=inc_lock_ref,
|
||||||
dec_lock_ref=lambda node, params: unlocks.append((node, params)),
|
dec_lock_ref=lambda node, params: unlocks.append((node, params)),
|
||||||
resolve_node_handle=lambda value: node if value == node_id else None,
|
|
||||||
)
|
)
|
||||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||||
|
|
||||||
wrapper.offload_nodes([node_id])
|
wrapper.offload_nodes([node_id])
|
||||||
|
|
||||||
assert locks == [node_id]
|
assert locks == [node_id]
|
||||||
assert node.external_cache_stored
|
assert not node.external_cache_stored
|
||||||
|
assert node.write_through_pending_id == node_id
|
||||||
assert not unlocks
|
assert not unlocks
|
||||||
|
|
||||||
linker.completed_offloads.append(False)
|
linker.completed_offloads.append(False)
|
||||||
@@ -172,9 +796,28 @@ def test_async_offload_pins_node_until_completion():
|
|||||||
wrapper.commit_completed_offloads(completed)
|
wrapper.commit_completed_offloads(completed)
|
||||||
|
|
||||||
assert not node.external_cache_stored
|
assert not node.external_cache_stored
|
||||||
|
assert node.write_through_pending_id is None
|
||||||
assert unlocks == [(node_id, lock_params)]
|
assert unlocks == [(node_id, lock_params)]
|
||||||
|
|
||||||
|
|
||||||
|
def test_offload_skips_node_already_stored_by_tree_core():
|
||||||
|
linker = _FakeLinker()
|
||||||
|
node = SimpleNamespace(
|
||||||
|
id=7,
|
||||||
|
external_cache_stored=True,
|
||||||
|
write_through_pending_id=None,
|
||||||
|
)
|
||||||
|
cache = _cache_for_wrapper(
|
||||||
|
tree_core=_FakeExternalTreeCore({node.id: node}),
|
||||||
|
inc_lock_ref=lambda node_id: pytest.fail("stored node must not be locked"),
|
||||||
|
)
|
||||||
|
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||||
|
|
||||||
|
wrapper.offload_nodes([node.id])
|
||||||
|
|
||||||
|
assert linker.queued_offloads == []
|
||||||
|
|
||||||
|
|
||||||
def test_async_load_pins_node_until_completion():
|
def test_async_load_pins_node_until_completion():
|
||||||
linker = _FakeLinker()
|
linker = _FakeLinker()
|
||||||
lock_params = object()
|
lock_params = object()
|
||||||
@@ -224,10 +867,6 @@ def test_release_request_cancels_queued_load():
|
|||||||
|
|
||||||
|
|
||||||
def test_failed_offload_rolls_back_split_fragments():
|
def test_failed_offload_rolls_back_split_fragments():
|
||||||
class _Component:
|
|
||||||
def build_external_linker_transfer(self, phase, node, keys):
|
|
||||||
return PoolTransfer(name=PoolName.KV, keys=["page"])
|
|
||||||
|
|
||||||
linker = _FakeLinker()
|
linker = _FakeLinker()
|
||||||
lock_params = object()
|
lock_params = object()
|
||||||
unlocks = []
|
unlocks = []
|
||||||
@@ -243,20 +882,10 @@ def test_failed_offload_rolls_back_split_fragments():
|
|||||||
)
|
)
|
||||||
nodes = {child.id: child, parent.id: parent}
|
nodes = {child.id: child, parent.id: parent}
|
||||||
|
|
||||||
def mark_pending(node_ids, ack_id):
|
|
||||||
for node_id in node_ids:
|
|
||||||
nodes[node_id].write_through_pending_id = ack_id
|
|
||||||
return list(node_ids)
|
|
||||||
|
|
||||||
cache = _cache_for_wrapper(
|
cache = _cache_for_wrapper(
|
||||||
tree_core=SimpleNamespace(
|
tree_core=_FakeExternalTreeCore(nodes),
|
||||||
enable_external_cache_linker=False,
|
|
||||||
mark_write_through_pending=mark_pending,
|
|
||||||
),
|
|
||||||
_components_tuple=(_Component(),),
|
|
||||||
inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=lambda: lock_params),
|
inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=lambda: lock_params),
|
||||||
dec_lock_ref=lambda node_id, params: unlocks.append((node_id, params)),
|
dec_lock_ref=lambda node_id, params: unlocks.append((node_id, params)),
|
||||||
resolve_node_handle=nodes.__getitem__,
|
|
||||||
)
|
)
|
||||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||||
wrapper.offload_nodes([child.id])
|
wrapper.offload_nodes([child.id])
|
||||||
@@ -299,10 +928,6 @@ def test_split_action_retargets_pending_external_offload():
|
|||||||
|
|
||||||
|
|
||||||
def test_reset_quiesces_backend_before_releasing_pending_locks():
|
def test_reset_quiesces_backend_before_releasing_pending_locks():
|
||||||
class _Component:
|
|
||||||
def build_external_linker_transfer(self, phase, node, keys):
|
|
||||||
return PoolTransfer(name=PoolName.KV, keys=["page"])
|
|
||||||
|
|
||||||
events = []
|
events = []
|
||||||
|
|
||||||
class _QuiescentFakeLinker(_FakeLinker):
|
class _QuiescentFakeLinker(_FakeLinker):
|
||||||
@@ -317,16 +942,9 @@ def test_reset_quiesces_backend_before_releasing_pending_locks():
|
|||||||
write_through_pending_id=None,
|
write_through_pending_id=None,
|
||||||
)
|
)
|
||||||
cache = _cache_for_wrapper(
|
cache = _cache_for_wrapper(
|
||||||
tree_core=SimpleNamespace(
|
tree_core=_FakeExternalTreeCore({node.id: node}),
|
||||||
enable_external_cache_linker=False,
|
|
||||||
mark_write_through_pending=lambda node_ids, ack_id: (
|
|
||||||
setattr(node, "write_through_pending_id", ack_id) or list(node_ids)
|
|
||||||
),
|
|
||||||
),
|
|
||||||
_components_tuple=(_Component(),),
|
|
||||||
inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=object),
|
inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=object),
|
||||||
dec_lock_ref=lambda node_id, params: events.append(("unlock", node_id)),
|
dec_lock_ref=lambda node_id, params: events.append(("unlock", node_id)),
|
||||||
resolve_node_handle=lambda node_id: node,
|
|
||||||
)
|
)
|
||||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||||
wrapper._queue_load("rid", node.id, [object()])
|
wrapper._queue_load("rid", node.id, [object()])
|
||||||
|
|||||||
@@ -8117,6 +8117,7 @@ class _InsertWalkSuite(CustomTestCase):
|
|||||||
|
|
||||||
_rid = 0
|
_rid = 0
|
||||||
_make_req = UnifiedRadixCacheSuite._make_req
|
_make_req = UnifiedRadixCacheSuite._make_req
|
||||||
|
_apply_match_to_req = UnifiedRadixCacheSuite._apply_match_to_req
|
||||||
_alloc = UnifiedRadixCacheSuite._alloc
|
_alloc = UnifiedRadixCacheSuite._alloc
|
||||||
_insert = UnifiedRadixCacheSuite._insert
|
_insert = UnifiedRadixCacheSuite._insert
|
||||||
_init_hicache = UnifiedRadixCacheSuite._init_hicache
|
_init_hicache = UnifiedRadixCacheSuite._init_hicache
|
||||||
|
|||||||
@@ -88,6 +88,11 @@ class UnifiedTreeCoreInspectionInterface(UnifiedTreeCoreInterface):
|
|||||||
"""The node's pending write-through id, if any."""
|
"""The node's pending write-through id, if any."""
|
||||||
...
|
...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def is_external_cache_stored(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node is known to be stored in the external cache."""
|
||||||
|
...
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def is_node_in_device_lru(
|
def is_node_in_device_lru(
|
||||||
self, node_id: NodeId, component_type: ComponentType
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
|
|||||||
@@ -81,6 +81,10 @@ class UnifiedTreeCoreInspector(UnifiedTreeCore, UnifiedTreeCoreInspectionInterfa
|
|||||||
"""The node's pending write-through id, if any."""
|
"""The node's pending write-through id, if any."""
|
||||||
return self.node_by_id(node_id).write_through_pending_id
|
return self.node_by_id(node_id).write_through_pending_id
|
||||||
|
|
||||||
|
def is_external_cache_stored(self, node_id: NodeId) -> bool:
|
||||||
|
"""Whether the node is known to be stored in the external cache."""
|
||||||
|
return self.node_by_id(node_id).external_cache_stored
|
||||||
|
|
||||||
def is_node_in_device_lru(
|
def is_node_in_device_lru(
|
||||||
self, node_id: NodeId, component_type: ComponentType
|
self, node_id: NodeId, component_type: ComponentType
|
||||||
) -> bool:
|
) -> bool:
|
||||||
|
|||||||
Reference in New Issue
Block a user