Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions python/sglang/srt/mem_cache/cache_init_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -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, ...] = ()
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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,
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
6 changes: 5 additions & 1 deletion python/sglang/srt/mem_cache/unified_radix_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
127 changes: 127 additions & 0 deletions test/registered/unit/mem_cache/test_lmcache_component_cursors.py
Original file line number Diff line number Diff line change
@@ -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"]))
14 changes: 12 additions & 2 deletions test/registered/unit/mem_cache/test_streaming_session_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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]
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down
49 changes: 49 additions & 0 deletions test/registered/unit/mem_cache/test_tree_core_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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()
Loading