[HiCache] Replace skip_lock_node_ids with a segment lock protocol (#36848)
This commit is contained in:
@@ -317,6 +317,8 @@ class DecodeRequest:
|
||||
prefix_match: Optional[DecodePrefixMatch] = None
|
||||
hicache_restored_kv_indices: Optional[torch.Tensor] = None
|
||||
hicache_restored_node: Any = None
|
||||
# Receipt for the inc_lock_ref held on hicache_restored_node.
|
||||
hicache_restore_lock_receipt: Optional[DecLockRefParams] = None
|
||||
hicache_load_consumer_index: int = -1
|
||||
hicache_restore_status: HiCacheRestoreResult = HiCacheRestoreResult.PENDING
|
||||
|
||||
@@ -426,12 +428,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
)
|
||||
|
||||
def _release_matched_prefix_lock(self, req: Req) -> None:
|
||||
params = DecLockRefParams(swa_uuid_for_lock=req.swa_uuid_for_lock)
|
||||
if req.swa_prefix_lock_released:
|
||||
self.tree_cache.dec_lock_ref(req.last_node, params, skip_swa=True)
|
||||
self.tree_cache.dec_lock_ref(req.last_node, req.lock_receipt, skip_swa=True)
|
||||
req.swa_prefix_lock_released = False
|
||||
else:
|
||||
self.tree_cache.dec_lock_ref(req.last_node, params)
|
||||
self.tree_cache.dec_lock_ref(req.last_node, req.lock_receipt)
|
||||
|
||||
def _reclaim_swa_tail_capacity(
|
||||
self, swa_tail_len: int, req_id: str
|
||||
@@ -676,9 +677,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
include_req=True,
|
||||
)
|
||||
# Keep aggregated scheduling semantics while preserving the SWA lock
|
||||
# boundary needed for the matching dec_lock_ref.
|
||||
lock_result = self.tree_cache.inc_lock_ref(result.last_device_node)
|
||||
req.swa_uuid_for_lock = lock_result.swa_uuid_for_lock
|
||||
# boundary needed for the matching dec_lock_ref; the full receipt
|
||||
# travels on the req so every later release mirrors this acquire.
|
||||
req.lock_receipt = self.tree_cache.inc_lock_ref(
|
||||
result.last_device_node
|
||||
).to_dec_params()
|
||||
return self._build_decode_prefix_match(req, result)
|
||||
|
||||
def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]:
|
||||
@@ -1239,8 +1242,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
and hasattr(self.tree_cache, "dec_swa_lock_only")
|
||||
):
|
||||
self.tree_cache.dec_swa_lock_only(
|
||||
decode_req.req.last_node,
|
||||
decode_req.req.swa_uuid_for_lock,
|
||||
decode_req.req.last_node, decode_req.req.lock_receipt
|
||||
)
|
||||
decode_req.req.swa_prefix_lock_released = True
|
||||
|
||||
|
||||
@@ -11,7 +11,9 @@ import torch
|
||||
|
||||
from sglang.srt.disaggregation.base import KVPoll
|
||||
from sglang.srt.managers.schedule_policy import match_prefix_for_req
|
||||
from sglang.srt.mem_cache.base_prefix_cache import InitLoadBackParams
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
InitLoadBackParams,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.disaggregation.decode import DecodeRequest
|
||||
@@ -184,8 +186,12 @@ class DecodeHiCacheTransferMixin:
|
||||
):
|
||||
self.tree_cache.release_aborted_request(decode_req.req.rid)
|
||||
if decode_req.hicache_restored_node is not None:
|
||||
self.tree_cache.dec_lock_ref(decode_req.hicache_restored_node)
|
||||
self.tree_cache.dec_lock_ref(
|
||||
decode_req.hicache_restored_node,
|
||||
decode_req.hicache_restore_lock_receipt,
|
||||
)
|
||||
decode_req.hicache_restored_node = None
|
||||
decode_req.hicache_restore_lock_receipt = None
|
||||
|
||||
def _try_hicache_queue_load_back(self, dr: DecodeRequest) -> bool:
|
||||
"""Queue one L2->L1 load_back op for ``dr``; True iff a DMA was queued.
|
||||
@@ -218,6 +224,12 @@ class DecodeHiCacheTransferMixin:
|
||||
req=dr.req,
|
||||
)
|
||||
)
|
||||
# The rematch repointed req.last_node to feed init_load_back's device
|
||||
# boundary, but the prealloc lock and the receipt on the req still
|
||||
# belong to pm.last_device_node; restore the pairing so any release
|
||||
# before the commit hands over the restored lock hits the right node
|
||||
# (the receipt's anchor makes a mispaired release assert).
|
||||
dr.req.last_node = pm.last_device_node
|
||||
# Failback: total coverage < required prefix means device alloc likely failed.
|
||||
if len(rematch.device_indices) + len(new_indices) < pm.decode_prefix_len:
|
||||
logger.warning(
|
||||
@@ -238,7 +250,9 @@ class DecodeHiCacheTransferMixin:
|
||||
[rematch.device_indices[pm.l1_prefix_len :], new_indices]
|
||||
)
|
||||
dr.hicache_restored_node = restored_node
|
||||
self.tree_cache.inc_lock_ref(restored_node)
|
||||
dr.hicache_restore_lock_receipt = self.tree_cache.inc_lock_ref(
|
||||
restored_node
|
||||
).to_dec_params()
|
||||
|
||||
if len(new_indices) == 0:
|
||||
# Whole prefix already on device; no DMA needed.
|
||||
@@ -303,7 +317,17 @@ class DecodeHiCacheTransferMixin:
|
||||
if prefix_match is None or not prefix_match.needs_local_restore:
|
||||
return
|
||||
|
||||
self.tree_cache.dec_lock_ref(prefix_match.last_device_node)
|
||||
req = decode_req.req
|
||||
restored_node = decode_req.hicache_restored_node
|
||||
restored_lock_receipt = decode_req.hicache_restore_lock_receipt
|
||||
assert restored_node is not None
|
||||
assert restored_lock_receipt is not None
|
||||
# Release preallocation before installing the restored lock receipt.
|
||||
self.tree_cache.dec_lock_ref(
|
||||
prefix_match.last_device_node,
|
||||
req.lock_receipt,
|
||||
skip_swa=req.swa_prefix_lock_released,
|
||||
)
|
||||
|
||||
self.tree_cache.req_to_token_pool.write(
|
||||
(
|
||||
@@ -312,7 +336,12 @@ class DecodeHiCacheTransferMixin:
|
||||
),
|
||||
decode_req.hicache_restored_kv_indices,
|
||||
)
|
||||
decode_req.req.prefix_indices = torch.cat(
|
||||
req.prefix_indices = torch.cat(
|
||||
[prefix_match.prefix_indices, decode_req.hicache_restored_kv_indices]
|
||||
)
|
||||
decode_req.req.last_node = decode_req.hicache_restored_node
|
||||
req.last_node = restored_node
|
||||
req.lock_receipt = restored_lock_receipt
|
||||
req.swa_prefix_lock_released = False
|
||||
# Prevent abort cleanup from releasing the transferred lock.
|
||||
decode_req.hicache_restored_node = None
|
||||
decode_req.hicache_restore_lock_receipt = None
|
||||
|
||||
@@ -107,6 +107,7 @@ from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
BasePrefixCache,
|
||||
DecLockRefParams,
|
||||
MatchPrefixParams,
|
||||
zero_match_result,
|
||||
)
|
||||
@@ -1125,13 +1126,11 @@ class Req(ReqDllmMixin):
|
||||
self.storage_prefetch_retry_pending = False
|
||||
self.storage_prefetch_retry_wait_polls = 0
|
||||
self.storage_prefetch_retry_attempts = 0
|
||||
# The node to lock until for swa radix tree lock ref
|
||||
self.swa_uuid_for_lock: Optional[int] = None
|
||||
# Receipt of the tree lock held on last_node (anchor, SWA boundary,
|
||||
# skipped components); every release replays it unchanged.
|
||||
self.lock_receipt: DecLockRefParams = DecLockRefParams()
|
||||
# Whether the prefill-time SWA tree lock has been released early
|
||||
self.swa_prefix_lock_released: bool = False
|
||||
# per-component nodes this req skipped locking (e.g. mamba on the decode
|
||||
# hold, already COW'd), so their dec releases only what it took.
|
||||
self.skip_lock_node_ids: dict = {}
|
||||
|
||||
# Whether or not if it is chunked. It increments whenever
|
||||
# it is chunked, and decrement whenever chunked request is
|
||||
@@ -1823,10 +1822,9 @@ class Req(ReqDllmMixin):
|
||||
self.last_node = None
|
||||
self.kv.cache_protected_len = 0
|
||||
self.num_matched_prefix_tokens = 0
|
||||
self.swa_uuid_for_lock = None
|
||||
self.lock_receipt = DecLockRefParams()
|
||||
self.swa_prefix_lock_released = False
|
||||
self.swa_branching_seqlen = None
|
||||
self.skip_lock_node_ids = {}
|
||||
self.extend_range = None
|
||||
self.dllm_initialized = False
|
||||
self.is_retracted = True
|
||||
@@ -3675,14 +3673,12 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
if (
|
||||
release_leaf_lock
|
||||
and not req.swa_prefix_lock_released
|
||||
and req.swa_uuid_for_lock is not None
|
||||
and req.lock_receipt.swa_uuid_for_lock is not None
|
||||
and req.last_node is not None
|
||||
and req.decode_batch_idx >= sliding_window_size
|
||||
):
|
||||
self.tree_cache.dec_swa_lock_only(
|
||||
req.last_node,
|
||||
req.swa_uuid_for_lock,
|
||||
skip_lock_node_ids=req.skip_lock_node_ids,
|
||||
req.last_node, req.lock_receipt
|
||||
)
|
||||
req.swa_prefix_lock_released = True
|
||||
elif self.forward_mode.is_extend() and self.tree_cache.is_chunk_cache():
|
||||
|
||||
@@ -1016,12 +1016,8 @@ class PrefillAdder:
|
||||
self._account_prefill_cache_admission(req, prefix_len)
|
||||
|
||||
def _req_inc_lock_ref(self, req: Req):
|
||||
result = self.tree_cache.inc_lock_ref(req.last_node)
|
||||
if self.is_hybrid_swa:
|
||||
req.swa_uuid_for_lock = result.swa_uuid_for_lock
|
||||
# match locks this node's components, so clear any stale skip set
|
||||
# carried from a previous scheduling of this req.
|
||||
req.skip_lock_node_ids = {}
|
||||
# Persist the release receipt.
|
||||
req.lock_receipt = self.tree_cache.inc_lock_ref(req.last_node).to_dec_params()
|
||||
|
||||
def add_dllm_staging_req(self, req: Req):
|
||||
assert self.dllm_config is not None
|
||||
@@ -1121,9 +1117,8 @@ class PrefillAdder:
|
||||
try:
|
||||
result = self.tree_cache.inc_lock_ref(last_node)
|
||||
if self.tree_cache.is_tree_cache():
|
||||
# init_load_back may revive SWA/Mamba tombstones while this
|
||||
# temporary admission lock is held. Release must mirror the
|
||||
# exact nodes skipped at acquire time.
|
||||
# Replay the acquire's receipt (SWA boundary uuid, mamba flag)
|
||||
# so release takes back exactly what this temporary lock took.
|
||||
dec_lock_params = result.to_dec_params()
|
||||
yield None
|
||||
finally:
|
||||
|
||||
@@ -174,8 +174,9 @@ class DynamicChunkSizer:
|
||||
# Walk the same match -> lock -> alloc lifecycle as a scheduled
|
||||
# request so release_kv_cache can release it symmetrically.
|
||||
req.init_next_round_input(self.tree_cache)
|
||||
lock = self.tree_cache.inc_lock_ref(req.last_node)
|
||||
req.swa_uuid_for_lock = lock.swa_uuid_for_lock
|
||||
req.lock_receipt = self.tree_cache.inc_lock_ref(
|
||||
req.last_node
|
||||
).to_dec_params()
|
||||
req.set_extend_range(
|
||||
len(req.prefix_indices), len(req.full_untruncated_fill_ids)
|
||||
)
|
||||
|
||||
@@ -131,39 +131,44 @@ class EvictResult:
|
||||
|
||||
@dataclasses.dataclass
|
||||
class IncLockRefResult:
|
||||
"""Result of an inc_lock_ref operation."""
|
||||
"""Receipt returned by ``inc_lock_ref``.
|
||||
|
||||
``node_id`` is the anchor the lock was taken on; a release replays the
|
||||
receipt on that node only. The SWA UUID marks the segment boundary;
|
||||
``None`` means root. ``skipped_lock_components`` records the components
|
||||
the acquire left untaken, so the release leaves them untouched.
|
||||
"""
|
||||
|
||||
delta: Optional[int] = None
|
||||
node_id: Optional[int] = None
|
||||
swa_uuid_for_lock: Optional[int] = None
|
||||
swa_uuid_for_host_lock: Optional[int] = None
|
||||
# Component nodes that were tombstones at acquire time. Replaying this set
|
||||
# at release prevents a short-lived lock from consuming a later load-back or
|
||||
# request lock after that tombstone becomes a valid device value.
|
||||
skip_lock_node_ids: dict[ComponentType, set[int]] = dataclasses.field(
|
||||
default_factory=dict
|
||||
)
|
||||
skipped_lock_components: tuple[ComponentType, ...] = ()
|
||||
|
||||
def to_dec_params(self) -> DecLockRefParams:
|
||||
"""Convert to the corresponding DecLockRefParams for dec_lock_ref."""
|
||||
return DecLockRefParams(
|
||||
node_id=self.node_id,
|
||||
swa_uuid_for_lock=self.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock=self.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids={
|
||||
component_type: set(node_ids)
|
||||
for component_type, node_ids in self.skip_lock_node_ids.items()
|
||||
},
|
||||
skipped_lock_components=tuple(self.skipped_lock_components),
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class DecLockRefParams:
|
||||
"""Parameters for dec_lock_ref operation."""
|
||||
"""Receipt required by unified-tree ``dec_lock_ref``.
|
||||
|
||||
Fields default to nothing-acquired, so a lost receipt under-releases (a
|
||||
leak the sanity checks report) instead of releasing another holder's
|
||||
lock. ``node_id`` is ``None`` only for receipts that never came from a
|
||||
unified-tree acquire (legacy caches, session sentinels).
|
||||
"""
|
||||
|
||||
node_id: Optional[int] = None
|
||||
swa_uuid_for_lock: Optional[int] = None
|
||||
swa_uuid_for_host_lock: Optional[int] = None
|
||||
skip_lock_node_ids: dict[ComponentType, set[int]] = dataclasses.field(
|
||||
default_factory=dict
|
||||
)
|
||||
skipped_lock_components: tuple[ComponentType, ...] = ()
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
|
||||
@@ -152,9 +152,23 @@ def _cache_actions_from_tagged(actions: Sequence[tuple]) -> list[CacheAction]:
|
||||
def _inc_lock_ref_result_from_binding(result) -> IncLockRefResult:
|
||||
return IncLockRefResult(
|
||||
delta=result.delta,
|
||||
node_id=result.node_id,
|
||||
swa_uuid_for_lock=result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock=result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids=_skip_lock_node_ids_from_binding(result.skip_lock_node_ids),
|
||||
skipped_lock_components=tuple(
|
||||
ComponentType(ct) for ct in result.skipped_lock_components
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _dec_lock_ref_params_to_binding(bindings_module, params: DecLockRefParams):
|
||||
"""Build the binding's params from the module that owns the core's binding
|
||||
(the inspection build is a distinct extension module with its own types)."""
|
||||
return bindings_module.DecLockRefParamsBinding(
|
||||
node_id=params.node_id,
|
||||
swa_uuid_for_lock=params.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock=params.swa_uuid_for_host_lock,
|
||||
skipped_lock_components=[int(ct) for ct in params.skipped_lock_components],
|
||||
)
|
||||
|
||||
|
||||
@@ -245,26 +259,6 @@ def _match_result_from_binding(result) -> MatchResult:
|
||||
)
|
||||
|
||||
|
||||
def _skip_lock_node_ids_from_binding(
|
||||
skip_lock_node_ids: dict[int, set[int]],
|
||||
) -> dict[ComponentType, set[int]]:
|
||||
"""Rekey the binding's component-value skip map by ComponentType."""
|
||||
return {
|
||||
ComponentType(component): set(node_ids)
|
||||
for component, node_ids in skip_lock_node_ids.items()
|
||||
}
|
||||
|
||||
|
||||
def _skip_lock_node_ids_to_binding(
|
||||
skip_lock_node_ids: dict[ComponentType, set[int]],
|
||||
) -> dict[int, set[int]]:
|
||||
"""Rekey a ComponentType skip map by the binding's component values."""
|
||||
return {
|
||||
int(component): set(node_ids)
|
||||
for component, node_ids in skip_lock_node_ids.items()
|
||||
}
|
||||
|
||||
|
||||
def _tracker_to_binding(tracker: dict[ComponentType, int]) -> dict[int, int]:
|
||||
"""Rekey a ComponentType tracker by the binding's component values."""
|
||||
return {int(component): freed for component, freed in tracker.items()}
|
||||
@@ -333,6 +327,12 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
raise ValueError(
|
||||
"Rust TreeCore does not support --radix-eviction-policy-config"
|
||||
)
|
||||
if ComponentType.SWA in self.tree_components and (
|
||||
params.sliding_window_size is None or params.sliding_window_size <= 0
|
||||
):
|
||||
raise ValueError(
|
||||
"the SWA tree component requires a positive sliding_window_size"
|
||||
)
|
||||
|
||||
self._page_size = params.page_size
|
||||
self.is_eagle = (
|
||||
@@ -415,45 +415,29 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
skip_lock_components: Sequence[ComponentType] = (),
|
||||
) -> IncLockRefResult:
|
||||
result = self._binding.inc_lock_ref(
|
||||
node_id, [int(component) for component in skip_lock_components]
|
||||
node_id, [int(ct) for ct in skip_lock_components]
|
||||
)
|
||||
return _inc_lock_ref_result_from_binding(result)
|
||||
|
||||
def dec_lock_ref(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
params: Optional[DecLockRefParams] = None,
|
||||
params: DecLockRefParams,
|
||||
skip_swa: bool = False,
|
||||
) -> DecLockRefResult:
|
||||
binding_params = (
|
||||
self._bindings.DecLockRefParamsBinding(
|
||||
swa_uuid_for_lock=params.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock=params.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids=_skip_lock_node_ids_to_binding(
|
||||
params.skip_lock_node_ids
|
||||
),
|
||||
)
|
||||
if params is not None
|
||||
else None
|
||||
self._binding.dec_lock_ref(
|
||||
node_id, _dec_lock_ref_params_to_binding(self._bindings, params), skip_swa
|
||||
)
|
||||
self._binding.dec_lock_ref(node_id, binding_params, skip_swa)
|
||||
return DecLockRefResult()
|
||||
|
||||
def dec_swa_lock_only(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
swa_uuid_for_lock: Optional[int],
|
||||
skip_lock_node_ids: Optional[dict] = None,
|
||||
params: DecLockRefParams,
|
||||
) -> DecSwaLockOnlyResult:
|
||||
result = DecSwaLockOnlyResult()
|
||||
new_device_frees, new_host_frees = self._binding.dec_swa_lock_only(
|
||||
node_id,
|
||||
swa_uuid_for_lock,
|
||||
(
|
||||
_skip_lock_node_ids_to_binding(skip_lock_node_ids)
|
||||
if skip_lock_node_ids
|
||||
else None
|
||||
),
|
||||
node_id, _dec_lock_ref_params_to_binding(self._bindings, params)
|
||||
)
|
||||
for component, tensors in new_device_frees.items():
|
||||
result.device_frees[ComponentType(component)].extend(tensors)
|
||||
@@ -503,30 +487,14 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
|
||||
def inc_host_lock_ref(self, node_id: NodeId) -> IncLockRefResult:
|
||||
result = self._binding.inc_host_lock_ref(node_id)
|
||||
return IncLockRefResult(
|
||||
delta=result.delta,
|
||||
swa_uuid_for_lock=result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock=result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids=_skip_lock_node_ids_from_binding(
|
||||
result.skip_lock_node_ids
|
||||
),
|
||||
)
|
||||
return _inc_lock_ref_result_from_binding(result)
|
||||
|
||||
def dec_host_lock_ref(
|
||||
self, node_id: NodeId, params: Optional[DecLockRefParams] = None
|
||||
self, node_id: NodeId, params: DecLockRefParams
|
||||
) -> DecLockRefResult:
|
||||
binding_params = (
|
||||
self._bindings.DecLockRefParamsBinding(
|
||||
swa_uuid_for_lock=params.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock=params.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids=_skip_lock_node_ids_to_binding(
|
||||
params.skip_lock_node_ids
|
||||
),
|
||||
)
|
||||
if params is not None
|
||||
else None
|
||||
self._binding.dec_host_lock_ref(
|
||||
node_id, _dec_lock_ref_params_to_binding(self._bindings, params)
|
||||
)
|
||||
self._binding.dec_host_lock_ref(node_id, binding_params)
|
||||
return DecLockRefResult()
|
||||
|
||||
def evictable_size(self) -> int:
|
||||
|
||||
@@ -501,9 +501,7 @@ class SWARadixCache(BasePrefixCache):
|
||||
|
||||
# Remove req slot release the cache lock
|
||||
self.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=req.swa_uuid_for_lock),
|
||||
skip_swa=req.swa_prefix_lock_released,
|
||||
req.last_node, req.lock_receipt, skip_swa=req.swa_prefix_lock_released
|
||||
)
|
||||
req.swa_prefix_lock_released = False
|
||||
|
||||
@@ -563,13 +561,10 @@ class SWARadixCache(BasePrefixCache):
|
||||
req.kv.cache_protected_len = len(new_indices)
|
||||
|
||||
self.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=req.swa_uuid_for_lock),
|
||||
skip_swa=req.swa_prefix_lock_released,
|
||||
req.last_node, req.lock_receipt, skip_swa=req.swa_prefix_lock_released
|
||||
)
|
||||
req.swa_prefix_lock_released = False
|
||||
result = self.inc_lock_ref(new_last_node)
|
||||
swa_uuid_for_lock = result.swa_uuid_for_lock
|
||||
lock_receipt = self.inc_lock_ref(new_last_node).to_dec_params()
|
||||
|
||||
# `req.prefix_indices` will be used in `PrefillAdder::add_chunked_req` later
|
||||
if len(new_indices) < len(kv_indices):
|
||||
@@ -579,7 +574,7 @@ class SWARadixCache(BasePrefixCache):
|
||||
else:
|
||||
req.prefix_indices = new_indices
|
||||
req.last_node = new_last_node
|
||||
req.swa_uuid_for_lock = swa_uuid_for_lock
|
||||
req.lock_receipt = lock_receipt
|
||||
|
||||
def pretty_print(self) -> None:
|
||||
self._print_helper(self.root_node, 0)
|
||||
@@ -806,13 +801,14 @@ class SWARadixCache(BasePrefixCache):
|
||||
def dec_swa_lock_only(
|
||||
self,
|
||||
node: TreeNode,
|
||||
swa_uuid_for_lock: Optional[int] = None,
|
||||
skip_lock_node_ids: Optional[dict] = None, # unused, signature parity only
|
||||
params: DecLockRefParams,
|
||||
):
|
||||
"""
|
||||
Decrement only the swa_lock_ref (and swa_protected_size_) along the chain
|
||||
[node, swa_uuid_for_lock], inclusive. The full_lock_ref is left untouched
|
||||
so the caller's full-cache protection is preserved.
|
||||
[node, receipt boundary uuid], inclusive. The full_lock_ref is left
|
||||
untouched so the caller's full-cache protection is preserved. Of the
|
||||
receipt this cache consumes only ``swa_uuid_for_lock``; it has no
|
||||
lower-priority components to drop.
|
||||
|
||||
Used to early-release the SWA portion of a request's tree lock once the
|
||||
request's decode position has advanced past the sliding window, so the
|
||||
@@ -826,13 +822,14 @@ class SWARadixCache(BasePrefixCache):
|
||||
as `swa_tombstone=True`. The full kv stays alive until the full-side
|
||||
lock drops; future prefix-matches stop before this tombstoned leaf.
|
||||
|
||||
Caller must ensure this is invoked at most once per (node, swa_uuid_for_lock)
|
||||
Caller must ensure this is invoked at most once per (node, boundary uuid)
|
||||
pair (track via e.g. `Req.swa_prefix_lock_released`). When the request
|
||||
finally releases its full lock via `dec_lock_ref`, pass `skip_swa=True`
|
||||
to avoid touching SWA state again.
|
||||
"""
|
||||
if self.disable:
|
||||
return
|
||||
swa_uuid_for_lock = params.swa_uuid_for_lock
|
||||
|
||||
while node != self.root_node:
|
||||
assert not node.swa_tombstone, (
|
||||
|
||||
@@ -180,34 +180,51 @@ Lock a node to protect it (and its ancestors) from eviction.
|
||||
| Aspect | Detail |
|
||||
|--------|--------|
|
||||
| **Purpose** | Called when a request begins using a cached prefix — prevents eviction of nodes it depends on |
|
||||
| **Inputs** | `node` — the last matched node (deepest) |
|
||||
| **Output** | `IncLockRefResult(swa_uuid_for_lock)` |
|
||||
| **Mutation** | Increments `lock_ref` per component along the path; moves tokens from evictable to protected size counters |
|
||||
| **Inputs** | `node` — the last matched node (deepest); `skip_lock_components` names components to leave untaken (the decode hold passes `(MAMBA,)`) |
|
||||
| **Output** | `IncLockRefResult(node_id, swa_uuid_for_lock, skipped_lock_components)` — the receipt the matching release must replay: the anchor node, the SWA boundary, and the skipped set |
|
||||
| **Mutation** | Increments `lock_ref` per component along its contiguous segment; moves data-bearing tokens from evictable to protected size counters |
|
||||
| **Complexity** | **O(D)** — Full: node to root; SWA: up to window boundary O(min(D, W)); Mamba: O(1).|
|
||||
|
||||
**Algorithm detail:** Calls `acquire_component_lock()` for each component.
|
||||
**Algorithm detail:** Calls `acquire_component_lock()` for each component. A lock
|
||||
covers a contiguous node segment and counts **every** node in it — tombstones
|
||||
included (they carry no tokens, so sizes only move for data-bearing nodes).
|
||||
|
||||
| Component | Strategy |
|
||||
|-----------|----------|
|
||||
| Full | **Path-lock**: walks from node to root, `lock_ref += 1` on every ancestor. On first lock (`lock_ref: 0→1`), moves tokens from `component_evictable_size_` to `component_protected_size_`. |
|
||||
| SWA | **Window-lock**: walks upward, accumulating SWA value lengths until `sliding_window_size` is filled. Records a `component_uuid` at the boundary node for `dec_lock_ref` to know where to stop. |
|
||||
| Mamba | **Single-node lock**: only `lock_ref += 1` on the node itself (mamba state is per-leaf, not per-path). |
|
||||
| SWA | **Segment-lock**: walks upward, `lock_ref += 1` on every node (tombstones included), accumulating position coverage (`len(key)`) until `sliding_window_size` is filled. Always stamps a boundary `component_uuid` at the last locked node; a `None` uuid in the receipt means the walk reached the root. |
|
||||
| Mamba | **Single-node lock**: only `lock_ref += 1` on the node itself (mamba state is per-leaf, not per-path). Taken unless the acquire lists it in `skip_lock_components`; the receipt records the skipped set. The core names no component: it drives whatever the tree registered through the same interface. |
|
||||
|
||||
---
|
||||
|
||||
### `dec_lock_ref(node, params?) → DecLockRefResult`
|
||||
### `dec_lock_ref(node, params, skip_swa=False) → DecLockRefResult`
|
||||
|
||||
Unlock a previously locked node path.
|
||||
Unlock a previously locked node path by replaying the acquire's receipt.
|
||||
|
||||
| Aspect | Detail |
|
||||
|--------|--------|
|
||||
| **Purpose** | Called when a request finishes — releases eviction protection |
|
||||
| **Inputs** | `node`, optional `params.swa_uuid_for_lock` for SWA boundary detection |
|
||||
| **Inputs** | `node`; required `params` receipt (`node_id` anchor, `swa_uuid_for_lock` boundary, `skipped_lock_components`); `skip_swa=True` after an earlier `dec_swa_lock_only`. A receipt whose anchor is not `node` is a protocol violation (assert): a mispaired release would otherwise walk another holder's segment. |
|
||||
| **Output** | `DecLockRefResult()` |
|
||||
| **Mutation** | Decrements `lock_ref` per component; moves tokens from protected back to evictable when `lock_ref` reaches 0 |
|
||||
| **Mutation** | Decrements `lock_ref` per component along the same segment the acquire counted; moves tokens from protected back to evictable when `lock_ref` reaches 0 |
|
||||
| **Complexity** | **O(D)** — symmetric to `inc_lock_ref` |
|
||||
|
||||
**Algorithm detail:** Calls `release_component_lock()` for each component. Full walks to root; SWA walks up until matching `component_uuid`; Mamba decrements single node.
|
||||
**Algorithm detail:** Releases auxiliary components before Full; every walk
|
||||
refreshes the evictable-leaf membership of each node whose last lock it drops,
|
||||
so the order is not load-bearing for the leaf sets. Full walks to root; SWA
|
||||
stops at the receipt boundary; components in `skipped_lock_components` are
|
||||
left alone. `skip_swa=True` also skips lower-priority components already
|
||||
released by `dec_swa_lock_only`. The host-side `dec_host_lock_ref` takes the
|
||||
same required receipt.
|
||||
|
||||
---
|
||||
|
||||
### `dec_swa_lock_only(node, params) → DecSwaLockOnlyResult`
|
||||
|
||||
Early-release only the SWA portion of a lock (decode advanced past the
|
||||
window), plus strictly-lower-priority co-located locks (e.g. Mamba) the
|
||||
receipt proves were taken. The eventual full release must pass
|
||||
`skip_swa=True`. At most once per (node, boundary uuid) pair.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -281,9 +281,11 @@ class FullComponent(TreeComponent):
|
||||
root = self.tree_core.root_node
|
||||
cur = node
|
||||
|
||||
# Skip the bottom evicted segment
|
||||
# The bottom device-evicted segment is locked too (no ledger move —
|
||||
# nothing is on device); a load-back that materializes a value under
|
||||
# lock credits protected directly.
|
||||
while cur is not root and cur.component_data[ct].value is None:
|
||||
result.skip_lock_node_ids.setdefault(ct, set()).add(cur.id)
|
||||
cur.component_data[ct].lock_ref += 1
|
||||
cur = cur.parent
|
||||
|
||||
# Lock the device-on segment up to root
|
||||
@@ -307,7 +309,7 @@ class FullComponent(TreeComponent):
|
||||
def release_component_lock(
|
||||
self,
|
||||
node: UnifiedTreeNode,
|
||||
params: Optional[DecLockRefParams],
|
||||
params: DecLockRefParams,
|
||||
lock_host: bool = False,
|
||||
) -> None:
|
||||
ct = self.component_type
|
||||
@@ -315,7 +317,6 @@ class FullComponent(TreeComponent):
|
||||
cd = node.component_data[ct]
|
||||
if cd.host_lock_ref == 0:
|
||||
return
|
||||
# Mirror of `acquire`. write_back uses a pure counter.
|
||||
if cd.host_value is None and not self.tree_core.is_write_back:
|
||||
return
|
||||
cd.host_lock_ref -= 1
|
||||
@@ -323,17 +324,13 @@ class FullComponent(TreeComponent):
|
||||
return
|
||||
|
||||
root = self.tree_core.root_node
|
||||
skip_lock_node_ids = params.skip_lock_node_ids.get(ct, ()) if params else ()
|
||||
cur = node
|
||||
while cur != root:
|
||||
if cur.id in skip_lock_node_ids:
|
||||
cur = cur.parent
|
||||
continue
|
||||
cd = cur.component_data[ct]
|
||||
assert cd.value is not None
|
||||
assert cd.lock_ref > 0
|
||||
|
||||
if cd.lock_ref == 1:
|
||||
assert cd.lock_ref > 0, (
|
||||
f"FULL segment release hit lock_ref=0 on node {cur.id}"
|
||||
)
|
||||
if cd.lock_ref == 1 and cd.value is not None:
|
||||
key_len = len(cd.value)
|
||||
self.tree_core.component_evictable_size_[ct] += key_len
|
||||
self.tree_core.component_protected_size_[ct] -= key_len
|
||||
@@ -422,8 +419,12 @@ class FullComponent(TreeComponent):
|
||||
n_len = len(cd.host_value)
|
||||
cd.value = device_indices[offset : offset + n_len].clone()
|
||||
offset += n_len
|
||||
# Full uses leaf sets, not LRU
|
||||
self.tree_core.component_evictable_size_[ct] += n_len
|
||||
# Full uses leaf sets, not LRU. A value materialized under
|
||||
# lock is protected; the last release moves it to evictable.
|
||||
if cd.lock_ref > 0:
|
||||
self.tree_core.component_protected_size_[ct] += n_len
|
||||
else:
|
||||
self.tree_core.component_evictable_size_[ct] += n_len
|
||||
self.tree_core._update_evictable_leaf_sets(n)
|
||||
|
||||
self.tree_core._update_evictable_leaf_sets(node)
|
||||
|
||||
@@ -233,14 +233,8 @@ class MambaComponent(TreeComponent):
|
||||
self._emit_excess_path_states_eviction(node, cache_actions)
|
||||
return
|
||||
if node.component_data[self.component_type].value is None:
|
||||
node.component_data[self.component_type].value = params.mamba_value
|
||||
# move from host LRU to device LRU
|
||||
host_lru = self.tree_core.host_lru_lists[self.component_type]
|
||||
if host_lru.in_list(node):
|
||||
host_lru.remove_node(node)
|
||||
self.tree_core.lru_lists[self.component_type].insert_mru(node)
|
||||
self.tree_core.component_evictable_size_[self.component_type] += len(
|
||||
params.mamba_value
|
||||
self.tree_core.set_component_device_value(
|
||||
node.id, self.component_type, params.mamba_value
|
||||
)
|
||||
node.last_access_time = get_and_increase_time_counter()
|
||||
self._emit_excess_path_states_eviction(node, cache_actions)
|
||||
@@ -441,19 +435,17 @@ class MambaComponent(TreeComponent):
|
||||
return result
|
||||
cd = node.component_data[ct]
|
||||
value = cd.host_value if lock_host else cd.value
|
||||
# A node in skip_lock_node_ids was a tombstone when this lock was acquired.
|
||||
if value is None:
|
||||
result.skip_lock_node_ids.setdefault(ct, set()).add(node.id)
|
||||
return result
|
||||
|
||||
# Tombstones are counted too; ledger/LRU track only data-bearing
|
||||
# nodes (a value materialized under lock is credited to protected
|
||||
# at the materialization site).
|
||||
if lock_host:
|
||||
if cd.host_lock_ref == 0:
|
||||
if cd.host_lock_ref == 0 and value is not None:
|
||||
host_lru = self.tree_core.host_lru_lists[ct]
|
||||
if host_lru.in_list(node):
|
||||
host_lru.remove_node(node)
|
||||
cd.host_lock_ref += 1
|
||||
else:
|
||||
if cd.lock_ref == 0:
|
||||
if cd.lock_ref == 0 and value is not None:
|
||||
vlen = len(value)
|
||||
self.tree_core.component_evictable_size_[ct] -= vlen
|
||||
self.tree_core.component_protected_size_[ct] += vlen
|
||||
@@ -463,32 +455,36 @@ class MambaComponent(TreeComponent):
|
||||
def release_component_lock(
|
||||
self,
|
||||
node: UnifiedTreeNode,
|
||||
params: Optional[DecLockRefParams],
|
||||
params: DecLockRefParams,
|
||||
lock_host: bool = False,
|
||||
) -> None:
|
||||
ct = self.component_type
|
||||
if node is self.tree_core.root_node:
|
||||
return
|
||||
cd = node.component_data[ct]
|
||||
skip_lock_node_ids = params.skip_lock_node_ids.get(ct, ()) if params else ()
|
||||
if node.id in skip_lock_node_ids:
|
||||
return
|
||||
|
||||
value = cd.host_value if lock_host else cd.value
|
||||
if lock_host:
|
||||
assert cd.host_lock_ref > 0, (
|
||||
f"Mamba release hit host_lock_ref=0 on node {node.id}"
|
||||
)
|
||||
cd.host_lock_ref -= 1
|
||||
if cd.host_lock_ref == 0 and cd.value is None and cd.host_value is not None:
|
||||
host_lru = self.tree_core.host_lru_lists[ct]
|
||||
if not host_lru.in_list(node):
|
||||
host_lru.insert_mru(node)
|
||||
if cd.host_lock_ref == 0:
|
||||
if cd.value is None and cd.host_value is not None:
|
||||
host_lru = self.tree_core.host_lru_lists[ct]
|
||||
if not host_lru.in_list(node):
|
||||
host_lru.insert_mru(node)
|
||||
self.tree_core._update_evictable_leaf_sets(node)
|
||||
return
|
||||
|
||||
if cd.lock_ref > 0:
|
||||
if cd.lock_ref == 1:
|
||||
vlen = len(value)
|
||||
self.tree_core.component_evictable_size_[ct] += vlen
|
||||
self.tree_core.component_protected_size_[ct] -= vlen
|
||||
cd.lock_ref -= 1
|
||||
assert cd.lock_ref > 0, f"Mamba release hit lock_ref=0 on node {node.id}"
|
||||
if cd.lock_ref == 1 and value is not None:
|
||||
vlen = len(value)
|
||||
self.tree_core.component_evictable_size_[ct] += vlen
|
||||
self.tree_core.component_protected_size_[ct] -= vlen
|
||||
cd.lock_ref -= 1
|
||||
if cd.lock_ref == 0:
|
||||
self.tree_core._update_evictable_leaf_sets(node)
|
||||
|
||||
def _alloc_mamba_slot(self) -> torch.Tensor:
|
||||
"""Allocate one mamba pool slot, evicting if necessary."""
|
||||
@@ -809,15 +805,11 @@ class MambaComponent(TreeComponent):
|
||||
return
|
||||
transfer = transfers[0]
|
||||
if transfer.device_indices is not None:
|
||||
cd = node.component_data[ct]
|
||||
cd.value = transfer.device_indices.clone()
|
||||
count = len(cd.value)
|
||||
# Move from host LRU to device LRU
|
||||
host_lru = self.tree_core.host_lru_lists[ct]
|
||||
if host_lru.in_list(node):
|
||||
host_lru.remove_node(node)
|
||||
self.tree_core.lru_lists[ct].insert_mru(node)
|
||||
self.tree_core.component_evictable_size_[ct] += count
|
||||
# The materialization primitive owns the ledger/LRU moves,
|
||||
# including crediting protected when restored under lock.
|
||||
self.tree_core.set_component_device_value(
|
||||
node.id, ct, transfer.device_indices.clone()
|
||||
)
|
||||
|
||||
elif phase == CacheTransferPhase.PREFETCH:
|
||||
if not transfers:
|
||||
|
||||
@@ -74,6 +74,9 @@ class SWAComponent(TreeComponent):
|
||||
), (
|
||||
f"SWAComponent requires SWATokenToKVPoolAllocator, got {type(params.token_to_kv_pool_allocator)}"
|
||||
)
|
||||
if params.sliding_window_size is None or params.sliding_window_size <= 0:
|
||||
raise ValueError("SWAComponent requires a positive sliding_window_size")
|
||||
|
||||
super().__init__(cache, params)
|
||||
self._session_leaf_covered_len: dict[str, dict[UnifiedTreeNode, int]] = {}
|
||||
self.sliding_window_size = params.sliding_window_size
|
||||
@@ -388,9 +391,9 @@ class SWAComponent(TreeComponent):
|
||||
|
||||
full_cd = node.component_data[BASE_COMPONENT_TYPE]
|
||||
swa_evicted_seqlen = params.swa_evicted_seqlen
|
||||
assert node.component_data[self.component_type].lock_ref == 0, (
|
||||
f"tombstone {self.component_type} lock_ref should be 0, node {node.id}"
|
||||
)
|
||||
# A locked tombstone is legal (segment locks count every node); the
|
||||
# full-value swap below is safe because full lock_ref >= swa
|
||||
# lock_ref, so a locked-SWA node always takes the Recover branch.
|
||||
assert swa_evicted_seqlen % self.tree_core.page_size == 0, (
|
||||
f"{self.component_type}: swa_evicted_seqlen must be page-aligned, {swa_evicted_seqlen=}"
|
||||
)
|
||||
@@ -464,9 +467,6 @@ class SWAComponent(TreeComponent):
|
||||
ct = self.component_type
|
||||
if node.component_data[ct].value is not None:
|
||||
return
|
||||
assert node.component_data[ct].lock_ref == 0, (
|
||||
f"tombstone {ct} lock_ref should be 0 on unevict, node {node.id}"
|
||||
)
|
||||
swa_evicted_seqlen = params.swa_evicted_seqlen
|
||||
assert swa_evicted_seqlen % self.tree_core.page_size == 0, (
|
||||
f"{ct}: swa_evicted_seqlen must be page-aligned, {swa_evicted_seqlen=}"
|
||||
@@ -576,6 +576,9 @@ class SWAComponent(TreeComponent):
|
||||
new_parent.component_data[self.component_type].lock_ref = child.component_data[
|
||||
self.component_type
|
||||
].lock_ref
|
||||
new_parent.component_data[
|
||||
self.component_type
|
||||
].host_lock_ref = child.component_data[self.component_type].host_lock_ref
|
||||
new_parent.component_data[
|
||||
self.component_type
|
||||
].session_ref = child.component_data[self.component_type].session_ref
|
||||
@@ -610,6 +613,8 @@ class SWAComponent(TreeComponent):
|
||||
parent_swa_data.metadata["host_uuid"] = host_uuid
|
||||
|
||||
host_lru = self.tree_core.host_lru_lists[self.component_type]
|
||||
# Host-locked halves stay out of the host LRU: in-flight IO
|
||||
# holds them, and host acquire removed the node at 0->1.
|
||||
if (
|
||||
new_parent.component_data[self.component_type].value is None
|
||||
and parent_swa_data.host_lock_ref == 0
|
||||
@@ -622,11 +627,16 @@ class SWAComponent(TreeComponent):
|
||||
):
|
||||
host_lru.insert_mru(child)
|
||||
|
||||
# parent inherits the swa_uuid from child for swa lock ref
|
||||
# The window-boundary uuids mark the node's older edge, which the
|
||||
# split moves to the parent — both tiers migrate with it.
|
||||
new_parent.component_data[self.component_type].metadata["uuid"] = (
|
||||
child.component_data[self.component_type].metadata.get("uuid")
|
||||
)
|
||||
child.component_data[self.component_type].metadata.pop("uuid", None)
|
||||
new_parent.component_data[self.component_type].metadata["host_uuid"] = (
|
||||
child.component_data[self.component_type].metadata.get("host_uuid")
|
||||
)
|
||||
child.component_data[self.component_type].metadata.pop("host_uuid", None)
|
||||
|
||||
def evict_component(
|
||||
self,
|
||||
@@ -754,10 +764,20 @@ class SWAComponent(TreeComponent):
|
||||
result: IncLockRefResult,
|
||||
lock_host: bool = False,
|
||||
) -> IncLockRefResult:
|
||||
"""Lock the contiguous segment covering the trailing window.
|
||||
|
||||
Every node in [node, boundary] is counted, tombstones included, so
|
||||
the paired release decrements the same contiguous segment with no
|
||||
carried skip state. Coverage is position-based (len(cur.key)); the
|
||||
boundary node is always uuid-stamped, so a release without a uuid
|
||||
means the segment reached the root. Ledger/LRU transitions track
|
||||
only data-bearing nodes; a value materialized later under lock is
|
||||
credited to protected by set_component_device_value.
|
||||
"""
|
||||
ct = self.component_type
|
||||
root = self.tree_core.root_node
|
||||
sliding_window_size = self.sliding_window_size
|
||||
swa_lock_size = 0
|
||||
covered = 0
|
||||
swa_uuid = None
|
||||
uuid_key = "host_uuid" if lock_host else "uuid"
|
||||
lru = (
|
||||
@@ -766,33 +786,25 @@ class SWAComponent(TreeComponent):
|
||||
else self.tree_core.lru_lists[ct]
|
||||
)
|
||||
|
||||
# Tombstoned nodes (cd.value is None) have no SWA chunk to protect
|
||||
# skip them and keep walking up. This path is hit when HiCache
|
||||
# backs up a FULL present internal node whose SWA was already evicted.
|
||||
cur = node
|
||||
while cur != root and swa_lock_size < sliding_window_size:
|
||||
while cur != root and covered < sliding_window_size:
|
||||
comp = cur.component_data[ct]
|
||||
value = comp.host_value if lock_host else comp.value
|
||||
if value is None:
|
||||
result.skip_lock_node_ids.setdefault(ct, set()).add(cur.id)
|
||||
cur = cur.parent
|
||||
continue
|
||||
|
||||
ref = comp.host_lock_ref if lock_host else comp.lock_ref
|
||||
if ref == 0:
|
||||
if ref == 0 and value is not None:
|
||||
if lock_host:
|
||||
if lru.in_list(cur):
|
||||
lru.remove_node(cur)
|
||||
else:
|
||||
key_len = len(cur.key)
|
||||
key_len = len(value)
|
||||
self.tree_core.component_evictable_size_[ct] -= key_len
|
||||
self.tree_core.component_protected_size_[ct] += key_len
|
||||
if lock_host:
|
||||
comp.host_lock_ref = ref + 1
|
||||
else:
|
||||
comp.lock_ref = ref + 1
|
||||
swa_lock_size += len(value)
|
||||
if swa_lock_size >= sliding_window_size:
|
||||
covered += len(cur.key)
|
||||
if covered >= sliding_window_size:
|
||||
if comp.metadata.get(uuid_key) is None:
|
||||
comp.metadata[uuid_key] = next_component_uuid()
|
||||
swa_uuid = comp.metadata[uuid_key]
|
||||
@@ -807,45 +819,47 @@ class SWAComponent(TreeComponent):
|
||||
def release_component_lock(
|
||||
self,
|
||||
node: UnifiedTreeNode,
|
||||
params: Optional[DecLockRefParams],
|
||||
params: DecLockRefParams,
|
||||
lock_host: bool = False,
|
||||
) -> None:
|
||||
ct = self.component_type
|
||||
root = self.tree_core.root_node
|
||||
swa_uuid_for_lock = (
|
||||
(params.swa_uuid_for_host_lock if lock_host else params.swa_uuid_for_lock)
|
||||
if params
|
||||
else None
|
||||
params.swa_uuid_for_host_lock if lock_host else params.swa_uuid_for_lock
|
||||
)
|
||||
skip_lock_node_ids = params.skip_lock_node_ids.get(ct, ()) if params else ()
|
||||
dec_swa = True
|
||||
uuid_key = "host_uuid" if lock_host else "uuid"
|
||||
|
||||
# A node in skip_lock_node_ids was a tombstone when this lock was acquired.
|
||||
cur = node
|
||||
while cur != root and dec_swa:
|
||||
comp = cur.component_data[ct]
|
||||
if cur.id in skip_lock_node_ids:
|
||||
cur = cur.parent
|
||||
continue
|
||||
ref = comp.host_lock_ref if lock_host else comp.lock_ref
|
||||
if ref == 0:
|
||||
cur = cur.parent
|
||||
continue
|
||||
if ref == 1:
|
||||
# Acquire counted every segment node and splits copy refs, so a
|
||||
# zero here means the release does not mirror its acquire.
|
||||
assert ref > 0, (
|
||||
f"SWA segment release hit {'host_' if lock_host else ''}"
|
||||
f"lock_ref=0 on node {cur.id}"
|
||||
)
|
||||
value = comp.host_value if lock_host else comp.value
|
||||
if ref == 1 and value is not None:
|
||||
if lock_host:
|
||||
if comp.value is None and comp.host_value is not None:
|
||||
if comp.value is None:
|
||||
host_lru = self.tree_core.host_lru_lists[ct]
|
||||
if not host_lru.in_list(cur):
|
||||
host_lru.insert_mru(cur)
|
||||
else:
|
||||
key_len = len(comp.value)
|
||||
key_len = len(value)
|
||||
self.tree_core.component_evictable_size_[ct] += key_len
|
||||
self.tree_core.component_protected_size_[ct] -= key_len
|
||||
if lock_host:
|
||||
comp.host_lock_ref = ref - 1
|
||||
else:
|
||||
comp.lock_ref = ref - 1
|
||||
if ref == 1:
|
||||
# This may have been the last lock holding the node out of
|
||||
# the evictable-leaf sets; refresh it here rather than rely
|
||||
# on the Full walk running after this one.
|
||||
self.tree_core._update_evictable_leaf_sets(cur)
|
||||
if swa_uuid_for_lock and comp.metadata.get(uuid_key) == swa_uuid_for_lock:
|
||||
dec_swa = False
|
||||
cur = cur.parent
|
||||
@@ -857,15 +871,14 @@ class SWAComponent(TreeComponent):
|
||||
device_frees: dict[ComponentType, list[torch.Tensor]],
|
||||
host_frees: dict[ComponentType, list[torch.Tensor]],
|
||||
) -> None:
|
||||
"""Early-release the SWA lock along [node, swa_uuid_for_lock] while
|
||||
leaving Full and Mamba locks intact.
|
||||
"""Early-release the SWA lock along [node, swa_uuid_for_lock]; this
|
||||
method touches only SWA state. The wrapping ``dec_swa_lock_only`` also
|
||||
drops strictly-lower-priority co-located locks (e.g. Mamba) per the
|
||||
receipt; the Full lock stays so the request's prefix is protected.
|
||||
|
||||
Called when a request's decode position has advanced past the sliding
|
||||
window — the SWA portion of the tree lock is no longer needed but the
|
||||
Full lock must stay so the request's prefix is protected.
|
||||
|
||||
Caller (UnifiedRadixCache.dec_swa_lock_only) must ensure this is
|
||||
invoked at most once per (node, swa_uuid_for_lock) pair.
|
||||
window. The caller must invoke this at most once per
|
||||
(node, swa_uuid_for_lock) pair.
|
||||
"""
|
||||
ct = self.component_type
|
||||
root = self.tree_core.root_node
|
||||
@@ -873,19 +886,16 @@ class SWAComponent(TreeComponent):
|
||||
cur = node
|
||||
while cur is not root:
|
||||
cd = cur.component_data[ct]
|
||||
# Acquire skips tombstoned nodes; release must skip them too. Same
|
||||
# for nodes with lock_ref == 0 — acquire never credited them.
|
||||
if cd.value is None or cd.lock_ref == 0:
|
||||
if swa_uuid_for_lock and cd.metadata.get("uuid") == swa_uuid_for_lock:
|
||||
break
|
||||
cur = cur.parent
|
||||
continue
|
||||
|
||||
assert cd.lock_ref > 0, (
|
||||
f"SWA window release hit lock_ref=0 on node {cur.id}"
|
||||
)
|
||||
cd.lock_ref -= 1
|
||||
if cd.lock_ref == 0:
|
||||
key_len = len(cur.key)
|
||||
self.tree_core.component_protected_size_[ct] -= key_len
|
||||
self.tree_core.component_evictable_size_[ct] += key_len
|
||||
self.tree_core._update_evictable_leaf_sets(cur)
|
||||
if cd.lock_ref == 0 and cd.value is not None:
|
||||
value_len = len(cd.value)
|
||||
self.tree_core.component_protected_size_[ct] -= value_len
|
||||
self.tree_core.component_evictable_size_[ct] += value_len
|
||||
if self.tree_core._is_device_leaf(cur):
|
||||
self.tree_core._evict_component_and_detach_lru(
|
||||
cur,
|
||||
|
||||
@@ -584,7 +584,7 @@ class TreeComponent(ABC):
|
||||
def release_component_lock(
|
||||
self,
|
||||
node: UnifiedTreeNode,
|
||||
params: Optional[DecLockRefParams],
|
||||
params: DecLockRefParams,
|
||||
lock_host: bool = False,
|
||||
) -> None:
|
||||
"""Decrement component lock refs, un-protecting nodes.
|
||||
|
||||
@@ -622,32 +622,66 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
self, node_id: NodeId, skip_lock_components: Sequence[ComponentType] = ()
|
||||
) -> IncLockRefResult:
|
||||
node = self.node_by_id(node_id)
|
||||
result = IncLockRefResult()
|
||||
skipped = tuple(skip_lock_components)
|
||||
# The receipt records the anchor and what was locked; the paired dec
|
||||
# replays exactly that.
|
||||
result = IncLockRefResult(node_id=node.id, skipped_lock_components=skipped)
|
||||
for component in self.components:
|
||||
if component.component_type in skip_lock_components:
|
||||
# Leave this component's value evictable and record every
|
||||
# non-root node (incl tombstones) so the matching dec skips a
|
||||
# lock we never took, which may be another req's on a shared node.
|
||||
if node is not self.root_node:
|
||||
result.skip_lock_node_ids.setdefault(
|
||||
component.component_type, set()
|
||||
).add(node.id)
|
||||
if component.component_type in skipped:
|
||||
continue
|
||||
result = component.acquire_component_lock(node=node, result=result)
|
||||
self._update_evictable_leaf_sets(node)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _assert_receipt_anchor(node: UnifiedTreeNode, params: DecLockRefParams) -> None:
|
||||
"""A receipt releases only the node its acquire returned; a mispaired
|
||||
node would silently release (or steal) another holder's segment."""
|
||||
assert params.node_id is None or params.node_id == node.id, (
|
||||
f"lock receipt anchored on node {params.node_id} released on node {node.id}"
|
||||
)
|
||||
|
||||
def _release_components(
|
||||
self,
|
||||
node: UnifiedTreeNode,
|
||||
params: DecLockRefParams,
|
||||
*,
|
||||
lock_host: bool = False,
|
||||
skip_swa_and_below: bool = False,
|
||||
) -> None:
|
||||
"""Release each component this receipt acquired. Auxiliaries go first
|
||||
so Full, whose walk refreshes leaf membership on every node it
|
||||
unlocks, sees their final refs; the auxiliary walks also refresh the
|
||||
nodes they unlock, so the order is not load-bearing for the sets."""
|
||||
swa_priority = None
|
||||
if skip_swa_and_below:
|
||||
swa_component = self.components_by_type.get(ComponentType.SWA)
|
||||
if swa_component is not None:
|
||||
swa_priority = swa_component.eviction_priority(is_leaf=False)
|
||||
for component in reversed(self.components):
|
||||
ct = component.component_type
|
||||
if ct in params.skipped_lock_components:
|
||||
continue
|
||||
if swa_priority is not None and (
|
||||
ct == ComponentType.SWA
|
||||
or component.eviction_priority(is_leaf=False) < swa_priority
|
||||
):
|
||||
continue
|
||||
component.release_component_lock(
|
||||
node=node, params=params, lock_host=lock_host
|
||||
)
|
||||
|
||||
def dec_lock_ref(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
params: Optional[DecLockRefParams] = None,
|
||||
params: DecLockRefParams,
|
||||
skip_swa: bool = False,
|
||||
) -> DecLockRefResult:
|
||||
node = self.node_by_id(node_id)
|
||||
for component in self.components:
|
||||
if skip_swa and component.component_type == ComponentType.SWA:
|
||||
continue
|
||||
component.release_component_lock(node=node, params=params)
|
||||
self._assert_receipt_anchor(node, params)
|
||||
# After an SWA early release (dec_swa_lock_only), SWA and the
|
||||
# lower-priority components it dropped are already released.
|
||||
self._release_components(node, params, skip_swa_and_below=skip_swa)
|
||||
self._update_evictable_leaf_sets(node)
|
||||
# TODO: delta is not aggregated from components; no caller uses it yet.
|
||||
return DecLockRefResult()
|
||||
@@ -655,36 +689,33 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
def dec_swa_lock_only(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
swa_uuid_for_lock: Optional[int],
|
||||
skip_lock_node_ids: Optional[dict] = None,
|
||||
params: DecLockRefParams,
|
||||
) -> DecSwaLockOnlyResult:
|
||||
"""Early-release the SWA portion of a request's tree lock, plus any
|
||||
strictly-lower-priority locks (e.g. Mamba) co-located on the node."""
|
||||
result = DecSwaLockOnlyResult()
|
||||
node = self.node_by_id(node_id)
|
||||
self._assert_receipt_anchor(node, params)
|
||||
swa_component = self.components_by_type.get(ComponentType.SWA)
|
||||
if swa_component is None:
|
||||
return result
|
||||
swa_component.release_window_lock(
|
||||
node, swa_uuid_for_lock, result.device_frees, result.host_frees
|
||||
node, params.swa_uuid_for_lock, result.device_frees, result.host_frees
|
||||
)
|
||||
|
||||
# Drop strictly-lower-priority locks (e.g. Mamba) co-located on the node,
|
||||
# honoring skip ids so we don't drop a lock a partial inc never took
|
||||
# (matters for FULL+SWA+MAMBA models, e.g. Inkling).
|
||||
# Drop strictly-lower-priority locks co-located on the node, skipping
|
||||
# any the paired inc never took (matters for FULL+SWA+MAMBA models).
|
||||
swa_priority = swa_component.eviction_priority(is_leaf=False)
|
||||
dec_params = DecLockRefParams(
|
||||
swa_uuid_for_lock=swa_uuid_for_lock,
|
||||
skip_lock_node_ids=skip_lock_node_ids or {},
|
||||
)
|
||||
for comp in self.components:
|
||||
for comp in reversed(self.components):
|
||||
if comp.component_type in params.skipped_lock_components:
|
||||
continue
|
||||
if comp.eviction_priority(is_leaf=False) < swa_priority:
|
||||
comp.release_component_lock(node, dec_params)
|
||||
comp.release_component_lock(node, params)
|
||||
return result
|
||||
|
||||
def inc_host_lock_ref(self, node_id: NodeId) -> IncLockRefResult:
|
||||
node = self.node_by_id(node_id)
|
||||
result = IncLockRefResult()
|
||||
result = IncLockRefResult(node_id=node.id)
|
||||
for component in self.components:
|
||||
result = component.acquire_component_lock(
|
||||
node=node, result=result, lock_host=True
|
||||
@@ -693,11 +724,11 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
return result
|
||||
|
||||
def dec_host_lock_ref(
|
||||
self, node_id: NodeId, params: Optional[DecLockRefParams] = None
|
||||
self, node_id: NodeId, params: DecLockRefParams
|
||||
) -> DecLockRefResult:
|
||||
node = self.node_by_id(node_id)
|
||||
for component in self.components:
|
||||
component.release_component_lock(node=node, params=params, lock_host=True)
|
||||
self._assert_receipt_anchor(node, params)
|
||||
self._release_components(node, params, lock_host=True)
|
||||
self._update_evictable_leaf_sets(node)
|
||||
return DecLockRefResult()
|
||||
|
||||
@@ -1260,7 +1291,10 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
assert cd.value is None
|
||||
n = len(fresh_value)
|
||||
cd.value = fresh_value.clone()
|
||||
self.component_evictable_size_[ct] += n
|
||||
if cd.lock_ref > 0:
|
||||
self.component_protected_size_[ct] += n
|
||||
else:
|
||||
self.component_evictable_size_[ct] += n
|
||||
self._update_evictable_leaf_sets(node)
|
||||
# A backuped node restored from fresh KV is a duplicate right away.
|
||||
self._update_duplicate_tracking(node)
|
||||
@@ -1887,7 +1921,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
return True
|
||||
|
||||
def _is_host_leaf(self, node: UnifiedTreeNode) -> bool:
|
||||
"""H-leaf: evicted, Full host value present, no children, unlocked, not root.
|
||||
"""H-leaf: evicted, Full host value present, no children, unlocked on
|
||||
both tiers, not root.
|
||||
|
||||
Only the Full (base) component host_value is required; auxiliary
|
||||
components are not mandatory for H-leaf membership. In-flight DMA
|
||||
@@ -1898,6 +1933,10 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
return False
|
||||
if any(cd.host_lock_ref > 0 for cd in node.component_data):
|
||||
return False
|
||||
# Segment locks count evicted nodes too: a device-locked candidate is
|
||||
# a live segment's anchor, and _evict_host_leaf would delete it.
|
||||
if any(cd.lock_ref > 0 for cd in node.component_data):
|
||||
return False
|
||||
if len(node.children) > 0:
|
||||
return False
|
||||
return True
|
||||
@@ -2258,12 +2297,18 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
# Full uses leaf sets, not LRU; its stores go through the insert paths.
|
||||
assert component_type != BASE_COMPONENT_TYPE
|
||||
node = self.node_by_id(node_id)
|
||||
node.component_data[component_type].value = value
|
||||
cd = node.component_data[component_type]
|
||||
cd.value = value
|
||||
host_lru = self.host_lru_lists[component_type]
|
||||
if host_lru.in_list(node):
|
||||
host_lru.remove_node(node)
|
||||
self.lru_lists[component_type].insert_mru(node)
|
||||
self.component_evictable_size_[component_type] += len(value)
|
||||
# A value materialized under lock is protected; the last release
|
||||
# moves it to evictable.
|
||||
if cd.lock_ref > 0:
|
||||
self.component_protected_size_[component_type] += len(value)
|
||||
else:
|
||||
self.component_evictable_size_[component_type] += len(value)
|
||||
|
||||
def get_component_device_value(
|
||||
self, node_id: NodeId, component_type: ComponentType
|
||||
@@ -2361,8 +2406,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
E(f"node {nid} {ct} host_lock_ref={cd.host_lock_ref}")
|
||||
if ct != FCT and fl < cd.lock_ref:
|
||||
E(f"node {nid} full_lock={fl} < {ct}_lock={cd.lock_ref}")
|
||||
if cd.value is None and cd.lock_ref > 0:
|
||||
E(f"node {nid} {ct} evicted but lock_ref={cd.lock_ref}")
|
||||
# Locked tombstones are legal: segment locks count every
|
||||
# node in [start, boundary], data-bearing or not.
|
||||
|
||||
# Collect expected leaf qualification (single pass)
|
||||
if self._is_device_leaf(node):
|
||||
|
||||
@@ -237,26 +237,27 @@ class UnifiedTreeCoreInterface(ABC):
|
||||
def inc_lock_ref(
|
||||
self, node_id: NodeId, skip_lock_components: Sequence[ComponentType] = ()
|
||||
) -> IncLockRefResult:
|
||||
"""Bump the reference count on a node's component locks, leaving any
|
||||
component in skip_lock_components evictable and recorded in the result."""
|
||||
"""Bump the reference count on a node's component locks. Components in
|
||||
``skip_lock_components`` are left untaken; the receipt records the
|
||||
anchor node and the skipped set so the paired release mirrors them."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def dec_lock_ref(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
params: Optional[DecLockRefParams] = None,
|
||||
params: DecLockRefParams,
|
||||
skip_swa: bool = False,
|
||||
) -> DecLockRefResult:
|
||||
"""Decrease the reference count on a node's component locks."""
|
||||
"""Decrease the reference count on a node's component locks. The
|
||||
receipt is required: a release must replay its acquire's evidence."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def dec_swa_lock_only(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
swa_uuid_for_lock: Optional[int],
|
||||
skip_lock_node_ids: Optional[dict] = None,
|
||||
params: DecLockRefParams,
|
||||
) -> DecSwaLockOnlyResult:
|
||||
"""Decrease only the SWA (and lower-priority co-located) reference
|
||||
counts; the result carries the freed slots."""
|
||||
@@ -313,7 +314,7 @@ class UnifiedTreeCoreInterface(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def dec_host_lock_ref(
|
||||
self, node_id: NodeId, params: Optional[DecLockRefParams] = None
|
||||
self, node_id: NodeId, params: DecLockRefParams
|
||||
) -> DecLockRefResult:
|
||||
"""Decrease the reference count on a node's host-side component locks."""
|
||||
...
|
||||
|
||||
@@ -878,7 +878,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
def dec_lock_ref(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
params: Optional[DecLockRefParams] = None,
|
||||
params: DecLockRefParams,
|
||||
skip_swa: bool = False,
|
||||
) -> DecLockRefResult:
|
||||
result = self.session.try_dec_lock_ref(node_id, params)
|
||||
@@ -889,28 +889,18 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
return self.tree_core.dec_lock_ref(node_id, params, skip_swa)
|
||||
|
||||
def _dec_req_lock(self, req: Req, *, skip_swa: bool = False) -> None:
|
||||
"""Release the tree lock a request holds on its last_node, honoring the
|
||||
components it skipped locking so it never drops a lock it never took."""
|
||||
self.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=req.swa_uuid_for_lock,
|
||||
skip_lock_node_ids=req.skip_lock_node_ids,
|
||||
),
|
||||
skip_swa=skip_swa,
|
||||
)
|
||||
"""Release the tree lock a request holds on its last_node with the
|
||||
receipt its acquire returned, so it never drops a lock it never took."""
|
||||
self.dec_lock_ref(req.last_node, req.lock_receipt, skip_swa=skip_swa)
|
||||
|
||||
def dec_swa_lock_only(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
swa_uuid_for_lock: Optional[int] = None,
|
||||
skip_lock_node_ids: Optional[dict] = None,
|
||||
params: DecLockRefParams,
|
||||
) -> None:
|
||||
if self.disable:
|
||||
return
|
||||
result = self.tree_core.dec_swa_lock_only(
|
||||
node_id, swa_uuid_for_lock, skip_lock_node_ids
|
||||
)
|
||||
result = self.tree_core.dec_swa_lock_only(node_id, params)
|
||||
self._free_values(result.device_frees, result.host_frees)
|
||||
|
||||
def inc_host_lock_ref(self, node_id: NodeId) -> IncLockRefResult:
|
||||
@@ -919,7 +909,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
return self.tree_core.inc_host_lock_ref(node_id)
|
||||
|
||||
def dec_host_lock_ref(
|
||||
self, node_id: NodeId, params: Optional[DecLockRefParams] = None
|
||||
self, node_id: NodeId, params: DecLockRefParams
|
||||
) -> DecLockRefResult:
|
||||
if self.disable:
|
||||
return DecLockRefResult()
|
||||
@@ -1099,20 +1089,20 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
new_indices[req.kv.cache_protected_len :],
|
||||
)
|
||||
|
||||
self._dec_req_lock(req)
|
||||
self._dec_req_lock(req, skip_swa=req.swa_prefix_lock_released)
|
||||
# Opt-in: leave the matched-prefix mamba evictable during decode (it is
|
||||
# already COW'd to the request's own slot, never read from this node again).
|
||||
# Safe only because any future COW source is the COWing request's own
|
||||
# admission-locked last_node (recorded only if still present, locked before
|
||||
# the next alloc) -- not this evictable node. A scheduler that matched a
|
||||
# whole batch before locking would break that. Off = original full lock.
|
||||
skip_lock_components = (
|
||||
(ComponentType.MAMBA,)
|
||||
if envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.get()
|
||||
else ()
|
||||
)
|
||||
lock_result = self.inc_lock_ref(
|
||||
new_last_node, skip_lock_components=skip_lock_components
|
||||
new_last_node,
|
||||
skip_lock_components=(
|
||||
(ComponentType.MAMBA,)
|
||||
if envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.get()
|
||||
else ()
|
||||
),
|
||||
)
|
||||
|
||||
# Update req fields
|
||||
@@ -1124,9 +1114,8 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
req.prefix_indices = new_indices
|
||||
req.kv.cache_protected_len = len(new_indices)
|
||||
req.last_node = new_last_node
|
||||
req.swa_uuid_for_lock = lock_result.swa_uuid_for_lock
|
||||
# carry the skip set so this node's dec releases only what we locked
|
||||
req.skip_lock_node_ids = lock_result.skip_lock_node_ids
|
||||
# Carry the receipt so this node's dec releases only what we locked.
|
||||
req.lock_receipt = lock_result.to_dec_params()
|
||||
# The rematch acquired a new SWA prefix lock.
|
||||
req.swa_prefix_lock_released = False
|
||||
|
||||
|
||||
@@ -48,18 +48,18 @@ class SessionSlot:
|
||||
|
||||
# First req's radix tree node (for dec_lock_ref on session close)
|
||||
last_node: Any = None
|
||||
swa_uuid_for_lock: Optional[str] = None
|
||||
# components the first req skipped locking on last_node, so release dec
|
||||
# releases only what it took (may share the node with another req).
|
||||
skip_lock_node_ids: dict = field(default_factory=dict)
|
||||
# Receipt of the first request's tree lock on last_node.
|
||||
lock_receipt: DecLockRefParams = field(default_factory=DecLockRefParams)
|
||||
# Whether the first request already released its SWA lock.
|
||||
swa_prefix_lock_released: bool = False
|
||||
|
||||
def save_from_req(self, req: Req, is_first: bool):
|
||||
"""Save KV state from a finishing request into this slot."""
|
||||
kv = req.detach_kv()
|
||||
if is_first:
|
||||
self.last_node = req.last_node
|
||||
self.swa_uuid_for_lock = req.swa_uuid_for_lock
|
||||
self.skip_lock_node_ids = req.skip_lock_node_ids
|
||||
self.lock_receipt = req.lock_receipt
|
||||
self.swa_prefix_lock_released = req.swa_prefix_lock_released
|
||||
# The slot takes over the request's KV record.
|
||||
self.kv = kv
|
||||
else:
|
||||
@@ -71,8 +71,8 @@ class SessionSlot:
|
||||
def restore_to_req(self, req: Req):
|
||||
"""Restore KV state from this slot into an incoming request."""
|
||||
req.kv = self.kv
|
||||
req.swa_uuid_for_lock = self.swa_uuid_for_lock
|
||||
req.skip_lock_node_ids = self.skip_lock_node_ids
|
||||
req.lock_receipt = self.lock_receipt
|
||||
req.swa_prefix_lock_released = self.swa_prefix_lock_released
|
||||
|
||||
# NOTE: the slot keeps sharing the record it just handed out. During
|
||||
# chunked prefill, a request may be rejected by
|
||||
@@ -290,8 +290,8 @@ class StreamingSession(BasePrefixCache):
|
||||
slot = SessionSlot(
|
||||
kv=kv,
|
||||
last_node=req.last_node,
|
||||
swa_uuid_for_lock=req.swa_uuid_for_lock,
|
||||
skip_lock_node_ids=req.skip_lock_node_ids,
|
||||
lock_receipt=req.lock_receipt,
|
||||
swa_prefix_lock_released=req.swa_prefix_lock_released,
|
||||
)
|
||||
self.slots[session_id] = slot
|
||||
else:
|
||||
@@ -396,13 +396,10 @@ class StreamingSession(BasePrefixCache):
|
||||
)
|
||||
|
||||
if lock_node is not None:
|
||||
self.inner.dec_lock_ref(
|
||||
lock_node,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=slot.swa_uuid_for_lock,
|
||||
skip_lock_node_ids=slot.skip_lock_node_ids,
|
||||
),
|
||||
)
|
||||
# skip_swa is an SWA-cache extension kwarg; a slot can only have
|
||||
# early-released when the inner cache supports SWA locks.
|
||||
skip = {"skip_swa": True} if slot.swa_prefix_lock_released else {}
|
||||
self.inner.dec_lock_ref(lock_node, slot.lock_receipt, **skip)
|
||||
|
||||
if slot.kv.holds_kv:
|
||||
self.free_kv_row(slot.kv, [(protected_len, slot.kv.kv_allocated_len)])
|
||||
|
||||
@@ -11,7 +11,8 @@ if TYPE_CHECKING:
|
||||
class ScriptedLockRefExhauster:
|
||||
def __init__(self, scheduler: Scheduler) -> None:
|
||||
self.scheduler = scheduler
|
||||
self._locked: List[Any] = []
|
||||
# (node, dec receipt) pairs; the receipt bounds the release walk.
|
||||
self._locked: List[tuple[Any, Any]] = []
|
||||
|
||||
def exhaust(self, *, leave_refs: int) -> None:
|
||||
tree_cache = self.scheduler.tree_cache
|
||||
@@ -24,17 +25,17 @@ class ScriptedLockRefExhauster:
|
||||
return
|
||||
|
||||
target = evictable[0]
|
||||
tree_cache.inc_lock_ref(to_node_handle(tree_cache, target))
|
||||
result = tree_cache.inc_lock_ref(to_node_handle(tree_cache, target))
|
||||
|
||||
newly_locked = [node for node in evictable if _node_lock_ref(node) > 0]
|
||||
if not newly_locked:
|
||||
return
|
||||
self._locked.append(target)
|
||||
self._locked.append((target, result.to_dec_params()))
|
||||
|
||||
def release(self) -> None:
|
||||
tree_cache = self.scheduler.tree_cache
|
||||
for node in self._locked:
|
||||
tree_cache.dec_lock_ref(to_node_handle(tree_cache, node))
|
||||
for node, dec_params in self._locked:
|
||||
tree_cache.dec_lock_ref(to_node_handle(tree_cache, node), dec_params)
|
||||
self._locked.clear()
|
||||
|
||||
def _evictable_nodes(self) -> List[Any]:
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
//! the rest from the `TreeComponent` defaults.
|
||||
|
||||
use std::cmp::Reverse;
|
||||
use std::collections::{BinaryHeap, HashMap, HashSet};
|
||||
use std::collections::{BinaryHeap, HashMap};
|
||||
|
||||
use tch::{Kind, Tensor};
|
||||
|
||||
@@ -277,8 +277,6 @@ impl<K: ChildKeyType> TreeComponent<K> for FullComponent {
|
||||
mut result: IncLockRefResult,
|
||||
lock_host: bool,
|
||||
) -> IncLockRefResult {
|
||||
let ct = FULL;
|
||||
|
||||
// Only the last host node needs to be protected.
|
||||
if lock_host {
|
||||
let node = tree_core.arena.node_mut(node_id);
|
||||
@@ -291,20 +289,17 @@ impl<K: ChildKeyType> TreeComponent<K> for FullComponent {
|
||||
return result;
|
||||
}
|
||||
|
||||
// Skip the bottom evicted segment, recording it for the matching release.
|
||||
let on_boundary = |node: &Node<K>| node.is_root() || node.has_device_value(FULL);
|
||||
// The bottom device-evicted segment is locked too (no ledger move —
|
||||
// nothing is on device); a load-back that materializes a value under
|
||||
// lock credits protected directly.
|
||||
let mut cur = node_id;
|
||||
let mut node = tree_core.arena.node(cur);
|
||||
if !on_boundary(node) {
|
||||
let skip_lock_node_ids = result.skip_lock_node_ids.entry(ct).or_default();
|
||||
loop {
|
||||
skip_lock_node_ids.insert(node.id);
|
||||
cur = node.parent();
|
||||
node = tree_core.arena.node(cur);
|
||||
if on_boundary(node) {
|
||||
break;
|
||||
}
|
||||
loop {
|
||||
let node = tree_core.arena.node_mut(cur);
|
||||
if node.is_root() || node.has_device_value(FULL) {
|
||||
break;
|
||||
}
|
||||
node.inc_device_lock_ref(FULL);
|
||||
cur = node.parent();
|
||||
}
|
||||
|
||||
// Lock the device-on segment up to the root.
|
||||
@@ -341,11 +336,9 @@ impl<K: ChildKeyType> TreeComponent<K> for FullComponent {
|
||||
&self,
|
||||
tree_core: &mut UnifiedTreeCore<K>,
|
||||
node_id: NodeIdx_,
|
||||
params: Option<&DecLockRefParams>,
|
||||
_params: &DecLockRefParams,
|
||||
lock_host: bool,
|
||||
) {
|
||||
let ct = FULL;
|
||||
|
||||
if lock_host {
|
||||
let node = tree_core.arena.node_mut(node_id);
|
||||
if node.host_lock_ref(FULL) == 0 {
|
||||
@@ -360,10 +353,6 @@ impl<K: ChildKeyType> TreeComponent<K> for FullComponent {
|
||||
return;
|
||||
}
|
||||
|
||||
let empty = HashSet::new();
|
||||
let skip_lock_node_ids = params
|
||||
.and_then(|p| p.skip_lock_node_ids.get(&ct))
|
||||
.unwrap_or(&empty);
|
||||
let mut cur = node_id;
|
||||
loop {
|
||||
let node = tree_core.arena.node_mut(cur);
|
||||
@@ -371,20 +360,12 @@ impl<K: ChildKeyType> TreeComponent<K> for FullComponent {
|
||||
break;
|
||||
}
|
||||
let parent = node.parent();
|
||||
if skip_lock_node_ids.contains(&node.id) {
|
||||
cur = parent;
|
||||
continue;
|
||||
}
|
||||
assert!(
|
||||
node.has_device_value(FULL),
|
||||
"release_component_lock: node {cur} has no FULL device value"
|
||||
);
|
||||
let old_lock_ref = node.device_lock_ref(FULL);
|
||||
assert!(
|
||||
old_lock_ref > 0,
|
||||
"release_component_lock: node {cur} is not locked"
|
||||
"FULL segment release hit lock_ref=0 on node {cur}"
|
||||
);
|
||||
let newly_unlocked_len = if old_lock_ref == 1 {
|
||||
let newly_unlocked_len = if old_lock_ref == 1 && node.has_device_value(FULL) {
|
||||
Some(node.device_value_len(FULL))
|
||||
} else {
|
||||
None
|
||||
@@ -393,6 +374,8 @@ impl<K: ChildKeyType> TreeComponent<K> for FullComponent {
|
||||
if let Some(key_len) = newly_unlocked_len {
|
||||
tree_core.dec_protected_size(FULL, key_len);
|
||||
tree_core.inc_evictable_size(FULL, key_len);
|
||||
}
|
||||
if old_lock_ref == 1 {
|
||||
tree_core.update_evictable_leaf_sets_(cur);
|
||||
}
|
||||
cur = parent;
|
||||
@@ -478,9 +461,16 @@ impl<K: ChildKeyType> TreeComponent<K> for FullComponent {
|
||||
let n_len = loaded.host_value_len(FULL) as i64;
|
||||
loaded
|
||||
.set_device_value(FULL, device_indices.narrow(0, offset, n_len).copy());
|
||||
let locked = loaded.device_lock_ref(FULL) > 0;
|
||||
offset += n_len;
|
||||
// Full uses leaf sets, not LRU.
|
||||
tree_core.inc_evictable_size(FULL, n_len as usize);
|
||||
// Full uses leaf sets, not LRU. A value materialized
|
||||
// under lock is protected; the last release moves it
|
||||
// to evictable.
|
||||
if locked {
|
||||
tree_core.inc_protected_size(FULL, n_len as usize);
|
||||
} else {
|
||||
tree_core.inc_evictable_size(FULL, n_len as usize);
|
||||
}
|
||||
tree_core.update_evictable_leaf_sets_(loaded_idx);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -178,16 +178,7 @@ impl<K: ChildKeyType> TreeComponent<K> for MambaComponent {
|
||||
return;
|
||||
}
|
||||
if !tree_core.arena.has_device_value(node_id, MAMBA) {
|
||||
// Tombstone refill: the node moves from the host LRU to the device LRU.
|
||||
tree_core
|
||||
.arena
|
||||
.set_device_value(node_id, MAMBA, mamba_value.shallow_clone());
|
||||
let host_lru = tree_core.host_lru_list_mut(MAMBA);
|
||||
if host_lru.in_list(Some(node_id)) {
|
||||
host_lru.remove_node(node_id);
|
||||
}
|
||||
tree_core.device_lru_list_mut(MAMBA).insert_mru(node_id);
|
||||
tree_core.inc_evictable_size(MAMBA, slot_len);
|
||||
tree_core.set_component_device_value_(node_id, MAMBA, mamba_value.shallow_clone());
|
||||
let tick = tree_core.arena.get_and_bump_access_counter();
|
||||
tree_core.arena.node_mut(node_id).last_access_counter = tick;
|
||||
self.emit_excess_path_states_eviction_(tree_core.arena.node(node_id).id, cache_actions);
|
||||
@@ -417,24 +408,19 @@ impl<K: ChildKeyType> TreeComponent<K> for MambaComponent {
|
||||
&self,
|
||||
tree_core: &mut UnifiedTreeCore<K>,
|
||||
node_id: NodeIdx_,
|
||||
mut result: IncLockRefResult,
|
||||
result: IncLockRefResult,
|
||||
lock_host: bool,
|
||||
) -> IncLockRefResult {
|
||||
let node = tree_core.arena.node(node_id);
|
||||
if node.is_root() {
|
||||
return result;
|
||||
}
|
||||
// A node in skip_lock_node_ids was a tombstone when this lock was acquired.
|
||||
if !Self::has_value(node, lock_host) {
|
||||
result
|
||||
.skip_lock_node_ids
|
||||
.entry(MAMBA)
|
||||
.or_default()
|
||||
.insert(node.id);
|
||||
return result;
|
||||
}
|
||||
// Tombstones are counted too; ledger/LRU track only data-bearing
|
||||
// nodes (a value materialized under lock is credited to protected
|
||||
// at the materialization site).
|
||||
let has_value = Self::has_value(node, lock_host);
|
||||
if lock_host {
|
||||
if node.host_lock_ref(MAMBA) == 0 {
|
||||
if node.host_lock_ref(MAMBA) == 0 && has_value {
|
||||
let host_lru = tree_core.host_lru_list_mut(MAMBA);
|
||||
if host_lru.in_list(Some(node_id)) {
|
||||
host_lru.remove_node(node_id);
|
||||
@@ -443,7 +429,7 @@ impl<K: ChildKeyType> TreeComponent<K> for MambaComponent {
|
||||
tree_core.arena.inc_host_lock_ref(node_id, MAMBA);
|
||||
} else {
|
||||
let value_len = node.device_value_len(MAMBA);
|
||||
if node.device_lock_ref(MAMBA) == 0 {
|
||||
if node.device_lock_ref(MAMBA) == 0 && has_value {
|
||||
tree_core.dec_evictable_size(MAMBA, value_len);
|
||||
tree_core.inc_protected_size(MAMBA, value_len);
|
||||
}
|
||||
@@ -457,43 +443,44 @@ impl<K: ChildKeyType> TreeComponent<K> for MambaComponent {
|
||||
&self,
|
||||
tree_core: &mut UnifiedTreeCore<K>,
|
||||
node_id: NodeIdx_,
|
||||
params: Option<&DecLockRefParams>,
|
||||
_params: &DecLockRefParams,
|
||||
lock_host: bool,
|
||||
) {
|
||||
if tree_core.arena.node(node_id).is_root() {
|
||||
return;
|
||||
}
|
||||
if let Some(params) = params
|
||||
&& params
|
||||
.skip_lock_node_ids
|
||||
.get(&MAMBA)
|
||||
.is_some_and(|ids| ids.contains(&tree_core.arena.node(node_id).id))
|
||||
{
|
||||
return;
|
||||
}
|
||||
if lock_host {
|
||||
let node = tree_core.arena.node_mut(node_id);
|
||||
assert!(
|
||||
node.host_lock_ref(MAMBA) > 0,
|
||||
"Mamba release hit host_lock_ref=0 on node {node_id}"
|
||||
);
|
||||
node.dec_host_lock_ref(MAMBA);
|
||||
if node.host_lock_ref(MAMBA) == 0
|
||||
&& !node.has_device_value(MAMBA)
|
||||
&& node.has_host_value(MAMBA)
|
||||
{
|
||||
let host_lru = tree_core.host_lru_list_mut(MAMBA);
|
||||
if !host_lru.in_list(Some(node_id)) {
|
||||
host_lru.insert_mru(node_id);
|
||||
if node.host_lock_ref(MAMBA) == 0 {
|
||||
if !node.has_device_value(MAMBA) && node.has_host_value(MAMBA) {
|
||||
let host_lru = tree_core.host_lru_list_mut(MAMBA);
|
||||
if !host_lru.in_list(Some(node_id)) {
|
||||
host_lru.insert_mru(node_id);
|
||||
}
|
||||
}
|
||||
tree_core.update_evictable_leaf_sets_(node_id);
|
||||
}
|
||||
return;
|
||||
}
|
||||
let node = tree_core.arena.node(node_id);
|
||||
let device_lock_ref = node.device_lock_ref(MAMBA);
|
||||
if device_lock_ref > 0 {
|
||||
if device_lock_ref == 1 {
|
||||
let value_len = node.device_value_len(MAMBA);
|
||||
tree_core.inc_evictable_size(MAMBA, value_len);
|
||||
tree_core.dec_protected_size(MAMBA, value_len);
|
||||
}
|
||||
tree_core.arena.dec_device_lock_ref(node_id, MAMBA);
|
||||
assert!(
|
||||
device_lock_ref > 0,
|
||||
"Mamba release hit lock_ref=0 on node {node_id}"
|
||||
);
|
||||
if device_lock_ref == 1 && node.has_device_value(MAMBA) {
|
||||
let value_len = node.device_value_len(MAMBA);
|
||||
tree_core.inc_evictable_size(MAMBA, value_len);
|
||||
tree_core.dec_protected_size(MAMBA, value_len);
|
||||
}
|
||||
tree_core.arena.dec_device_lock_ref(node_id, MAMBA);
|
||||
if device_lock_ref == 1 {
|
||||
tree_core.update_evictable_leaf_sets_(node_id);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -613,16 +600,9 @@ impl<K: ChildKeyType> TreeComponent<K> for MambaComponent {
|
||||
return;
|
||||
};
|
||||
if let Some(device_indices) = &transfer.device_indices {
|
||||
let node = tree_core.arena.node_mut(node_id);
|
||||
node.set_device_value(MAMBA, device_indices.copy());
|
||||
let count = node.device_value_len(MAMBA);
|
||||
// Move from host LRU to device LRU
|
||||
let host_lru = tree_core.host_lru_list_mut(MAMBA);
|
||||
if host_lru.in_list(Some(node_id)) {
|
||||
host_lru.remove_node(node_id);
|
||||
}
|
||||
tree_core.device_lru_list_mut(MAMBA).insert_mru(node_id);
|
||||
tree_core.inc_evictable_size(MAMBA, count);
|
||||
// The materialization primitive owns the ledger/LRU moves,
|
||||
// including crediting protected when restored under lock.
|
||||
tree_core.set_component_device_value_(node_id, MAMBA, device_indices.copy());
|
||||
}
|
||||
}
|
||||
// The python elif chain has no BACKUP_STORAGE arm.
|
||||
|
||||
@@ -344,7 +344,7 @@ pub trait TreeComponent<K: ChildKeyType> {
|
||||
&self,
|
||||
tree_core: &mut UnifiedTreeCore<K>,
|
||||
node_id: NodeIdx_,
|
||||
params: Option<&DecLockRefParams>,
|
||||
params: &DecLockRefParams,
|
||||
lock_host: bool,
|
||||
);
|
||||
|
||||
@@ -466,6 +466,49 @@ pub const BASE_COMPONENT_TYPE: ComponentType = ComponentType::Full;
|
||||
/// Slots per tier — the arrays are sized to this, not the enabled subset.
|
||||
pub const NUM_COMPONENT_TYPES: usize = ComponentType::Mamba as usize + 1;
|
||||
|
||||
/// A set of component types (bitmask over `ComponentType::idx`), e.g. the
|
||||
/// components an `inc_lock_ref` left untaken.
|
||||
#[derive(Copy, Clone, Default, PartialEq, Eq, Debug)]
|
||||
pub struct ComponentSet(u8);
|
||||
|
||||
impl ComponentSet {
|
||||
pub const EMPTY: ComponentSet = ComponentSet(0);
|
||||
|
||||
/// The set holding exactly one component.
|
||||
pub const fn of(component_type: ComponentType) -> ComponentSet {
|
||||
ComponentSet(1 << component_type.idx())
|
||||
}
|
||||
|
||||
pub fn insert(&mut self, component_type: ComponentType) {
|
||||
self.0 |= 1 << component_type.idx();
|
||||
}
|
||||
|
||||
pub const fn contains(self, component_type: ComponentType) -> bool {
|
||||
self.0 & (1 << component_type.idx()) != 0
|
||||
}
|
||||
|
||||
pub const fn is_empty(self) -> bool {
|
||||
self.0 == 0
|
||||
}
|
||||
|
||||
/// The members, in component-index order.
|
||||
pub fn iter(self) -> impl Iterator<Item = ComponentType> {
|
||||
(0..NUM_COMPONENT_TYPES)
|
||||
.filter(move |idx| self.0 & (1 << idx) != 0)
|
||||
.map(ComponentType::from_idx)
|
||||
}
|
||||
}
|
||||
|
||||
impl FromIterator<ComponentType> for ComponentSet {
|
||||
fn from_iter<I: IntoIterator<Item = ComponentType>>(iter: I) -> Self {
|
||||
let mut set = ComponentSet::EMPTY;
|
||||
for component_type in iter {
|
||||
set.insert(component_type);
|
||||
}
|
||||
set
|
||||
}
|
||||
}
|
||||
|
||||
impl ComponentType {
|
||||
/// Index into a per-component array.
|
||||
pub const fn idx(self) -> usize {
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
//! SWA values arrive pool-resolved; the full->SWA index translation happens at
|
||||
//! the cache boundary.
|
||||
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::collections::HashMap;
|
||||
|
||||
use tch::{Kind, Tensor};
|
||||
|
||||
@@ -34,10 +34,15 @@ impl SwaComponent {
|
||||
impl SwaComponent {
|
||||
/// Build the driver from the tree's init params.
|
||||
pub fn new(params: &CacheInitParams) -> Self {
|
||||
let sliding_window_size = params
|
||||
.swa_sliding_window_size
|
||||
.expect("the SWA component requires swa_sliding_window_size");
|
||||
assert!(
|
||||
sliding_window_size > 0,
|
||||
"swa_sliding_window_size must be positive"
|
||||
);
|
||||
SwaComponent {
|
||||
sliding_window_size: params
|
||||
.swa_sliding_window_size
|
||||
.expect("the SWA component requires swa_sliding_window_size"),
|
||||
sliding_window_size,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -408,11 +413,9 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
||||
}
|
||||
|
||||
let swa_evicted_seqlen = params.swa_evicted_seqlen;
|
||||
assert_eq!(
|
||||
node.device_lock_ref(SWA),
|
||||
0,
|
||||
"tombstone Swa lock_ref should be 0, node {node_id}"
|
||||
);
|
||||
// A locked tombstone is legal (segment locks count every node); the
|
||||
// full-value swap below is safe because full lock_ref >= swa
|
||||
// lock_ref, so a locked-SWA node always takes the Recover branch.
|
||||
assert_eq!(
|
||||
swa_evicted_seqlen % tree_core.page_size,
|
||||
0,
|
||||
@@ -495,11 +498,6 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
||||
if node.has_device_value(SWA) {
|
||||
return;
|
||||
}
|
||||
assert_eq!(
|
||||
node.device_lock_ref(SWA),
|
||||
0,
|
||||
"tombstone Swa lock_ref should be 0 on unevict, node {node_id}"
|
||||
);
|
||||
let swa_evicted_seqlen = params.swa_evicted_seqlen;
|
||||
assert_eq!(
|
||||
swa_evicted_seqlen % tree_core.page_size,
|
||||
@@ -587,26 +585,33 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
||||
let (new_parent, child) = tree_core.arena.node_pair_mut(new_parent_id, child_id);
|
||||
let split_len = new_parent.key.atom_len() as i64;
|
||||
new_parent.copy_device_lock_ref(SWA, child);
|
||||
new_parent.copy_host_lock_ref(SWA, child);
|
||||
if child.has_device_value(SWA) {
|
||||
Node::redistribute_child_device_value(new_parent, child, SWA, split_len);
|
||||
}
|
||||
if child.has_host_value(SWA) {
|
||||
Node::redistribute_child_host_value(new_parent, child, SWA, split_len);
|
||||
// Device-tombstoned sides park in the host LRU.
|
||||
let parent_is_tombstone = !new_parent.has_device_value(SWA);
|
||||
let child_is_tombstone = !child.has_device_value(SWA);
|
||||
// Device-tombstoned sides park in the host LRU. Host-locked
|
||||
// halves stay out of it: in-flight IO holds them, and host
|
||||
// acquire removed the node at 0->1.
|
||||
let parent_parks =
|
||||
!new_parent.has_device_value(SWA) && new_parent.host_lock_ref(SWA) == 0;
|
||||
let child_parks = !child.has_device_value(SWA) && child.host_lock_ref(SWA) == 0;
|
||||
let host_lru = tree_core.host_lru_list_mut(SWA);
|
||||
if parent_is_tombstone {
|
||||
if parent_parks {
|
||||
host_lru.insert_mru(new_parent_id);
|
||||
}
|
||||
if child_is_tombstone && !host_lru.in_list(Some(child_id)) {
|
||||
if child_parks && !host_lru.in_list(Some(child_id)) {
|
||||
host_lru.insert_mru(child_id);
|
||||
}
|
||||
}
|
||||
|
||||
// parent inherits the swa_uuid from child for swa lock ref
|
||||
// The window-boundary uuids mark the node's older edge, which the
|
||||
// split moves to the parent — both tiers migrate with it.
|
||||
let swa_uuid = tree_core.arena.node_mut(child_id).swa_uuid.take();
|
||||
tree_core.arena.node_mut(new_parent_id).swa_uuid = swa_uuid;
|
||||
let swa_host_uuid = tree_core.arena.node_mut(child_id).swa_host_uuid.take();
|
||||
tree_core.arena.node_mut(new_parent_id).swa_host_uuid = swa_host_uuid;
|
||||
}
|
||||
|
||||
fn evict_component(
|
||||
@@ -1033,46 +1038,44 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
||||
mut result: IncLockRefResult,
|
||||
lock_host: bool,
|
||||
) -> IncLockRefResult {
|
||||
let ct = SWA;
|
||||
// Lock the contiguous segment covering the trailing window.
|
||||
//
|
||||
// Every node in [node, boundary] is counted, tombstones included, so
|
||||
// the paired release decrements the same contiguous segment with no
|
||||
// carried skip state. Coverage is position-based (key length); the
|
||||
// boundary node is always uuid-stamped, so a release without a uuid
|
||||
// means the segment reached the root. Ledger/LRU transitions track
|
||||
// only data-bearing nodes; a value materialized later under lock is
|
||||
// credited to protected by set_component_device_value.
|
||||
let sliding_window_size = self.sliding_window_size;
|
||||
let mut swa_lock_size = 0;
|
||||
let mut covered = 0;
|
||||
let mut swa_uuid = None;
|
||||
|
||||
// Tombstoned nodes (cd.value is None) have no SWA chunk to protect
|
||||
// skip them and keep walking up. This path is hit when HiCache
|
||||
// backs up a FULL present internal node whose SWA was already evicted.
|
||||
let mut cur = node_id;
|
||||
loop {
|
||||
let node = tree_core.arena.node_mut(cur);
|
||||
if node.is_root() || swa_lock_size >= sliding_window_size {
|
||||
if node.is_root() || covered >= sliding_window_size {
|
||||
break;
|
||||
}
|
||||
let parent = node.parent();
|
||||
if !Self::has_value(node, lock_host) {
|
||||
result
|
||||
.skip_lock_node_ids
|
||||
.entry(ct)
|
||||
.or_default()
|
||||
.insert(node.id);
|
||||
cur = parent;
|
||||
continue;
|
||||
}
|
||||
let key_len = node.key.atom_len();
|
||||
let has_value = Self::has_value(node, lock_host);
|
||||
let value_len = Self::value_len(node, lock_host);
|
||||
let newly_locked = Self::lock_ref(node, lock_host) == 0;
|
||||
Self::inc_lock_ref(node, lock_host);
|
||||
swa_lock_size += Self::value_len(node, lock_host);
|
||||
if newly_locked {
|
||||
if newly_locked && has_value {
|
||||
if lock_host {
|
||||
let host_lru = tree_core.host_lru_list_mut(SWA);
|
||||
if host_lru.in_list(Some(cur)) {
|
||||
host_lru.remove_node(cur);
|
||||
}
|
||||
} else {
|
||||
tree_core.dec_evictable_size(SWA, key_len);
|
||||
tree_core.inc_protected_size(SWA, key_len);
|
||||
tree_core.dec_evictable_size(SWA, value_len);
|
||||
tree_core.inc_protected_size(SWA, value_len);
|
||||
}
|
||||
}
|
||||
if swa_lock_size >= sliding_window_size {
|
||||
covered += key_len;
|
||||
if covered >= sliding_window_size {
|
||||
swa_uuid = Some(Self::ensure_swa_uuid(tree_core, cur, lock_host));
|
||||
}
|
||||
cur = parent;
|
||||
@@ -1090,23 +1093,15 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
||||
&self,
|
||||
tree_core: &mut UnifiedTreeCore<K>,
|
||||
node_id: NodeIdx_,
|
||||
params: Option<&DecLockRefParams>,
|
||||
params: &DecLockRefParams,
|
||||
lock_host: bool,
|
||||
) {
|
||||
let ct = SWA;
|
||||
let swa_uuid_for_lock = params.and_then(|p| {
|
||||
if lock_host {
|
||||
p.swa_uuid_for_host_lock
|
||||
} else {
|
||||
p.swa_uuid_for_lock
|
||||
}
|
||||
});
|
||||
let empty = HashSet::new();
|
||||
let skip_lock_node_ids = params
|
||||
.and_then(|p| p.skip_lock_node_ids.get(&ct))
|
||||
.unwrap_or(&empty);
|
||||
let swa_uuid_for_lock = if lock_host {
|
||||
params.swa_uuid_for_host_lock
|
||||
} else {
|
||||
params.swa_uuid_for_lock
|
||||
};
|
||||
|
||||
// A node in skip_lock_node_ids was a tombstone when this lock was acquired.
|
||||
let mut cur = node_id;
|
||||
loop {
|
||||
let node = tree_core.arena.node_mut(cur);
|
||||
@@ -1114,30 +1109,36 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
||||
break;
|
||||
}
|
||||
let parent = node.parent();
|
||||
if skip_lock_node_ids.contains(&node.id) {
|
||||
cur = parent;
|
||||
continue;
|
||||
}
|
||||
let lock_ref = Self::lock_ref(node, lock_host);
|
||||
if lock_ref == 0 {
|
||||
cur = parent;
|
||||
continue;
|
||||
}
|
||||
if lock_ref == 1 {
|
||||
// Acquire counted every segment node and splits copy refs, so a
|
||||
// zero here means the release does not mirror its acquire.
|
||||
assert!(
|
||||
lock_ref > 0,
|
||||
"SWA segment release hit {}lock_ref=0 on node {cur}",
|
||||
if lock_host { "host_" } else { "" }
|
||||
);
|
||||
let has_value = Self::has_value(node, lock_host);
|
||||
let value_len = Self::value_len(node, lock_host);
|
||||
if lock_ref == 1 && has_value {
|
||||
if lock_host {
|
||||
if !node.has_device_value(SWA) && node.has_host_value(SWA) {
|
||||
if !node.has_device_value(SWA) {
|
||||
let host_lru = tree_core.host_lru_list_mut(SWA);
|
||||
if !host_lru.in_list(Some(cur)) {
|
||||
host_lru.insert_mru(cur);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let key_len = node.device_value_len(SWA);
|
||||
tree_core.inc_evictable_size(SWA, key_len);
|
||||
tree_core.dec_protected_size(SWA, key_len);
|
||||
tree_core.inc_evictable_size(SWA, value_len);
|
||||
tree_core.dec_protected_size(SWA, value_len);
|
||||
}
|
||||
}
|
||||
Self::dec_lock_ref(tree_core.arena.node_mut(cur), lock_host);
|
||||
if lock_ref == 1 {
|
||||
// This may have been the last lock holding the node out of
|
||||
// the evictable-leaf sets; refresh it here rather than rely
|
||||
// on the Full walk running after this one.
|
||||
tree_core.update_evictable_leaf_sets_(cur);
|
||||
}
|
||||
if swa_uuid_for_lock.is_some()
|
||||
&& Self::swa_uuid(tree_core.arena.node(cur), lock_host) == swa_uuid_for_lock
|
||||
{
|
||||
@@ -1147,15 +1148,14 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
||||
}
|
||||
}
|
||||
|
||||
/// Early-release the SWA lock along [node, swa_uuid_for_lock] while
|
||||
/// leaving Full and Mamba locks intact.
|
||||
/// Early-release the SWA lock along [node, swa_uuid_for_lock]; this
|
||||
/// method touches only SWA state. The wrapping `dec_swa_lock_only` also
|
||||
/// drops strictly-lower-priority co-located locks (e.g. Mamba) per the
|
||||
/// receipt; the Full lock stays so the request's prefix is protected.
|
||||
///
|
||||
/// Called when a request's decode position has advanced past the sliding
|
||||
/// window — the SWA portion of the tree lock is no longer needed but the
|
||||
/// Full lock must stay so the request's prefix is protected.
|
||||
///
|
||||
/// Caller (UnifiedRadixCache.dec_swa_lock_only) must ensure this is
|
||||
/// invoked at most once per (node, swa_uuid_for_lock) pair.
|
||||
/// window. The caller must invoke this at most once per
|
||||
/// (node, swa_uuid_for_lock) pair.
|
||||
fn release_window_lock(
|
||||
&self,
|
||||
tree_core: &mut UnifiedTreeCore<K>,
|
||||
@@ -1172,21 +1172,20 @@ impl<K: ChildKeyType> TreeComponent<K> for SwaComponent {
|
||||
break;
|
||||
}
|
||||
let parent = node.parent();
|
||||
// Acquire skips tombstoned nodes; release must skip them too. Same
|
||||
// for nodes with lock_ref == 0 — acquire never credited them.
|
||||
if !node.has_device_value(SWA) || node.device_lock_ref(SWA) == 0 {
|
||||
if swa_uuid_for_lock.is_some() && node.swa_uuid == swa_uuid_for_lock {
|
||||
break;
|
||||
}
|
||||
cur = parent;
|
||||
continue;
|
||||
}
|
||||
|
||||
assert!(
|
||||
node.device_lock_ref(SWA) > 0,
|
||||
"SWA window release hit lock_ref=0 on node {cur}"
|
||||
);
|
||||
let has_value = node.has_device_value(SWA);
|
||||
let value_len = node.device_value_len(SWA);
|
||||
node.dec_device_lock_ref(SWA);
|
||||
if node.device_lock_ref(SWA) == 0 {
|
||||
let key_len = node.key.atom_len();
|
||||
tree_core.dec_protected_size(SWA, key_len);
|
||||
tree_core.inc_evictable_size(SWA, key_len);
|
||||
let now_unlocked = node.device_lock_ref(SWA) == 0;
|
||||
if now_unlocked {
|
||||
tree_core.update_evictable_leaf_sets_(cur);
|
||||
}
|
||||
if now_unlocked && has_value {
|
||||
tree_core.dec_protected_size(SWA, value_len);
|
||||
tree_core.inc_evictable_size(SWA, value_len);
|
||||
if tree_core.is_evictable_device_leaf_(tree_core.arena.node(cur)) {
|
||||
tree_core.evict_component_and_detach_lru_(
|
||||
cur,
|
||||
|
||||
@@ -283,6 +283,12 @@ impl<K: ChildKeyType> Node<K> {
|
||||
self.set_lock_ref_(slot, src_node.lock_ref_(slot));
|
||||
}
|
||||
|
||||
/// Copy the component's host lock refcount from `src_node`.
|
||||
pub fn copy_host_lock_ref(&mut self, component_type: ComponentType, src_node: &Node<K>) {
|
||||
let slot = ValueSlotIdx::host(component_type);
|
||||
self.set_lock_ref_(slot, src_node.lock_ref_(slot));
|
||||
}
|
||||
|
||||
/// Split the component's device value between a new parent and the child.
|
||||
pub fn redistribute_child_device_value(
|
||||
parent_node: &mut Node<K>,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
//! Python bindings: the `mem_cache` extension module and its TreeCore adapter.
|
||||
|
||||
use std::borrow::Cow;
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use pyo3::buffer::PyBuffer;
|
||||
@@ -10,7 +10,7 @@ use pyo3::prelude::*;
|
||||
use pyo3::types::{PyBytes, PyDict, PyList};
|
||||
use tch::{Device, Kind, Tensor};
|
||||
|
||||
use crate::components::{ComponentType, FULL, MAMBA, SWA};
|
||||
use crate::components::{ComponentSet, ComponentType, FULL, MAMBA, SWA};
|
||||
use crate::node::ChildKeyType;
|
||||
use crate::node::{KeyNamespaceRef, NodeAccessError, NodeId, TreeCoreRuntimeError};
|
||||
use crate::unified_tree_core::KvCacheEvent;
|
||||
@@ -648,24 +648,27 @@ impl InsertResultBinding {
|
||||
#[pyclass(get_all, set_all)]
|
||||
#[derive(Clone, Default)]
|
||||
pub struct DecLockRefParamsBinding {
|
||||
pub node_id: Option<NodeId>,
|
||||
pub swa_uuid_for_lock: Option<i64>,
|
||||
pub swa_uuid_for_host_lock: Option<i64>,
|
||||
pub skip_lock_node_ids: HashMap<u8, HashSet<NodeId>>,
|
||||
pub skipped_lock_components: Vec<u8>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl DecLockRefParamsBinding {
|
||||
#[new]
|
||||
#[pyo3(signature = (swa_uuid_for_lock = None, swa_uuid_for_host_lock = None, skip_lock_node_ids = None))]
|
||||
#[pyo3(signature = (node_id = None, swa_uuid_for_lock = None, swa_uuid_for_host_lock = None, skipped_lock_components = Vec::new()))]
|
||||
fn new(
|
||||
node_id: Option<NodeId>,
|
||||
swa_uuid_for_lock: Option<i64>,
|
||||
swa_uuid_for_host_lock: Option<i64>,
|
||||
skip_lock_node_ids: Option<HashMap<u8, HashSet<NodeId>>>,
|
||||
skipped_lock_components: Vec<u8>,
|
||||
) -> Self {
|
||||
DecLockRefParamsBinding {
|
||||
node_id,
|
||||
swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: skip_lock_node_ids.unwrap_or_default(),
|
||||
skipped_lock_components,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -674,44 +677,49 @@ impl DecLockRefParamsBinding {
|
||||
/// Convert into the tree core's dec-lock params.
|
||||
fn to_dec_lock_ref_params(&self) -> PyResult<DecLockRefParams> {
|
||||
Ok(DecLockRefParams {
|
||||
node_id: self.node_id,
|
||||
swa_uuid_for_lock: self.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: self.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: self
|
||||
.skip_lock_node_ids
|
||||
.iter()
|
||||
.map(|(ct, node_ids)| {
|
||||
Ok::<_, PyErr>((parse_component_type(*ct)?, node_ids.clone()))
|
||||
})
|
||||
.collect::<PyResult<_>>()?,
|
||||
skipped_lock_components: component_set_from_py(&self.skipped_lock_components)?,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Python-visible inc_lock_ref result; hand skip_lock_node_ids back to the
|
||||
/// matching dec_lock_ref.
|
||||
/// Python-visible inc_lock_ref result; the receipt (anchor node, boundary
|
||||
/// uuids, skipped components) is handed back to the matching dec_lock_ref.
|
||||
#[pyclass(get_all)]
|
||||
pub struct IncLockRefResultBinding {
|
||||
delta: Option<usize>,
|
||||
node_id: Option<NodeId>,
|
||||
swa_uuid_for_lock: Option<i64>,
|
||||
swa_uuid_for_host_lock: Option<i64>,
|
||||
skip_lock_node_ids: HashMap<u8, HashSet<NodeId>>,
|
||||
skipped_lock_components: Vec<u8>,
|
||||
}
|
||||
|
||||
impl IncLockRefResultBinding {
|
||||
fn from_result(result: crate::unified_tree_core::IncLockRefResult) -> Self {
|
||||
Self {
|
||||
delta: result.delta,
|
||||
node_id: result.node_id,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result
|
||||
.skip_lock_node_ids
|
||||
.into_iter()
|
||||
.map(|(ct, node_ids)| (component_type_to_u8(ct), node_ids))
|
||||
skipped_lock_components: result
|
||||
.skipped_lock_components
|
||||
.iter()
|
||||
.map(|ct| ct.idx() as u8)
|
||||
.collect(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse Python component-type ids into a component set.
|
||||
fn component_set_from_py(component_types: &[u8]) -> PyResult<ComponentSet> {
|
||||
component_types
|
||||
.iter()
|
||||
.map(|ct| parse_component_type(*ct))
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Convert a Python component-keyed tracker into the core's counts.
|
||||
fn tracker_from_py(tracker: HashMap<u8, usize>) -> PyResult<HashMap<ComponentType, usize>> {
|
||||
tracker
|
||||
@@ -1074,23 +1082,17 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
|
||||
cache_actions_to_py(py, actions)
|
||||
}
|
||||
|
||||
/// Bump the reference count on a node's component locks.
|
||||
/// Bump the reference count on a node's component locks; the listed
|
||||
/// components are left untaken and recorded in the receipt.
|
||||
fn inc_lock_ref(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
node_id: NodeId,
|
||||
skip_lock_components: Option<Vec<u8>>,
|
||||
skip_lock_components: Vec<u8>,
|
||||
) -> PyResult<IncLockRefResultBinding> {
|
||||
let skip_lock_components = skip_lock_components
|
||||
.unwrap_or_default()
|
||||
.into_iter()
|
||||
.map(parse_component_type)
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
let skip = component_set_from_py(&skip_lock_components)?;
|
||||
let result = py
|
||||
.allow_threads(|| {
|
||||
self.core()
|
||||
.inc_lock_ref_with_skip(node_id, &skip_lock_components)
|
||||
})
|
||||
.allow_threads(|| self.core().inc_lock_ref(node_id, skip))
|
||||
.map_err(node_access_error)?;
|
||||
Ok(IncLockRefResultBinding::from_result(result))
|
||||
}
|
||||
@@ -1100,11 +1102,11 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
node_id: NodeId,
|
||||
params: Option<&DecLockRefParamsBinding>,
|
||||
params: &DecLockRefParamsBinding,
|
||||
skip_swa: bool,
|
||||
) -> PyResult<()> {
|
||||
let params = params.map(|p| p.to_dec_lock_ref_params()).transpose()?;
|
||||
py.allow_threads(|| self.core().dec_lock_ref(node_id, params.as_ref(), skip_swa))
|
||||
let params = params.to_dec_lock_ref_params()?;
|
||||
py.allow_threads(|| self.core().dec_lock_ref(node_id, ¶ms, skip_swa))
|
||||
.map_err(node_access_error)?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -1115,22 +1117,16 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
node_id: NodeId,
|
||||
swa_uuid_for_lock: Option<i64>,
|
||||
skip_lock_node_ids: Option<HashMap<u8, HashSet<NodeId>>>,
|
||||
params: &DecLockRefParamsBinding,
|
||||
) -> PyResult<(Py<PyDict>, Py<PyDict>)> {
|
||||
let skip_lock_node_ids = skip_lock_node_ids
|
||||
.unwrap_or_default()
|
||||
.into_iter()
|
||||
.map(|(ct, node_ids)| Ok((parse_component_type(ct)?, node_ids)))
|
||||
.collect::<PyResult<HashMap<_, _>>>()?;
|
||||
let params = params.to_dec_lock_ref_params()?;
|
||||
let (device_frees, host_frees) = py
|
||||
.allow_threads(|| {
|
||||
let mut device_frees = HashMap::new();
|
||||
let mut host_frees = HashMap::new();
|
||||
self.core().dec_swa_lock_only_with_skip(
|
||||
self.core().dec_swa_lock_only(
|
||||
node_id,
|
||||
swa_uuid_for_lock,
|
||||
Some(&skip_lock_node_ids),
|
||||
¶ms,
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)?;
|
||||
@@ -1733,16 +1729,7 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
|
||||
let result = py
|
||||
.allow_threads(|| self.core().inc_host_lock_ref(node_id))
|
||||
.map_err(node_access_error)?;
|
||||
Ok(IncLockRefResultBinding {
|
||||
delta: result.delta,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result
|
||||
.skip_lock_node_ids
|
||||
.into_iter()
|
||||
.map(|(ct, node_ids)| (component_type_to_u8(ct), node_ids))
|
||||
.collect(),
|
||||
})
|
||||
Ok(IncLockRefResultBinding::from_result(result))
|
||||
}
|
||||
|
||||
/// Decrease the reference count on a node's host-side component locks.
|
||||
@@ -1750,10 +1737,10 @@ impl<K: ChildKeyType + Send + Sync> TreeCoreBinding<K> {
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
node_id: NodeId,
|
||||
params: Option<&DecLockRefParamsBinding>,
|
||||
params: &DecLockRefParamsBinding,
|
||||
) -> PyResult<()> {
|
||||
let params = params.map(|p| p.to_dec_lock_ref_params()).transpose()?;
|
||||
py.allow_threads(|| self.core().dec_host_lock_ref(node_id, params.as_ref()))
|
||||
let params = params.to_dec_lock_ref_params()?;
|
||||
py.allow_threads(|| self.core().dec_host_lock_ref(node_id, ¶ms))
|
||||
.map_err(node_access_error)?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -2371,23 +2358,24 @@ macro_rules! tree_core_binding {
|
||||
}
|
||||
|
||||
/// Bump the reference count on a node's component locks.
|
||||
#[pyo3(signature = (node_id, skip_lock_components = None))]
|
||||
#[pyo3(signature = (node_id, skip_lock_components = Vec::new()))]
|
||||
fn inc_lock_ref(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
node_id: NodeId,
|
||||
skip_lock_components: Option<Vec<u8>>,
|
||||
skip_lock_components: Vec<u8>,
|
||||
) -> PyResult<IncLockRefResultBinding> {
|
||||
self.inner.inc_lock_ref(py, node_id, skip_lock_components)
|
||||
}
|
||||
|
||||
/// Decrease the reference count on a node's component locks.
|
||||
#[pyo3(signature = (node_id, params = None, skip_swa = false))]
|
||||
/// Decrease the reference count on a node's component locks. The
|
||||
/// receipt is required: a release must replay its acquire's evidence.
|
||||
#[pyo3(signature = (node_id, params, skip_swa = false))]
|
||||
fn dec_lock_ref(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
node_id: NodeId,
|
||||
params: Option<&DecLockRefParamsBinding>,
|
||||
params: &DecLockRefParamsBinding,
|
||||
skip_swa: bool,
|
||||
) -> PyResult<()> {
|
||||
self.inner.dec_lock_ref(py, node_id, params, skip_swa)
|
||||
@@ -2395,20 +2383,14 @@ macro_rules! tree_core_binding {
|
||||
|
||||
/// Early-release the SWA portion of a request's tree lock; returns this
|
||||
/// release's per-component (device_frees, host_frees).
|
||||
#[pyo3(signature = (node_id, swa_uuid_for_lock = None, skip_lock_node_ids = None))]
|
||||
#[pyo3(signature = (node_id, params))]
|
||||
fn dec_swa_lock_only(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
node_id: NodeId,
|
||||
swa_uuid_for_lock: Option<i64>,
|
||||
skip_lock_node_ids: Option<HashMap<u8, HashSet<NodeId>>>,
|
||||
params: &DecLockRefParamsBinding,
|
||||
) -> PyResult<(Py<PyDict>, Py<PyDict>)> {
|
||||
self.inner.dec_swa_lock_only(
|
||||
py,
|
||||
node_id,
|
||||
swa_uuid_for_lock,
|
||||
skip_lock_node_ids,
|
||||
)
|
||||
self.inner.dec_swa_lock_only(py, node_id, params)
|
||||
}
|
||||
|
||||
/// Store a component's device value on a node (the SWA rebuild write-back).
|
||||
@@ -2817,12 +2799,12 @@ macro_rules! tree_core_binding {
|
||||
}
|
||||
|
||||
/// Decrease the reference count on a node's host-side component locks.
|
||||
#[pyo3(signature = (node_id, params = None))]
|
||||
/// The receipt is required, as for dec_lock_ref.
|
||||
fn dec_host_lock_ref(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
node_id: NodeId,
|
||||
params: Option<&DecLockRefParamsBinding>,
|
||||
params: &DecLockRefParamsBinding,
|
||||
) -> PyResult<()> {
|
||||
self.inner.dec_host_lock_ref(py, node_id, params)
|
||||
}
|
||||
|
||||
@@ -70,7 +70,7 @@ impl TreeComponent<Vec<i64>> for DefaultComponentForTest {
|
||||
&self,
|
||||
tree_core: &mut UnifiedTreeCore<Vec<i64>>,
|
||||
node_id: NodeIdx_,
|
||||
params: Option<&DecLockRefParams>,
|
||||
params: &DecLockRefParams,
|
||||
lock_host: bool,
|
||||
) {
|
||||
unimplemented!()
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use super::*;
|
||||
use crate::components::FULL;
|
||||
use crate::components::{ComponentSet, FULL};
|
||||
use crate::node::NodeAccessError;
|
||||
use crate::test_utils::accumulate_step;
|
||||
use crate::unified_tree_core::CacheInitParams;
|
||||
@@ -625,10 +625,9 @@ fn inc_lock_ref_locks_the_device_path() {
|
||||
let mut tc = core();
|
||||
let (n1, n2) = lock_chain(&mut tc);
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(n2).id)
|
||||
.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(result.delta, Some(5));
|
||||
assert!(result.skip_lock_node_ids.is_empty());
|
||||
assert_eq!(tc.arena.device_lock_ref(n1, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(n2, FULL), 1);
|
||||
let state = tc.component_state(FULL);
|
||||
@@ -641,10 +640,10 @@ fn inc_lock_ref_locks_the_device_path() {
|
||||
fn inc_lock_ref_again_only_bumps_the_refs() {
|
||||
let mut tc = core();
|
||||
let (n1, n2) = lock_chain(&mut tc);
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(n2).id)
|
||||
.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(result.delta, Some(0));
|
||||
assert_eq!(tc.arena.device_lock_ref(n1, FULL), 2);
|
||||
@@ -659,10 +658,10 @@ fn inc_lock_ref_counts_only_newly_locked_nodes() {
|
||||
// n1 is already locked via its own path; locking n2 moves only n2's tokens.
|
||||
let mut tc = core();
|
||||
let (n1, n2) = lock_chain(&mut tc);
|
||||
tc.inc_lock_ref(tc.arena.node(n1).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n1).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(n2).id)
|
||||
.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(result.delta, Some(3));
|
||||
assert_eq!(tc.arena.device_lock_ref(n1, FULL), 2);
|
||||
@@ -673,8 +672,9 @@ fn inc_lock_ref_counts_only_newly_locked_nodes() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn inc_lock_ref_collects_the_evicted_bottom_segment() {
|
||||
// n2 and n3 are evicted (no device value): the walk records both and locks only n1.
|
||||
fn inc_lock_ref_counts_the_evicted_bottom_segment() {
|
||||
// n2 and n3 are evicted (no device value): counted in the segment with no
|
||||
// ledger move; only n1's tokens turn protected.
|
||||
let mut tc = core();
|
||||
let root = tc.arena.root();
|
||||
let n1 = tc
|
||||
@@ -709,16 +709,12 @@ fn inc_lock_ref_collects_the_evicted_bottom_segment() {
|
||||
tc.component_state_mut(FULL).evictable_size = 2;
|
||||
tc.evictable_device_leaves.add(n1);
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(n3).id)
|
||||
.inc_lock_ref(tc.arena.node(n3).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(result.delta, Some(2));
|
||||
assert_eq!(
|
||||
result.skip_lock_node_ids[&FULL],
|
||||
HashSet::from([tc.arena.node(n2).id, tc.arena.node(n3).id])
|
||||
);
|
||||
assert_eq!(tc.arena.device_lock_ref(n1, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(n2, FULL), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(n3, FULL), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(n2, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(n3, FULL), 1);
|
||||
// The locked ancestor leaves the D-leaf set.
|
||||
assert!(!tc.evictable_device_leaves.contains(n1));
|
||||
}
|
||||
@@ -728,15 +724,18 @@ fn lock_round_trips_on_a_root_anchor_are_noops() {
|
||||
let mut tc = core();
|
||||
let root = tc.arena.root();
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(root).id)
|
||||
.inc_lock_ref(tc.arena.node(root).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(result.delta, Some(0));
|
||||
assert!(result.skip_lock_node_ids.is_empty());
|
||||
// The protected root keeps its construction-time lock through the pair.
|
||||
assert_eq!(tc.arena.device_lock_ref(root, FULL), 1);
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(root).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -760,7 +759,7 @@ fn lock_walks_stop_at_the_root_of_a_salted_chain() {
|
||||
.set_device_value(n1, FULL, Tensor::from_slice(&[0i64, 1]));
|
||||
tc.component_state_mut(FULL).evictable_size = 2;
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(n1).id)
|
||||
.inc_lock_ref(tc.arena.node(n1).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(result.delta, Some(2));
|
||||
assert_eq!(tc.arena.device_lock_ref(n1, FULL), 1);
|
||||
@@ -769,7 +768,11 @@ fn lock_walks_stop_at_the_root_of_a_salted_chain() {
|
||||
// The release walk stops at the same boundary.
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n1).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -798,17 +801,20 @@ fn lock_walks_treat_a_present_but_empty_value_as_device_on() {
|
||||
Tensor::from_slice(&empty),
|
||||
);
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(n1).id)
|
||||
.inc_lock_ref(tc.arena.node(n1).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
// A present-but-empty value is device-on (Python `value is not None`):
|
||||
// locked, zero tokens moved.
|
||||
assert_eq!(result.delta, Some(0));
|
||||
assert_eq!(tc.arena.device_lock_ref(n1, FULL), 1);
|
||||
assert!(result.skip_lock_node_ids.is_empty());
|
||||
// The release side moves the same zero tokens back.
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n1).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -822,11 +828,15 @@ fn lock_walks_treat_a_present_but_empty_value_as_device_on() {
|
||||
fn dec_lock_ref_unlocks_and_restores_sizes() {
|
||||
let mut tc = core();
|
||||
let (n1, n2) = lock_chain(&mut tc);
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n2).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -875,19 +885,15 @@ fn dec_lock_ref_replays_the_skip_set() {
|
||||
.set_device_value(n1, FULL, Tensor::from_slice(&[0i64, 1]));
|
||||
tc.component_state_mut(FULL).evictable_size = 2;
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(n3).id)
|
||||
.inc_lock_ref(tc.arena.node(n3).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
let params = DecLockRefParams {
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
// The still-evicted n2 and n3 are skipped instead of tripping the lock asserts.
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n3).id,
|
||||
Some(¶ms),
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(tc.arena.node(n3).id, ¶ms, /* skip_swa = */ false)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.device_lock_ref(n1, FULL), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(n2, FULL), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(n3, FULL), 0);
|
||||
@@ -897,7 +903,7 @@ fn dec_lock_ref_replays_the_skip_set() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn temp_lock_skips_the_evicted_anchor_and_mirrors_on_release() {
|
||||
fn temp_lock_counts_the_evicted_anchor_and_mirrors_on_release() {
|
||||
// Chain root -> a -> y -> anchor with FULL device values; the anchor is evicted.
|
||||
let mut tc = core();
|
||||
let root = tc.arena.root();
|
||||
@@ -933,34 +939,32 @@ fn temp_lock_skips_the_evicted_anchor_and_mirrors_on_release() {
|
||||
tc.arena
|
||||
.set_device_value(y, FULL, Tensor::from_slice(&[0i64]));
|
||||
tc.component_state_mut(FULL).evictable_size = 3;
|
||||
// The temp lock records the evicted anchor and locks only its ancestors.
|
||||
// The temp lock counts the evicted anchor too (no ledger move).
|
||||
let temp_lock = tc
|
||||
.inc_lock_ref(tc.arena.node(anchor).id)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.device_lock_ref(anchor, FULL), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(y, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, FULL), 1);
|
||||
assert_eq!(
|
||||
temp_lock.skip_lock_node_ids[&FULL],
|
||||
HashSet::from([tc.arena.node(anchor).id])
|
||||
);
|
||||
// A load-back restores the anchor; the second acquire covers it.
|
||||
tc.arena
|
||||
.set_device_value(anchor, FULL, Tensor::from_slice(&[0i64]));
|
||||
let second_lock = tc
|
||||
.inc_lock_ref(tc.arena.node(anchor).id)
|
||||
.inc_lock_ref(tc.arena.node(anchor).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.device_lock_ref(anchor, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(y, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, FULL), 1);
|
||||
// A load-back restores the anchor under the held lock: credited to
|
||||
// protected, as the production commit does. The second acquire stacks.
|
||||
tc.arena
|
||||
.set_device_value(anchor, FULL, Tensor::from_slice(&[0i64]));
|
||||
tc.inc_protected_size(FULL, 1);
|
||||
let second_lock = tc
|
||||
.inc_lock_ref(tc.arena.node(anchor).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.device_lock_ref(anchor, FULL), 2);
|
||||
assert_eq!(tc.arena.device_lock_ref(y, FULL), 2);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, FULL), 2);
|
||||
// Releasing the temp lock mirrors its skip set: the anchor keeps its lock.
|
||||
// Each release takes back exactly its own refs.
|
||||
let temp_params = DecLockRefParams {
|
||||
skip_lock_node_ids: temp_lock.skip_lock_node_ids,
|
||||
skipped_lock_components: temp_lock.skipped_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(anchor).id,
|
||||
Some(&temp_params),
|
||||
&temp_params,
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -968,12 +972,12 @@ fn temp_lock_skips_the_evicted_anchor_and_mirrors_on_release() {
|
||||
assert_eq!(tc.arena.device_lock_ref(y, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, FULL), 1);
|
||||
let second_params = DecLockRefParams {
|
||||
skip_lock_node_ids: second_lock.skip_lock_node_ids,
|
||||
skipped_lock_components: second_lock.skipped_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(anchor).id,
|
||||
Some(&second_params),
|
||||
&second_params,
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -983,9 +987,9 @@ fn temp_lock_skips_the_evicted_anchor_and_mirrors_on_release() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "has no FULL device value")]
|
||||
fn dec_lock_ref_panics_without_replaying_the_skip_set() {
|
||||
// Dropping the acquire's skip set makes the release walk hit the tombstone.
|
||||
#[should_panic(expected = "FULL segment release hit lock_ref=0")]
|
||||
fn dec_lock_ref_panics_on_double_release() {
|
||||
// The second, unpaired release hits the already-unlocked segment.
|
||||
let mut tc = core();
|
||||
let root = tc.arena.root();
|
||||
let n1 = tc
|
||||
@@ -1009,11 +1013,25 @@ fn dec_lock_ref_panics_without_replaying_the_skip_set() {
|
||||
tc.arena
|
||||
.set_device_value(n1, FULL, Tensor::from_slice(&[0i64, 1]));
|
||||
tc.component_state_mut(FULL).evictable_size = 2;
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n2).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n2).id,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -1023,11 +1041,15 @@ fn dec_lock_ref_panics_without_replaying_the_skip_set() {
|
||||
fn dec_lock_ref_with_skip_swa_still_releases_full() {
|
||||
let mut tc = core();
|
||||
let (_n1, n2) = lock_chain(&mut tc);
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n2).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ true,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -1039,13 +1061,17 @@ fn nested_locks_release_pairwise() {
|
||||
// Two acquires then two releases: sizes move only on the outermost pair.
|
||||
let mut tc = core();
|
||||
let (_n1, n2) = lock_chain(&mut tc);
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n2).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -1055,7 +1081,11 @@ fn nested_locks_release_pairwise() {
|
||||
assert!(!tc.evictable_device_leaves.contains(n2));
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n2).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -1066,13 +1096,17 @@ fn nested_locks_release_pairwise() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "is not locked")]
|
||||
#[should_panic(expected = "FULL segment release hit lock_ref=0")]
|
||||
fn dec_lock_ref_panics_on_an_unlocked_node() {
|
||||
let mut tc = core();
|
||||
let (_n1, n2) = lock_chain(&mut tc);
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n2).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -1105,7 +1139,7 @@ fn inc_lock_ref_panics_on_an_evicted_ancestor() {
|
||||
tc.arena
|
||||
.set_device_value(n2, FULL, Tensor::from_slice(&[0i64, 1, 2]));
|
||||
tc.component_state_mut(FULL).evictable_size = 3;
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n2).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
}
|
||||
|
||||
@@ -1125,7 +1159,7 @@ fn inc_lock_ref_panics_when_evictable_size_is_unaccounted() {
|
||||
.unwrap();
|
||||
tc.arena
|
||||
.set_device_value(n1, FULL, Tensor::from_slice(&[0i64]));
|
||||
tc.inc_lock_ref(tc.arena.node(n1).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n1).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
}
|
||||
|
||||
@@ -1140,7 +1174,11 @@ fn dec_lock_ref_panics_on_protected_underflow() {
|
||||
.set_lock_ref_(ValueSlotIdx::device(FULL), 1);
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n2).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -1183,7 +1221,6 @@ fn inc_host_lock_ref_pins_the_backuped_anchor() {
|
||||
.inc_host_lock_ref(tc.arena.node(node).id)
|
||||
.expect("live test node");
|
||||
assert_eq!(result.delta, None);
|
||||
assert!(result.skip_lock_node_ids.is_empty());
|
||||
assert_eq!(tc.arena.host_lock_ref(node, FULL), 1);
|
||||
// The pinned anchor leaves the H-leaf set; the device tier is untouched.
|
||||
assert!(!tc.evictable_host_leaves.contains(node));
|
||||
@@ -1256,7 +1293,7 @@ fn host_lock_round_trips_on_a_root_anchor_are_noops() {
|
||||
.expect("live test node");
|
||||
assert_eq!(result.delta, None);
|
||||
assert_eq!(tc.arena.host_lock_ref(root, FULL), 0);
|
||||
tc.dec_host_lock_ref(tc.arena.node(root).id, /* params = */ None)
|
||||
tc.dec_host_lock_ref(tc.arena.node(root).id, &DecLockRefParams::default())
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(root, FULL), 0);
|
||||
}
|
||||
@@ -1281,7 +1318,7 @@ fn dec_host_lock_ref_unpins_and_restores_the_h_leaf_set() {
|
||||
tc.component_state_mut(FULL).evictable_size = 7;
|
||||
tc.inc_host_lock_ref(tc.arena.node(node).id)
|
||||
.expect("live test node");
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, /* params = */ None)
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, &DecLockRefParams::default())
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(node, FULL), 0);
|
||||
assert!(tc.evictable_host_leaves.contains(node));
|
||||
@@ -1294,7 +1331,7 @@ fn dec_host_lock_ref_unpins_and_restores_the_h_leaf_set() {
|
||||
fn dec_host_lock_ref_on_an_unlocked_anchor_is_a_noop() {
|
||||
let mut tc = core();
|
||||
let node = host_lock_anchor(&mut tc);
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, /* params = */ None)
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, &DecLockRefParams::default())
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(node, FULL), 0);
|
||||
}
|
||||
@@ -1307,7 +1344,7 @@ fn dec_host_lock_ref_keeps_the_counter_when_the_host_value_is_gone() {
|
||||
tc.inc_host_lock_ref(tc.arena.node(node).id)
|
||||
.expect("live test node");
|
||||
let _ = tc.arena.take_host_value(node, FULL);
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, /* params = */ None)
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, &DecLockRefParams::default())
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(node, FULL), 1);
|
||||
}
|
||||
@@ -1318,7 +1355,7 @@ fn host_lock_round_trip_under_write_back_is_a_pure_counter() {
|
||||
let (_n1, n2) = lock_chain(&mut tc);
|
||||
tc.inc_host_lock_ref(tc.arena.node(n2).id)
|
||||
.expect("live test node");
|
||||
tc.dec_host_lock_ref(tc.arena.node(n2).id, /* params = */ None)
|
||||
tc.dec_host_lock_ref(tc.arena.node(n2).id, &DecLockRefParams::default())
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(n2, FULL), 0);
|
||||
let state = tc.component_state(FULL);
|
||||
@@ -1346,7 +1383,10 @@ fn release_host_arm_updates_the_h_leaf_set_without_the_dispatcher() {
|
||||
tc.inc_host_lock_ref(tc.arena.node(node).id)
|
||||
.expect("live test node");
|
||||
FullComponent.release_component_lock(
|
||||
&mut tc, node, /* params = */ None, /* lock_host = */ true,
|
||||
&mut tc,
|
||||
node,
|
||||
&DecLockRefParams::default(),
|
||||
/* lock_host = */ true,
|
||||
);
|
||||
assert!(tc.evictable_host_leaves.contains(node));
|
||||
}
|
||||
@@ -1359,11 +1399,11 @@ fn nested_host_locks_release_pairwise() {
|
||||
.expect("live test node");
|
||||
tc.inc_host_lock_ref(tc.arena.node(node).id)
|
||||
.expect("live test node");
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, /* params = */ None)
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, &DecLockRefParams::default())
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(node, FULL), 1);
|
||||
assert!(!tc.evictable_host_leaves.contains(node));
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, /* params = */ None)
|
||||
tc.dec_host_lock_ref(tc.arena.node(node).id, &DecLockRefParams::default())
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(node, FULL), 0);
|
||||
assert!(tc.evictable_host_leaves.contains(node));
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use super::*;
|
||||
use crate::components::{FULL, MAMBA, SWA};
|
||||
use crate::components::{ComponentSet, FULL, MAMBA, SWA};
|
||||
use crate::test_utils::{accumulate_step, action_kinds};
|
||||
use crate::unified_lru_list::UnifiedLRUList;
|
||||
|
||||
@@ -292,7 +292,6 @@ fn device_lock_moves_the_slot_between_evictable_and_protected_once() {
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert!(result.skip_lock_node_ids.is_empty());
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 0);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 1);
|
||||
mamba.acquire_component_lock(
|
||||
@@ -303,25 +302,34 @@ fn device_lock_moves_the_slot_between_evictable_and_protected_once() {
|
||||
);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 1);
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 2);
|
||||
mamba.release_component_lock(&mut tc, a, None, /* lock_host = */ false);
|
||||
mamba.release_component_lock(
|
||||
&mut tc,
|
||||
a,
|
||||
&DecLockRefParams::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 1);
|
||||
mamba.release_component_lock(&mut tc, a, None, /* lock_host = */ false);
|
||||
mamba.release_component_lock(
|
||||
&mut tc,
|
||||
a,
|
||||
&DecLockRefParams::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 1);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 0);
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn skip_aware_lock_records_only_the_mamba_target() {
|
||||
fn lock_without_mamba_records_the_receipt_and_leaves_mamba_evictable() {
|
||||
let (mut tc, parent, leaf) = hybrid_lock_core();
|
||||
let leaf_handle = tc.arena.node(leaf).id;
|
||||
|
||||
let result = tc
|
||||
.inc_lock_ref_with_skip(leaf_handle, &[MAMBA])
|
||||
.inc_lock_ref(leaf_handle, ComponentSet::of(MAMBA))
|
||||
.expect("live test node");
|
||||
|
||||
assert_eq!(result.skip_lock_node_ids[&MAMBA].len(), 1);
|
||||
assert!(result.skip_lock_node_ids[&MAMBA].contains(&leaf_handle));
|
||||
assert!(result.skipped_lock_components.contains(MAMBA));
|
||||
assert_eq!(tc.arena.node(parent).device_lock_ref(MAMBA), 0);
|
||||
assert_eq!(tc.arena.node(leaf).device_lock_ref(MAMBA), 0);
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 2);
|
||||
@@ -329,36 +337,45 @@ fn skip_aware_lock_records_only_the_mamba_target() {
|
||||
assert_eq!(tc.arena.node(parent).device_lock_ref(FULL), 1);
|
||||
assert_eq!(tc.arena.node(leaf).device_lock_ref(FULL), 1);
|
||||
|
||||
tc.dec_lock_ref(
|
||||
leaf_handle,
|
||||
Some(&DecLockRefParams {
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
..Default::default()
|
||||
}),
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
// The receipt replays exactly what was taken: FULL only.
|
||||
let params = DecLockRefParams {
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
tc.dec_lock_ref(leaf_handle, ¶ms, /* skip_swa = */ false)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.node(parent).device_lock_ref(FULL), 0);
|
||||
assert_eq!(tc.arena.node(leaf).device_lock_ref(FULL), 0);
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 2);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn swa_only_release_honors_a_skipped_mamba_target() {
|
||||
fn swa_only_release_spares_another_holders_mamba_lock() {
|
||||
let (mut tc, _parent, leaf) = hybrid_lock_core();
|
||||
let leaf_handle = tc.arena.node(leaf).id;
|
||||
let owner = tc.inc_lock_ref(leaf_handle).expect("live test node");
|
||||
let skipped = tc
|
||||
.inc_lock_ref_with_skip(leaf_handle, &[MAMBA])
|
||||
let owner = tc
|
||||
.inc_lock_ref(leaf_handle, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
let holder = tc
|
||||
.inc_lock_ref(leaf_handle, ComponentSet::of(MAMBA))
|
||||
.expect("live test node");
|
||||
assert!(!owner.skipped_lock_components.contains(MAMBA));
|
||||
assert!(holder.skipped_lock_components.contains(MAMBA));
|
||||
assert_eq!(tc.arena.node(leaf).device_lock_ref(MAMBA), 1);
|
||||
|
||||
// The holder's early SWA release must not drop the owner's mamba lock.
|
||||
let holder_params = DecLockRefParams {
|
||||
swa_uuid_for_lock: holder.swa_uuid_for_lock,
|
||||
skipped_lock_components: holder.skipped_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
let mut device_frees = HashMap::new();
|
||||
let mut host_frees = HashMap::new();
|
||||
tc.dec_swa_lock_only_with_skip(
|
||||
tc.dec_swa_lock_only(
|
||||
leaf_handle,
|
||||
skipped.swa_uuid_for_lock,
|
||||
Some(&skipped.skip_lock_node_ids),
|
||||
&holder_params,
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)
|
||||
@@ -369,33 +386,22 @@ fn swa_only_release_honors_a_skipped_mamba_target() {
|
||||
assert_eq!(tc.arena.node(leaf).device_lock_ref(MAMBA), 1);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 1);
|
||||
|
||||
let skipped_params = DecLockRefParams {
|
||||
swa_uuid_for_lock: skipped.swa_uuid_for_lock,
|
||||
skip_lock_node_ids: skipped.skip_lock_node_ids,
|
||||
..Default::default()
|
||||
};
|
||||
tc.dec_lock_ref(
|
||||
leaf_handle,
|
||||
Some(&skipped_params),
|
||||
/* skip_swa = */ true,
|
||||
)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(leaf_handle, &holder_params, /* skip_swa = */ true)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.node(leaf).device_lock_ref(MAMBA), 1);
|
||||
let owner_params = DecLockRefParams {
|
||||
swa_uuid_for_lock: owner.swa_uuid_for_lock,
|
||||
skip_lock_node_ids: owner.skip_lock_node_ids,
|
||||
skipped_lock_components: owner.skipped_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
tc.dec_lock_ref(
|
||||
leaf_handle,
|
||||
Some(&owner_params),
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(leaf_handle, &owner_params, /* skip_swa = */ false)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.node(leaf).device_lock_ref(MAMBA), 0);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tombstone_lock_is_recorded_and_replayed_at_release() {
|
||||
fn tombstone_lock_is_counted_with_no_ledger_move() {
|
||||
let mut tc = mamba_core(/* page_size = */ 1);
|
||||
let [a] = chain::<1>(&mut tc);
|
||||
let mamba = mamba_component();
|
||||
@@ -405,14 +411,15 @@ fn tombstone_lock_is_recorded_and_replayed_at_release() {
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert!(result.skip_lock_node_ids[&MAMBA].contains(&tc.arena.node(a).id));
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 0);
|
||||
// The replayed skip set keeps the release from touching the node.
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 1);
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 0);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 0);
|
||||
// The paired release decrements the counted tombstone, ledger untouched.
|
||||
let params = DecLockRefParams {
|
||||
skip_lock_node_ids: result.skip_lock_node_ids.clone(),
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
..DecLockRefParams::default()
|
||||
};
|
||||
mamba.release_component_lock(&mut tc, a, Some(¶ms), /* lock_host = */ false);
|
||||
mamba.release_component_lock(&mut tc, a, ¶ms, /* lock_host = */ false);
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 0);
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 0);
|
||||
}
|
||||
@@ -428,8 +435,12 @@ fn root_locks_are_noops() {
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert!(result.skip_lock_node_ids.is_empty());
|
||||
mamba.release_component_lock(&mut tc, root, None, /* lock_host = */ false);
|
||||
mamba.release_component_lock(
|
||||
&mut tc,
|
||||
root,
|
||||
&DecLockRefParams::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 0);
|
||||
}
|
||||
|
||||
@@ -448,7 +459,12 @@ fn host_lock_detaches_and_reattaches_the_host_lru() {
|
||||
);
|
||||
assert!(!tc.host_lru_list(MAMBA).in_list(Some(a)));
|
||||
assert_eq!(tc.arena.node(a).host_lock_ref(MAMBA), 1);
|
||||
mamba.release_component_lock(&mut tc, a, None, /* lock_host = */ true);
|
||||
mamba.release_component_lock(
|
||||
&mut tc,
|
||||
a,
|
||||
&DecLockRefParams::default(),
|
||||
/* lock_host = */ true,
|
||||
);
|
||||
assert!(tc.host_lru_list(MAMBA).in_list(Some(a)));
|
||||
assert_eq!(tc.arena.node(a).host_lock_ref(MAMBA), 0);
|
||||
}
|
||||
@@ -466,7 +482,12 @@ fn host_unlock_skips_the_lru_for_device_backed_nodes() {
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ true,
|
||||
);
|
||||
mamba.release_component_lock(&mut tc, a, None, /* lock_host = */ true);
|
||||
mamba.release_component_lock(
|
||||
&mut tc,
|
||||
a,
|
||||
&DecLockRefParams::default(),
|
||||
/* lock_host = */ true,
|
||||
);
|
||||
assert!(!tc.host_lru_list(MAMBA).in_list(Some(a)));
|
||||
}
|
||||
|
||||
@@ -1820,7 +1841,7 @@ fn branching_from_a_host_full_hit_is_reusable_after_insert() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn skip_set_release_after_a_restore_and_relock_keeps_the_new_lock() {
|
||||
fn release_after_a_restore_and_relock_keeps_the_other_lock() {
|
||||
let mut tc = mamba_core(/* page_size = */ 1);
|
||||
let [a] = chain::<1>(&mut tc);
|
||||
let mamba = mamba_component();
|
||||
@@ -1830,28 +1851,33 @@ fn skip_set_release_after_a_restore_and_relock_keeps_the_new_lock() {
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert!(first.skip_lock_node_ids[&MAMBA].contains(&tc.arena.node(a).id));
|
||||
// The tombstone is restored and a second request locks it before the
|
||||
// first release replays its skip set.
|
||||
set_mamba_device(&mut tc, a, 7);
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 1);
|
||||
// The tombstone is restored under the held lock (credited to protected)
|
||||
// and a second request stacks its own lock on it.
|
||||
tc.set_component_device_value_(a, MAMBA, Tensor::from_slice(&[7i64]));
|
||||
let _ = mamba.acquire_component_lock(
|
||||
&mut tc,
|
||||
a,
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 1);
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 2);
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 0);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 1);
|
||||
let params = DecLockRefParams {
|
||||
skip_lock_node_ids: first.skip_lock_node_ids.clone(),
|
||||
skipped_lock_components: first.skipped_lock_components,
|
||||
..DecLockRefParams::default()
|
||||
};
|
||||
mamba.release_component_lock(&mut tc, a, Some(¶ms), /* lock_host = */ false);
|
||||
// The replayed skip keeps the restored node's fresh lock intact.
|
||||
mamba.release_component_lock(&mut tc, a, ¶ms, /* lock_host = */ false);
|
||||
// The first release takes back exactly its own ref.
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 1);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 1);
|
||||
mamba.release_component_lock(&mut tc, a, None, /* lock_host = */ false);
|
||||
mamba.release_component_lock(
|
||||
&mut tc,
|
||||
a,
|
||||
&DecLockRefParams::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
assert_eq!(tc.arena.node(a).device_lock_ref(MAMBA), 0);
|
||||
assert_eq!(tc.evictable_size_(MAMBA), 1);
|
||||
assert_eq!(tc.protected_size_(MAMBA), 0);
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use super::*;
|
||||
use crate::components::{FULL, MAMBA, SWA};
|
||||
use crate::components::{ComponentSet, FULL, MAMBA, SWA};
|
||||
use crate::test_utils::{accumulate_step, action_kinds};
|
||||
use crate::unified_tree_core::CacheInitParams;
|
||||
|
||||
@@ -766,18 +766,32 @@ fn insert_overlap_recovers_a_tombstone_inside_the_window() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "tombstone Swa lock_ref should be 0, node")]
|
||||
fn insert_overlap_panics_on_a_locked_swa_tombstone() {
|
||||
fn insert_overlap_recovers_a_locked_swa_tombstone() {
|
||||
let mut tc = swa_core(/* window = */ 8, /* page_size = */ 1);
|
||||
tc.insert(&insert_params_swa(&vec![1, 2, 3], &[10, 11, 12], 0, 0));
|
||||
let root = tc.arena.root();
|
||||
let leaf = child_of(&tc, root, &[1]);
|
||||
// The rebuild is deferred, so the leaf is still an SWA tombstone; a raw
|
||||
// lock on it breaks the tombstones-are-unlocked contract.
|
||||
// A segment lock may hold an SWA tombstone; the co-held FULL lock (the
|
||||
// full >= swa protocol invariant) forces the Recover branch, so the
|
||||
// locked full stays on the node.
|
||||
tc.arena
|
||||
.node_mut(leaf)
|
||||
.set_lock_ref_(ValueSlotIdx::device(SWA), 1);
|
||||
tc.insert(&insert_params_swa(&vec![1, 2, 3], &[20, 21, 22], 0, 0));
|
||||
tc.arena
|
||||
.node_mut(leaf)
|
||||
.set_lock_ref_(ValueSlotIdx::device(FULL), 1);
|
||||
let result = tc.insert(&insert_params_swa(&vec![1, 2, 3], &[20, 21, 22], 0, 0));
|
||||
assert!(
|
||||
tc.arena
|
||||
.device_value(leaf, FULL)
|
||||
.equal(&Tensor::from_slice(&[10i64, 11, 12]))
|
||||
);
|
||||
assert!(
|
||||
result
|
||||
.cache_actions
|
||||
.iter()
|
||||
.any(|action| matches!(action, CacheAction::RecoverSwaWithLockedFull { .. }))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1576,7 +1590,7 @@ fn acquire_lock_reuses_the_stamped_uuid_and_shifts_sizes_once() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn acquire_lock_skips_tombstones_and_records_them() {
|
||||
fn acquire_lock_counts_tombstones_toward_the_window() {
|
||||
let mut tc = swa_core(/* window = */ 2, /* page_size = */ 1);
|
||||
let [a, b, c] = chain(&mut tc);
|
||||
store_swa_device(&mut tc, a);
|
||||
@@ -1587,14 +1601,13 @@ fn acquire_lock_skips_tombstones_and_records_them() {
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
// The valueless b is recorded and skipped; the window fills at a.
|
||||
// The valueless b is counted too (no ledger move); position-based
|
||||
// coverage fills the window at b, so a stays outside the segment.
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 1);
|
||||
assert_eq!(result.skip_lock_node_ids[&SWA].len(), 1);
|
||||
assert!(result.skip_lock_node_ids[&SWA].contains(&tc.arena.node(b).id));
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 0);
|
||||
assert!(result.swa_uuid_for_lock.is_some());
|
||||
assert_eq!(node_swa_uuid(&tc, a), result.swa_uuid_for_lock);
|
||||
assert_eq!(node_swa_uuid(&tc, b), result.swa_uuid_for_lock);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1624,11 +1637,14 @@ fn inc_lock_ref_runs_full_and_swa_walks_together() {
|
||||
store_swa_device(&mut tc, b);
|
||||
store_swa_device(&mut tc, c);
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(c).id)
|
||||
.inc_lock_ref(tc.arena.node(c).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
// FULL sees a valueless path (skip segment only); SWA locks its window.
|
||||
// FULL counts its valueless bottom segment (no ledger move); SWA locks
|
||||
// its window.
|
||||
assert_eq!(result.delta, Some(0));
|
||||
assert_eq!(result.skip_lock_node_ids[&FULL].len(), 3);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, FULL), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, FULL), 1);
|
||||
assert!(result.swa_uuid_for_lock.is_some());
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 1);
|
||||
@@ -1657,10 +1673,10 @@ fn inc_host_lock_ref_runs_full_and_swa_host_arms_together() {
|
||||
// The release replays the acquire's uuid and unwinds both arms.
|
||||
let params = DecLockRefParams {
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
tc.dec_host_lock_ref(tc.arena.node(c).id, Some(¶ms))
|
||||
tc.dec_host_lock_ref(tc.arena.node(c).id, ¶ms)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(c, FULL), 0);
|
||||
assert_eq!(tc.arena.host_lock_ref(c, SWA), 0);
|
||||
@@ -1690,10 +1706,10 @@ fn dec_host_lock_ref_with_the_inner_uuid_leaves_an_outer_window_pinned() {
|
||||
// window's lock above the boundary survives.
|
||||
let params = DecLockRefParams {
|
||||
swa_uuid_for_host_lock: inner.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: inner.skip_lock_node_ids,
|
||||
skipped_lock_components: inner.skipped_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
tc.dec_host_lock_ref(tc.arena.node(c).id, Some(¶ms))
|
||||
tc.dec_host_lock_ref(tc.arena.node(c).id, ¶ms)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.host_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.host_lock_ref(b, SWA), 1);
|
||||
@@ -1758,7 +1774,7 @@ fn acquire_host_lock_reuses_the_stamped_uuid_and_skips_unlisted_nodes() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn acquire_host_lock_skips_host_tombstones_and_records_them() {
|
||||
fn acquire_host_lock_counts_host_tombstones_toward_the_window() {
|
||||
let mut tc = swa_core(/* window = */ 2, /* page_size = */ 1);
|
||||
let [a, b, c] = chain(&mut tc);
|
||||
set_swa_host(&mut tc, a);
|
||||
@@ -1769,12 +1785,12 @@ fn acquire_host_lock_skips_host_tombstones_and_records_them() {
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ true,
|
||||
);
|
||||
// The host-valueless b is counted too; position-based coverage fills
|
||||
// the window at b, so a stays outside the segment.
|
||||
assert_eq!(tc.arena.host_lock_ref(c, SWA), 1);
|
||||
assert_eq!(tc.arena.host_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.host_lock_ref(a, SWA), 1);
|
||||
assert_eq!(result.skip_lock_node_ids[&SWA].len(), 1);
|
||||
assert!(result.skip_lock_node_ids[&SWA].contains(&tc.arena.node(b).id));
|
||||
assert_eq!(node_swa_host_uuid(&tc, a), result.swa_uuid_for_host_lock);
|
||||
assert_eq!(tc.arena.host_lock_ref(b, SWA), 1);
|
||||
assert_eq!(tc.arena.host_lock_ref(a, SWA), 0);
|
||||
assert_eq!(node_swa_host_uuid(&tc, b), result.swa_uuid_for_host_lock);
|
||||
assert!(result.swa_uuid_for_host_lock.is_some());
|
||||
}
|
||||
|
||||
@@ -1996,11 +2012,12 @@ fn release_lock_returns_the_window_to_evictable() {
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ false);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ false);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 0);
|
||||
@@ -2029,11 +2046,12 @@ fn release_lock_keeps_sizes_while_other_locks_remain() {
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: first.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: first.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: first.skip_lock_node_ids,
|
||||
skipped_lock_components: first.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ false);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ false);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 1);
|
||||
assert_eq!(tc.swa_evictable_size(), 1);
|
||||
@@ -2056,11 +2074,12 @@ fn release_lock_replays_the_tombstone_skips() {
|
||||
// b gained a device value AFTER the acquire recorded it as a tombstone.
|
||||
store_swa_device(&mut tc, b);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ false);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ false);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 0);
|
||||
@@ -2087,11 +2106,12 @@ fn release_lock_stops_at_the_window_uuid() {
|
||||
.node_mut(a)
|
||||
.set_lock_ref_(ValueSlotIdx::device(SWA), 1);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ false);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ false);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 1);
|
||||
@@ -2120,11 +2140,12 @@ fn release_host_lock_stops_at_the_host_uuid_boundary() {
|
||||
/* lock_host = */ true,
|
||||
);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ true);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ true);
|
||||
assert_eq!(tc.arena.host_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.host_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.host_lock_ref(a, SWA), 1);
|
||||
@@ -2132,34 +2153,27 @@ fn release_host_lock_stops_at_the_host_uuid_boundary() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn release_lock_without_params_passes_over_an_unlocked_middle_node() {
|
||||
#[should_panic(expected = "SWA segment release hit lock_ref=0")]
|
||||
fn release_lock_without_the_boundary_uuid_dies_at_the_segment_edge() {
|
||||
let mut tc = swa_core(/* window = */ 1, /* page_size = */ 1);
|
||||
let [a, b, c] = chain(&mut tc);
|
||||
store_swa_device(&mut tc, a);
|
||||
store_swa_device(&mut tc, b);
|
||||
let [_a, _b, c] = chain(&mut tc);
|
||||
store_swa_device(&mut tc, c);
|
||||
let swa = swa_component(1);
|
||||
// The 1-atom window locks only the acquired node: c and a, never b.
|
||||
// The 1-atom window locks only the acquired node c.
|
||||
let _ = swa.acquire_component_lock(
|
||||
&mut tc,
|
||||
c,
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
let _ = swa.acquire_component_lock(
|
||||
// A receipt-less release overshoots the boundary into unlocked territory
|
||||
// and dies there instead of silently stealing whatever it crosses.
|
||||
swa.release_component_lock(
|
||||
&mut tc,
|
||||
a,
|
||||
IncLockRefResult::default(),
|
||||
c,
|
||||
&DecLockRefParams::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
swa.release_component_lock(
|
||||
&mut tc, c, /* params = */ None, /* lock_host = */ false,
|
||||
);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 0);
|
||||
assert_eq!(tc.swa_evictable_size(), 3);
|
||||
assert_eq!(tc.swa_protected_size(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -2178,11 +2192,12 @@ fn release_host_lock_reparks_tombstoned_host_nodes() {
|
||||
/* lock_host = */ true,
|
||||
);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ true);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ true);
|
||||
assert_eq!(tc.arena.host_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.host_lock_ref(b, SWA), 0);
|
||||
assert!(tc.host_lru_list(SWA).in_list(Some(c)));
|
||||
@@ -2198,19 +2213,16 @@ fn inc_then_dec_lock_ref_roundtrips_with_dec_params() {
|
||||
store_swa_device(&mut tc, b);
|
||||
store_swa_device(&mut tc, c);
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(c).id)
|
||||
.inc_lock_ref(tc.arena.node(c).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(c).id,
|
||||
Some(¶ms),
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(tc.arena.node(c).id, ¶ms, /* skip_swa = */ false)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.swa_evictable_size(), 3);
|
||||
@@ -2230,13 +2242,17 @@ fn dec_swa_lock_only_releases_swa_while_full_stays_locked() {
|
||||
// Fund FULL's evictable counter for its lock walk (raw slot sets skip it).
|
||||
tc.component_state_mut(FULL).evictable_size = 3;
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(c).id)
|
||||
.inc_lock_ref(tc.arena.node(c).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
let mut device_frees = HashMap::new();
|
||||
let mut host_frees = HashMap::new();
|
||||
tc.dec_swa_lock_only(
|
||||
tc.arena.node(c).id,
|
||||
result.swa_uuid_for_lock,
|
||||
&DecLockRefParams {
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)
|
||||
@@ -2273,7 +2289,11 @@ fn dec_swa_lock_only_evicts_a_fully_unlocked_device_leaf() {
|
||||
let mut host_frees = HashMap::new();
|
||||
tc.dec_swa_lock_only(
|
||||
tc.arena.node(c).id,
|
||||
result.swa_uuid_for_lock,
|
||||
&DecLockRefParams {
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)
|
||||
@@ -2296,7 +2316,11 @@ fn dec_swa_lock_only_is_a_noop_without_the_swa_component() {
|
||||
let mut host_frees = HashMap::new();
|
||||
tc.dec_swa_lock_only(
|
||||
tc.arena.node(root).id,
|
||||
None,
|
||||
&DecLockRefParams {
|
||||
swa_uuid_for_lock: None,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)
|
||||
@@ -2374,11 +2398,12 @@ fn release_lock_skip_set_leaves_a_relocked_tombstone_credited() {
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: first.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: first.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: first.skip_lock_node_ids,
|
||||
skipped_lock_components: first.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ false);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ false);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 1);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 0);
|
||||
@@ -2387,28 +2412,27 @@ fn release_lock_skip_set_leaves_a_relocked_tombstone_credited() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn release_lock_passes_over_uncredited_nodes_without_params() {
|
||||
#[should_panic(expected = "SWA segment release hit lock_ref=0")]
|
||||
fn double_release_with_one_receipt_dies_loud() {
|
||||
let mut tc = swa_core(/* window = */ 2, /* page_size = */ 1);
|
||||
let [a, b, c] = chain(&mut tc);
|
||||
store_swa_device(&mut tc, a);
|
||||
store_swa_device(&mut tc, b);
|
||||
store_swa_device(&mut tc, c);
|
||||
let swa = swa_component(2);
|
||||
let _ = swa.acquire_component_lock(
|
||||
let result = swa.acquire_component_lock(
|
||||
&mut tc,
|
||||
c,
|
||||
IncLockRefResult::default(),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
// No params: the walk crosses the never-credited a up to the root.
|
||||
swa.release_component_lock(
|
||||
&mut tc, c, /* params = */ None, /* lock_host = */ false,
|
||||
);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 0);
|
||||
assert_eq!(tc.swa_evictable_size(), 3);
|
||||
assert_eq!(tc.swa_protected_size(), 0);
|
||||
let params = DecLockRefParams {
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
..Default::default()
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ false);
|
||||
// Consuming the same receipt twice dies at the first unlocked node.
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -2435,7 +2459,11 @@ fn dec_swa_lock_only_releases_the_window_exactly_once() {
|
||||
let mut host_frees = HashMap::new();
|
||||
tc.dec_swa_lock_only(
|
||||
tc.arena.node(c).id,
|
||||
first.swa_uuid_for_lock,
|
||||
&DecLockRefParams {
|
||||
swa_uuid_for_lock: first.swa_uuid_for_lock,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)
|
||||
@@ -2447,7 +2475,11 @@ fn dec_swa_lock_only_releases_the_window_exactly_once() {
|
||||
assert_eq!(tc.swa_protected_size(), 2);
|
||||
tc.dec_swa_lock_only(
|
||||
tc.arena.node(c).id,
|
||||
first.swa_uuid_for_lock,
|
||||
&DecLockRefParams {
|
||||
swa_uuid_for_lock: first.swa_uuid_for_lock,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)
|
||||
@@ -2483,7 +2515,11 @@ fn dec_swa_lock_only_leaves_out_of_window_swa_locks_alone() {
|
||||
let mut host_frees = HashMap::new();
|
||||
tc.dec_swa_lock_only(
|
||||
tc.arena.node(c).id,
|
||||
result.swa_uuid_for_lock,
|
||||
&DecLockRefParams {
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)
|
||||
@@ -2496,7 +2532,8 @@ fn dec_swa_lock_only_leaves_out_of_window_swa_locks_alone() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn release_window_lock_passes_over_an_unlocked_valued_node() {
|
||||
#[should_panic(expected = "SWA window release hit lock_ref=0")]
|
||||
fn release_window_lock_without_the_uuid_dies_past_the_boundary() {
|
||||
let mut tc = swa_core(/* window = */ 2, /* page_size = */ 1);
|
||||
let [a, b, c] = chain(&mut tc);
|
||||
store_swa_device(&mut tc, a);
|
||||
@@ -2511,11 +2548,9 @@ fn release_window_lock_passes_over_an_unlocked_valued_node() {
|
||||
);
|
||||
let mut device_frees = HashMap::new();
|
||||
let mut host_frees = HashMap::new();
|
||||
// No uuid bound: the walk crosses the valued-but-unlocked a to the root.
|
||||
// Without the boundary uuid the walk crosses the segment edge into the
|
||||
// unlocked a and dies there instead of stealing.
|
||||
swa.release_window_lock(&mut tc, c, None, &mut device_frees, &mut host_frees);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(a, SWA), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -2533,7 +2568,8 @@ fn release_window_lock_passes_over_a_mid_chain_tombstone_without_a_uuid() {
|
||||
);
|
||||
let mut device_frees = HashMap::new();
|
||||
let mut host_frees = HashMap::new();
|
||||
// No uuid bound: the walk crosses the mid-chain tombstone b and releases a.
|
||||
// Walked-to-root acquire (window > chain): the uuid-less release counts
|
||||
// back through the mid-chain tombstone b and releases a.
|
||||
swa.release_window_lock(&mut tc, c, None, &mut device_frees, &mut host_frees);
|
||||
assert_eq!(tc.arena.device_lock_ref(c, SWA), 0);
|
||||
assert_eq!(tc.arena.device_lock_ref(b, SWA), 0);
|
||||
@@ -2558,11 +2594,12 @@ fn release_host_lock_does_not_repark_a_node_whose_host_value_was_taken() {
|
||||
// device value either, so the release has nothing to park.
|
||||
let _ = tc.arena.take_host_value(a, SWA);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, a, Some(¶ms), /* lock_host = */ true);
|
||||
swa.release_component_lock(&mut tc, a, ¶ms, /* lock_host = */ true);
|
||||
assert_eq!(tc.arena.host_lock_ref(a, SWA), 0);
|
||||
assert!(!tc.host_lru_list(SWA).in_list(Some(a)));
|
||||
}
|
||||
@@ -2584,11 +2621,12 @@ fn release_host_lock_skips_reparking_device_valued_nodes() {
|
||||
/* lock_host = */ true,
|
||||
);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ true);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ true);
|
||||
// Device-valued nodes never re-park in the host LRU on host release.
|
||||
assert!(!tc.host_lru_list(SWA).in_list(Some(c)));
|
||||
assert!(!tc.host_lru_list(SWA).in_list(Some(b)));
|
||||
@@ -2612,11 +2650,12 @@ fn release_host_lock_leaves_an_already_listed_node_listed() {
|
||||
// Something re-listed b while the lock was held (e.g. a split re-park).
|
||||
tc.host_lru_list_mut(SWA).insert_mru(b);
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: result.skip_lock_node_ids,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
swa.release_component_lock(&mut tc, c, Some(¶ms), /* lock_host = */ true);
|
||||
swa.release_component_lock(&mut tc, c, ¶ms, /* lock_host = */ true);
|
||||
assert!(tc.host_lru_list(SWA).in_list(Some(b)));
|
||||
assert!(tc.host_lru_list(SWA).in_list(Some(c)));
|
||||
let _ = a;
|
||||
@@ -2890,6 +2929,29 @@ fn redistribute_on_node_split_moves_the_swa_uuid_to_the_parent() {
|
||||
assert_eq!(node_swa_uuid(&tc, node), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn redistribute_on_node_split_preserves_host_lock_state() {
|
||||
let mut tc = swa_core(/* window = */ 4, /* page_size = */ 1);
|
||||
tc.insert(&insert_params_swa(&vec![1, 2], &[10, 11], 0, 0));
|
||||
let root = tc.arena.root();
|
||||
let node = child_of(&tc, root, &[1]);
|
||||
// A host-locked tombstone mid-IO: both halves must stay pinned and out
|
||||
// of the host LRU, and the host boundary uuid moves to the parent.
|
||||
tc.arena
|
||||
.set_host_value(node, SWA, Tensor::from_slice(&[70i64, 71]));
|
||||
tc.arena
|
||||
.node_mut(node)
|
||||
.set_lock_ref_(ValueSlotIdx::host(SWA), 1);
|
||||
tc.arena.node_mut(node).swa_host_uuid = Some(9);
|
||||
let (parent, _) = tc.split_node_(node, /* split_len = */ 1);
|
||||
assert_eq!(tc.arena.host_lock_ref(parent, SWA), 1);
|
||||
assert_eq!(tc.arena.host_lock_ref(node, SWA), 1);
|
||||
assert_eq!(node_swa_host_uuid(&tc, parent), Some(9));
|
||||
assert_eq!(node_swa_host_uuid(&tc, node), None);
|
||||
assert!(!tc.host_lru_list(SWA).in_list(Some(parent)));
|
||||
assert!(!tc.host_lru_list(SWA).in_list(Some(node)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn finalize_window_arithmetic_at_page_boundaries() {
|
||||
let mut tc = swa_core(/* window = */ 4, /* page_size = */ 2);
|
||||
@@ -3277,17 +3339,24 @@ fn reinsert_rejects_a_page_misaligned_boundary() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "tombstone Swa lock_ref should be 0 on unevict")]
|
||||
fn reinsert_rejects_a_locked_tombstone() {
|
||||
fn reinsert_rebuilds_a_locked_tombstone() {
|
||||
let mut tc = swa_core(/* window = */ 8, /* page_size = */ 1);
|
||||
tc.insert(&insert_params_swa(&vec![1, 2], &[10, 11], 0, 0));
|
||||
let root = tc.arena.root();
|
||||
let node = child_of(&tc, root, &[1]);
|
||||
evict_full(&mut tc, node, /* remaining_size = */ 0);
|
||||
// Segment locks count evicted nodes, so a locked tombstone is legal and
|
||||
// the re-insert rebuilds its SWA from the fresh KV.
|
||||
tc.arena
|
||||
.node_mut(node)
|
||||
.set_lock_ref_(ValueSlotIdx::device(SWA), 1);
|
||||
tc.insert(&insert_params_swa(&vec![1, 2], &[20, 21], 0, 0));
|
||||
let result = tc.insert(&insert_params_swa(&vec![1, 2], &[20, 21], 0, 0));
|
||||
assert!(
|
||||
result
|
||||
.cache_actions
|
||||
.iter()
|
||||
.any(|action| matches!(action, CacheAction::SwaRebuild { .. }))
|
||||
);
|
||||
}
|
||||
|
||||
fn set_full_host(tc: &mut UnifiedTreeCore<Vec<i64>>, node: NodeIdx_) {
|
||||
@@ -4856,13 +4925,16 @@ fn deep_swa_tree_survives_backup_evict_and_load_back_rounds() {
|
||||
)
|
||||
.expect("live test node");
|
||||
assert!(actions.is_empty());
|
||||
let lock = tc.inc_lock_ref(anchor).expect("live test node");
|
||||
let lock = tc
|
||||
.inc_lock_ref(anchor, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: lock.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: lock.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: lock.skip_lock_node_ids,
|
||||
skipped_lock_components: lock.skipped_lock_components,
|
||||
};
|
||||
tc.dec_lock_ref(anchor, Some(¶ms), /* skip_swa = */ false)
|
||||
tc.dec_lock_ref(anchor, ¶ms, /* skip_swa = */ false)
|
||||
.expect("live test node");
|
||||
tc.finish_load_back(anchor).expect("live test node");
|
||||
}
|
||||
@@ -4960,3 +5032,104 @@ fn recovered_swa_span_evicts_before_the_window_leaf() {
|
||||
assert!(tc.arena.has_device_value(leaf, FULL));
|
||||
tc.sanity_check(&[], &[]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn aux_release_refreshes_the_leaf_set_whatever_the_release_order() {
|
||||
let mut tc = swa_core(/* window = */ 1, /* page_size = */ 1);
|
||||
tc.insert(&insert_params_swa(&vec![1], &[10], 0, 0));
|
||||
let leaf = child_of(&tc, tc.arena.root(), &[1]);
|
||||
store_swa_device(&mut tc, leaf);
|
||||
assert!(tc.evictable_device_leaves.contains(leaf));
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(leaf).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert!(!tc.evictable_device_leaves.contains(leaf));
|
||||
let params = DecLockRefParams {
|
||||
node_id: result.node_id,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: None,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
};
|
||||
// Full first: its walk still sees the SWA lock, so the leaf stays out.
|
||||
crate::components::FullComponent
|
||||
.release_component_lock(&mut tc, leaf, ¶ms, /* lock_host = */ false);
|
||||
assert!(!tc.evictable_device_leaves.contains(leaf));
|
||||
// The SWA release drops the last lock and must readmit the leaf itself.
|
||||
swa_component(1).release_component_lock(&mut tc, leaf, ¶ms, /* lock_host = */ false);
|
||||
assert!(tc.evictable_device_leaves.contains(leaf));
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "lock receipt anchored on node")]
|
||||
fn dec_lock_ref_rejects_a_receipt_from_another_node() {
|
||||
let mut tc = swa_core(/* window = */ 2, /* page_size = */ 1);
|
||||
let [a, _b, c] = chain(&mut tc);
|
||||
store_swa_device(&mut tc, c);
|
||||
let result = tc
|
||||
.inc_lock_ref(tc.arena.node(c).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
let params = DecLockRefParams {
|
||||
node_id: result.node_id,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: None,
|
||||
skipped_lock_components: result.skipped_lock_components,
|
||||
};
|
||||
// Same receipt, wrong anchor: the walk would otherwise release a's
|
||||
// segment, which this holder never locked.
|
||||
tc.dec_lock_ref(tc.arena.node(a).id, ¶ms, /* skip_swa = */ false)
|
||||
.expect("live test node");
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "lock receipt anchored on node")]
|
||||
fn dec_host_lock_ref_rejects_a_receipt_from_another_node() {
|
||||
let mut tc = swa_core(/* window = */ 2, /* page_size = */ 1);
|
||||
let [a, _b, c] = chain(&mut tc);
|
||||
set_swa_host(&mut tc, c);
|
||||
tc.host_lru_list_mut(SWA).insert_mru(c);
|
||||
let result = tc
|
||||
.inc_host_lock_ref(tc.arena.node(c).id)
|
||||
.expect("live test node");
|
||||
let params = DecLockRefParams {
|
||||
node_id: result.node_id,
|
||||
swa_uuid_for_lock: None,
|
||||
swa_uuid_for_host_lock: result.swa_uuid_for_host_lock,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
};
|
||||
tc.dec_host_lock_ref(tc.arena.node(a).id, ¶ms)
|
||||
.expect("live test node");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn receipt_anchor_follows_the_locked_node_through_a_split() {
|
||||
let mut tc = swa_core(/* window = */ 1, /* page_size = */ 1);
|
||||
tc.insert(&insert_params_swa(&vec![1, 2], &[10, 11], 0, 0));
|
||||
let leaf = child_of(&tc, tc.arena.root(), &[1]);
|
||||
store_swa_device(&mut tc, leaf);
|
||||
let leaf_id = tc.arena.node(leaf).id;
|
||||
let result = tc
|
||||
.inc_lock_ref(leaf_id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(result.node_id, Some(leaf_id));
|
||||
// Diverge inside the node: the split keeps the id on the deeper half.
|
||||
tc.insert(&insert_params_swa(&vec![1, 3], &[12, 13], 0, 0));
|
||||
let params = DecLockRefParams {
|
||||
node_id: result.node_id,
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: None,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
};
|
||||
tc.dec_lock_ref(leaf_id, ¶ms, /* skip_swa = */ false)
|
||||
.expect("live test node");
|
||||
assert_eq!(
|
||||
tc.arena
|
||||
.device_lock_ref(tc.arena.resolve(leaf_id).expect("live test node"), SWA),
|
||||
0
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "swa_sliding_window_size must be positive")]
|
||||
fn new_panics_on_a_zero_sliding_window_size() {
|
||||
SwaComponent::new(&swa_params_with_window(0));
|
||||
}
|
||||
|
||||
@@ -3,7 +3,7 @@ use std::sync::Mutex;
|
||||
use tch::Tensor;
|
||||
|
||||
use super::*;
|
||||
use crate::components::{FULL, MAMBA, SWA};
|
||||
use crate::components::{ComponentSet, FULL, MAMBA, SWA};
|
||||
use crate::node::{NodeAccessError, ValueSlotIdx};
|
||||
use crate::test_utils::{accumulate_step, action_kinds};
|
||||
|
||||
@@ -91,7 +91,7 @@ impl TreeComponent<Vec<i64>> for RecordingComponentForTest {
|
||||
&self,
|
||||
_tree_core: &mut UnifiedTreeCore<Vec<i64>>,
|
||||
_node_id: NodeIdx_,
|
||||
_params: Option<&DecLockRefParams>,
|
||||
_params: &DecLockRefParams,
|
||||
_lock_host: bool,
|
||||
) {
|
||||
unimplemented!()
|
||||
@@ -213,7 +213,7 @@ impl TreeComponent<Vec<i64>> for CountingComponentForTest {
|
||||
&self,
|
||||
_tree_core: &mut UnifiedTreeCore<Vec<i64>>,
|
||||
_node_id: NodeIdx_,
|
||||
_params: Option<&DecLockRefParams>,
|
||||
_params: &DecLockRefParams,
|
||||
_lock_host: bool,
|
||||
) {
|
||||
unimplemented!()
|
||||
@@ -293,11 +293,11 @@ impl TreeComponent<Vec<i64>> for LowPriorityComponentForTest {
|
||||
&self,
|
||||
_tree_core: &mut UnifiedTreeCore<Vec<i64>>,
|
||||
_node_id: NodeIdx_,
|
||||
params: Option<&DecLockRefParams>,
|
||||
params: &DecLockRefParams,
|
||||
lock_host: bool,
|
||||
) {
|
||||
assert!(!lock_host);
|
||||
assert!(params.is_some_and(|p| p.swa_uuid_for_lock.is_some()));
|
||||
assert!(params.swa_uuid_for_lock.is_some());
|
||||
panic!("low-priority release dispatched");
|
||||
}
|
||||
}
|
||||
@@ -368,7 +368,7 @@ impl TreeComponent<Vec<i64>> for SwaComponentForTest {
|
||||
&self,
|
||||
_tree_core: &mut UnifiedTreeCore<Vec<i64>>,
|
||||
_node_id: NodeIdx_,
|
||||
_params: Option<&DecLockRefParams>,
|
||||
_params: &DecLockRefParams,
|
||||
_lock_host: bool,
|
||||
) {
|
||||
unimplemented!()
|
||||
@@ -481,7 +481,7 @@ impl TreeComponent<Vec<i64>> for SwaEvictionComponentForTest {
|
||||
&self,
|
||||
_tree_core: &mut UnifiedTreeCore<Vec<i64>>,
|
||||
_node_id: NodeIdx_,
|
||||
_params: Option<&DecLockRefParams>,
|
||||
_params: &DecLockRefParams,
|
||||
_lock_host: bool,
|
||||
) {
|
||||
unimplemented!()
|
||||
@@ -503,7 +503,7 @@ fn locked_anchor_for_dispatch(tc: &mut UnifiedTreeCore<Vec<i64>>) -> NodeIdx_ {
|
||||
tc.arena
|
||||
.set_device_value(n1, FULL, Tensor::from_slice(&[0i64, 1]));
|
||||
tc.component_state_mut(FULL).evictable_size = 2;
|
||||
tc.inc_lock_ref(tc.arena.node(n1).id)
|
||||
tc.inc_lock_ref(tc.arena.node(n1).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
n1
|
||||
}
|
||||
@@ -516,7 +516,11 @@ fn dec_lock_ref_skip_swa_skips_the_swa_component() {
|
||||
// The skipped Swa driver is never dispatched, so its stub cannot panic.
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(n1).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ true,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -541,7 +545,7 @@ fn inc_lock_ref_reaches_every_component() {
|
||||
tc.arena
|
||||
.set_device_value(n1, FULL, Tensor::from_slice(&[0i64, 1]));
|
||||
tc.component_state_mut(FULL).evictable_size = 2;
|
||||
let _ = tc.inc_lock_ref(tc.arena.node(n1).id);
|
||||
let _ = tc.inc_lock_ref(tc.arena.node(n1).id, ComponentSet::EMPTY);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -552,7 +556,11 @@ fn dec_lock_ref_without_skip_swa_reaches_every_component() {
|
||||
tc.register_component_(Arc::new(SwaComponentForTest));
|
||||
let _ = tc.dec_lock_ref(
|
||||
tc.arena.node(n1).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
);
|
||||
}
|
||||
@@ -630,7 +638,11 @@ fn dec_swa_lock_only_dispatches_lower_priority_releases() {
|
||||
let mut host_frees = HashMap::new();
|
||||
let _ = tc.dec_swa_lock_only(
|
||||
tc.arena.node(root).id,
|
||||
Some(7),
|
||||
&DecLockRefParams {
|
||||
swa_uuid_for_lock: Some(7),
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
);
|
||||
@@ -671,7 +683,11 @@ fn dec_swa_lock_only_returns_device_frees_in_the_device_dict() {
|
||||
let mut host_frees = HashMap::new();
|
||||
tc.dec_swa_lock_only(
|
||||
tc.arena.node(a).id,
|
||||
result.swa_uuid_for_lock,
|
||||
&DecLockRefParams {
|
||||
swa_uuid_for_lock: result.swa_uuid_for_lock,
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
&mut device_frees,
|
||||
&mut host_frees,
|
||||
)
|
||||
@@ -3755,11 +3771,15 @@ fn commit_load_back_reattaches_device_slices_and_restores_the_match() {
|
||||
assert_eq!(tc.full_evictable_size(), 4);
|
||||
// The orchestrator re-locks the loaded path right after commit; that lock walk
|
||||
// also re-evaluates the parent's transient D-leaf membership.
|
||||
tc.inc_lock_ref(tc.arena.node(child).id)
|
||||
tc.inc_lock_ref(tc.arena.node(child).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
tc.dec_lock_ref(
|
||||
tc.arena.node(child).id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -5261,7 +5281,7 @@ fn stale_handle_returns_err_after_its_node_is_freed() {
|
||||
tc.evict_device_leaf(leaf, /* is_write_back = */ false)
|
||||
.expect("live test node");
|
||||
assert!(matches!(
|
||||
tc.inc_lock_ref(leaf),
|
||||
tc.inc_lock_ref(leaf, ComponentSet::EMPTY),
|
||||
Err(NodeAccessError { node_id }) if node_id == leaf
|
||||
));
|
||||
}
|
||||
@@ -6164,7 +6184,7 @@ fn reset_restores_a_fresh_tree() {
|
||||
..insert_params(&vec![7, 8], &[20, 21])
|
||||
});
|
||||
let matched = tc.match_prefix(&match_params(&vec![1, 2, 3]));
|
||||
tc.inc_lock_ref(matched.best_match_node_id)
|
||||
tc.inc_lock_ref(matched.best_match_node_id, ComponentSet::EMPTY)
|
||||
.expect("live match node");
|
||||
assert_eq!(tc.protected_size(), 3);
|
||||
// Seed aux LRU, host LRU, and host-leaf state so the reset must clear each.
|
||||
@@ -6216,7 +6236,7 @@ fn size_accessors_mirror_the_full_component_state() {
|
||||
assert_eq!(tc.protected_size(), 0);
|
||||
assert_eq!(tc.component_evictable_size(FULL), 3);
|
||||
let matched = tc.match_prefix(&match_params(&vec![1, 2, 3]));
|
||||
tc.inc_lock_ref(matched.best_match_node_id)
|
||||
tc.inc_lock_ref(matched.best_match_node_id, ComponentSet::EMPTY)
|
||||
.expect("live match node");
|
||||
assert_eq!(tc.protected_size(), 3);
|
||||
assert_eq!(tc.full_protected_size(), 3);
|
||||
@@ -6306,7 +6326,7 @@ fn walk_for_kv_canary_chains_slots_across_namespaces() {
|
||||
fn walk_for_kv_canary_unlocked_only_skips_locked_nodes_but_keeps_the_chain() {
|
||||
let mut tc = core();
|
||||
let (a, _b) = matched_chain(&mut tc);
|
||||
tc.inc_lock_ref(tc.arena.node(a).id)
|
||||
tc.inc_lock_ref(tc.arena.node(a).id, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(
|
||||
sorted_canary_rows(tc.walk_for_kv_canary(true, false)),
|
||||
@@ -6659,13 +6679,18 @@ fn sanity_check_passes_on_a_healthy_tree() {
|
||||
let leaf = tc
|
||||
.match_prefix(&match_params(&vec![1, 2, 9]))
|
||||
.best_match_node_id;
|
||||
tc.inc_lock_ref(leaf).expect("live test node");
|
||||
tc.inc_lock_ref(leaf, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
tc.sanity_check(&[(1, leaf)], &[(2, leaf)]);
|
||||
tc.dec_lock_ref(
|
||||
tc.arena
|
||||
.node(tc.arena.resolve(leaf).expect("live test node"))
|
||||
.id,
|
||||
/* params = */ None,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live test node");
|
||||
@@ -6774,16 +6799,28 @@ fn sanity_check_detects_an_evicted_parent_prefix() {
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "evicted but lock_ref")]
|
||||
fn sanity_check_detects_a_locked_tombstone() {
|
||||
fn sanity_check_accepts_a_locked_tombstone() {
|
||||
// Segment locks count evicted nodes, so a device-locked tombstone is a
|
||||
// legal state the checker must not flag.
|
||||
let mut tc = sane_tree();
|
||||
// write_back spares the tombstone's ancestors the backup-chain rule.
|
||||
tc.is_write_back = true;
|
||||
let leaf = tc
|
||||
.match_prefix(&match_params(&vec![1, 2, 9]))
|
||||
.best_match_node_id;
|
||||
tc.inc_lock_ref(leaf).expect("live test node");
|
||||
let _ = tc
|
||||
.arena
|
||||
.take_device_value(tc.arena.resolve(leaf).expect("live test node"), FULL);
|
||||
let leaf_idx = tc.arena.resolve(leaf).expect("live test node");
|
||||
// Tombstone the leaf consistently first (host copy, ledger, leaf sets),
|
||||
// then lock through it: the bottom segment counts the tombstone.
|
||||
tc.arena
|
||||
.set_host_value(leaf_idx, FULL, Tensor::from_slice(&[9i64]));
|
||||
let taken = tc.arena.take_device_value(leaf_idx, FULL);
|
||||
tc.dec_evictable_size(FULL, taken.size()[0] as usize);
|
||||
tc.update_evictable_leaf_sets_(leaf_idx);
|
||||
let parent_idx = tc.arena.node(leaf_idx).parent();
|
||||
tc.update_evictable_leaf_sets_(parent_idx);
|
||||
tc.inc_lock_ref(leaf, ComponentSet::EMPTY)
|
||||
.expect("live test node");
|
||||
assert_eq!(tc.arena.device_lock_ref(leaf_idx, FULL), 1);
|
||||
tc.sanity_check(&[], &[]);
|
||||
}
|
||||
|
||||
@@ -8049,13 +8086,16 @@ fn run_random_op_sequence(mut tc: UnifiedTreeCore<Vec<i64>>, page: usize, mamba:
|
||||
2 => {
|
||||
// Balanced lock round trip on whatever the key matches.
|
||||
let anchor = tc.match_prefix(&match_params(&key)).best_match_node_id;
|
||||
let lock = tc.inc_lock_ref(anchor).expect("live match anchor");
|
||||
let lock = tc
|
||||
.inc_lock_ref(anchor, ComponentSet::EMPTY)
|
||||
.expect("live match anchor");
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: lock.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: lock.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: lock.skip_lock_node_ids,
|
||||
skipped_lock_components: lock.skipped_lock_components,
|
||||
};
|
||||
tc.dec_lock_ref(anchor, Some(¶ms), /* skip_swa = */ false)
|
||||
tc.dec_lock_ref(anchor, ¶ms, /* skip_swa = */ false)
|
||||
.expect("live match anchor");
|
||||
}
|
||||
_ => {
|
||||
@@ -8063,7 +8103,9 @@ fn run_random_op_sequence(mut tc: UnifiedTreeCore<Vec<i64>>, page: usize, mamba:
|
||||
let matched = tc.match_prefix(&match_params(&key));
|
||||
let anchor = matched.best_match_node_id;
|
||||
let matched_len = matched.device_indices.numel() as usize;
|
||||
let lock = tc.inc_lock_ref(anchor).expect("live match anchor");
|
||||
let lock = tc
|
||||
.inc_lock_ref(anchor, ComponentSet::EMPTY)
|
||||
.expect("live match anchor");
|
||||
tc.insert(&sequence_insert_params(
|
||||
&key,
|
||||
matched_len,
|
||||
@@ -8072,11 +8114,12 @@ fn run_random_op_sequence(mut tc: UnifiedTreeCore<Vec<i64>>, page: usize, mamba:
|
||||
mamba,
|
||||
));
|
||||
let params = DecLockRefParams {
|
||||
node_id: None,
|
||||
swa_uuid_for_lock: lock.swa_uuid_for_lock,
|
||||
swa_uuid_for_host_lock: lock.swa_uuid_for_host_lock,
|
||||
skip_lock_node_ids: lock.skip_lock_node_ids,
|
||||
skipped_lock_components: lock.skipped_lock_components,
|
||||
};
|
||||
tc.dec_lock_ref(anchor, Some(¶ms), /* skip_swa = */ false)
|
||||
tc.dec_lock_ref(anchor, ¶ms, /* skip_swa = */ false)
|
||||
.expect("live match anchor");
|
||||
}
|
||||
}
|
||||
@@ -8224,10 +8267,17 @@ fn a_zero_length_match_anchors_at_the_root() {
|
||||
.best_match_node_id;
|
||||
assert_eq!(anchor, tc.root_node_handle(Some("salted")));
|
||||
// The root handle stays valid across a full namespace eviction.
|
||||
tc.inc_lock_ref(anchor).expect("live root");
|
||||
tc.inc_lock_ref(anchor, ComponentSet::EMPTY)
|
||||
.expect("live root");
|
||||
drain_full_device(&mut tc);
|
||||
tc.dec_lock_ref(
|
||||
anchor, /* params = */ None, /* skip_swa = */ false,
|
||||
anchor,
|
||||
/* params = */
|
||||
&DecLockRefParams {
|
||||
skipped_lock_components: ComponentSet::EMPTY,
|
||||
..Default::default()
|
||||
},
|
||||
/* skip_swa = */ false,
|
||||
)
|
||||
.expect("live root");
|
||||
assert!(tc.arena.resolve(anchor).is_ok());
|
||||
|
||||
@@ -8,7 +8,9 @@ use std::sync::Arc;
|
||||
use sha2::{Digest, Sha256};
|
||||
use tch::{Device, Kind, Tensor};
|
||||
|
||||
use crate::components::{self, FullComponent, MambaComponent, SwaComponent, TreeComponent};
|
||||
use crate::components::{
|
||||
self, ComponentSet, FullComponent, MambaComponent, SwaComponent, TreeComponent,
|
||||
};
|
||||
use crate::components::{
|
||||
BASE_COMPONENT_TYPE, ComponentType, FULL, MAMBA, NUM_COMPONENT_TYPES, SWA,
|
||||
};
|
||||
@@ -34,28 +36,41 @@ fn next_coexist_reclaim_digest(current: i64, node_id: NodeId, component_idx: usi
|
||||
// ---- interface types ----
|
||||
|
||||
/// Result of `inc_lock_ref`, handed back to the matching `dec_lock_ref`.
|
||||
///
|
||||
/// The receipt a release needs is per-component lock evidence: the SWA
|
||||
/// segment boundary uuid (None means the segment reached the root) and
|
||||
/// whether the single-node Mamba lock was taken (the decode hold opts
|
||||
/// out). Locks count every node in their contiguous segment, so no
|
||||
/// per-node skip state exists. Receipt fields default to nothing-acquired;
|
||||
/// `inc_lock_ref` stamps what it actually took.
|
||||
#[derive(Default)]
|
||||
pub struct IncLockRefResult {
|
||||
/// Tokens newly protected (moved out of evictable) by this lock.
|
||||
pub delta: Option<usize>,
|
||||
/// The node the lock was taken on; a release replays the receipt there only.
|
||||
pub node_id: Option<NodeId>,
|
||||
/// SWA lock-window uuid minted/reused by the device lock walk.
|
||||
pub swa_uuid_for_lock: Option<i64>,
|
||||
/// SWA lock-window uuid minted/reused by the host lock walk.
|
||||
pub swa_uuid_for_host_lock: Option<i64>,
|
||||
/// Per-component nodes that were tombstones at acquire time; replayed at
|
||||
/// release so the unlock skips them.
|
||||
pub skip_lock_node_ids: HashMap<ComponentType, HashSet<NodeId>>,
|
||||
/// Components the acquire left untaken; the release skips them too.
|
||||
pub skipped_lock_components: ComponentSet,
|
||||
}
|
||||
|
||||
/// Params for `dec_lock_ref`.
|
||||
/// Params for `dec_lock_ref`. Receipt fields default to nothing-acquired so
|
||||
/// a lost receipt under-releases (a leak sanity checks report) instead of
|
||||
/// releasing a lock another holder owns.
|
||||
#[derive(Default)]
|
||||
pub struct DecLockRefParams {
|
||||
/// The node the matching acquire locked; None only for receipts that did
|
||||
/// not come from this core (a mispaired anchor is a protocol violation).
|
||||
pub node_id: Option<NodeId>,
|
||||
/// SWA lock-window uuid the device unlock stops at, from the matching acquire.
|
||||
pub swa_uuid_for_lock: Option<i64>,
|
||||
/// SWA lock-window uuid the host unlock stops at, from the matching acquire.
|
||||
pub swa_uuid_for_host_lock: Option<i64>,
|
||||
/// Per-component nodes the unlock walk skips (from the matching acquire).
|
||||
pub skip_lock_node_ids: HashMap<ComponentType, HashSet<NodeId>>,
|
||||
/// Components the matching acquire left untaken.
|
||||
pub skipped_lock_components: ComponentSet,
|
||||
}
|
||||
|
||||
/// Result of `dec_lock_ref`.
|
||||
@@ -792,58 +807,90 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
||||
self.swa_uuid_counter
|
||||
}
|
||||
|
||||
/// Bump the reference count on a node's component locks.
|
||||
pub fn inc_lock_ref(&mut self, node_id: NodeId) -> Result<IncLockRefResult, NodeAccessError> {
|
||||
self.inc_lock_ref_with_skip(node_id, &[])
|
||||
}
|
||||
|
||||
/// Bump component locks, leaving explicitly skipped target components evictable.
|
||||
pub fn inc_lock_ref_with_skip(
|
||||
/// Bump the reference count on a node's component locks. Components in
|
||||
/// `skip_lock_components` are left untaken; the receipt records the anchor
|
||||
/// node and the skipped set so the paired release mirrors them.
|
||||
pub fn inc_lock_ref(
|
||||
&mut self,
|
||||
node_id: NodeId,
|
||||
skip_lock_components: &[ComponentType],
|
||||
skip_lock_components: ComponentSet,
|
||||
) -> Result<IncLockRefResult, NodeAccessError> {
|
||||
let node_id = self.arena.resolve(node_id)?;
|
||||
let node = self.arena.node(node_id);
|
||||
let node_handle = node.id;
|
||||
let is_root = node.is_root();
|
||||
let mut result = IncLockRefResult::default();
|
||||
let node_idx = self.arena.resolve(node_id)?;
|
||||
let mut result = IncLockRefResult {
|
||||
node_id: Some(self.arena.node(node_idx).id),
|
||||
skipped_lock_components: skip_lock_components,
|
||||
..Default::default()
|
||||
};
|
||||
for i in 0..self.components.len() {
|
||||
let component_type = self.components[i].component_type();
|
||||
if skip_lock_components.contains(&component_type) {
|
||||
if !is_root {
|
||||
result
|
||||
.skip_lock_node_ids
|
||||
.entry(component_type)
|
||||
.or_default()
|
||||
.insert(node_handle);
|
||||
}
|
||||
let component = Arc::clone(&self.components[i]);
|
||||
if skip_lock_components.contains(component.component_type()) {
|
||||
continue;
|
||||
}
|
||||
let component = Arc::clone(&self.components[i]);
|
||||
result = component
|
||||
.acquire_component_lock(self, node_id, result, /* lock_host = */ false);
|
||||
.acquire_component_lock(self, node_idx, result, /* lock_host = */ false);
|
||||
}
|
||||
self.update_evictable_leaf_sets_(node_id);
|
||||
self.update_evictable_leaf_sets_(node_idx);
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// Decrease the reference count on a node's component locks.
|
||||
/// A receipt releases only the node its acquire returned; a mispaired
|
||||
/// node would silently release (or steal) another holder's segment.
|
||||
fn assert_receipt_anchor_(&self, node_idx: NodeIdx_, params: &DecLockRefParams) {
|
||||
if let Some(anchor) = params.node_id {
|
||||
let node_handle = self.arena.node(node_idx).id;
|
||||
assert!(
|
||||
anchor == node_handle,
|
||||
"lock receipt anchored on node {anchor} released on node {node_handle}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// Release each component this receipt acquired. Auxiliaries go first so
|
||||
/// Full, whose walk refreshes leaf membership on every node it unlocks,
|
||||
/// sees their final refs; the auxiliary walks also refresh the nodes they
|
||||
/// unlock, so the order is not load-bearing for the sets.
|
||||
fn release_components_(
|
||||
&mut self,
|
||||
node_idx: NodeIdx_,
|
||||
params: &DecLockRefParams,
|
||||
lock_host: bool,
|
||||
skip_swa_and_below: bool,
|
||||
) {
|
||||
let swa_priority = if skip_swa_and_below {
|
||||
self.try_component_by_type_(SWA)
|
||||
.map(|swa| swa.eviction_priority(/* is_leaf = */ false))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
for i in (0..self.components.len()).rev() {
|
||||
let component = Arc::clone(&self.components[i]);
|
||||
let ct = component.component_type();
|
||||
if params.skipped_lock_components.contains(ct) {
|
||||
continue;
|
||||
}
|
||||
if let Some(swa_priority) = swa_priority
|
||||
&& (ct == SWA || component.eviction_priority(/* is_leaf = */ false) < swa_priority)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
component.release_component_lock(self, node_idx, params, lock_host);
|
||||
}
|
||||
}
|
||||
|
||||
/// Decrease the reference count on a node's component locks. The receipt
|
||||
/// is required: a release must replay its acquire's evidence. After an SWA
|
||||
/// early release (`dec_swa_lock_only`), `skip_swa` leaves SWA and the
|
||||
/// lower-priority components it already dropped alone.
|
||||
pub fn dec_lock_ref(
|
||||
&mut self,
|
||||
node_id: NodeId,
|
||||
params: Option<&DecLockRefParams>,
|
||||
params: &DecLockRefParams,
|
||||
skip_swa: bool,
|
||||
) -> Result<DecLockRefResult, NodeAccessError> {
|
||||
let node_id = self.arena.resolve(node_id)?;
|
||||
for i in 0..self.components.len() {
|
||||
if skip_swa && self.components[i].component_type() == SWA {
|
||||
continue;
|
||||
}
|
||||
let component = Arc::clone(&self.components[i]);
|
||||
component.release_component_lock(self, node_id, params, /* lock_host = */ false);
|
||||
}
|
||||
self.update_evictable_leaf_sets_(node_id);
|
||||
let node_idx = self.arena.resolve(node_id)?;
|
||||
self.assert_receipt_anchor_(node_idx, params);
|
||||
self.release_components_(node_idx, params, /* lock_host = */ false, skip_swa);
|
||||
self.update_evictable_leaf_sets_(node_idx);
|
||||
// TODO: delta is not aggregated from components; no caller uses it yet.
|
||||
Ok(DecLockRefResult::default())
|
||||
}
|
||||
@@ -853,50 +900,37 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
||||
pub fn dec_swa_lock_only(
|
||||
&mut self,
|
||||
node_id: NodeId,
|
||||
swa_uuid_for_lock: Option<i64>,
|
||||
params: &DecLockRefParams,
|
||||
device_frees: &mut HashMap<ComponentType, Vec<Tensor>>,
|
||||
host_frees: &mut HashMap<ComponentType, Vec<Tensor>>,
|
||||
) -> Result<(), NodeAccessError> {
|
||||
self.dec_swa_lock_only_with_skip(
|
||||
node_id,
|
||||
swa_uuid_for_lock,
|
||||
/* skip_lock_node_ids = */ None,
|
||||
device_frees,
|
||||
host_frees,
|
||||
)
|
||||
}
|
||||
|
||||
/// Skip-aware variant used when an acquire deliberately omitted a component.
|
||||
pub fn dec_swa_lock_only_with_skip(
|
||||
&mut self,
|
||||
node_id: NodeId,
|
||||
swa_uuid_for_lock: Option<i64>,
|
||||
skip_lock_node_ids: Option<&HashMap<ComponentType, HashSet<NodeId>>>,
|
||||
device_frees: &mut HashMap<ComponentType, Vec<Tensor>>,
|
||||
host_frees: &mut HashMap<ComponentType, Vec<Tensor>>,
|
||||
) -> Result<(), NodeAccessError> {
|
||||
let node_id = self.arena.resolve(node_id)?;
|
||||
let node_idx = self.arena.resolve(node_id)?;
|
||||
self.assert_receipt_anchor_(node_idx, params);
|
||||
let Some(swa) = self.try_component_by_type_(SWA) else {
|
||||
return Ok(());
|
||||
};
|
||||
swa.release_window_lock(self, node_id, swa_uuid_for_lock, device_frees, host_frees);
|
||||
swa.release_window_lock(
|
||||
self,
|
||||
node_idx,
|
||||
params.swa_uuid_for_lock,
|
||||
device_frees,
|
||||
host_frees,
|
||||
);
|
||||
|
||||
// Drop strictly-lower-priority locks (e.g. Mamba) co-located on the node.
|
||||
// Drop strictly-lower-priority locks co-located on the node, skipping
|
||||
// any the paired inc never took.
|
||||
let swa_priority = swa.eviction_priority(/* is_leaf = */ false);
|
||||
let dec_params = DecLockRefParams {
|
||||
swa_uuid_for_lock,
|
||||
skip_lock_node_ids: skip_lock_node_ids.cloned().unwrap_or_default(),
|
||||
..Default::default()
|
||||
};
|
||||
for i in 0..self.components.len() {
|
||||
for i in (0..self.components.len()).rev() {
|
||||
let component = Arc::clone(&self.components[i]);
|
||||
if params
|
||||
.skipped_lock_components
|
||||
.contains(component.component_type())
|
||||
{
|
||||
continue;
|
||||
}
|
||||
if component.eviction_priority(/* is_leaf = */ false) < swa_priority {
|
||||
component.release_component_lock(
|
||||
self,
|
||||
node_id,
|
||||
Some(&dec_params),
|
||||
/* lock_host = */ false,
|
||||
);
|
||||
component
|
||||
.release_component_lock(self, node_idx, params, /* lock_host = */ false);
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
@@ -925,29 +959,31 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
||||
&mut self,
|
||||
node_id: NodeId,
|
||||
) -> Result<IncLockRefResult, NodeAccessError> {
|
||||
let node_id = self.arena.resolve(node_id)?;
|
||||
let mut result = IncLockRefResult::default();
|
||||
let node_idx = self.arena.resolve(node_id)?;
|
||||
let mut result = IncLockRefResult {
|
||||
node_id: Some(self.arena.node(node_idx).id),
|
||||
..Default::default()
|
||||
};
|
||||
for i in 0..self.components.len() {
|
||||
let component = Arc::clone(&self.components[i]);
|
||||
result = component
|
||||
.acquire_component_lock(self, node_id, result, /* lock_host = */ true);
|
||||
.acquire_component_lock(self, node_idx, result, /* lock_host = */ true);
|
||||
}
|
||||
self.update_evictable_leaf_sets_(node_id);
|
||||
self.update_evictable_leaf_sets_(node_idx);
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// Decrease the reference count on a node's host-side component locks.
|
||||
/// The receipt is required, as for `dec_lock_ref`.
|
||||
pub fn dec_host_lock_ref(
|
||||
&mut self,
|
||||
node_id: NodeId,
|
||||
params: Option<&DecLockRefParams>,
|
||||
params: &DecLockRefParams,
|
||||
) -> Result<DecLockRefResult, NodeAccessError> {
|
||||
let node_id = self.arena.resolve(node_id)?;
|
||||
for i in 0..self.components.len() {
|
||||
let component = Arc::clone(&self.components[i]);
|
||||
component.release_component_lock(self, node_id, params, /* lock_host = */ true);
|
||||
}
|
||||
self.update_evictable_leaf_sets_(node_id);
|
||||
let node_idx = self.arena.resolve(node_id)?;
|
||||
self.assert_receipt_anchor_(node_idx, params);
|
||||
self.release_components_(node_idx, params, /* lock_host = */ true, false);
|
||||
self.update_evictable_leaf_sets_(node_idx);
|
||||
Ok(DecLockRefResult::default())
|
||||
}
|
||||
|
||||
@@ -1877,7 +1913,14 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
||||
pub fn unevict_node_on_insert_(&mut self, node_id: NodeIdx_, fresh_value: &Tensor) {
|
||||
self.arena
|
||||
.set_device_value(node_id, FULL, fresh_value.copy());
|
||||
self.inc_evictable_size(FULL, fresh_value.size()[0] as usize);
|
||||
let tokens = fresh_value.size()[0] as usize;
|
||||
// A value materialized under lock is protected; the last release
|
||||
// moves it to evictable.
|
||||
if self.arena.device_lock_ref(node_id, FULL) > 0 {
|
||||
self.inc_protected_size(FULL, tokens);
|
||||
} else {
|
||||
self.inc_evictable_size(FULL, tokens);
|
||||
}
|
||||
self.update_evictable_leaf_sets_(node_id);
|
||||
self.update_full_coexisting_host_tracking_(node_id);
|
||||
if let Some(parent_id) = self.arena.node(node_id).try_parent() {
|
||||
@@ -2644,6 +2687,11 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
||||
if node.is_host_locked() {
|
||||
return false;
|
||||
}
|
||||
// Segment locks count evicted nodes too: a device-locked candidate is
|
||||
// a live segment's anchor, and evict_host_leaf_ would delete it.
|
||||
if node.is_device_locked() {
|
||||
return false;
|
||||
}
|
||||
if !node.children.is_empty() {
|
||||
return false;
|
||||
}
|
||||
@@ -3702,7 +3750,13 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
||||
host_lru.remove_node(node_id);
|
||||
}
|
||||
self.device_lru_list_mut(component_type).insert_mru(node_id);
|
||||
self.inc_evictable_size(component_type, tokens);
|
||||
// A value materialized under lock is protected; the last release
|
||||
// moves it to evictable.
|
||||
if self.arena.device_lock_ref(node_id, component_type) > 0 {
|
||||
self.inc_protected_size(component_type, tokens);
|
||||
} else {
|
||||
self.inc_evictable_size(component_type, tokens);
|
||||
}
|
||||
}
|
||||
|
||||
/// The component's device value on the node, or None if evicted.
|
||||
@@ -3931,12 +3985,8 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
||||
device_state.lock_ref
|
||||
));
|
||||
}
|
||||
if device_state.value.is_none() && device_state.lock_ref > 0 {
|
||||
errors.push(format!(
|
||||
"node {node_id} {ct:?} evicted but lock_ref={}",
|
||||
device_state.lock_ref
|
||||
));
|
||||
}
|
||||
// Locked tombstones are legal: segment locks count every
|
||||
// node in [start, boundary], data-bearing or not.
|
||||
}
|
||||
|
||||
// Collect expected leaf qualification (single pass)
|
||||
|
||||
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
|
||||
|
||||
@@ -40,7 +41,7 @@ def _make_req(
|
||||
req.kv = ReqKvInfo(req_pool_idx=req_pool_idx)
|
||||
req.skip_radix_cache_insert = False
|
||||
req.last_node = None
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.session = None
|
||||
req.return_logprob = False
|
||||
req.logprob_start_len = -1
|
||||
|
||||
@@ -35,13 +35,17 @@ from unittest.mock import MagicMock
|
||||
import torch
|
||||
|
||||
from sglang.srt.disaggregation.decode import DecodePreallocQueue
|
||||
from sglang.srt.disaggregation.decode_hicache_mixin import DecodePrefixMatch
|
||||
from sglang.srt.disaggregation.decode_hicache_mixin import (
|
||||
DecodeHiCacheTransferMixin,
|
||||
DecodePrefixMatch,
|
||||
)
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
DecLockRefParams,
|
||||
InsertParams,
|
||||
MatchPrefixParams,
|
||||
)
|
||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey
|
||||
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
|
||||
from sglang.srt.utils.common import Range
|
||||
|
||||
|
||||
@@ -387,7 +391,7 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
req.last_node = object()
|
||||
req.finished_reason = None
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = 123
|
||||
req.lock_receipt = DecLockRefParams(swa_uuid_for_lock=123)
|
||||
req.swa_prefix_lock_released = False
|
||||
req.pd_rebootstrap_in_progress = False
|
||||
req.sampling_params.max_new_tokens = 16
|
||||
@@ -462,7 +466,10 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
self.assertEqual(preallocated, [])
|
||||
self.assertEqual(failed, [])
|
||||
queue._pre_alloc.assert_not_called()
|
||||
queue.tree_cache.dec_swa_lock_only.assert_called_once_with(req.last_node, 123)
|
||||
queue.tree_cache.dec_swa_lock_only.assert_called_once_with(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=123),
|
||||
)
|
||||
queue.tree_cache.dec_lock_ref.assert_called_once_with(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=123),
|
||||
@@ -472,6 +479,51 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
queue._swa_tail_len.assert_called_once_with(8)
|
||||
queue._allocatable_token_budgets.assert_called_once()
|
||||
|
||||
def test_hicache_restore_commit_hands_over_lock_with_receipt(self):
|
||||
"""The hicache-restore commit must release the prealloc lock with the
|
||||
req's receipt, honoring a prior early SWA release (skip_swa), and hand
|
||||
the restored node's lock to the req atomically: receipt fields move
|
||||
with last_node, the early-release flag resets (the restored lock is
|
||||
fresh), and the decode_req drops ownership so a post-commit abort
|
||||
cannot release the restored lock a second time."""
|
||||
q = DecodeHiCacheTransferMixin.__new__(DecodeHiCacheTransferMixin)
|
||||
q.tree_cache = MagicMock()
|
||||
|
||||
req = MagicMock()
|
||||
req.req_pool_idx = 0
|
||||
req.lock_receipt = DecLockRefParams(swa_uuid_for_lock=123)
|
||||
req.swa_prefix_lock_released = True # SWA tail-prealloc released early
|
||||
|
||||
prealloc_node = object()
|
||||
restored_node = object()
|
||||
decode_req = MagicMock()
|
||||
decode_req.req = req
|
||||
decode_req.prefix_match = DecodePrefixMatch(
|
||||
prefix_indices=torch.arange(4, dtype=torch.int64),
|
||||
l2_host_hit_length=4,
|
||||
l3_storage_hit_length=0,
|
||||
last_device_node=prealloc_node,
|
||||
)
|
||||
decode_req.hicache_restored_node = restored_node
|
||||
decode_req.hicache_restore_lock_receipt = DecLockRefParams(
|
||||
swa_uuid_for_lock=456, skipped_lock_components=(ComponentType.MAMBA,)
|
||||
)
|
||||
decode_req.hicache_restored_kv_indices = torch.arange(4, 8, dtype=torch.int64)
|
||||
|
||||
q._commit_hicache_local_restore_to_req(decode_req)
|
||||
|
||||
q.tree_cache.dec_lock_ref.assert_called_once_with(
|
||||
prealloc_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=123),
|
||||
skip_swa=True,
|
||||
)
|
||||
self.assertIs(req.last_node, restored_node)
|
||||
self.assertEqual(req.lock_receipt.swa_uuid_for_lock, 456)
|
||||
self.assertIn(ComponentType.MAMBA, req.lock_receipt.skipped_lock_components)
|
||||
self.assertFalse(req.swa_prefix_lock_released)
|
||||
self.assertIsNone(decode_req.hicache_restored_node)
|
||||
self.assertIsNone(decode_req.hicache_restore_lock_receipt)
|
||||
|
||||
def test_repeated_incremental_no_leak(self):
|
||||
"""Multiple incremental transfers shouldn't leak lock_refs."""
|
||||
cache, req_to_token = _make_cache_with_pools()
|
||||
|
||||
@@ -16,6 +16,7 @@ from types import SimpleNamespace
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
DecLockRefParams,
|
||||
EvictParams,
|
||||
IncLockRefResult,
|
||||
)
|
||||
@@ -182,12 +183,12 @@ class _RecordingComp:
|
||||
|
||||
class TestDecSwaLockSkip(unittest.TestCase):
|
||||
"""dec_swa_lock_only early-releases SWA plus co-located lower-tier (Mamba)
|
||||
locks. On a full-only-locked node (decode skip) it must thread the skip set
|
||||
into that lower-tier release, else it drops a mamba lock it never took --
|
||||
another request's, on a shared FULL+SWA+MAMBA node (Inkling). Guards the
|
||||
contract without booting a 3-component model."""
|
||||
locks. On a node whose acquire skipped Mamba (decode hold), the release
|
||||
must skip it too, else it drops a mamba lock it never took -- another
|
||||
request's, on a shared FULL+SWA+MAMBA node (Inkling). Guards the contract
|
||||
without booting a 3-component model."""
|
||||
|
||||
def test_threads_skip_ids_into_lower_tier_release(self):
|
||||
def _run(self, skipped_lock_components):
|
||||
# internal-node priority: full=2 > swa=1 > mamba=0
|
||||
full = _RecordingComp(ComponentType.FULL, 2)
|
||||
swa = _RecordingComp(ComponentType.SWA, 1)
|
||||
@@ -197,23 +198,27 @@ class TestDecSwaLockSkip(unittest.TestCase):
|
||||
components=(full, swa, mamba),
|
||||
components_by_type={ComponentType.SWA: swa},
|
||||
node_by_id=lambda node_id: node,
|
||||
_assert_receipt_anchor=UnifiedTreeCore._assert_receipt_anchor,
|
||||
)
|
||||
|
||||
UnifiedTreeCore.dec_swa_lock_only(
|
||||
tree_core,
|
||||
node.id,
|
||||
swa_uuid_for_lock=None,
|
||||
skip_lock_node_ids={ComponentType.MAMBA: {7}},
|
||||
DecLockRefParams(skipped_lock_components=skipped_lock_components),
|
||||
)
|
||||
return full, mamba
|
||||
|
||||
# mamba (below swa) is released, honoring the skip set
|
||||
self.assertEqual(len(mamba.released), 1)
|
||||
self.assertEqual(
|
||||
mamba.released[0].skip_lock_node_ids.get(ComponentType.MAMBA), {7}
|
||||
)
|
||||
def test_unlocked_mamba_is_not_released(self):
|
||||
full, mamba = self._run(skipped_lock_components=(ComponentType.MAMBA,))
|
||||
# mamba took no lock at acquire, so the early release skips it too
|
||||
self.assertEqual(mamba.released, [])
|
||||
# full (above swa) is never touched
|
||||
self.assertEqual(full.released, [])
|
||||
|
||||
def test_lower_tier_released_when_locked(self):
|
||||
full, mamba = self._run(skipped_lock_components=())
|
||||
self.assertEqual(len(mamba.released), 1)
|
||||
self.assertEqual(full.released, [])
|
||||
|
||||
|
||||
class TestMambaDonatedAllocRatio(unittest.TestCase):
|
||||
def test_prefill_peak_ratio2_exhausts_pool(self):
|
||||
|
||||
@@ -76,10 +76,10 @@ def test_lock_moves_tokens_between_evictable_and_protected():
|
||||
InsertParams(key=_key([1, 2]), value=torch.tensor([10, 11], dtype=torch.int64)),
|
||||
)
|
||||
matched = core.match_prefix(MatchPrefixParams(key=_key([1, 2])))
|
||||
core.inc_lock_ref(matched.best_match_node)
|
||||
lock = core.inc_lock_ref(matched.best_match_node)
|
||||
assert core.protected_size() == 2
|
||||
assert core.evictable_size() == 0
|
||||
core.dec_lock_ref(matched.best_match_node)
|
||||
core.dec_lock_ref(matched.best_match_node, lock.to_dec_params())
|
||||
assert core.evictable_size() == 2
|
||||
|
||||
|
||||
|
||||
@@ -26,6 +26,7 @@ from sglang.srt.disaggregation.kv_events import (
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
DecLockRefParams,
|
||||
InsertParams,
|
||||
InsertResult,
|
||||
MatchPrefixParams,
|
||||
@@ -230,8 +231,10 @@ def test_stale_handle_operations_raise_key_error_without_poisoning_the_core():
|
||||
|
||||
operations = {
|
||||
"inc_lock_ref": lambda: core.inc_lock_ref(stale_root),
|
||||
"dec_lock_ref": lambda: core.dec_lock_ref(stale_root),
|
||||
"dec_swa_lock_only": lambda: core.dec_swa_lock_only(stale_root, None),
|
||||
"dec_lock_ref": lambda: core.dec_lock_ref(stale_root, DecLockRefParams()),
|
||||
"dec_swa_lock_only": lambda: core.dec_swa_lock_only(
|
||||
stale_root, DecLockRefParams()
|
||||
),
|
||||
"evict_device_leaf": lambda: core.evict_device_leaf(stale_root, False),
|
||||
"drop_subtree_no_host": lambda: core.drop_subtree_no_host(stale_root),
|
||||
"demote": lambda: core.demote(stale_root),
|
||||
@@ -267,7 +270,9 @@ def test_stale_handle_operations_raise_key_error_without_poisoning_the_core():
|
||||
stale_root, {}, {}
|
||||
),
|
||||
"inc_host_lock_ref": lambda: core.inc_host_lock_ref(stale_root),
|
||||
"dec_host_lock_ref": lambda: core.dec_host_lock_ref(stale_root),
|
||||
"dec_host_lock_ref": lambda: core.dec_host_lock_ref(
|
||||
stale_root, DecLockRefParams()
|
||||
),
|
||||
"mark_write_through_pending": lambda: core.mark_write_through_pending(
|
||||
[stale_root], stale_root
|
||||
),
|
||||
@@ -416,10 +421,10 @@ def test_lock_and_unlock_move_tokens_between_protected_and_evictable():
|
||||
_insert(core, [1, 2, 3], [10, 11, 12])
|
||||
_insert(core, [1, 2, 3, 4, 5], [20, 21, 22, 13, 14])
|
||||
matched = core.match_prefix(MatchPrefixParams(key=_key([1, 2, 3, 4, 5])))
|
||||
core.inc_lock_ref(matched.best_match_node)
|
||||
lock = core.inc_lock_ref(matched.best_match_node)
|
||||
assert core.protected_size() == 5
|
||||
assert core.evictable_size() == 0
|
||||
core.dec_lock_ref(matched.best_match_node)
|
||||
core.dec_lock_ref(matched.best_match_node, lock.to_dec_params())
|
||||
assert core.protected_size() == 0
|
||||
assert core.evictable_size() == 5
|
||||
|
||||
@@ -874,8 +879,8 @@ def test_host_lock_refs_round_trip():
|
||||
_insert(core, [1], [10])
|
||||
leaf = core.match_prefix(MatchPrefixParams(key=_key([1]))).best_match_node
|
||||
core.commit_backup(leaf, torch.tensor([100], dtype=torch.int64), {})
|
||||
core.inc_host_lock_ref(leaf)
|
||||
core.dec_host_lock_ref(leaf)
|
||||
host_lock = core.inc_host_lock_ref(leaf)
|
||||
core.dec_host_lock_ref(leaf, host_lock.to_dec_params())
|
||||
core.sanity_check([], [])
|
||||
|
||||
|
||||
@@ -1204,11 +1209,6 @@ def test_swa_requires_the_sliding_window_size():
|
||||
)
|
||||
|
||||
|
||||
def test_swa_without_a_window_is_rejected_through_the_adapter():
|
||||
with pytest.raises(ValueError, match="requires swa_sliding_window_size"):
|
||||
_tree_core(tree_components=(ComponentType.FULL, ComponentType.SWA))
|
||||
|
||||
|
||||
def test_enable_hicache_constructs():
|
||||
mem_cache.RustUnifiedTreeCoreBinding(
|
||||
mem_cache.TreeCoreInitParamsBinding(enable_hicache=True),
|
||||
@@ -1260,6 +1260,14 @@ def _swa_tree_core(window: int = 8, **params_overrides) -> RustUnifiedTreeCore:
|
||||
)
|
||||
|
||||
|
||||
def test_swa_core_rejects_a_missing_or_non_positive_window():
|
||||
"""A zero window can never fill, so no boundary uuid would ever be stamped;
|
||||
the adapter refuses it up front instead of letting the core misbehave later."""
|
||||
for window in (None, 0, -1):
|
||||
with pytest.raises(ValueError, match="positive sliding_window_size"):
|
||||
_swa_tree_core(window=window)
|
||||
|
||||
|
||||
def test_write_back_load_back_ignores_auxiliary_nodes_for_pending_ownership():
|
||||
core = _swa_tree_core(window=4)
|
||||
core.set_hicache_enabled()
|
||||
@@ -1583,20 +1591,20 @@ def test_skipped_mamba_lock_survives_swa_only_release_through_the_adapter():
|
||||
node = core.match_prefix(MatchPrefixParams(key=_key([1, 2]))).best_match_node
|
||||
|
||||
owner = core.inc_lock_ref(node)
|
||||
skipped = core.inc_lock_ref(node, skip_lock_components=(ComponentType.MAMBA,))
|
||||
assert skipped.skip_lock_node_ids == {ComponentType.MAMBA: {node}}
|
||||
holder = core.inc_lock_ref(node, skip_lock_components=(ComponentType.MAMBA,))
|
||||
assert ComponentType.MAMBA not in owner.skipped_lock_components
|
||||
assert ComponentType.MAMBA in holder.skipped_lock_components
|
||||
assert core.mamba_protected_size() == 1
|
||||
|
||||
released = core.dec_swa_lock_only(
|
||||
node,
|
||||
skipped.swa_uuid_for_lock,
|
||||
skip_lock_node_ids=skipped.skip_lock_node_ids,
|
||||
)
|
||||
# The holder's receipt says it never took mamba: its early SWA release
|
||||
# must leave the owner's mamba lock alone.
|
||||
released = core.dec_swa_lock_only(node, holder.to_dec_params())
|
||||
assert dict(released.device_frees) == {}
|
||||
assert dict(released.host_frees) == {}
|
||||
assert core.mamba_protected_size() == 1
|
||||
|
||||
core.dec_lock_ref(node, skipped.to_dec_params(), skip_swa=True)
|
||||
core.dec_lock_ref(node, holder.to_dec_params(), skip_swa=True)
|
||||
assert core.mamba_protected_size() == 1
|
||||
core.dec_lock_ref(node, owner.to_dec_params())
|
||||
assert core.protected_size() == 0
|
||||
assert core.swa_protected_size() == 0
|
||||
@@ -1674,10 +1682,9 @@ def test_mamba_eviction_walk_frees_slots_through_the_adapter():
|
||||
assert torch.cat(device_frees[ComponentType.MAMBA]).tolist() == [7]
|
||||
assert core.mamba_evictable_size() == 1
|
||||
|
||||
# A pre-eviction node handle locked after the tombstoning lands in the
|
||||
# skip map, and the replay keeps the release off it.
|
||||
# A pre-eviction node handle still lock-round-trips: the segment lock
|
||||
# counts the tombstone and the paired release takes it back exactly.
|
||||
lock = core.inc_lock_ref(internal)
|
||||
assert internal in lock.skip_lock_node_ids[ComponentType.MAMBA]
|
||||
core.dec_lock_ref(internal, lock.to_dec_params())
|
||||
core.sanity_check([], [])
|
||||
|
||||
@@ -1892,8 +1899,6 @@ def test_component_device_value_round_trips():
|
||||
|
||||
|
||||
def test_lock_uuid_round_trips_through_dec_lock_ref():
|
||||
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
|
||||
|
||||
core = _swa_tree_core(window=2)
|
||||
first = _insert(core, [1, 2, 3], [10, 11, 12])
|
||||
# The window cap split the leaf: rebuild the in-window nodes' SWA values.
|
||||
@@ -1912,10 +1917,7 @@ def test_lock_uuid_round_trips_through_dec_lock_ref():
|
||||
assert core.swa_evictable_size() == 1
|
||||
core.dec_lock_ref(
|
||||
node,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=result.swa_uuid_for_lock,
|
||||
skip_lock_node_ids=result.skip_lock_node_ids,
|
||||
),
|
||||
DecLockRefParams(swa_uuid_for_lock=result.swa_uuid_for_lock),
|
||||
)
|
||||
# The uuid-bounded release returned the window to evictable.
|
||||
assert core.swa_protected_size() == 0
|
||||
@@ -1925,32 +1927,27 @@ def test_lock_uuid_round_trips_through_dec_lock_ref():
|
||||
assert again.swa_uuid_for_lock == result.swa_uuid_for_lock
|
||||
|
||||
|
||||
def test_swa_skip_map_crosses_the_binding_and_replays():
|
||||
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
|
||||
|
||||
def test_swa_tombstones_cross_the_binding_and_release_balanced():
|
||||
core = _swa_tree_core(window=8)
|
||||
_insert(core, [1, 2], [10, 11])
|
||||
second = _insert(core, [1, 2, 3, 4], [10, 11, 12, 13])
|
||||
leaf = second.cache_actions[-1].node_id
|
||||
# Only the leaf carries SWA; its ancestor is recorded as a tombstone skip.
|
||||
# Only the leaf carries SWA; the ancestor tombstone is counted too, and
|
||||
# the under-window walk reaches the root without stamping a uuid.
|
||||
core.set_component_device_value(
|
||||
leaf, ComponentType.SWA, torch.tensor([52, 53], dtype=torch.int64)
|
||||
)
|
||||
result = core.inc_lock_ref(leaf)
|
||||
assert result.skip_lock_node_ids[ComponentType.SWA]
|
||||
assert result.swa_uuid_for_lock is None
|
||||
assert core.swa_protected_size() == 2
|
||||
core.dec_lock_ref(
|
||||
leaf,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=result.swa_uuid_for_lock,
|
||||
skip_lock_node_ids=result.skip_lock_node_ids,
|
||||
),
|
||||
DecLockRefParams(swa_uuid_for_lock=result.swa_uuid_for_lock),
|
||||
)
|
||||
assert core.swa_protected_size() == 0
|
||||
|
||||
|
||||
def test_dec_swa_lock_only_frees_flow_after_the_full_release():
|
||||
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
|
||||
|
||||
core = _swa_tree_core(window=2)
|
||||
first = _insert(core, [1, 2], [10, 11])
|
||||
node = first.cache_actions[0].node_id
|
||||
@@ -1958,17 +1955,15 @@ def test_dec_swa_lock_only_frees_flow_after_the_full_release():
|
||||
node, ComponentType.SWA, torch.tensor([50, 51], dtype=torch.int64)
|
||||
)
|
||||
result = core.inc_lock_ref(node)
|
||||
# A non-None boundary: the window fills at the locked node itself.
|
||||
assert result.swa_uuid_for_lock is not None
|
||||
# The FULL lock releases first (skip_swa), then the early window release
|
||||
# finds a fully unlocked device leaf and evicts it in place.
|
||||
core.dec_lock_ref(
|
||||
node,
|
||||
DecLockRefParams(skip_lock_node_ids=result.skip_lock_node_ids),
|
||||
skip_swa=True,
|
||||
)
|
||||
core.dec_lock_ref(node, result.to_dec_params(), skip_swa=True)
|
||||
device_frees: dict = {}
|
||||
host_frees: dict = {}
|
||||
_accumulate_step(
|
||||
core.dec_swa_lock_only(node, result.swa_uuid_for_lock),
|
||||
core.dec_swa_lock_only(node, result.to_dec_params()),
|
||||
{},
|
||||
device_frees,
|
||||
host_frees,
|
||||
@@ -1977,7 +1972,7 @@ def test_dec_swa_lock_only_frees_flow_after_the_full_release():
|
||||
assert core.get_component_device_value(node, ComponentType.SWA) is None
|
||||
|
||||
|
||||
def test_dec_swa_lock_only_returns_the_window_frees():
|
||||
def test_dec_swa_lock_only_releases_once_and_a_repeat_dies_loud():
|
||||
core = _swa_tree_core(window=2)
|
||||
first = _insert(core, [1, 2, 3], [10, 11, 12])
|
||||
for action in first.cache_actions:
|
||||
@@ -1991,22 +1986,19 @@ def test_dec_swa_lock_only_returns_the_window_frees():
|
||||
device_frees: dict = {}
|
||||
host_frees: dict = {}
|
||||
_accumulate_step(
|
||||
core.dec_swa_lock_only(node, result.swa_uuid_for_lock),
|
||||
core.dec_swa_lock_only(node, result.to_dec_params()),
|
||||
{},
|
||||
device_frees,
|
||||
host_frees,
|
||||
)
|
||||
# The FULL lock still protects the path: the SWA release frees nothing and
|
||||
# the rebuilt values survive; a repeat release is a no-op.
|
||||
# The FULL lock still protects the path: the SWA release frees nothing
|
||||
# and the rebuilt values survive.
|
||||
assert device_frees == {}
|
||||
assert core.get_component_device_value(node, ComponentType.SWA) is not None
|
||||
_accumulate_step(
|
||||
core.dec_swa_lock_only(node, result.swa_uuid_for_lock),
|
||||
{},
|
||||
device_frees,
|
||||
host_frees,
|
||||
)
|
||||
assert device_frees == {}
|
||||
# A repeat release of the same window is a protocol violation and dies
|
||||
# at the segment instead of silently walking it.
|
||||
with pytest.raises(BaseException, match="SWA window release hit lock_ref=0"):
|
||||
core.dec_swa_lock_only(node, result.to_dec_params())
|
||||
|
||||
|
||||
def test_swa_rebuild_applies_through_the_python_allocator():
|
||||
@@ -2035,7 +2027,7 @@ def test_recover_with_locked_full_applies_through_the_python_allocator():
|
||||
# The decode advanced past the window: the SWA lock releases early, then
|
||||
# window eviction tombstones the SWA slot under the FULL lock (the state a
|
||||
# locked-full overlap recovers from); its frees return to the allocator.
|
||||
cache.dec_swa_lock_only(node, lock.swa_uuid_for_lock)
|
||||
cache.dec_swa_lock_only(node, lock.to_dec_params())
|
||||
tracker = {ComponentType.FULL: 0, ComponentType.SWA: 0}
|
||||
device_frees: dict = {}
|
||||
host_frees: dict = {}
|
||||
|
||||
@@ -4,7 +4,7 @@ import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT, ReqKvInfo
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import MatchResult
|
||||
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams, MatchResult
|
||||
from sglang.srt.session.streaming_session import SessionSlot, StreamingSession
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
@@ -54,6 +54,7 @@ class _FakeInnerCache:
|
||||
self.match_results = list(match_results or [])
|
||||
self.dec_lock_ref_calls = []
|
||||
self.dec_lock_ref_params = []
|
||||
self.dec_lock_ref_skip_swa = []
|
||||
|
||||
def cache_finished_req(self, *args, **kwargs):
|
||||
raise AssertionError("Streaming requests should not delegate to inner cache")
|
||||
@@ -66,6 +67,7 @@ class _FakeInnerCache:
|
||||
def dec_lock_ref(self, node, *args, **kwargs):
|
||||
self.dec_lock_ref_calls.append(node)
|
||||
self.dec_lock_ref_params.append(args[0] if args else kwargs.get("params"))
|
||||
self.dec_lock_ref_skip_swa.append(kwargs.get("skip_swa", False))
|
||||
|
||||
def supports_mamba(self):
|
||||
return False
|
||||
@@ -97,9 +99,9 @@ class _FakeReq:
|
||||
self.extra_key = None
|
||||
self.cache_salt = None
|
||||
self.last_node = None
|
||||
self.swa_uuid_for_lock = None
|
||||
self.skip_lock_node_ids = {}
|
||||
self.swa_branching_seqlen = None
|
||||
self.lock_receipt = DecLockRefParams()
|
||||
self.swa_prefix_lock_released = False
|
||||
self.to_finish = None
|
||||
self.finished_reason = None
|
||||
self.finished_len = None
|
||||
@@ -232,13 +234,11 @@ def test_nth_mid_abort_nukes_session_slot():
|
||||
assert req.kv.req_pool_idx is None
|
||||
|
||||
|
||||
def test_release_session_threads_mamba_skip_ids():
|
||||
"""release_session must forward the slot's skip_lock_node_ids to
|
||||
def test_release_session_threads_mamba_lock_receipt():
|
||||
"""release_session must forward the slot's mamba lock receipt to
|
||||
dec_lock_ref. The first req's last_node may be full-only-locked (mamba
|
||||
skipped at inc), so without the skip set the release would drop a mamba
|
||||
not taken at inc), so without the receipt the release would drop a mamba
|
||||
lock the session never took -- another request's, on a shared node."""
|
||||
from sglang.srt.mem_cache.unified_cache.components import ComponentType
|
||||
|
||||
req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128)
|
||||
req_to_token_pool = _FakeReqToTokenPool(req_to_token)
|
||||
allocator = _FakeAllocator()
|
||||
@@ -255,7 +255,6 @@ def test_release_session_threads_mamba_skip_ids():
|
||||
cache_protected_len=0,
|
||||
),
|
||||
last_node=lock_node,
|
||||
skip_lock_node_ids={ComponentType.MAMBA: {42}},
|
||||
)
|
||||
|
||||
tree_cache.release_session("session-a")
|
||||
@@ -263,7 +262,39 @@ def test_release_session_threads_mamba_skip_ids():
|
||||
assert inner.dec_lock_ref_calls == [lock_node]
|
||||
params = inner.dec_lock_ref_params[0]
|
||||
assert params is not None
|
||||
assert params.skip_lock_node_ids.get(ComponentType.MAMBA) == {42}
|
||||
assert params.skipped_lock_components == ()
|
||||
assert inner.dec_lock_ref_skip_swa == [False]
|
||||
|
||||
|
||||
def test_release_session_skips_swa_after_early_release():
|
||||
"""A slot saved from a req that early-released its SWA lock
|
||||
(swa_prefix_lock_released) must release with skip_swa, or the session
|
||||
close double-releases the SWA segment."""
|
||||
req_to_token = torch.arange(256, dtype=torch.int32).reshape(2, 128)
|
||||
req_to_token_pool = _FakeReqToTokenPool(req_to_token)
|
||||
allocator = _FakeAllocator()
|
||||
inner = _FakeInnerCache(req_to_token_pool, allocator, page_size=1)
|
||||
tree_cache = StreamingSession(inner)
|
||||
|
||||
lock_node = SimpleNamespace(id=42)
|
||||
tree_cache.slots["session-a"] = SessionSlot(
|
||||
kv=ReqKvInfo(
|
||||
req_pool_idx=0,
|
||||
kv_committed_len=50,
|
||||
kv_allocated_len=50,
|
||||
swa_evicted_seqlen=0,
|
||||
cache_protected_len=0,
|
||||
),
|
||||
last_node=lock_node,
|
||||
lock_receipt=DecLockRefParams(node_id=42, swa_uuid_for_lock=7),
|
||||
swa_prefix_lock_released=True,
|
||||
)
|
||||
|
||||
tree_cache.release_session("session-a")
|
||||
|
||||
assert inner.dec_lock_ref_calls == [lock_node]
|
||||
assert inner.dec_lock_ref_params[0].swa_uuid_for_lock == 7
|
||||
assert inner.dec_lock_ref_skip_swa == [True]
|
||||
|
||||
|
||||
def test_session_slot_does_not_restore_swa_branching_seqlen():
|
||||
|
||||
@@ -20,6 +20,7 @@ import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo, ScheduleBatch
|
||||
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
|
||||
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||
from sglang.srt.mem_cache.common import free_swa_out_of_window_slots
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
@@ -111,7 +112,7 @@ def _make_req(req_pool_idx, token_ids, cache_protected_len, tree):
|
||||
extra_key=None,
|
||||
cache_salt=None,
|
||||
last_node=tree.root_node,
|
||||
swa_uuid_for_lock=None,
|
||||
lock_receipt=DecLockRefParams(),
|
||||
swa_prefix_lock_released=False,
|
||||
prefix_indices=torch.tensor([], dtype=torch.int64, device=tree.device),
|
||||
_kv_committed_len=len(token_ids),
|
||||
|
||||
@@ -158,7 +158,7 @@ class TestSWALockReleaseLifecycle(CustomTestCase):
|
||||
self.assertFalse(leaf.swa_tombstone)
|
||||
self.assertTrue(tree.swa_lru_list.in_list(leaf))
|
||||
|
||||
tree.dec_swa_lock_only(leaf, swa_uuid_for_lock=swa_uuid)
|
||||
tree.dec_swa_lock_only(leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
|
||||
|
||||
self.assertTrue(leaf.swa_tombstone)
|
||||
self.assertFalse(tree.swa_lru_list.in_list(leaf))
|
||||
@@ -196,7 +196,7 @@ class TestSWALockReleaseLifecycle(CustomTestCase):
|
||||
swa_evictable_before = tree.swa_evictable_size_
|
||||
swa_avail_before = allocator.swa_available_size()
|
||||
|
||||
tree.dec_swa_lock_only(leaf_a, swa_uuid_for_lock=swa_uuid)
|
||||
tree.dec_swa_lock_only(leaf_a, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
|
||||
|
||||
self.assertFalse(internal.swa_tombstone)
|
||||
self.assertTrue(tree.swa_lru_list.in_list(internal))
|
||||
@@ -221,7 +221,7 @@ class TestSWALockReleaseLifecycle(CustomTestCase):
|
||||
inc_res = tree.inc_lock_ref(leaf)
|
||||
swa_uuid = inc_res.swa_uuid_for_lock
|
||||
|
||||
tree.dec_swa_lock_only(leaf, swa_uuid_for_lock=swa_uuid)
|
||||
tree.dec_swa_lock_only(leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
|
||||
self.assertTrue(leaf.swa_tombstone)
|
||||
self.assertEqual(leaf.full_lock_ref, 1)
|
||||
|
||||
@@ -308,7 +308,7 @@ class TestSWALockReleaseLifecycle(CustomTestCase):
|
||||
inc_res = tree.inc_lock_ref(leaf)
|
||||
swa_uuid = inc_res.swa_uuid_for_lock
|
||||
|
||||
tree.dec_swa_lock_only(leaf, swa_uuid_for_lock=swa_uuid)
|
||||
tree.dec_swa_lock_only(leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
|
||||
self.assertTrue(leaf.swa_tombstone)
|
||||
|
||||
swa_evictable_before_delete = tree.swa_evictable_size_
|
||||
@@ -354,7 +354,9 @@ class TestSWALockReleaseLifecycle(CustomTestCase):
|
||||
swa_avail_before = allocator.swa_available_size()
|
||||
full_avail_before = allocator.full_available_size()
|
||||
|
||||
tree.dec_swa_lock_only(leaf, swa_uuid_for_lock=swa_uuid)
|
||||
tree.dec_swa_lock_only(
|
||||
leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid)
|
||||
)
|
||||
|
||||
self.assertTrue(leaf.swa_tombstone)
|
||||
self.assertFalse(tree.swa_lru_list.in_list(leaf))
|
||||
@@ -406,7 +408,7 @@ class TestSWALockReleaseLifecycle(CustomTestCase):
|
||||
swa_evictable_before = tree.swa_evictable_size_
|
||||
swa_avail_before = allocator.swa_available_size()
|
||||
|
||||
tree.dec_swa_lock_only(leaf_a, swa_uuid_for_lock=swa_uuid)
|
||||
tree.dec_swa_lock_only(leaf_a, DecLockRefParams(swa_uuid_for_lock=swa_uuid))
|
||||
|
||||
# Leaf side: tombstoned and pages freed.
|
||||
self.assertTrue(leaf_a.swa_tombstone)
|
||||
|
||||
@@ -794,7 +794,7 @@ class TestSWA(unittest.TestCase):
|
||||
req.extra_key = None
|
||||
req.cache_salt = None
|
||||
req.last_node = tree.root_node
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.kv.swa_evicted_seqlen = 0
|
||||
req.kv.cache_protected_len = 1
|
||||
# Intentionally mismatch to ensure code does not use len(prefix_indices).
|
||||
@@ -832,7 +832,7 @@ class TestSWA(unittest.TestCase):
|
||||
req2.extra_key = None
|
||||
req2.cache_salt = None
|
||||
req2.last_node = tree.root_node
|
||||
req2.swa_uuid_for_lock = None
|
||||
req2.lock_receipt = DecLockRefParams()
|
||||
req2.kv.swa_evicted_seqlen = 0
|
||||
req2.kv.cache_protected_len = 1
|
||||
req2.prefix_indices = torch.tensor([21, 22, 23, 24, 25], device=tree.device)
|
||||
@@ -1322,7 +1322,7 @@ class TestCacheUnfinishedReqEvictedPrefix(CustomTestCase):
|
||||
req.cache_salt = None
|
||||
req.kv.cache_protected_len = 0
|
||||
req.last_node = tree.root_node
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.prefix_indices = torch.empty(0, dtype=torch.int64, device=tree.device)
|
||||
req.kv.swa_evicted_seqlen = evicted
|
||||
|
||||
|
||||
@@ -25,7 +25,6 @@ from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
DecLockRefParams,
|
||||
EvictParams,
|
||||
InsertParams,
|
||||
MatchPrefixParams,
|
||||
@@ -588,7 +587,7 @@ def bench_lock_unlock(
|
||||
lr = env.tree.inc_lock_ref(node)
|
||||
env.tree.dec_lock_ref(
|
||||
node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)),
|
||||
lr.to_dec_params(),
|
||||
)
|
||||
|
||||
warmup = min(20, num_pairs // 10)
|
||||
@@ -633,9 +632,7 @@ def bench_cache_finished(
|
||||
if v is None:
|
||||
env.tree.dec_lock_ref(
|
||||
node,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)
|
||||
),
|
||||
lr.to_dec_params(),
|
||||
)
|
||||
continue
|
||||
kv_indices = torch.cat([mr.device_indices, v])
|
||||
@@ -652,8 +649,8 @@ def bench_cache_finished(
|
||||
req.last_node = node
|
||||
req.kv.cache_protected_len = matched_len
|
||||
req.kv.kv_committed_len = len(seq)
|
||||
if hasattr(lr, "swa_uuid_for_lock"):
|
||||
req.swa_uuid_for_lock = lr.swa_uuid_for_lock
|
||||
if hasattr(lr, "to_dec_params"):
|
||||
req.lock_receipt = lr.to_dec_params()
|
||||
env.rtp.req_to_token[req.kv.req_pool_idx, : len(kv_indices)] = kv_indices
|
||||
req_items.append(req)
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ import unittest
|
||||
from array import array
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass, replace
|
||||
from types import SimpleNamespace
|
||||
from typing import Optional
|
||||
from unittest import mock
|
||||
|
||||
@@ -25,7 +26,7 @@ from sglang.srt.disaggregation.kv_events import (
|
||||
StorageMedium,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req, ReqKvInfo
|
||||
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
@@ -98,6 +99,7 @@ from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
from sglang.srt.session.streaming_session import SessionSlot
|
||||
from sglang.srt.utils import get_device
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
@@ -1446,9 +1448,7 @@ class UnifiedRadixCacheSuite:
|
||||
# Unlock -> should now be evictable
|
||||
cache.dec_lock_ref(
|
||||
m.last_device_node,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=getattr(lock_result, "swa_uuid_for_lock", None)
|
||||
),
|
||||
lock_result.to_dec_params(),
|
||||
)
|
||||
result = cache.evict(EvictParams(num_tokens=len(seq_a)))
|
||||
self.assertGreaterEqual(result.num_tokens_evicted, len(seq_a))
|
||||
@@ -1539,7 +1539,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_committed_len = kv_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
req.full_untruncated_fill_ids = array("q", input_ids + output_ids)
|
||||
req.set_extend_range(
|
||||
@@ -1580,7 +1580,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_allocated_len = kv_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
if self.cfg.has_mamba:
|
||||
req.kv.mamba_last_track_seqlen = kv_len
|
||||
@@ -1624,7 +1624,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_committed_len = kv_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.swa_prefix_lock_released = True
|
||||
req.extra_key = None
|
||||
req.full_untruncated_fill_ids = array("q", tokens)
|
||||
@@ -1659,7 +1659,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_committed_len = kv_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
if self.cfg.has_mamba:
|
||||
req.kv.mamba_last_track_seqlen = kv_len
|
||||
@@ -1673,7 +1673,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
|
||||
req.lock_receipt,
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -1698,7 +1698,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_committed_len = len(tokens)
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
req.kv.swa_evicted_seqlen = evicted_len
|
||||
|
||||
@@ -1716,7 +1716,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
|
||||
req.lock_receipt,
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -1800,7 +1800,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_committed_len = kv_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
req.full_untruncated_fill_ids = array("q", input_ids)
|
||||
req.set_extend_range(
|
||||
@@ -1925,7 +1925,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_committed_len = kv_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
req.kv.swa_evicted_seqlen = 0
|
||||
|
||||
@@ -1955,7 +1955,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
|
||||
req.lock_receipt,
|
||||
)
|
||||
cache.dec_lock_ref(last_device_node, lock_result.to_dec_params())
|
||||
cache.sanity_check()
|
||||
@@ -2021,7 +2021,7 @@ class UnifiedRadixCacheSuite:
|
||||
1,
|
||||
"Mamba locked before release",
|
||||
)
|
||||
cache.dec_swa_lock_only(node_a, lock_result.swa_uuid_for_lock)
|
||||
cache.dec_swa_lock_only(node_a, lock_result.to_dec_params())
|
||||
self.assertEqual(_device_lock_ref(cache, node_a, ComponentType.SWA), 0)
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node_a, ComponentType.MAMBA),
|
||||
@@ -2071,7 +2071,58 @@ class UnifiedRadixCacheSuite:
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
cache.dec_lock_ref(node_a, DecLockRefParams(swa_uuid_for_lock=None))
|
||||
cache.dec_lock_ref(
|
||||
node_a, DecLockRefParams(swa_uuid_for_lock=None), skip_swa=True
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
def test_mamba_opt_out_holder_cannot_release_another_holders_mamba_lock(self):
|
||||
"""Holder A takes the mamba lock; holder B opts out (skip_lock_components=(ComponentType.MAMBA,))
|
||||
on the same node. B's early SWA release and final release must leave
|
||||
A's mamba lock intact -- a lost/defaulted receipt on B's side used to
|
||||
decrement A's lock without tripping any assert."""
|
||||
if not self.cfg.has_swa or not self.cfg.has_mamba:
|
||||
self.skipTest("requires SWA and Mamba components")
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
|
||||
seq = self._make_seq(
|
||||
1, (self.cfg.sliding_window_size // self.cfg.page_size) + 4
|
||||
)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq)
|
||||
node = cache.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq)))
|
||||
).last_device_node
|
||||
self.assertIsNotNone(_device_value(cache, node, ComponentType.MAMBA))
|
||||
|
||||
lock_a = cache.inc_lock_ref(node)
|
||||
lock_b = cache.inc_lock_ref(node, skip_lock_components=(ComponentType.MAMBA,))
|
||||
self.assertNotIn(ComponentType.MAMBA, lock_a.skipped_lock_components)
|
||||
self.assertIn(ComponentType.MAMBA, lock_b.skipped_lock_components)
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node, ComponentType.MAMBA), 1, "only A holds mamba"
|
||||
)
|
||||
|
||||
# B: early SWA release, then final release -- both replay B's receipt.
|
||||
cache.dec_swa_lock_only(node, lock_b.to_dec_params())
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node, ComponentType.MAMBA),
|
||||
1,
|
||||
"B's early release spares A",
|
||||
)
|
||||
cache.dec_lock_ref(node, lock_b.to_dec_params(), skip_swa=True)
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node, ComponentType.MAMBA),
|
||||
1,
|
||||
"B's final release spares A",
|
||||
)
|
||||
|
||||
cache.dec_swa_lock_only(node, lock_a.to_dec_params())
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node, ComponentType.MAMBA),
|
||||
0,
|
||||
"A's release drops its own lock",
|
||||
)
|
||||
cache.dec_lock_ref(node, lock_a.to_dec_params(), skip_swa=True)
|
||||
cache.sanity_check()
|
||||
|
||||
def test_swa_early_release_drops_co_located_mamba_lock(self):
|
||||
@@ -2106,11 +2157,9 @@ class UnifiedRadixCacheSuite:
|
||||
# Early SWA release (decode advanced past the window), via the public
|
||||
# path the scheduler calls. The leaf's SWA is tombstoned and the
|
||||
# co-located lower-tier Mamba lock must drop in the same release.
|
||||
cache.dec_swa_lock_only(node_a, lock_result.swa_uuid_for_lock)
|
||||
cache.dec_swa_lock_only(node_a, lock_result.to_dec_params())
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node_a, ComponentType.SWA),
|
||||
0,
|
||||
"SWA early-released",
|
||||
_device_lock_ref(cache, node_a, ComponentType.SWA), 0, "SWA early-released"
|
||||
)
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node_a, ComponentType.MAMBA),
|
||||
@@ -2140,14 +2189,11 @@ class UnifiedRadixCacheSuite:
|
||||
skipped = cache.inc_lock_ref(
|
||||
node_a, skip_lock_components=(ComponentType.MAMBA,)
|
||||
)
|
||||
self.assertEqual(skipped.skip_lock_node_ids, {ComponentType.MAMBA: {node_a}})
|
||||
self.assertNotIn(ComponentType.MAMBA, owner.skipped_lock_components)
|
||||
self.assertIn(ComponentType.MAMBA, skipped.skipped_lock_components)
|
||||
self.assertEqual(_device_lock_ref(cache, node_a, ComponentType.MAMBA), 1)
|
||||
|
||||
cache.dec_swa_lock_only(
|
||||
node_a,
|
||||
skipped.swa_uuid_for_lock,
|
||||
skip_lock_node_ids=skipped.skip_lock_node_ids,
|
||||
)
|
||||
cache.dec_swa_lock_only(node_a, skipped.to_dec_params())
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node_a, ComponentType.MAMBA),
|
||||
1,
|
||||
@@ -2312,7 +2358,7 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertGreaterEqual(_device_lock_ref(cache, node_a, ComponentType.MAMBA), 1)
|
||||
self.assertGreaterEqual(_device_lock_ref(cache, node_a, ComponentType.FULL), 1)
|
||||
|
||||
cache.dec_swa_lock_only(node_a, lock_result.swa_uuid_for_lock)
|
||||
cache.dec_swa_lock_only(node_a, lock_result.to_dec_params())
|
||||
self.assertEqual(
|
||||
_device_lock_ref(cache, node_a, ComponentType.SWA), 0, "SWA released"
|
||||
)
|
||||
@@ -2386,7 +2432,7 @@ class UnifiedRadixCacheSuite:
|
||||
cache.sanity_check()
|
||||
cache.dec_lock_ref(
|
||||
leaf,
|
||||
DecLockRefParams(swa_uuid_for_lock=lock_result.swa_uuid_for_lock),
|
||||
lock_result.to_dec_params(),
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -2548,7 +2594,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_committed_len = pre_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
|
||||
swa_avail_before = allocator.swa_attn_allocator.available_size()
|
||||
@@ -2575,7 +2621,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
|
||||
req.lock_receipt,
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -2637,7 +2683,7 @@ class UnifiedRadixCacheSuite:
|
||||
req.kv.kv_committed_len = pre_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
|
||||
with envs.SGLANG_OPT_UNIFIED_CACHE_FREE_OUT_OF_WINDOW_SLOTS.override(True):
|
||||
@@ -2651,7 +2697,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
|
||||
req.lock_receipt,
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -2715,18 +2761,26 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertIsNotNone(_device_value(cache, node, ComponentType.FULL))
|
||||
self.assertIsNotNone(_device_value(cache, node, aux))
|
||||
|
||||
lock_result = cache.inc_lock_ref(node)
|
||||
self.assertGreater(_device_lock_ref(cache, node, ComponentType.FULL), 0)
|
||||
self.assertGreater(_device_lock_ref(cache, node, aux), 0)
|
||||
|
||||
# Reach the "FULL locked, aux unlocked" state the way production does:
|
||||
# mamba via the decode-hold opt-out (skip_lock_components=(ComponentType.MAMBA,)), SWA via the
|
||||
# early window release (its own first-class op).
|
||||
aux_len = len(_device_value(cache, node, aux))
|
||||
cache.tree_core.set_component_protected_size(
|
||||
aux, cache.tree_core.component_protected_size(aux) - aux_len
|
||||
)
|
||||
cache.tree_core.set_component_evictable_size(
|
||||
aux, cache.tree_core.component_evictable_size(aux) + aux_len
|
||||
)
|
||||
cache.tree_core.set_component_device_lock_ref(node, aux, 0)
|
||||
if aux == ComponentType.MAMBA:
|
||||
lock_result = cache.inc_lock_ref(
|
||||
node, skip_lock_components=(ComponentType.MAMBA,)
|
||||
)
|
||||
else:
|
||||
lock_result = cache.inc_lock_ref(node)
|
||||
self.assertGreater(_device_lock_ref(cache, node, aux), 0)
|
||||
cache.dec_swa_lock_only(
|
||||
node,
|
||||
lock_result.to_dec_params(),
|
||||
)
|
||||
# FULL still locked -> not a device leaf -> no inline evict; the
|
||||
# value stays evictable for the explicit aux eviction below.
|
||||
self.assertIsNotNone(_device_value(cache, node, aux))
|
||||
self.assertGreater(_device_lock_ref(cache, node, ComponentType.FULL), 0)
|
||||
self.assertEqual(_device_lock_ref(cache, node, aux), 0)
|
||||
self.assertFalse(cache.tree_core.is_device_evictable_leaf(node))
|
||||
|
||||
evict_params = EvictParams(num_tokens=0)
|
||||
@@ -2747,7 +2801,8 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
node,
|
||||
DecLockRefParams(swa_uuid_for_lock=lock_result.swa_uuid_for_lock),
|
||||
lock_result.to_dec_params(),
|
||||
skip_swa=(aux == ComponentType.SWA),
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -2789,9 +2844,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
m_base.last_device_node,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=getattr(lock_result, "swa_uuid_for_lock", None)
|
||||
),
|
||||
lock_result.to_dec_params(),
|
||||
)
|
||||
# After unlock, base should be in evictable_device_leaves
|
||||
self.assertTrue(
|
||||
@@ -2910,7 +2963,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
m.last_device_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)),
|
||||
lr.to_dec_params(),
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -2973,7 +3026,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
m.last_device_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)),
|
||||
lr.to_dec_params(),
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -5094,7 +5147,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
cache.dec_lock_ref(
|
||||
m.last_device_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)),
|
||||
lr.to_dec_params(),
|
||||
)
|
||||
m = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", leaf))))
|
||||
self.assertGreaterEqual(len(m.device_indices), len(base))
|
||||
@@ -6070,9 +6123,7 @@ class UnifiedRadixCacheSuite:
|
||||
finally:
|
||||
cache.dec_lock_ref(
|
||||
parent,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=getattr(lock_result, "swa_uuid_for_lock", None)
|
||||
),
|
||||
lock_result.to_dec_params(),
|
||||
)
|
||||
self.assertTrue(cache.tree_core.is_full_device_evicted(leaf))
|
||||
self.assertTrue(cache.tree_core.is_backuped(leaf))
|
||||
@@ -7050,7 +7101,7 @@ class UnifiedRadixCacheSuite:
|
||||
)
|
||||
|
||||
temp_lock = cache.inc_lock_ref(leaf)
|
||||
self.assertEqual(_device_lock_ref(cache, tombstone, ComponentType.SWA), 0)
|
||||
self.assertEqual(_device_lock_ref(cache, tombstone, ComponentType.SWA), 1)
|
||||
|
||||
xfer = cache.tree_core.build_hicache_transfers(
|
||||
ComponentType.SWA, leaf, CacheTransferPhase.LOAD_BACK
|
||||
@@ -7069,7 +7120,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
load_back_lock = cache.inc_lock_ref(leaf)
|
||||
request_lock = cache.inc_lock_ref(leaf)
|
||||
self.assertEqual(_device_lock_ref(cache, tombstone, ComponentType.SWA), 2)
|
||||
self.assertEqual(_device_lock_ref(cache, tombstone, ComponentType.SWA), 3)
|
||||
|
||||
cache.dec_lock_ref(leaf, temp_lock.to_dec_params())
|
||||
self.assertEqual(_device_lock_ref(cache, tombstone, ComponentType.SWA), 2)
|
||||
@@ -7167,14 +7218,12 @@ class UnifiedRadixCacheSuite:
|
||||
self._release_ongoing_load_back_locks(cache)
|
||||
cache.sanity_check()
|
||||
|
||||
def test_hicache_full_temp_lock_skips_evicted_anchor_and_mirrors_on_release(
|
||||
def test_hicache_full_temp_lock_covers_evicted_anchor_and_mirrors_on_release(
|
||||
self,
|
||||
):
|
||||
"""Acquire records the evicted anchor in skip_lock_node_ids (phase 1)
|
||||
and locks device-on ancestors only (phase 2). After load_back
|
||||
restores the anchor, a second acquire covers it; releasing the
|
||||
first must mirror the skip so the anchor's lock_ref is not
|
||||
decremented twice.
|
||||
"""Segment locks count the evicted anchor too (no skip receipts), so
|
||||
a value restored mid-hold stays correctly attributed: each release
|
||||
takes back exactly its own ref regardless of interleaved holders.
|
||||
"""
|
||||
if self._skip_unsupported_hicache_test():
|
||||
return
|
||||
@@ -7189,25 +7238,37 @@ class UnifiedRadixCacheSuite:
|
||||
self._simulate_backup_tree(cache)
|
||||
|
||||
anchor_value = _device_value(cache, anchor, ComponentType.FULL)
|
||||
# Simulate the anchor's FULL device eviction: drop the value and take
|
||||
# its tokens out of the evictable ledger, as a real evict would.
|
||||
cache.tree_core.set_component_device_value_raw(anchor, ComponentType.FULL, None)
|
||||
cache.tree_core.set_component_evictable_size(
|
||||
ComponentType.FULL,
|
||||
cache.tree_core.component_evictable_size(ComponentType.FULL)
|
||||
- len(anchor_value),
|
||||
)
|
||||
|
||||
self.assertEqual(_device_lock_ref(cache, anchor, ComponentType.FULL), 0)
|
||||
self.assertEqual(_device_lock_ref(cache, y, ComponentType.FULL), 0)
|
||||
self.assertEqual(_device_lock_ref(cache, a, ComponentType.FULL), 0)
|
||||
|
||||
temp_lock = cache.inc_lock_ref(anchor)
|
||||
self.assertEqual(_device_lock_ref(cache, anchor, ComponentType.FULL), 0)
|
||||
self.assertEqual(_device_lock_ref(cache, anchor, ComponentType.FULL), 1)
|
||||
self.assertEqual(_device_lock_ref(cache, y, ComponentType.FULL), 1)
|
||||
self.assertEqual(_device_lock_ref(cache, a, ComponentType.FULL), 1)
|
||||
self.assertIn(ComponentType.FULL, temp_lock.skip_lock_node_ids)
|
||||
self.assertIn(anchor, temp_lock.skip_lock_node_ids[ComponentType.FULL])
|
||||
|
||||
# Restore the value mid-hold: a value materialized under lock is
|
||||
# protected until the last release, exactly as a load-back credits it.
|
||||
cache.tree_core.set_component_device_value_raw(
|
||||
anchor, ComponentType.FULL, anchor_value
|
||||
)
|
||||
cache.tree_core.set_component_protected_size(
|
||||
ComponentType.FULL,
|
||||
cache.tree_core.component_protected_size(ComponentType.FULL)
|
||||
+ len(anchor_value),
|
||||
)
|
||||
|
||||
second_lock = cache.inc_lock_ref(anchor)
|
||||
self.assertEqual(_device_lock_ref(cache, anchor, ComponentType.FULL), 1)
|
||||
self.assertEqual(_device_lock_ref(cache, anchor, ComponentType.FULL), 2)
|
||||
self.assertEqual(_device_lock_ref(cache, y, ComponentType.FULL), 2)
|
||||
self.assertEqual(_device_lock_ref(cache, a, ComponentType.FULL), 2)
|
||||
|
||||
@@ -7249,7 +7310,7 @@ class UnifiedRadixCacheSuite:
|
||||
)
|
||||
|
||||
temp_lock = cache.inc_lock_ref(node)
|
||||
self.assertEqual(_device_lock_ref(cache, node, ComponentType.MAMBA), 0)
|
||||
self.assertEqual(_device_lock_ref(cache, node, ComponentType.MAMBA), 1)
|
||||
|
||||
xfer = cache.tree_core.build_hicache_transfers(
|
||||
ComponentType.MAMBA, node, CacheTransferPhase.LOAD_BACK
|
||||
@@ -7266,7 +7327,7 @@ class UnifiedRadixCacheSuite:
|
||||
|
||||
load_back_lock = cache.inc_lock_ref(node)
|
||||
request_lock = cache.inc_lock_ref(node)
|
||||
self.assertEqual(_device_lock_ref(cache, node, ComponentType.MAMBA), 2)
|
||||
self.assertEqual(_device_lock_ref(cache, node, ComponentType.MAMBA), 3)
|
||||
|
||||
cache.dec_lock_ref(node, temp_lock.to_dec_params())
|
||||
self.assertEqual(_device_lock_ref(cache, node, ComponentType.MAMBA), 2)
|
||||
@@ -7274,6 +7335,9 @@ class UnifiedRadixCacheSuite:
|
||||
cache.dec_lock_ref(node, load_back_lock.to_dec_params())
|
||||
cache.dec_lock_ref(node, request_lock.to_dec_params())
|
||||
self.assertEqual(_device_lock_ref(cache, node, ComponentType.MAMBA), 0)
|
||||
# The commit ran under held locks: the restored value must have been
|
||||
# credited to protected, or the ledger drifts on the final release.
|
||||
cache.sanity_check()
|
||||
|
||||
def test_hicache_mixed_backup_evict_insert(self):
|
||||
"""Complex scenario: backup some, evict, insert new, verify invariants."""
|
||||
@@ -7337,9 +7401,7 @@ class UnifiedRadixCacheSuite:
|
||||
finally:
|
||||
cache.dec_lock_ref(
|
||||
parent,
|
||||
DecLockRefParams(
|
||||
swa_uuid_for_lock=getattr(lr, "swa_uuid_for_lock", None)
|
||||
),
|
||||
lr.to_dec_params(),
|
||||
)
|
||||
|
||||
self.assertTrue(
|
||||
@@ -7627,7 +7689,7 @@ class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase):
|
||||
req.kv.kv_committed_len = len(tokens)
|
||||
req.kv.kv_allocated_len = len(tokens)
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
req.kv.mamba_last_track_seqlen = len(tokens)
|
||||
return req
|
||||
@@ -8140,7 +8202,7 @@ class TestResumableInsertWalk(_InsertWalkSuite):
|
||||
|
||||
# Fill the host pool below len(top) free, keeping the on-path H-leaf
|
||||
# the oldest host entry and pinning the unbacked path root.
|
||||
cache.inc_lock_ref(top)
|
||||
top_lock = cache.inc_lock_ref(top)
|
||||
host_pool = cache.cache_controller.mem_pool_host
|
||||
start = 1000
|
||||
top_len = _node_key_length(cache, top)
|
||||
@@ -8158,7 +8220,7 @@ class TestResumableInsertWalk(_InsertWalkSuite):
|
||||
cache.writing_check(write_back=True)
|
||||
cache.evict(EvictParams(num_tokens=count))
|
||||
self.assertTrue(cache.tree_core.is_full_device_evicted(filler))
|
||||
cache.dec_lock_ref(top)
|
||||
cache.dec_lock_ref(top, top_lock.to_dec_params())
|
||||
|
||||
# The crossing backup evicts exactly the on-path H-leaf, then the
|
||||
# remaining suffix is recreated as a fresh leaf.
|
||||
@@ -8527,11 +8589,13 @@ class TestResumableInsertWalkSWA(_InsertWalkSuite):
|
||||
|
||||
lock_result = cache.inc_lock_ref(node)
|
||||
self.assertGreaterEqual(_device_lock_ref(cache, node, ComponentType.SWA), 1)
|
||||
cache.dec_swa_lock_only(node, lock_result.swa_uuid_for_lock)
|
||||
cache.dec_swa_lock_only(node, lock_result.to_dec_params())
|
||||
self.assertEqual(_device_lock_ref(cache, node, ComponentType.SWA), 0)
|
||||
self.assertGreaterEqual(_device_lock_ref(cache, node, ComponentType.FULL), 1)
|
||||
|
||||
cache.dec_lock_ref(node, DecLockRefParams(swa_uuid_for_lock=None))
|
||||
cache.dec_lock_ref(
|
||||
node, DecLockRefParams(swa_uuid_for_lock=None), skip_swa=True
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
|
||||
@@ -8661,7 +8725,7 @@ class TestReturnedValuesDrain(_InsertWalkSuite):
|
||||
(
|
||||
"dec_swa_lock_only",
|
||||
lambda: make(DecSwaLockOnlyResult),
|
||||
lambda: cache.dec_swa_lock_only(node),
|
||||
lambda: cache.dec_swa_lock_only(node, DecLockRefParams()),
|
||||
None,
|
||||
),
|
||||
]
|
||||
@@ -9182,7 +9246,7 @@ class TestSWAWindowUnderBigramKey(CustomTestCase):
|
||||
req.kv.kv_committed_len = seq_len
|
||||
req.last_node = cache.root_node_handle()
|
||||
req.kv.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = None
|
||||
req.lock_receipt = DecLockRefParams()
|
||||
req.extra_key = None
|
||||
|
||||
with envs.SGLANG_OPT_UNIFIED_CACHE_FREE_OUT_OF_WINDOW_SLOTS.override(True):
|
||||
@@ -9204,7 +9268,7 @@ class TestSWAWindowUnderBigramKey(CustomTestCase):
|
||||
|
||||
cache.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
|
||||
req.lock_receipt,
|
||||
)
|
||||
cache.sanity_check()
|
||||
|
||||
@@ -9402,5 +9466,445 @@ class TestAnchorLockOutcomePolicy(CustomTestCase):
|
||||
cache.match_prefix.assert_called_once()
|
||||
|
||||
|
||||
@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA")
|
||||
class TestSegmentLockProtocol(_InsertWalkSuite):
|
||||
"""Segment-lock protocol regressions, replaying the production lock-theft
|
||||
failure classes (F1/F2) and the split hazards.
|
||||
|
||||
The protocol: a lock covers the contiguous node segment
|
||||
[start, boundary-uuid], counting every node (tombstones included), so a
|
||||
release needs only the receipt (anchor node, boundary uuid, skipped
|
||||
components) and any ref==0 met inside the segment is a hard protocol
|
||||
violation. The replays read the tree through the inspection interface, so
|
||||
they run unchanged against the Python and Rust cores.
|
||||
"""
|
||||
|
||||
cfg = CacheConfig(
|
||||
components=(ComponentType.FULL, ComponentType.SWA), sliding_window_size=8
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _swa_ref(cache, node_id):
|
||||
return _device_lock_ref(cache, node_id, ComponentType.SWA)
|
||||
|
||||
@staticmethod
|
||||
def _segment(cache, leaf_id, window):
|
||||
"""Node ids from leaf up to the position-based window boundary."""
|
||||
nodes, covered, cur = [], 0, leaf_id
|
||||
while not cache.tree_core.is_root(cur) and covered < window:
|
||||
nodes.append(cur)
|
||||
covered += cache.tree_core.get_node_key_length(cur)
|
||||
cur = _node_parent(cache, cur)
|
||||
return nodes
|
||||
|
||||
def _match_leaf(self, cache, seq):
|
||||
m = cache.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq))))
|
||||
return m.last_device_node
|
||||
|
||||
@staticmethod
|
||||
def _deepest(cache):
|
||||
"""Structurally deepest node — bypasses SWA match validation, which
|
||||
never adopts a holed window (simulates the stale-relock drift case)."""
|
||||
node_id = cache.root_node_handle()
|
||||
while True:
|
||||
children = _node_children(cache, node_id)
|
||||
if not children:
|
||||
return node_id
|
||||
node_id = children[0]
|
||||
|
||||
def _assert_protocol_violation(self, fn, fragment):
|
||||
"""The Python core asserts; the Rust core panics (a BaseException
|
||||
subclass at the PyO3 boundary). Either way the message names the
|
||||
violation and the operation never completes silently."""
|
||||
try:
|
||||
fn()
|
||||
except (KeyboardInterrupt, SystemExit):
|
||||
raise
|
||||
except BaseException as exc: # pyo3 PanicException derives from BaseException
|
||||
self.assertIn(fragment, str(exc))
|
||||
else:
|
||||
self.fail(f"protocol violation went unreported: {fragment}")
|
||||
|
||||
def test_rebuilt_tombstone_relock_release_no_theft(self):
|
||||
"""F1 attribution replay: A locks a window containing a tombstone; the
|
||||
tombstone is rebuilt and locked by B mid-hold; A's release must leave
|
||||
B's refs intact (the old skip-set protocol decremented B's lock)."""
|
||||
sw = self.cfg.sliding_window_size
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
seq = self._make_seq(1, 2 * sw)
|
||||
# SWA data only for the last sw//2 positions: the window has a hole.
|
||||
# As in production, the evicted prefix's SWA slots are released before
|
||||
# the insert so a later FULL free finds no live SWA peer.
|
||||
swa_evicted = len(seq) - sw // 2
|
||||
value = self._alloc(allocator, len(seq))
|
||||
allocator.free_swa(value[:swa_evicted])
|
||||
cache.insert(
|
||||
InsertParams(
|
||||
key=RadixKey(array("q", seq)),
|
||||
value=value,
|
||||
swa_evicted_seqlen=swa_evicted,
|
||||
)
|
||||
)
|
||||
leaf = self._deepest(cache)
|
||||
segment = self._segment(cache, leaf, sw)
|
||||
self.assertTrue(
|
||||
any(_device_value(cache, n, ComponentType.SWA) is None for n in segment),
|
||||
"fixture must place a tombstone inside the window",
|
||||
)
|
||||
|
||||
lock_a = cache.inc_lock_ref(leaf)
|
||||
# Count-everything: every segment node carries A's ref, tombstones
|
||||
# included, and the boundary uuid is always stamped.
|
||||
self.assertIsNotNone(lock_a.swa_uuid_for_lock)
|
||||
self.assertEqual(lock_a.node_id, leaf)
|
||||
for n in segment:
|
||||
self.assertEqual(self._swa_ref(cache, n), 1)
|
||||
cache.sanity_check()
|
||||
|
||||
# Rebuild the tombstones under A's lock (Recover path: FULL is
|
||||
# locked); the rebuilt values must be credited to protected.
|
||||
cache.insert(
|
||||
InsertParams(
|
||||
key=RadixKey(array("q", seq)),
|
||||
value=self._alloc(allocator, len(seq)),
|
||||
swa_evicted_seqlen=0,
|
||||
)
|
||||
)
|
||||
cache.sanity_check()
|
||||
leaf = self._deepest(cache)
|
||||
segment = self._segment(cache, leaf, sw)
|
||||
for n in segment:
|
||||
self.assertIsNotNone(_device_value(cache, n, ComponentType.SWA))
|
||||
|
||||
lock_b = cache.inc_lock_ref(leaf)
|
||||
for n in segment:
|
||||
self.assertEqual(self._swa_ref(cache, n), 2)
|
||||
|
||||
# THE regression: A's release takes back exactly A's refs.
|
||||
cache.dec_lock_ref(leaf, lock_a.to_dec_params())
|
||||
for n in segment:
|
||||
self.assertEqual(self._swa_ref(cache, n), 1)
|
||||
cache.sanity_check()
|
||||
|
||||
cache.dec_lock_ref(leaf, lock_b.to_dec_params())
|
||||
for n in segment:
|
||||
self.assertEqual(self._swa_ref(cache, n), 0)
|
||||
cache.sanity_check()
|
||||
|
||||
def test_release_without_receipt_fails_loud(self):
|
||||
"""A release missing its boundary uuid must die at the segment edge
|
||||
(ref==0 assert) instead of silently walking to root stealing other
|
||||
holders' locks — the F1 failure made loud."""
|
||||
sw = self.cfg.sliding_window_size
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
seq = self._make_seq(1, 3 * sw)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq)
|
||||
leaf = self._match_leaf(cache, seq)
|
||||
lock = cache.inc_lock_ref(leaf)
|
||||
self.assertIsNotNone(lock.swa_uuid_for_lock)
|
||||
|
||||
self._assert_protocol_violation(
|
||||
lambda: cache.dec_lock_ref(leaf, DecLockRefParams(swa_uuid_for_lock=None)),
|
||||
"lock_ref=0",
|
||||
)
|
||||
|
||||
def test_release_on_another_node_fails_loud(self):
|
||||
"""The receipt anchors the lock on the node it was taken on; replaying
|
||||
it on a different node (the rematch-clobbered ``req.last_node`` class
|
||||
of bug) must assert instead of walking that node's segment."""
|
||||
sw = self.cfg.sliding_window_size
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
seq = self._make_seq(1, 2 * sw)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq)
|
||||
leaf = self._match_leaf(cache, seq)
|
||||
parent = _node_parent(cache, leaf)
|
||||
self.assertFalse(cache.tree_core.is_root(parent))
|
||||
lock = cache.inc_lock_ref(leaf)
|
||||
self.assertEqual(lock.node_id, leaf)
|
||||
|
||||
self._assert_protocol_violation(
|
||||
lambda: cache.dec_lock_ref(parent, lock.to_dec_params()),
|
||||
"lock receipt anchored on node",
|
||||
)
|
||||
|
||||
def test_double_release_fails_loud(self):
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
seq = self._make_seq(1, 2 * self.cfg.sliding_window_size)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq)
|
||||
leaf = self._match_leaf(cache, seq)
|
||||
lock = cache.inc_lock_ref(leaf)
|
||||
cache.dec_lock_ref(leaf, lock.to_dec_params())
|
||||
self._assert_protocol_violation(
|
||||
lambda: cache.dec_lock_ref(leaf, lock.to_dec_params()), "lock_ref=0"
|
||||
)
|
||||
|
||||
def test_finish_after_early_release_without_skip_swa_fails_loud(self):
|
||||
"""F2 replay: retraction-after-early-release used to run a second SWA
|
||||
walk that stole ancestors' locks; now it dies at the first node."""
|
||||
sw = self.cfg.sliding_window_size
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
seq = self._make_seq(1, 2 * sw)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq)
|
||||
leaf = self._match_leaf(cache, seq)
|
||||
lock = cache.inc_lock_ref(leaf)
|
||||
cache.dec_swa_lock_only(leaf, lock.to_dec_params())
|
||||
self._assert_protocol_violation(
|
||||
lambda: cache.dec_lock_ref(leaf, lock.to_dec_params()), "lock_ref=0"
|
||||
)
|
||||
|
||||
def test_finish_after_early_release_with_skip_swa(self):
|
||||
"""The correct F2 flow: skip_swa honors the early release."""
|
||||
sw = self.cfg.sliding_window_size
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
seq = self._make_seq(1, 2 * sw)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq)
|
||||
leaf = self._match_leaf(cache, seq)
|
||||
lock = cache.inc_lock_ref(leaf)
|
||||
cache.dec_swa_lock_only(leaf, lock.to_dec_params())
|
||||
cache.dec_lock_ref(leaf, lock.to_dec_params(), skip_swa=True)
|
||||
self.assertEqual(self._swa_ref(cache, leaf), 0)
|
||||
self.assertEqual(_device_lock_ref(cache, leaf, ComponentType.FULL), 0)
|
||||
cache.sanity_check()
|
||||
|
||||
def test_split_under_lock_releases_balanced(self):
|
||||
"""A mid-segment split mints a new node with copied refs and migrates
|
||||
the boundary uuid; the original receipt (its anchor stays on the
|
||||
deeper half) still releases exactly."""
|
||||
sw = self.cfg.sliding_window_size
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
seq = self._make_seq(1, 2 * sw)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq)
|
||||
leaf = self._match_leaf(cache, seq)
|
||||
lock = cache.inc_lock_ref(leaf)
|
||||
pre_segment = self._segment(cache, leaf, sw)
|
||||
|
||||
# Diverge inside the window to force a split of a locked node.
|
||||
fork = seq[: len(seq) - sw // 2] + self._make_seq(9000, sw)
|
||||
self._insert(cache, allocator, req_to_token_pool, fork)
|
||||
post_segment = self._segment(cache, leaf, sw)
|
||||
self.assertGreater(len(post_segment), len(pre_segment))
|
||||
for n in post_segment:
|
||||
self.assertEqual(self._swa_ref(cache, n), 1)
|
||||
|
||||
cache.dec_lock_ref(leaf, lock.to_dec_params())
|
||||
for n in post_segment:
|
||||
self.assertEqual(self._swa_ref(cache, n), 0)
|
||||
cache.sanity_check()
|
||||
|
||||
def test_aux_release_readmits_the_leaf_whatever_the_release_order(self):
|
||||
"""Each component's release refreshes the leaf sets of the nodes it
|
||||
unlocks, so a leaf whose last lock is an auxiliary one is readmitted
|
||||
even when Full released first. Component-level replay of the Python
|
||||
core; the Rust crate covers its own order in its unit tests."""
|
||||
if _selected_tree_core_test_backend() != "python":
|
||||
self.skipTest("drives Python component objects directly")
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
seq = self._make_seq(1, self.cfg.sliding_window_size)
|
||||
self._insert(cache, allocator, req_to_token_pool, seq)
|
||||
leaf = self._match_leaf(cache, seq)
|
||||
node = cache.tree_core.node_by_id(leaf)
|
||||
self.assertIn(node, cache.tree_core.evictable_device_leaves)
|
||||
|
||||
params = cache.inc_lock_ref(leaf).to_dec_params()
|
||||
self.assertNotIn(node, cache.tree_core.evictable_device_leaves)
|
||||
# Full first: its walk still sees the SWA lock, so the leaf stays out.
|
||||
cache.components[ComponentType.FULL].release_component_lock(node, params)
|
||||
self.assertNotIn(node, cache.tree_core.evictable_device_leaves)
|
||||
# The SWA release drops the last lock and must readmit the leaf itself.
|
||||
cache.components[ComponentType.SWA].release_component_lock(node, params)
|
||||
self.assertIn(node, cache.tree_core.evictable_device_leaves)
|
||||
cache.sanity_check()
|
||||
|
||||
|
||||
@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA")
|
||||
class TestSegmentLockFuzz(_InsertWalkSuite):
|
||||
"""Random lock/insert/evict interleavings with the tree's own ledger
|
||||
recomputation (sanity_check) as the per-step oracle, on both cores."""
|
||||
|
||||
cfg = CacheConfig(
|
||||
components=(ComponentType.FULL, ComponentType.SWA),
|
||||
sliding_window_size=8,
|
||||
kv_size=4096,
|
||||
max_num_reqs=256,
|
||||
)
|
||||
|
||||
def _lock_skips(self, rng):
|
||||
"""The decode hold opts the Mamba lock out; exercise both receipts."""
|
||||
if self.cfg.has_mamba and rng.random() < 0.5:
|
||||
return (ComponentType.MAMBA,)
|
||||
return ()
|
||||
|
||||
def _run_seed(self, seed: int, steps: int = 120):
|
||||
import random as _random
|
||||
|
||||
rng = _random.Random(seed)
|
||||
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
chains: list[list[int]] = []
|
||||
held: list[list] = [] # [node_id, receipt, released] entries
|
||||
|
||||
for step in range(steps):
|
||||
op = rng.random()
|
||||
try:
|
||||
if op < 0.35 or not chains:
|
||||
# Insert: fresh chain or extend/diverge an existing one.
|
||||
if chains and rng.random() < 0.6:
|
||||
base = rng.choice(chains)
|
||||
cut = rng.randrange(1, len(base) + 1)
|
||||
seq = base[:cut] + self._make_seq(
|
||||
1000 * (step + 1), rng.randrange(2, 12)
|
||||
)
|
||||
else:
|
||||
seq = self._make_seq(1000 * (step + 1), rng.randrange(4, 20))
|
||||
if allocator.available_size() < len(seq):
|
||||
cache.evict(EvictParams(num_tokens=len(seq) * 2))
|
||||
if allocator.available_size() < len(seq):
|
||||
continue
|
||||
swa_evict = rng.randrange(0, len(seq)) if rng.random() < 0.3 else 0
|
||||
value = self._alloc(allocator, len(seq))
|
||||
# Release the evicted prefix's SWA peers first, as the
|
||||
# scheduler does before inserting a window-trimmed request.
|
||||
allocator.free_swa(value[:swa_evict])
|
||||
params = InsertParams(
|
||||
key=RadixKey(array("q", seq)),
|
||||
value=value,
|
||||
swa_evicted_seqlen=swa_evict,
|
||||
)
|
||||
if self.cfg.has_mamba:
|
||||
req = self._make_req(req_to_token_pool)
|
||||
params.mamba_value = req.kv.mamba_pool_idx.unsqueeze(0)
|
||||
cache.insert(params)
|
||||
chains.append(seq)
|
||||
elif op < 0.6:
|
||||
# Lock a random chain's current deepest device node.
|
||||
seq = rng.choice(chains)
|
||||
m = cache.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq)))
|
||||
)
|
||||
node_id = m.last_device_node
|
||||
if cache.tree_core.is_root(node_id):
|
||||
continue
|
||||
receipt = cache.inc_lock_ref(
|
||||
node_id, skip_lock_components=self._lock_skips(rng)
|
||||
)
|
||||
held.append([node_id, receipt, False])
|
||||
elif op < 0.8 and held:
|
||||
# Full release of a random held lock.
|
||||
idx = rng.randrange(len(held))
|
||||
node_id, receipt, released = held.pop(idx)
|
||||
cache.dec_lock_ref(
|
||||
node_id, receipt.to_dec_params(), skip_swa=released
|
||||
)
|
||||
elif op < 0.9 and held:
|
||||
# Early SWA release of a random not-yet-released lock.
|
||||
idx = rng.randrange(len(held))
|
||||
node_id, receipt, released = held[idx]
|
||||
if released or receipt.swa_uuid_for_lock is None:
|
||||
continue
|
||||
cache.dec_swa_lock_only(
|
||||
node_id,
|
||||
receipt.to_dec_params(),
|
||||
)
|
||||
held[idx][2] = True
|
||||
else:
|
||||
cache.evict(
|
||||
EvictParams(
|
||||
num_tokens=rng.randrange(0, 32),
|
||||
swa_num_tokens=rng.randrange(0, 32),
|
||||
mamba_num=rng.randrange(0, 4) if self.cfg.has_mamba else 0,
|
||||
)
|
||||
)
|
||||
except AssertionError:
|
||||
raise
|
||||
cache.sanity_check()
|
||||
|
||||
# Drain remaining locks; the tree must come back exactly balanced.
|
||||
for node_id, receipt, released in held:
|
||||
cache.dec_lock_ref(node_id, receipt.to_dec_params(), skip_swa=released)
|
||||
cache.sanity_check()
|
||||
|
||||
def test_fuzz_seed0(self):
|
||||
self._run_seed(0)
|
||||
|
||||
def test_fuzz_seed1(self):
|
||||
self._run_seed(1)
|
||||
|
||||
def test_fuzz_seed2(self):
|
||||
self._run_seed(2)
|
||||
|
||||
|
||||
@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA")
|
||||
class TestSegmentLockFuzzWithMamba(TestSegmentLockFuzz):
|
||||
"""The FULL+SWA+MAMBA (Inkling) shape: the Mamba opt-out receipt, the
|
||||
lower-priority cascade on early SWA release, and Mamba evictions all join
|
||||
the interleavings."""
|
||||
|
||||
cfg = CacheConfig(
|
||||
components=(ComponentType.FULL, ComponentType.SWA, ComponentType.MAMBA),
|
||||
sliding_window_size=8,
|
||||
kv_size=4096,
|
||||
max_num_reqs=256,
|
||||
mamba_cache_size=512,
|
||||
)
|
||||
|
||||
|
||||
class TestStreamingSessionLockLifecycle(CustomTestCase):
|
||||
"""A streaming session must persist swa_prefix_lock_released: closing or
|
||||
aborting a session whose first turn early-released its SWA lock must not
|
||||
release the SWA segment a second time."""
|
||||
|
||||
cfg = CacheConfig(
|
||||
page_size=1,
|
||||
components=(ComponentType.FULL, ComponentType.SWA),
|
||||
sliding_window_size=4,
|
||||
kv_size=64,
|
||||
max_context_len=64,
|
||||
)
|
||||
|
||||
def _lock_and_early_release(self, cache, allocator):
|
||||
tokens = array("q", range(1, 9))
|
||||
value = allocator.alloc(len(tokens))
|
||||
cache.insert(InsertParams(key=RadixKey(tokens), value=value))
|
||||
match = cache.match_prefix(MatchPrefixParams(key=RadixKey(tokens)))
|
||||
node = match.last_device_node
|
||||
lock = cache.inc_lock_ref(node)
|
||||
self.assertIsNotNone(lock.swa_uuid_for_lock)
|
||||
cache.dec_swa_lock_only(node, lock.to_dec_params())
|
||||
return node, lock
|
||||
|
||||
def _streaming_req(self, node, lock, *, session):
|
||||
# No KV row is held: the slot only carries the tree lock receipt.
|
||||
kv = ReqKvInfo()
|
||||
return SimpleNamespace(
|
||||
kv=kv,
|
||||
detach_kv=lambda: kv,
|
||||
last_node=node,
|
||||
lock_receipt=lock.to_dec_params(),
|
||||
swa_prefix_lock_released=True,
|
||||
session=session,
|
||||
finished_reason=None,
|
||||
)
|
||||
|
||||
def test_close_after_early_release_releases_swa_once(self):
|
||||
cache, allocator, _ = build_fixture(self.cfg)
|
||||
node, lock = self._lock_and_early_release(cache, allocator)
|
||||
req = self._streaming_req(node, lock, session=None)
|
||||
slot = SessionSlot()
|
||||
cache.session.slots["s"] = slot
|
||||
slot.save_from_req(req, is_first=True)
|
||||
cache.session.release_session("s")
|
||||
cache.sanity_check()
|
||||
|
||||
def test_first_req_mid_abort_after_early_release(self):
|
||||
cache, allocator, pool = build_fixture(self.cfg)
|
||||
node, lock = self._lock_and_early_release(cache, allocator)
|
||||
session = SimpleNamespace(
|
||||
session_id="s2", streaming=True, abort_req=lambda: None
|
||||
)
|
||||
req = self._streaming_req(node, lock, session=session)
|
||||
req.finished_reason = FINISH_ABORT()
|
||||
self.assertTrue(cache.session.try_cache_finished_req(req))
|
||||
cache.sanity_check()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user