[HiCache] Replace skip_lock_node_ids with a segment lock protocol (#36848)
This commit is contained in:
@@ -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