[sgl] wire SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW on Unified Cache. (#28161)
This commit is contained in:
@@ -514,6 +514,49 @@ class SWAComponent(TreeComponent):
|
||||
dec_swa = False
|
||||
cur = cur.parent
|
||||
|
||||
def release_window_lock(
|
||||
self,
|
||||
node: UnifiedTreeNode,
|
||||
swa_uuid_for_lock: Optional[int] = None,
|
||||
) -> None:
|
||||
"""Early-release the SWA lock along [node, swa_uuid_for_lock] while
|
||||
leaving Full and Mamba locks intact.
|
||||
|
||||
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.
|
||||
"""
|
||||
ct = self.component_type
|
||||
root = self.cache.root_node
|
||||
|
||||
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
|
||||
|
||||
cd.lock_ref -= 1
|
||||
if cd.lock_ref == 0:
|
||||
key_len = len(cur.key)
|
||||
self.cache.component_protected_size_[ct] -= key_len
|
||||
self.cache.component_evictable_size_[ct] += key_len
|
||||
if self.cache._is_device_leaf(cur):
|
||||
self.cache._evict_component_and_detach_lru(
|
||||
cur, self, target=EvictLayer.DEVICE
|
||||
)
|
||||
|
||||
if swa_uuid_for_lock and cd.metadata.get("uuid") == swa_uuid_for_lock:
|
||||
break
|
||||
cur = cur.parent
|
||||
|
||||
def prepare_for_caching_req(
|
||||
self,
|
||||
req: Req,
|
||||
|
||||
@@ -638,7 +638,10 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
return result
|
||||
|
||||
def dec_lock_ref(
|
||||
self, node: Any, params: Optional[DecLockRefParams] = None
|
||||
self,
|
||||
node: Any,
|
||||
params: Optional[DecLockRefParams] = None,
|
||||
skip_swa: bool = False,
|
||||
) -> DecLockRefResult:
|
||||
result = self.session.try_dec_lock_ref(node, params)
|
||||
if result is not None:
|
||||
@@ -646,12 +649,36 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
if self.disable:
|
||||
return DecLockRefResult()
|
||||
for component in self._components_tuple:
|
||||
if skip_swa and component.component_type == ComponentType.SWA:
|
||||
continue
|
||||
component.release_component_lock(node=node, params=params)
|
||||
|
||||
self._update_evictable_leaf_sets(node)
|
||||
# TODO: delta is not aggregated from components; no caller uses it yet.
|
||||
return DecLockRefResult()
|
||||
|
||||
def dec_swa_lock_only(
|
||||
self,
|
||||
node: UnifiedTreeNode,
|
||||
swa_uuid_for_lock: Optional[int] = None,
|
||||
) -> None:
|
||||
"""Early-release the SWA portion of a request's tree lock, plus any
|
||||
strictly-lower-priority locks (e.g. Mamba) co-located on `node`.
|
||||
"""
|
||||
if self.disable:
|
||||
return
|
||||
swa_component = self.components.get(ComponentType.SWA)
|
||||
if swa_component is None:
|
||||
return
|
||||
swa_component.release_window_lock(node, swa_uuid_for_lock)
|
||||
|
||||
# Drop strictly-lower-priority locks (e.g. Mamba) co-located on `node`.
|
||||
swa_priority = swa_component.eviction_priority(is_leaf=False)
|
||||
dec_params = DecLockRefParams(swa_uuid_for_lock=swa_uuid_for_lock)
|
||||
for comp in self._components_tuple:
|
||||
if comp.eviction_priority(is_leaf=False) < swa_priority:
|
||||
comp.release_component_lock(node, dec_params)
|
||||
|
||||
def inc_host_lock_ref(self, node: Any) -> IncLockRefResult:
|
||||
if self.disable:
|
||||
return IncLockRefResult()
|
||||
@@ -741,6 +768,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
self.dec_lock_ref(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)),
|
||||
skip_swa=getattr(req, "swa_prefix_lock_released", False),
|
||||
)
|
||||
|
||||
# cleanup
|
||||
@@ -1248,6 +1276,18 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
if comp.eviction_priority(is_leaf) <= trigger_priority:
|
||||
if comp is not trigger and comp.node_has_component_data(node, target):
|
||||
cd = node.component_data[comp.component_type]
|
||||
# A comp whose TRUE internal priority outranks the trigger
|
||||
# is only in this loop because leaf-collapse flattened
|
||||
# priorities; a lock on it is a legit pin and must be
|
||||
# spared. A lock on a strictly-lower-priority tier is a
|
||||
# real strand — fall through to the assert below.
|
||||
if comp.eviction_priority(
|
||||
is_leaf=False
|
||||
) >= trigger.eviction_priority(is_leaf=False):
|
||||
if EvictLayer.DEVICE in target and cd.lock_ref != 0:
|
||||
continue
|
||||
if EvictLayer.HOST in target and cd.host_lock_ref != 0:
|
||||
continue
|
||||
if EvictLayer.DEVICE in target:
|
||||
assert cd.lock_ref == 0
|
||||
if EvictLayer.HOST in target:
|
||||
|
||||
@@ -27,6 +27,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
DecLockRefParams,
|
||||
EvictParams,
|
||||
EvictResult,
|
||||
IncLockRefResult,
|
||||
InitLoadBackParams,
|
||||
InsertParams,
|
||||
MatchPrefixParams,
|
||||
@@ -1261,6 +1262,265 @@ class UnifiedRadixCacheSuite:
|
||||
self.assertGreaterEqual(result.swa_num_tokens_evicted, 0)
|
||||
tree.sanity_check()
|
||||
|
||||
def test_leaf_transition_swa_evict_spares_locked_full(self):
|
||||
if not self.cfg.has_swa or not self.cfg.has_mamba:
|
||||
self.skipTest("requires SWA and Mamba components")
|
||||
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
|
||||
n_short = (self.cfg.sliding_window_size // self.cfg.page_size) + 4
|
||||
seq_a = self._make_seq(1, n_short)
|
||||
seq_ab = seq_a + self._make_seq(7000, 2)
|
||||
self._insert(tree, allocator, req_to_token_pool, seq_a)
|
||||
self._insert(tree, allocator, req_to_token_pool, seq_ab)
|
||||
|
||||
m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_a))))
|
||||
node_a = m.last_device_node
|
||||
self.assertGreater(len(node_a.children), 0, "A must be internal")
|
||||
|
||||
swa_cd = node_a.component_data[ComponentType.SWA]
|
||||
mamba_cd = node_a.component_data[ComponentType.MAMBA]
|
||||
full_cd = node_a.component_data[ComponentType.FULL]
|
||||
|
||||
# A request locks A, then decodes past the window → early-release the SWA
|
||||
# portion. On this internal node, dec_swa_lock_only also drops the
|
||||
# strictly-lower-tier Mamba lock (the co-located Mamba is useless once SWA
|
||||
# is gone), leaving only the Full path-lock held. This is what guarantees
|
||||
# the later SWA-eviction cascade never meets a legitimately-locked Mamba.
|
||||
lock_result = tree.inc_lock_ref(node_a)
|
||||
self.assertGreaterEqual(mamba_cd.lock_ref, 1, "Mamba locked before release")
|
||||
tree.dec_swa_lock_only(node_a, lock_result.swa_uuid_for_lock)
|
||||
self.assertEqual(swa_cd.lock_ref, 0)
|
||||
self.assertEqual(
|
||||
mamba_cd.lock_ref, 0, "dec_swa_lock_only drops the lower-tier Mamba lock"
|
||||
)
|
||||
self.assertGreaterEqual(full_cd.lock_ref, 1)
|
||||
self.assertTrue(tree.lru_lists[ComponentType.SWA].in_list(node_a))
|
||||
|
||||
# Evict the child branch (Full/device eviction only) → A becomes a
|
||||
# Full-locked leaf with its now-unlocked SWA still in the SWA LRU. We do
|
||||
# NOT tombstone aux at the leaf-transition; the held Full pins the node.
|
||||
tree.evict(EvictParams(num_tokens=len(seq_ab)))
|
||||
|
||||
self.assertEqual(len(node_a.children), 0, "A should now be a leaf")
|
||||
self.assertGreaterEqual(full_cd.lock_ref, 1, "Full must stay locked")
|
||||
self.assertTrue(
|
||||
tree.lru_lists[ComponentType.SWA].in_list(node_a),
|
||||
"A's unlocked SWA stays in the LRU (not tombstoned at transition)",
|
||||
)
|
||||
|
||||
# SWA eviction now selects A (it is in the SWA LRU). A's Full lock is a
|
||||
# higher-or-equal internal tier, so the cascade skips it and spares the
|
||||
# Full KV. The unlocked lower-tier Mamba is cascaded as part of the atomic
|
||||
# leaf teardown. This used to assert `cd.lock_ref == 0` on the locked Full.
|
||||
tree.evict(EvictParams(num_tokens=0, swa_num_tokens=len(seq_a)))
|
||||
self.assertGreaterEqual(full_cd.lock_ref, 1, "Full must remain locked")
|
||||
self.assertIsNotNone(full_cd.value, "Full KV must survive the SWA cascade")
|
||||
self.assertIsNone(swa_cd.value, "A's SWA was freed by its own eviction")
|
||||
tree.sanity_check()
|
||||
|
||||
tree.dec_lock_ref(node_a, DecLockRefParams(swa_uuid_for_lock=None))
|
||||
tree.sanity_check()
|
||||
|
||||
def test_swa_early_release_drops_co_located_mamba_lock(self):
|
||||
if not self.cfg.has_swa or not self.cfg.has_mamba:
|
||||
self.skipTest("requires SWA and Mamba components")
|
||||
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
|
||||
n_short = (self.cfg.sliding_window_size // self.cfg.page_size) + 4
|
||||
seq_a = self._make_seq(1, n_short)
|
||||
self._insert(tree, allocator, req_to_token_pool, seq_a)
|
||||
node_a = tree.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq_a)))
|
||||
).last_device_node
|
||||
self.assertEqual(len(node_a.children), 0, "A must be a leaf")
|
||||
|
||||
swa_cd = node_a.component_data[ComponentType.SWA]
|
||||
mamba_cd = node_a.component_data[ComponentType.MAMBA]
|
||||
full_cd = node_a.component_data[ComponentType.FULL]
|
||||
self.assertIsNotNone(mamba_cd.value, "A must hold a Mamba checkpoint")
|
||||
|
||||
# Natural lock acquisition — records inc_lock_ref in the lock trace.
|
||||
lock_result = tree.inc_lock_ref(node_a)
|
||||
self.assertGreaterEqual(swa_cd.lock_ref, 1, "SWA locked")
|
||||
self.assertGreaterEqual(mamba_cd.lock_ref, 1, "Mamba locked")
|
||||
self.assertGreaterEqual(full_cd.lock_ref, 1, "Full locked")
|
||||
|
||||
# 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.
|
||||
tree.dec_swa_lock_only(node_a, lock_result.swa_uuid_for_lock)
|
||||
self.assertEqual(swa_cd.lock_ref, 0, "SWA early-released")
|
||||
self.assertEqual(
|
||||
mamba_cd.lock_ref,
|
||||
0,
|
||||
"Mamba lock must drop on early SWA release",
|
||||
)
|
||||
self.assertGreaterEqual(full_cd.lock_ref, 1, "Full stays locked")
|
||||
|
||||
def test_cascade_evict_asserts_on_locked_internal_mamba(self):
|
||||
if not self.cfg.has_swa or not self.cfg.has_mamba:
|
||||
self.skipTest("requires SWA and Mamba components")
|
||||
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
|
||||
n_short = (self.cfg.sliding_window_size // self.cfg.page_size) + 4
|
||||
seq_a = self._make_seq(1, n_short)
|
||||
seq_ab = seq_a + self._make_seq(7000, 2)
|
||||
self._insert(tree, allocator, req_to_token_pool, seq_a)
|
||||
self._insert(tree, allocator, req_to_token_pool, seq_ab)
|
||||
|
||||
node_a = tree.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq_a)))
|
||||
).last_device_node
|
||||
self.assertGreater(len(node_a.children), 0, "A must be internal")
|
||||
|
||||
mamba_cd = node_a.component_data[ComponentType.MAMBA]
|
||||
full_cd = node_a.component_data[ComponentType.FULL]
|
||||
self.assertIsNotNone(mamba_cd.value, "A must hold a Mamba checkpoint")
|
||||
|
||||
# Lock ONLY Mamba — a stranded lower-priority lock that no supported path
|
||||
# produces. The cascade must surface it rather than silently skip.
|
||||
tree.components[ComponentType.MAMBA].acquire_component_lock(
|
||||
node_a, IncLockRefResult()
|
||||
)
|
||||
self.assertGreaterEqual(mamba_cd.lock_ref, 1, "Mamba locked")
|
||||
self.assertEqual(full_cd.lock_ref, 0, "Full unlocked")
|
||||
|
||||
tracker = {ct: 0 for ct in tree.tree_components}
|
||||
tree._evict_component_and_detach_lru(
|
||||
node_a,
|
||||
tree.components[ComponentType.SWA],
|
||||
target=EvictLayer.DEVICE,
|
||||
tracker=tracker,
|
||||
)
|
||||
# No higher-or-equal tier pins the node, so even with early-release on
|
||||
# the stranded Mamba lock must trip the hard-invariant assert.
|
||||
with self.assertRaises(AssertionError):
|
||||
tree._cascade_evict(node_a, tree.components[ComponentType.SWA], tracker)
|
||||
|
||||
# Clean up the forced lock so teardown/sanity is consistent.
|
||||
tree.components[ComponentType.MAMBA].release_component_lock(
|
||||
node_a, DecLockRefParams(swa_uuid_for_lock=None)
|
||||
)
|
||||
|
||||
def test_cascade_evict_asserts_on_locked_leaf_mamba(self):
|
||||
if not self.cfg.has_swa or not self.cfg.has_mamba:
|
||||
self.skipTest("requires SWA and Mamba components")
|
||||
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
|
||||
n_short = (self.cfg.sliding_window_size // self.cfg.page_size) + 4
|
||||
seq_a = self._make_seq(1, n_short)
|
||||
self._insert(tree, allocator, req_to_token_pool, seq_a)
|
||||
|
||||
node_a = tree.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq_a)))
|
||||
).last_device_node
|
||||
self.assertEqual(len(node_a.children), 0, "A must be a leaf")
|
||||
|
||||
mamba_cd = node_a.component_data[ComponentType.MAMBA]
|
||||
full_cd = node_a.component_data[ComponentType.FULL]
|
||||
self.assertIsNotNone(mamba_cd.value, "A must hold a Mamba checkpoint")
|
||||
|
||||
# Lock ONLY Mamba (Full stays unlocked) — a stranded lower-tier lock.
|
||||
tree.components[ComponentType.MAMBA].acquire_component_lock(
|
||||
node_a, IncLockRefResult()
|
||||
)
|
||||
self.assertGreaterEqual(mamba_cd.lock_ref, 1, "Mamba locked")
|
||||
self.assertEqual(full_cd.lock_ref, 0, "Full unlocked")
|
||||
|
||||
tracker = {ct: 0 for ct in tree.tree_components}
|
||||
tree._evict_component_and_detach_lru(
|
||||
node_a,
|
||||
tree.components[ComponentType.SWA],
|
||||
target=EvictLayer.DEVICE,
|
||||
tracker=tracker,
|
||||
)
|
||||
# No higher-or-equal tier pins the node, so even with early-release on
|
||||
# the stranded Mamba lock must trip the hard-invariant assert.
|
||||
with self.assertRaises(AssertionError):
|
||||
tree._cascade_evict(node_a, tree.components[ComponentType.SWA], tracker)
|
||||
|
||||
# Clean up the forced lock so teardown/sanity is consistent.
|
||||
tree.components[ComponentType.MAMBA].release_component_lock(
|
||||
node_a, DecLockRefParams(swa_uuid_for_lock=None)
|
||||
)
|
||||
|
||||
def test_dec_swa_lock_only_hicache_child_on_host_treated_as_device_leaf(self):
|
||||
if not self.cfg.has_swa or not self.cfg.has_mamba:
|
||||
self.skipTest("requires SWA and Mamba components")
|
||||
|
||||
tree, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||
|
||||
n_short = (self.cfg.sliding_window_size // self.cfg.page_size) + 4
|
||||
seq_a = self._make_seq(1, n_short)
|
||||
seq_ab = seq_a + self._make_seq(7000, 2)
|
||||
self._insert(tree, allocator, req_to_token_pool, seq_a)
|
||||
self._insert(tree, allocator, req_to_token_pool, seq_ab)
|
||||
|
||||
node_a = tree.match_prefix(
|
||||
MatchPrefixParams(key=RadixKey(array("q", seq_a)))
|
||||
).last_device_node
|
||||
self.assertGreater(len(node_a.children), 0, "A must have children")
|
||||
|
||||
self.assertFalse(
|
||||
tree._is_device_leaf(node_a),
|
||||
"A is not a device-leaf while child holds Full on device",
|
||||
)
|
||||
|
||||
self._simulate_backup(tree, node_a)
|
||||
self.assertTrue(node_a.backuped, "node_a must be backuped (invariant a)")
|
||||
|
||||
def _collect_descendants(node):
|
||||
out = []
|
||||
for c in list(node.children.values()):
|
||||
out.extend(_collect_descendants(c))
|
||||
out.append(c)
|
||||
return out
|
||||
|
||||
descendants = _collect_descendants(node_a)
|
||||
self.assertGreater(len(descendants), 0)
|
||||
for desc in descendants:
|
||||
self._simulate_backup(tree, desc)
|
||||
self.assertTrue(desc.backuped, "desc must be backuped before demote")
|
||||
tracker = {ct: 0 for ct in tree.tree_components}
|
||||
tree._evict_to_host(desc, tracker)
|
||||
self.assertTrue(desc.evicted, "desc should be D->H demoted")
|
||||
self.assertIsNone(desc.component_data[ComponentType.FULL].value)
|
||||
|
||||
self.assertTrue(
|
||||
tree._is_device_leaf(node_a),
|
||||
"A is a HiCache device-leaf (no child with Full on device)",
|
||||
)
|
||||
self.assertGreater(len(node_a.children), 0, "A still has tree-children")
|
||||
self.assertIn(node_a, tree.evictable_device_leaves)
|
||||
tree.sanity_check()
|
||||
|
||||
lock_result = tree.inc_lock_ref(node_a)
|
||||
swa_cd = node_a.component_data[ComponentType.SWA]
|
||||
mamba_cd = node_a.component_data[ComponentType.MAMBA]
|
||||
full_cd = node_a.component_data[ComponentType.FULL]
|
||||
self.assertGreaterEqual(swa_cd.lock_ref, 1)
|
||||
self.assertGreaterEqual(mamba_cd.lock_ref, 1)
|
||||
self.assertGreaterEqual(full_cd.lock_ref, 1)
|
||||
|
||||
tree.dec_swa_lock_only(node_a, lock_result.swa_uuid_for_lock)
|
||||
self.assertEqual(swa_cd.lock_ref, 0, "SWA released")
|
||||
self.assertEqual(mamba_cd.lock_ref, 0, "Mamba dropped by dec_swa_lock_only")
|
||||
self.assertGreaterEqual(full_cd.lock_ref, 1, "Full kept by contract")
|
||||
self.assertIsNotNone(
|
||||
swa_cd.value,
|
||||
"SWA slot stays under contract (lazy reclaim by drive_eviction)",
|
||||
)
|
||||
self.assertTrue(
|
||||
tree.lru_lists[ComponentType.SWA].in_list(node_a),
|
||||
"SWA stays in LRU for drive_eviction to pick later",
|
||||
)
|
||||
|
||||
tree.dec_lock_ref(
|
||||
node_a, DecLockRefParams(swa_uuid_for_lock=None), skip_swa=True
|
||||
)
|
||||
self.assertTrue(tree._is_device_leaf(node_a))
|
||||
tree.sanity_check()
|
||||
|
||||
def test_swa_evict_full_leaf_cascades_all(self):
|
||||
if not self.cfg.has_swa:
|
||||
self.skipTest("requires SWA component")
|
||||
|
||||
Reference in New Issue
Block a user