From 5609f8e509f5e0b245a291a90a457c9140fa7c19 Mon Sep 17 00:00:00 2001 From: Ke Bao Date: Sun, 19 Jul 2026 00:40:02 +0800 Subject: [PATCH] Reset only the used mamba state on radix cache hit (#31643) --- .../sglang/srt/mem_cache/mamba_radix_cache.py | 69 +++++++++++-------- .../unit/mem_cache/test_mamba_unittest.py | 59 ++++++++++++++++ 2 files changed, 98 insertions(+), 30 deletions(-) diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index 73623430d..921f5af2c 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -19,7 +19,6 @@ limitations under the License. The radix tree data structure for managing the hybrid (full and Mamba) KV cache. """ -import heapq from array import array from collections import defaultdict from typing import TYPE_CHECKING, List, Optional, Tuple @@ -83,7 +82,11 @@ class TreeNode: self.full_lock_ref = 0 self.mamba_lock_ref = 0 # last access time is only used for sanity check. LRU is maintained by the lru list. + # `last_access_time` tracks the full LRU (whole matched path is reused as prefix); + # `mamba_last_access_time` tracks the mamba LRU, which only touches the single state + # actually consumed per access, so the two orders diverge and need separate stamps. self.last_access_time = get_last_access_time() + self.mamba_last_access_time = self.last_access_time self.hit_count = 0 self.host_ref_counter = 0 @@ -172,10 +175,12 @@ class LRUList: self.prv = "mamba_prev" self.nxt = "mamba_next" self.lock_ref = "mamba_lock_ref" + self.time_attr = "mamba_last_access_time" else: self.prv = "prev" self.nxt = "next" self.lock_ref = "full_lock_ref" + self.time_attr = "last_access_time" # Initialize dummy head and tail nodes self.head = TreeNode() # Most recently used side self.tail = TreeNode() # Least recently used side @@ -226,6 +231,8 @@ class LRUList: assert ( not self.mamba or node.mamba_value is not None ), f"Resetting mamba tombstone node in mamba lru list: {node.id=}" + if self.mamba: + node.mamba_last_access_time = get_last_access_time() self._remove_node(node) self._add_node(node) @@ -255,6 +262,8 @@ class LRUList: assert ( node.id not in self.cache ), f"Inserting node {node.id=} already in lru list, existing node: {self.cache[node.id].id=}" + if self.mamba: + node.mamba_last_access_time = get_last_access_time() self.cache[node.id] = node self._add_node(node) @@ -330,21 +339,20 @@ class LRUList: msg = f"{self.mamba=} LRU list: " x_lru = self._get_lru() while x_lru is not None and x_lru.id in self.cache: - msg += f"[{x_lru.id}] {x_lru.last_access_time:f} -> " + msg += f"[{x_lru.id}] {getattr(x_lru, self.time_attr):f} -> " x_lru = getattr(x_lru, self.prv) print(msg) if not tree_cache: return - msg = f"{self.mamba=} Nodes (sorted by last_access_time): " + msg = f"{self.mamba=} Nodes (sorted by {self.time_attr}): " if self.mamba: nodes = tree_cache._collect_nontombstone_nodes() else: nodes = tree_cache._collect_all_nodes() - heapq.heapify(nodes) - while len(nodes): - x = heapq.heappop(nodes) - msg += f"[{x.id}] {x.last_access_time:f} -> " + nodes.sort(key=lambda n: getattr(n, self.time_attr)) + for x in nodes: + msg += f"[{x.id}] {getattr(x, self.time_attr):f} -> " print(msg) # Note: this is expensive, only use for debug @@ -364,8 +372,8 @@ class LRUList: # Note: this is expensive, only use for debug or idle check def sanity_check(self, tree_cache: MambaRadixCache): """ - Check if the lru list is valid by rebuilding the lru list from the tree, heapifying it, and - checking if the lru list is valid. + Check the lru list is valid by rebuilding it from the tree, sorting by this list's + access-time stamp, and checking the order matches the linked list. """ try: if self.mamba: @@ -374,16 +382,16 @@ class LRUList: nodes = tree_cache._collect_all_nodes() total_nodes = len(nodes) total_lru = len(self.cache) - # heapify based on last_access_time - heapq.heapify(nodes) + # rebuild expected order from this list's own access-time stamp (full and mamba + # lists have independent recency, so they use different stamps) + nodes.sort(key=lambda n: getattr(n, self.time_attr)) # the root node is not in the lru list assert len(nodes) == ( total_lru + (0 if self.mamba else 1) ), f"len(nodes): {len(nodes)}, total_lru: {total_lru}" x_lru = self._get_lru() - while len(nodes): - x = heapq.heappop(nodes) + for x in nodes: if x == tree_cache.root_node: # root node is not in the lru list continue @@ -393,7 +401,7 @@ class LRUList: assert ( x == x_lru - ), f"Incorrect LRU list, {self.mamba=}, x: {x.id=} != x_lru: {x_lru.id=}, {x.last_access_time=}, {x_lru.last_access_time=}" + ), f"Incorrect LRU list, {self.mamba=}, x: {x.id=} != x_lru: {x_lru.id=}, {getattr(x, self.time_attr)=}, {getattr(x_lru, self.time_attr)=}" assert ( x_lru.full_lock_ref == 0 ), f"x_lru should not be locked when idle, {x_lru.full_lock_ref=}, {x_lru.id=}" @@ -1104,11 +1112,16 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): cow_mamba = params.cow_mamba req = params.req - # update time for matched nodes, and make nodes closer to root to be least recently used - # this allows mamba to evict nodes closer to root first + # Full KV of the whole matched path is reused as prefix, so refresh the entire + # chain (nodes closer to root end up least recently used, evicted first). node_update = last_node self.full_lru_list.reset_node_and_parents_mru(node_update, self.root_node) - self.mamba_lru_list.reset_node_and_parents_mru(node_update, self.root_node) + # Mamba only consumes last_node's state (cf. inc_lock_ref, which locks just this + # node's mamba_value). Refreshing ancestors would keep a whole session's states + # adjacent in the mamba LRU and evict cold sessions wholesale; touch only the used + # state so older leaves survive. + if last_node is not self.root_node and last_node.mamba_value is not None: + self.mamba_lru_list.reset_node_mru(last_node) # This last_access_time is for sanity check, can be deleted after validation in production cur_time = get_last_access_time() @@ -1170,12 +1183,13 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): new_node.key = child.key[:split_len] new_node.value = child.value[:split_len].clone() - # child time should be later than parent's time for mamba tombstone + # child time should be later than the new parent's time in the full LRU child.last_access_time = get_last_access_time() + # A split does not change the set of live mamba states (child keeps its value, + # new_node is a mamba tombstone), so the mamba LRU is left untouched — only the + # full LRU reorders around the new intermediate node. self.full_lru_list.remove_node(child) - if child.mamba_value is not None: - self.mamba_lru_list.remove_node(child) child.parent = new_node child.key = child.key[split_len:] child.value = child.value[split_len:].clone() @@ -1184,12 +1198,10 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): child.hash_value, split_len, self.page_size ) - # insert the new node and child into the lru lists, insert + # insert the new node and child into the full lru list, insert # parent first so that parent is after child in the lru list self.full_lru_list.insert_mru(new_node) self.full_lru_list.insert_mru(child) - if child.mamba_value is not None: - self.mamba_lru_list.insert_mru(child) return new_node def _insert_helper( @@ -1201,14 +1213,14 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): chunked: bool = False, prev_prefix_len: int = 0, ) -> Tuple[int, bool]: - # Update the last access time from root to leaf, so that - # mamba will tombstone the node closer to root first + # Refresh the full LRU from root to leaf (the whole path is reused as prefix). + # The mamba states of these existing nodes were not recomputed this insert, so + # the mamba LRU is left untouched here; only genuinely new mamba states (the new + # leaf / a revived tombstone below) are inserted. assert mamba_value is not None, "Mamba value should not be None here." node.last_access_time = get_last_access_time() if node != self.root_node: self.full_lru_list.reset_node_mru(node) - if node.mamba_value is not None: - self.mamba_lru_list.reset_node_mru(node) if len(key) == 0: return 0, True @@ -1219,8 +1231,6 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): node = node.children[child_key] node.last_access_time = get_last_access_time() self.full_lru_list.reset_node_mru(node) - if node.mamba_value is not None: - self.mamba_lru_list.reset_node_mru(node) prefix_len = node.key.match(key, page_size=self.page_size) if prev_prefix_len < total_prefix_length + prefix_len: @@ -1260,7 +1270,6 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): else: # mamba value already exists mamba_value_exist = True self.full_lru_list.reset_node_mru(node) - self.mamba_lru_list.reset_node_mru(node) node.last_access_time = get_last_access_time() return total_prefix_length, mamba_value_exist diff --git a/test/registered/unit/mem_cache/test_mamba_unittest.py b/test/registered/unit/mem_cache/test_mamba_unittest.py index 01a830342..599e06953 100755 --- a/test/registered/unit/mem_cache/test_mamba_unittest.py +++ b/test/registered/unit/mem_cache/test_mamba_unittest.py @@ -387,6 +387,65 @@ class TestMamba(unittest.TestCase): print(available_and_evictable_str(tree)) tree.sanity_check() + def test_mamba_lru_match_refreshes_only_used_node(self): + """A prefix-cache hit must refresh only the matched leaf's mamba state in + the mamba LRU, not its ancestors. Whole-chain refresh clustered a session's + states adjacently, so under mamba-pool pressure eviction dropped whole cold + sessions instead of the intermediate states reuse never needs. Guards against + reverting the mamba list to reset_node_and_parents_mru. + """ + tree, allocator, req_to_token_pool, make_dummy_req = ( + self._setup_tree_and_allocator() + ) + + def insert(token_ids): + req = make_dummy_req() + kv = allocator.alloc(len(token_ids)) + tree.insert( + InsertParams( + key=RadixKey(array("q", token_ids)), + value=kv, + mamba_value=req.mamba_pool_idx.unsqueeze(0), + ) + ) + + def match_leaf(token_ids): + return tree.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", token_ids))) + ).last_device_node + + def mamba_lru_mru_to_lru(): + lst = tree.mamba_lru_list + order, x = [], getattr(lst.head, lst.nxt) + while x is not None and x is not lst.tail and x.id in lst.cache: + order.append(x) + x = getattr(x, lst.nxt) + return order + + # Two independent sessions, each a 2-node mamba chain: + # root -> a1 -> b1 and root -> a2 -> b2 + insert([1, 2, 3]) + insert([1, 2, 3, 4, 5, 6]) + insert([7, 8, 9]) + insert([7, 8, 9, 10, 11, 12]) + + b1 = match_leaf([1, 2, 3, 4, 5, 6]) + a1 = b1.parent + b2 = match_leaf([7, 8, 9, 10, 11, 12]) + a2 = b2.parent + # Session 2 was matched last, so session 1's ancestor a1 is older than a2. + order = mamba_lru_mru_to_lru() + self.assertGreater(order.index(a1), order.index(a2)) + + # Re-access session 1. Only its consumed leaf (b1) moves to MRU; its ancestor + # a1 must stay put -- whole-chain reset would bump a1 right behind b1, making + # it newer than a2. + self.assertIs(match_leaf([1, 2, 3, 4, 5, 6]), b1) + order = mamba_lru_mru_to_lru() + self.assertIs(order[0], b1) + self.assertGreater(order.index(a1), order.index(a2)) + tree.sanity_check() + def test_mamba_radix_cache_kv_events(self): tree, allocator, _, make_dummy_req = self._setup_tree_and_allocator( enable_kv_cache_events=True