diff --git a/python/sglang/srt/mem_cache/unified_cache_components/swa_component.py b/python/sglang/srt/mem_cache/unified_cache_components/swa_component.py index 54303981c..30a783cca 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/swa_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/swa_component.py @@ -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, diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 4df8109be..9514172d8 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -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: diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 2ec8a57b0..6cbb56e5e 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -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")