From 6965fe0eecb832749d1cf324e66a4b01fdc16065 Mon Sep 17 00:00:00 2001 From: Bi Xue Date: Mon, 1 Jun 2026 04:35:18 -0700 Subject: [PATCH] [sgl] Window-aware LRU refresh for SWA prefix cache in unified cache (#26615) --- .../unified_cache_components/__init__.py | 2 + .../unified_cache_components/swa_component.py | 26 ++ .../tree_component.py | 30 ++ .../srt/mem_cache/unified_radix_cache.py | 37 ++- .../test_unified_radix_cache_unittest.py | 297 ++++++++++++++++++ 5 files changed, 388 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/mem_cache/unified_cache_components/__init__.py b/python/sglang/srt/mem_cache/unified_cache_components/__init__.py index 2d848b3ec..144b27e51 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/__init__.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/__init__.py @@ -8,6 +8,7 @@ from sglang.srt.mem_cache.unified_cache_components.tree_component import ( ComponentData, ComponentType, EvictLayer, + LRURefreshPhase, TreeComponent, get_and_increase_time_counter, next_component_uuid, @@ -20,6 +21,7 @@ __all__ = [ "EvictLayer", "FullComponent", "CacheTransferPhase", + "LRURefreshPhase", "MambaComponent", "SWAComponent", "TreeComponent", 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 0ba006079..415be271f 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 @@ -19,6 +19,7 @@ from sglang.srt.mem_cache.unified_cache_components.tree_component import ( CacheTransferPhase, ComponentType, EvictLayer, + LRURefreshPhase, TreeComponent, next_component_uuid, ) @@ -60,6 +61,31 @@ class SWAComponent(TreeComponent): full_indices ) + def refresh_lru( + self, + phase: LRURefreshPhase, + node: UnifiedTreeNode, + root_node: UnifiedTreeNode, + ) -> None: + match phase: + case LRURefreshPhase.WALKDOWN: + # Walk-down would refresh every visited ancestor to MRU, + # but most are outside the active sliding window and must + # stay evictable. Window-bounded refresh runs at + # MATCH_END / INSERT_END instead. + return + case LRURefreshPhase.MATCH_END | LRURefreshPhase.INSERT_END: + self.cache.lru_lists[ + self.component_type + ].reset_node_and_window_ancestors_mru( + node, + root_node, + self.sliding_window_size + self.cache.page_size, + self.node_has_component_data, + ) + case _: + raise ValueError(f"Unknown LRURefreshPhase: {phase}") + def _restore_device_value(self, node: UnifiedTreeNode, value: torch.Tensor) -> None: ct = self.component_type node.component_data[ct].value = value diff --git a/python/sglang/srt/mem_cache/unified_cache_components/tree_component.py b/python/sglang/srt/mem_cache/unified_cache_components/tree_component.py index 2b9b03b88..91be6d4f6 100644 --- a/python/sglang/srt/mem_cache/unified_cache_components/tree_component.py +++ b/python/sglang/srt/mem_cache/unified_cache_components/tree_component.py @@ -83,6 +83,13 @@ class CacheTransferPhase(str, Enum): PREFETCH = "prefetch" # Storage→H +class LRURefreshPhase(str, Enum): + + WALKDOWN = "walkdown" # touching a node while walking through the tree + MATCH_END = "match_end" # end of a successful prefix match + INSERT_END = "insert_end" # after a new/updated leaf is committed + + def get_and_increase_time_counter() -> float64: global _LAST_ACCESS_TIME_COUNTER_FLOAT ret = _LAST_ACCESS_TIME_COUNTER_FLOAT @@ -115,6 +122,29 @@ class TreeComponent(ABC): value = node.component_data[self.component_type].value return len(value) if value is not None else 0 + def refresh_lru( + self, + phase: LRURefreshPhase, + node: UnifiedTreeNode, + root_node: UnifiedTreeNode, + ) -> None: + ct = self.component_type + match phase: + case LRURefreshPhase.WALKDOWN: + if node.component_data[ct].value is None: + return + self.cache.lru_lists[ct].reset_node_mru(node) + case LRURefreshPhase.MATCH_END: + self.cache.lru_lists[ct].reset_node_and_parents_mru( + node, root_node, self.node_has_component_data + ) + case LRURefreshPhase.INSERT_END: + # WALKDOWN already refreshed every node on the insert path + # (including the new leaf), so there is nothing more to do. + return + case _: + raise ValueError(f"Unknown LRURefreshPhase: {phase}") + @abstractmethod def create_match_validator( self, match_device_only: bool = False diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 93b6924fb..1bde47c5d 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -44,6 +44,7 @@ from sglang.srt.mem_cache.unified_cache_components import ( ComponentType, EvictLayer, FullComponent, + LRURefreshPhase, MambaComponent, SWAComponent, TreeComponent, @@ -186,6 +187,24 @@ class UnifiedLRUList: prev_node = node node = node.parent + def reset_node_and_window_ancestors_mru( + self, + node: UnifiedTreeNode, + root_node: UnifiedTreeNode, + window_size: int, + should_include, + ): + prev_node = self.head + accumulated = 0 + while node != root_node and accumulated < window_size: + if should_include(node): + assert node.id in self.cache + self._remove_node(node) + self._add_node_after(prev_node, node) + prev_node = node + accumulated += len(node.key) + node = node.parent + def in_list(self, node: Optional[UnifiedTreeNode]): return node is not None and node.id in self.cache @@ -800,9 +819,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): for comp in self._components_tuple: if comp.component_type == BASE_COMPONENT_TYPE: continue # Full uses last_access_time, not LRU - self.lru_lists[comp.component_type].reset_node_and_parents_mru( - node_update, self.root_node, comp.node_has_component_data - ) + comp.refresh_lru(LRURefreshPhase.MATCH_END, node_update, self.root_node) cur_time = get_and_increase_time_counter() while node_update: @@ -878,7 +895,10 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): def _touch_node(self, node: UnifiedTreeNode): node.last_access_time = get_and_increase_time_counter() if node != self.root_node: - self._for_each_component_lru(node, UnifiedLRUList.reset_node_mru) + for comp in self._components_tuple: + if comp.component_type == BASE_COMPONENT_TYPE: + continue + comp.refresh_lru(LRURefreshPhase.WALKDOWN, node, self.root_node) def _add_new_node( self, @@ -1014,6 +1034,15 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): params=params, result=result, ) + + if target_node is not self.root_node: + for component in self._components_tuple: + if component.component_type == BASE_COMPONENT_TYPE: + continue + component.refresh_lru( + LRURefreshPhase.INSERT_END, target_node, self.root_node + ) + if is_new_leaf: self._inc_hit_count(target_node, params.chunked) return result 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 66167eb96..9eb4cdc6c 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 @@ -46,6 +46,7 @@ from sglang.srt.mem_cache.unified_cache_components.tree_component import ( ) from sglang.srt.mem_cache.unified_radix_cache import ( COMPONENT_REGISTRY, + UnifiedLRUList, UnifiedRadixCache, UnifiedTreeNode, ) @@ -1113,6 +1114,200 @@ class UnifiedRadixCacheSuite: ) tree.sanity_check() + def _swa_lru_order(self, tree): + lru = tree.lru_lists[ComponentType.SWA] + pt = lru._pt + nodes: list = [] + cur = lru.head.lru_next[pt] + while cur is not lru.tail: + nodes.append(cur) + cur = cur.lru_next[pt] + return nodes + + def _swa_pinning_cfg_supported(self) -> bool: + if not self.cfg.has_swa or self.cfg.has_mamba: + return False + cushion = self.cfg.sliding_window_size + self.cfg.page_size + pages_per_node = 8 + side_pages = 5 + if pages_per_node * self.cfg.page_size < cushion: + return False + chain_inserts_pages = 4 * pages_per_node + side_pages + if self.cfg.kv_size < chain_inserts_pages * self.cfg.page_size: + return False + return True + + def test_swa_lru_walk_down_does_not_refresh_ancestors_during_insert(self): + if not self._swa_pinning_cfg_supported(): + self.skipTest("requires SWA-only config with node size >= cushion") + tree, allocator, req_to_token_pool = build_fixture(self.cfg) + + seq_a = self._make_seq(1, 8) + seq_ab = seq_a + self._make_seq(100, 8) + seq_abc = seq_ab + self._make_seq(200, 8) + self._insert(tree, allocator, req_to_token_pool, seq_a) + self._insert(tree, allocator, req_to_token_pool, seq_ab) + self._insert(tree, allocator, req_to_token_pool, seq_abc) + + seq_side = self._make_seq(900, 5) + self._insert(tree, allocator, req_to_token_pool, seq_side) + + pre = self._swa_lru_order(tree) + self.assertEqual(len(pre), 4) + side_node, c_node, b_node, a_node = pre + + seq_abcd = seq_abc + self._make_seq(300, 8) + self._insert(tree, allocator, req_to_token_pool, seq_abcd) + + post = self._swa_lru_order(tree) + # new leaf E exists now, length 5 + self.assertEqual(len(post), 5) + # side branch must still appear BEFORE B and A in MRU->LRU order: + # bounded refresh on new leaf E (size=8 >= cushion=5) refreshes only E. + side_pos = post.index(side_node) + b_pos = post.index(b_node) + a_pos = post.index(a_node) + self.assertLess( + side_pos, + b_pos, + f"side branch must remain ahead of B (no walk-down refresh); " + f"post={[n.id for n in post]}, side={side_node.id}, B={b_node.id}", + ) + self.assertLess( + side_pos, + a_pos, + f"side branch must remain ahead of A (no walk-down refresh); " + f"post={[n.id for n in post]}, side={side_node.id}, A={a_node.id}", + ) + tree.sanity_check() + + def test_swa_lru_match_only_refreshes_window_cushion(self): + if not self._swa_pinning_cfg_supported(): + self.skipTest("requires SWA-only config with node size >= cushion") + tree, allocator, req_to_token_pool = build_fixture(self.cfg) + + seq_a = self._make_seq(1, 8) + seq_ab = seq_a + self._make_seq(100, 8) + seq_abc = seq_ab + self._make_seq(200, 8) + self._insert(tree, allocator, req_to_token_pool, seq_a) + self._insert(tree, allocator, req_to_token_pool, seq_ab) + self._insert(tree, allocator, req_to_token_pool, seq_abc) + + seq_side = self._make_seq(900, 5) + self._insert(tree, allocator, req_to_token_pool, seq_side) + + pre = self._swa_lru_order(tree) + self.assertEqual(len(pre), 4) + side_node, c_node, b_node, a_node = pre + + m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_abc)))) + self.assertEqual(len(m.device_indices), len(seq_abc)) + + post = self._swa_lru_order(tree) + self.assertIs(post[0], c_node, "C (last matched node) must be MRU") + self.assertIs( + post[-1], + a_node, + "Oldest out-of-cushion ancestor A must remain at LRU tail; " + f"got post={[n.id for n in post]}, " + f"pre={[n.id for n in pre]}", + ) + self.assertIn( + side_node, + post[:2], + "Side branch must NOT be pushed below ancestors after deep match; " + f"got post={[n.id for n in post]}", + ) + tree.sanity_check() + + def test_swa_lru_old_ancestors_evict_first_under_pressure(self): + if not self._swa_pinning_cfg_supported(): + self.skipTest("requires SWA-only config with node size >= cushion") + tree, allocator, req_to_token_pool = build_fixture(self.cfg) + + seq_a = self._make_seq(1, 8) + seq_ab = seq_a + self._make_seq(100, 8) + seq_abc = seq_ab + self._make_seq(200, 8) + self._insert(tree, allocator, req_to_token_pool, seq_a) + self._insert(tree, allocator, req_to_token_pool, seq_ab) + self._insert(tree, allocator, req_to_token_pool, seq_abc) + + seq_side = self._make_seq(900, 5) + self._insert(tree, allocator, req_to_token_pool, seq_side) + + m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_abc)))) + self.assertEqual(len(m.device_indices), len(seq_abc)) + + m_side_before = tree.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", seq_side))) + ) + self.assertEqual(len(m_side_before.device_indices), len(seq_side)) + + tree.evict(EvictParams(num_tokens=0, swa_num_tokens=self.cfg.page_size)) + + m_side_after = tree.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", seq_side))) + ) + self.assertEqual( + len(m_side_after.device_indices), + len(seq_side), + "Side branch SWA must survive eviction; oldest ancestors (A) " + "should be evicted first under bounded SWA LRU refresh.", + ) + tree.sanity_check() + + def test_swa_lru_cushion_bound_is_sliding_window_plus_page_size(self): + if not self._swa_pinning_cfg_supported(): + self.skipTest("requires SWA-only config with node size >= cushion") + tree, allocator, req_to_token_pool = build_fixture(self.cfg) + + seq_a = self._make_seq(1, 8) + seq_ab = seq_a + self._make_seq(100, 8) + seq_abc = seq_ab + self._make_seq(200, 8) + self._insert(tree, allocator, req_to_token_pool, seq_a) + self._insert(tree, allocator, req_to_token_pool, seq_ab) + self._insert(tree, allocator, req_to_token_pool, seq_abc) + + seq_side = self._make_seq(900, 5) + self._insert(tree, allocator, req_to_token_pool, seq_side) + + pre = self._swa_lru_order(tree) + self.assertEqual(len(pre), 4) + side_node, c_node, b_node, a_node = pre + + m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_abc)))) + self.assertEqual(len(m.device_indices), len(seq_abc)) + post = self._swa_lru_order(tree) + + cushion = self.cfg.sliding_window_size + self.cfg.page_size + self.assertGreaterEqual(len(c_node.key), cushion) + self.assertIs(post[0], c_node, "C alone exhausts cushion → only C refreshed") + # B and A: untouched ordering relative to each other AND to side_node + b_pos = post.index(b_node) + a_pos = post.index(a_node) + side_pos = post.index(side_node) + self.assertLess(side_pos, b_pos, "B was below side in pre, must stay below") + self.assertLess(b_pos, a_pos, "A was below B in pre, must stay below") + tree.sanity_check() + + def test_swa_sanity_check_passes_after_deep_match(self): + if not self._swa_pinning_cfg_supported(): + self.skipTest("requires SWA-only config with node size >= cushion") + tree, allocator, req_to_token_pool = build_fixture(self.cfg) + + seq_a = self._make_seq(1, 8) + seq_ab = seq_a + self._make_seq(100, 8) + seq_abc = seq_ab + self._make_seq(200, 8) + self._insert(tree, allocator, req_to_token_pool, seq_a) + self._insert(tree, allocator, req_to_token_pool, seq_ab) + self._insert(tree, allocator, req_to_token_pool, seq_abc) + self._insert(tree, allocator, req_to_token_pool, self._make_seq(900, 5)) + + for _ in range(3): + m = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", seq_abc)))) + self.assertEqual(len(m.device_indices), len(seq_abc)) + tree.sanity_check() + def test_tombstone_cleanup_respects_locked_parent(self): tree, _, _ = build_fixture(self.cfg) parent = UnifiedTreeNode(self.cfg.components) @@ -2960,6 +3155,108 @@ class UnifiedRadixCacheSuite: tree.sanity_check() +class UnifiedLRUListBoundedRefreshTest(CustomTestCase): + + components = (ComponentType.FULL, ComponentType.SWA) + + def _make_node(self, key_len: int) -> UnifiedTreeNode: + n = UnifiedTreeNode(self.components) + n.key = RadixKey(list(range(key_len))) + return n + + def _build_chain(self, key_lens: list[int]) -> tuple: + root = self._make_node(0) + nodes = [] + parent = root + for kl in key_lens: + n = self._make_node(kl) + n.parent = parent + nodes.append(n) + parent = n + return root, nodes + + def _lru_order(self, lru: UnifiedLRUList) -> list: + pt = lru._pt + out = [] + cur = lru.head.lru_next[pt] + while cur is not lru.tail: + out.append(cur) + cur = cur.lru_next[pt] + return out + + def test_bounded_refresh_stops_after_accumulated_meets_window(self): + root, [a, b, c, d] = self._build_chain([2, 2, 2, 2]) + lru = UnifiedLRUList(ComponentType.SWA, self.components) + for n in (a, b, c, d): + lru.insert_mru(n) + self.assertEqual(self._lru_order(lru), [d, c, b, a]) + + # window=5, page_size=1 implicit; nodes are size 2 each + # Walking up from D: visit D(acc=2<5) -> visit C(acc=4<5) -> visit + # B(acc=6>=5, refresh and stop). A is NOT touched. + lru.reset_node_and_window_ancestors_mru( + d, root, window_size=5, should_include=lambda _n: True + ) + # Expected MRU->LRU: D, C, B (refreshed in walk-up order), A (untouched) + self.assertEqual(self._lru_order(lru), [d, c, b, a]) + + def test_bounded_refresh_skips_non_included(self): + root, [a, b, c, d] = self._build_chain([2, 2, 2, 2]) + lru = UnifiedLRUList(ComponentType.SWA, self.components) + for n in (a, c, d): # b excluded from LRU (simulated tombstone) + lru.insert_mru(n) + self.assertEqual(self._lru_order(lru), [d, c, a]) + + included = {a, c, d} + lru.reset_node_and_window_ancestors_mru( + d, root, window_size=5, should_include=lambda n: n in included + ) + # D, C refreshed; B contributes 2 to acc (now 6 >= 5) but is skipped; + # A is not visited because the walk stops at B. + self.assertEqual(self._lru_order(lru), [d, c, a]) + + def test_bounded_refresh_visits_only_until_window_filled(self): + root, [a, b, c, d] = self._build_chain([3, 3, 3, 3]) + lru = UnifiedLRUList(ComponentType.SWA, self.components) + # MRU->LRU: A, B, C, D (deepest is at the LRU tail, oldest) + for n in (d, c, b, a): + lru.insert_mru(n) + self.assertEqual(self._lru_order(lru), [a, b, c, d]) + + # window=5: walking up from D visits D(acc=3<5) and C(acc=6>=5, stop). + # Order expected: [D, C, A, B]. Why: D and C move to head in that + # order (D first, then C right after D). A and B keep their relative + # positions (they were the surviving prefix [A, B] before). + lru.reset_node_and_window_ancestors_mru( + d, root, window_size=5, should_include=lambda _n: True + ) + self.assertEqual(self._lru_order(lru), [d, c, a, b]) + + def test_bounded_refresh_stops_at_root(self): + root, [a, b] = self._build_chain([1, 1]) + lru = UnifiedLRUList(ComponentType.SWA, self.components) + for n in (a, b): + lru.insert_mru(n) + self.assertEqual(self._lru_order(lru), [b, a]) + + # Big window — refresh walks A and B both, then hits root and stops. + lru.reset_node_and_window_ancestors_mru( + b, root, window_size=1000, should_include=lambda _n: True + ) + self.assertEqual(self._lru_order(lru), [b, a]) + + def test_bounded_refresh_window_zero_is_noop(self): + root, [a, b, c] = self._build_chain([2, 2, 2]) + lru = UnifiedLRUList(ComponentType.SWA, self.components) + for n in (a, b, c): + lru.insert_mru(n) + before = self._lru_order(lru) + lru.reset_node_and_window_ancestors_mru( + c, root, window_size=0, should_include=lambda _n: True + ) + self.assertEqual(self._lru_order(lru), before) + + _CONFIGS: list[CacheConfig] = [ CacheConfig(page_size=1, components=(ComponentType.FULL,)), CacheConfig(page_size=4, components=(ComponentType.FULL,)),