diff --git a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py index 3c2988854..0a3ae9160 100644 --- a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py @@ -632,6 +632,11 @@ class HiMambaRadixCache(MambaRadixCache): break if node.parent.full_lock_ref > 0 or node.parent.mamba_lock_ref > 0: break + if ( + node.parent.host_ref_counter > 0 + or node.parent.host_mamba_ref_counter > 0 + ): + break parent = node.parent diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index d201093f2..1f2c2017e 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -664,6 +664,7 @@ class MambaRadixCache(BasePrefixCache): value=page_aligned_kv_indices, mamba_value=mamba_value_forked, prev_prefix_len=req.cache_protected_len, + chunked=chunked, ) ) new_prefix_len, mamba_exist = result.prefix_len, result.mamba_exist diff --git a/test/registered/unit/mem_cache/test_mamba_unittest.py b/test/registered/unit/mem_cache/test_mamba_unittest.py index 6158f622b..4d4450ea5 100755 --- a/test/registered/unit/mem_cache/test_mamba_unittest.py +++ b/test/registered/unit/mem_cache/test_mamba_unittest.py @@ -13,7 +13,8 @@ from sglang.srt.mem_cache.base_prefix_cache import ( ) from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.common import available_and_evictable_str -from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache +from sglang.srt.mem_cache.hi_mamba_radix_cache import HiMambaRadixCache +from sglang.srt.mem_cache.mamba_radix_cache import LRUList, MambaRadixCache, TreeNode from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.sampling.sampling_params import SamplingParams @@ -392,6 +393,53 @@ class TestMamba(unittest.TestCase): return tree, allocator, req_to_token_pool, make_dummy_req + def test_hi_mamba_tombstone_cleanup_respects_host_ref(self): + tree = object.__new__(HiMambaRadixCache) + root = TreeNode() + parent = TreeNode() + deleted = TreeNode() + + root.key = RadixKey([]) + parent.key = RadixKey([1]) + deleted.key = RadixKey([2]) + parent.parent = root + deleted.parent = parent + parent.value = torch.tensor([1], dtype=torch.int64) + parent.protect_host() + root.children[parent.key.child_key(1)] = parent + + class RecordingCacheController: + def __init__(self): + self.device_evictions = [] + self.host_evictions = [] + + def evict_device(self, value): + self.device_evictions.append(value) + + def evict_host(self, value): + self.host_evictions.append(value) + + tree.root_node = root + tree.page_size = 1 + tree.full_lru_list = LRUList(mamba=False) + tree.full_lru_list.insert_mru(parent) + tree.cache_controller = RecordingCacheController() + tree.full_evictable_size_ = len(parent.value) + tree.evictable_full_device_leaves = {parent} + tree.evictable_full_host_leaves = set() + + result_node, full_evicted, mamba_evicted = ( + tree._iteratively_delete_tombstone_leaf(deleted) + ) + + self.assertIs(result_node, deleted) + self.assertEqual(full_evicted, 0) + self.assertEqual(mamba_evicted, 0) + self.assertIs(root.children[parent.key.child_key(1)], parent) + self.assertTrue(tree.full_lru_list.in_list(parent)) + self.assertEqual(tree.cache_controller.device_evictions, []) + self.assertEqual(tree.cache_controller.host_evictions, []) + def test_mamba_pool_cpu_offload(self): """MambaPool.get_cpu_copy / load_cpu_copy round-trips conv and temporal state.""" _, _, req_to_token_pool, _ = self._setup_tree_and_allocator() 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 114a06e1b..069205311 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 @@ -28,7 +28,10 @@ from sglang.srt.mem_cache.memory_pool import ( from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool, SWATokenToKVPoolAllocator from sglang.srt.mem_cache.unified_cache_components.tree_component import ComponentType -from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache +from sglang.srt.mem_cache.unified_radix_cache import ( + UnifiedRadixCache, + UnifiedTreeNode, +) from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ( ServerArgs, @@ -871,6 +874,30 @@ class UnifiedRadixCacheSuite: ) tree.sanity_check() + def test_tombstone_cleanup_respects_locked_parent(self): + tree, _, _ = build_fixture(self.cfg) + parent = UnifiedTreeNode(self.cfg.components) + deleted = UnifiedTreeNode(self.cfg.components) + + parent.key = RadixKey(self._make_seq(1, 1)) + deleted.key = RadixKey(self._make_seq(1000, 1)) + parent.parent = tree.root_node + deleted.parent = parent + parent.component_data[ComponentType.FULL].value = torch.arange( + self.cfg.page_size, dtype=torch.int64, device=tree.device + ) + parent.component_data[ComponentType.FULL].lock_ref = 1 + parent_key = parent.key.child_key(tree.page_size) + tree.root_node.children[parent_key] = parent + + tracker = {ct: 0 for ct in tree.tree_components} + + tree._iteratively_delete_tombstone_leaf(deleted, tracker) + + self.assertIn(parent_key, tree.root_node.children) + self.assertIs(tree.root_node.children[parent_key], parent) + self.assertTrue(all(evicted == 0 for evicted in tracker.values())) + def test_internal_readonly_does_not_modify_tree(self): """Verify readonly match does not modify tree structure (no split).""" if self.cfg.page_size > 1 or self.cfg.has_mamba or self.cfg.has_swa: