diff --git a/python/sglang/srt/mem_cache/cache_init_params.py b/python/sglang/srt/mem_cache/cache_init_params.py index fbfca879a415..0aa5d1e4d6d9 100644 --- a/python/sglang/srt/mem_cache/cache_init_params.py +++ b/python/sglang/srt/mem_cache/cache_init_params.py @@ -55,5 +55,6 @@ class CacheInitParams: component_registry_override: Optional[dict[ComponentType, type[TreeComponent]]] = ( None ) + tree_core_backend: Optional[str] = dataclasses.field(default=None, kw_only=True) mtp_draft_device_pools: tuple[object, ...] = () diff --git a/python/sglang/srt/mem_cache/storage/lmcache/lmcache_unified_radix_cache.py b/python/sglang/srt/mem_cache/storage/lmcache/lmcache_unified_radix_cache.py index 159ff9ad6d90..bf526ec0e467 100644 --- a/python/sglang/srt/mem_cache/storage/lmcache/lmcache_unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/storage/lmcache/lmcache_unified_radix_cache.py @@ -769,7 +769,10 @@ def _prepare_external_slots_for_insert(self, req: Req) -> None: external_tokens - self.lmcache_connector.aligned_swa_window_size(), 0, ) - req.kv.swa_evicted_seqlen = max(req.kv.swa_evicted_seqlen, swa_missing_end) + req.kv.set_evicted_seqlen( + ComponentType.SWA, + max(req.kv.get_evicted_seqlen(ComponentType.SWA), swa_missing_end), + ) def _publish_external_loaded_prefix(self, req: Req, *, token_ids_len: int) -> None: """Publish retrieved KV and immutable Mamba state into the device tree.""" @@ -807,7 +810,7 @@ def _publish_external_loaded_prefix(self, req: Req, *, token_ids_len: int) -> No value=kv_indices[:total_hit].to(dtype=torch.int64, copy=True), mamba_value=checkpoint, prev_prefix_len=prev_prefix_len, - swa_evicted_seqlen=req.kv.swa_evicted_seqlen, + component_evicted_seqlens=req.kv.component_evicted_seqlens.copy(), chunked=True, priority=req.priority or 0, ) diff --git a/python/sglang/srt/mem_cache/unified_cache/component_type.py b/python/sglang/srt/mem_cache/unified_cache/component_type.py index f09a63f0c42f..77ae5420de02 100644 --- a/python/sglang/srt/mem_cache/unified_cache/component_type.py +++ b/python/sglang/srt/mem_cache/unified_cache/component_type.py @@ -10,6 +10,7 @@ class ComponentType(int, Enum): SWA = 1 MAMBA = 2 C128 = 3 + AUXILIARY_SWA = 4 def __str__(self) -> str: # keep human-readable logging return self.name.lower() diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index e8665f5cbeed..19c96520fa35 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -204,7 +204,11 @@ def __init__( ) # The TreeCore owns the tree member-var state (structure, LRUs, sizes, # evictable leaves) and drives the components' tree-level hooks. - self._tree_core_backend = envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.get() + self._tree_core_backend = ( + params.tree_core_backend + if params.tree_core_backend is not None + else envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.get() + ) self.tree_core = create_tree_core( name=self._tree_core_backend, params=params, diff --git a/test/registered/unit/mem_cache/test_lmcache_component_cursors.py b/test/registered/unit/mem_cache/test_lmcache_component_cursors.py new file mode 100644 index 000000000000..904e3fdfc249 --- /dev/null +++ b/test/registered/unit/mem_cache/test_lmcache_component_cursors.py @@ -0,0 +1,127 @@ +import importlib.util +import sys +from types import ModuleType, SimpleNamespace +from unittest import mock + +import pytest +import torch + +from sglang.srt.managers.schedule_batch import ReqKvInfo +from sglang.srt.mem_cache.base_prefix_cache import ( + IncLockRefResult, + InsertResult, + MatchResult, +) +from sglang.srt.mem_cache.radix_cache import RadixKey +from sglang.srt.mem_cache.unified_cache.component_type import ComponentType +from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +@pytest.fixture +def cache_class(): + # Exercise the real cache module without installing the optional MP service. + metadata = ModuleType("lmcache.integration.sglang.lmcache_mp_metadata") + metadata.LMCacheExternalFlow = mock.Mock() + metadata.LMCachePendingStore = mock.Mock() + connector = ModuleType("lmcache.integration.sglang.unified_lmcache_mp_connector") + connector.UnifiedLMCacheMPConnector = mock.Mock() + with mock.patch.dict( + sys.modules, {metadata.__name__: metadata, connector.__name__: connector} + ): + spec = importlib.util.find_spec( + "sglang.srt.mem_cache.storage.lmcache.lmcache_unified_radix_cache" + ) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module.LMCacheUnifiedRadixCache + + +def _cache_and_request(cache_class, *, swa_enabled, swa_cursor=0): + cache = cache_class.__new__(cache_class) + cache.is_swa_enabled = swa_enabled + cache.lmcache_connector = SimpleNamespace(aligned_swa_window_size=lambda: 4) + flow = SimpleNamespace( + load=SimpleNamespace(local_hit_tokens=2, device_indices=torch.arange(2, 8)), + loaded_skip_tokens=0, + total_hit=8, + key=RadixKey(list(range(8))), + prefix_published=False, + mamba_value=None, + free_mamba_after_load=False, + request_mamba_value=None, + allocated_request_mamba_for_load=False, + load_req=None, + ) + cache._external_flows = {"request": flow} + req = SimpleNamespace( + rid="request", + kv=ReqKvInfo( + req_pool_idx=0, + kv_allocated_len=12, + cache_protected_len=8, + component_evicted_seqlens={ComponentType.AUXILIARY_SWA: 6}, + ), + priority=0, + last_node=None, + ) + if swa_enabled: + req.kv.set_evicted_seqlen(ComponentType.SWA, swa_cursor) + return cache, req, flow + + +@pytest.mark.parametrize("swa_cursor", [0, 6]) +def test_external_swa_load_preserves_independent_cursors(cache_class, swa_cursor): + cache, req, flow = _cache_and_request( + cache_class, swa_enabled=True, swa_cursor=swa_cursor + ) + + cache._publish_external_loaded_prefix(req, token_ids_len=12) + + assert flow.prefix_published + assert req.kv.cache_protected_len == 2 + assert req.kv.component_evicted_seqlens == { + ComponentType.SWA: max(swa_cursor, 4), + ComponentType.AUXILIARY_SWA: 6, + } + + +@pytest.mark.parametrize("swa_enabled", [False, True]) +def test_mamba_publication_snapshots_component_cursors(cache_class, swa_enabled): + cache, req, flow = _cache_and_request(cache_class, swa_enabled=swa_enabled) + flow.mamba_value = torch.tensor([1], dtype=torch.int64) + cache.req_to_token_pool = mock.Mock() + cache.req_to_token_pool.req_to_token = torch.arange(12).reshape(1, 12) + cache.insert = mock.Mock(return_value=InsertResult(prefix_len=8)) + cache.inc_lock_ref = mock.Mock(return_value=IncLockRefResult(node_id=42)) + matched = MatchResult( + device_indices=torch.arange(8), + last_device_node=42, + last_host_node=None, + best_match_node=None, + ) + + with mock.patch.object(UnifiedRadixCache, "match_prefix", return_value=matched): + cache._publish_external_loaded_prefix(req, token_ids_len=12) + + insert_params = cache.insert.call_args.args[0] + expected = {ComponentType.AUXILIARY_SWA: 6} + if swa_enabled: + expected[ComponentType.SWA] = 4 + assert insert_params.component_evicted_seqlens == expected + assert ( + insert_params.component_evicted_seqlens is not req.kv.component_evicted_seqlens + ) + req.kv.set_evicted_seqlen(ComponentType.AUXILIARY_SWA, 99) + assert insert_params.get_evicted_seqlen(ComponentType.AUXILIARY_SWA) == 6 + assert req.kv.cache_protected_len == 8 + assert req.lock_receipt.node_id == 42 + assert flow.prefix_published + assert flow.mamba_value is None + assert req.prefix_indices.tolist() == list(range(12)) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/mem_cache/test_streaming_session_unit.py b/test/registered/unit/mem_cache/test_streaming_session_unit.py index a41d58d65cb3..3813c84585fe 100644 --- a/test/registered/unit/mem_cache/test_streaming_session_unit.py +++ b/test/registered/unit/mem_cache/test_streaming_session_unit.py @@ -125,6 +125,7 @@ def test_session_slot_round_trip_preserves_component_state(): req.kv.mamba_last_track_idx = 0 req.kv.mamba_last_track_seqlen = 3 req.kv.set_evicted_seqlen(ComponentType.SWA, 2) + req.kv.set_evicted_seqlen(ComponentType.AUXILIARY_SWA, 3) slot = SessionSlot() slot.save_from_req(req, is_first=True) @@ -135,7 +136,10 @@ def test_session_slot_round_trip_preserves_component_state(): assert next_req.kv.mamba_next_track_idx == 1 assert next_req.kv.mamba_last_track_idx == 0 assert next_req.kv.mamba_last_track_seqlen == 3 - assert next_req.kv.component_evicted_seqlens == {ComponentType.SWA: 2} + assert next_req.kv.component_evicted_seqlens == { + ComponentType.SWA: 2, + ComponentType.AUXILIARY_SWA: 3, + } assert req.kv.component_evicted_seqlens == {} next_req.kv.mark_kv_released() assert next_req.kv.is_kv_released @@ -261,6 +265,8 @@ def test_release_session_preserves_component_lock_receipt(uuid): ) acquired.set_lock_uuid(ComponentType.SWA, uuid) acquired.set_lock_uuid(ComponentType.SWA, 19, lock_host=True) + acquired.set_lock_uuid(ComponentType.AUXILIARY_SWA, 23) + acquired.set_lock_uuid(ComponentType.AUXILIARY_SWA, None, lock_host=True) tree_cache.slots["session-a"] = SessionSlot( kv=ReqKvInfo( req_pool_idx=0, @@ -274,6 +280,8 @@ def test_release_session_preserves_component_lock_receipt(uuid): acquired.set_lock_uuid(ComponentType.SWA, 99) acquired.set_lock_uuid(ComponentType.SWA, 99, lock_host=True) + acquired.set_lock_uuid(ComponentType.AUXILIARY_SWA, 99) + acquired.set_lock_uuid(ComponentType.AUXILIARY_SWA, 99, lock_host=True) tree_cache.release_session("session-a") assert inner.dec_lock_ref_calls == [lock_node] @@ -282,6 +290,8 @@ def test_release_session_preserves_component_lock_receipt(uuid): assert params.skipped_lock_components == (ComponentType.MAMBA,) assert params.get_lock_uuid(ComponentType.SWA) == uuid assert params.get_lock_uuid(ComponentType.SWA, lock_host=True) == 19 + assert params.get_lock_uuid(ComponentType.AUXILIARY_SWA) == 23 + assert params.get_lock_uuid(ComponentType.AUXILIARY_SWA, lock_host=True) is None for lock_host in (False, True): with pytest.raises(KeyError): params.get_lock_uuid(ComponentType.MAMBA, lock_host=lock_host) @@ -378,7 +388,7 @@ def test_trim_overshoot_postcondition(): @pytest.mark.parametrize("operation", ["trim", "match"]) -@pytest.mark.parametrize("component", [ComponentType.SWA, ComponentType.FULL]) +@pytest.mark.parametrize("component", [ComponentType.SWA, ComponentType.AUXILIARY_SWA]) def test_session_rewind_keeps_component_cursors_page_aligned(operation, component): """Rewinding below any component cursor must free whole pages.""" page_size = 16 diff --git a/test/registered/unit/mem_cache/test_tree_core_registry.py b/test/registered/unit/mem_cache/test_tree_core_registry.py index 11c8b9d5917d..ccb9319dcf1f 100644 --- a/test/registered/unit/mem_cache/test_tree_core_registry.py +++ b/test/registered/unit/mem_cache/test_tree_core_registry.py @@ -81,6 +81,10 @@ class _StubMambaComponent(_StubFullComponent): component_type = ComponentType.MAMBA +class _StubAuxiliarySWAComponent(_StubFullComponent): + component_type = ComponentType.AUXILIARY_SWA + + class TreeCoreRegistryTest(CustomTestCase): def setUp(self): self._registry_snapshot = dict(_TREE_CORE_REGISTRY) @@ -196,6 +200,51 @@ def test_env_var_routes_construction_to_the_selected_backend(self): component = cache.components[ComponentType.FULL] self.assertIs(component.tree_core, core) + def test_backend_override_is_instance_local(self): + for backend, default_backend, expected_backend in ( + ("python", "unregistered-test-core", "python"), + (None, "python", "python"), + (None, "unregistered-test-core", None), + ("", "python", None), + ("unregistered-test-core", "python", None), + ): + with ( + self.subTest(backend=backend, default_backend=default_backend), + envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.override(default_backend), + ): + params = self._cache_params(tree_core_backend=backend) + if expected_backend is None: + with self.assertRaisesRegex(ValueError, "is not registered"): + UnifiedRadixCache(params) + else: + cache = UnifiedRadixCache(params) + self.assertIsInstance(cache.tree_core, UnifiedTreeCore) + self.assertEqual(cache._tree_core_backend, expected_backend) + self.assertEqual( + envs.SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND.get(), default_backend + ) + + def test_auxiliary_component_uses_its_own_node_storage(self): + cache = UnifiedRadixCache( + self._cache_params( + tree_components=(ComponentType.FULL, ComponentType.AUXILIARY_SWA), + component_registry_override={ + ComponentType.FULL: _StubFullComponent, + ComponentType.AUXILIARY_SWA: _StubAuxiliarySWAComponent, + }, + tree_core_backend="python", + ) + ) + node = cache.tree_core.root_node + auxiliary = node.component_data[ComponentType.AUXILIARY_SWA] + full = node.component_data[ComponentType.FULL] + self.assertIsNot(auxiliary, full) + auxiliary.metadata["boundary"] = 7 + self.assertNotIn("boundary", full.metadata) + self.assertIs( + cache.components[ComponentType.AUXILIARY_SWA].tree_core, cache.tree_core + ) + if __name__ == "__main__": unittest.main()