From 6e0b7f35adaa7056ae662b16b0ad1eace2b38e5c Mon Sep 17 00:00:00 2001 From: Jialin Ouyang Date: Wed, 20 May 2026 10:05:04 -0700 Subject: [PATCH] [radix cache] pluggable RadixCache factory (--radix-cache-backend) (#25101) --- .../sglang/srt/mem_cache/kv_cache_builder.py | 97 +----- python/sglang/srt/mem_cache/registry.py | 194 ++++++++++++ python/sglang/srt/server_args.py | 14 + .../unit/mem_cache/test_registry.py | 291 ++++++++++++++++++ 4 files changed, 516 insertions(+), 80 deletions(-) create mode 100644 python/sglang/srt/mem_cache/registry.py create mode 100644 test/registered/unit/mem_cache/test_registry.py diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 5e74ca0e6..821eccdde 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -27,9 +27,8 @@ from sglang.srt.configs.model_config import ModelImpl from sglang.srt.environ import envs from sglang.srt.managers.mm_utils import init_mm_embedding_cache from sglang.srt.mem_cache.cache_init_params import CacheInitParams -from sglang.srt.mem_cache.radix_cache import RadixCache +from sglang.srt.mem_cache.registry import TreeCacheBuildContext, create_tree_cache from sglang.srt.model_loader.utils import get_resolved_model_impl -from sglang.srt.session.streaming_session import StreamingSession if TYPE_CHECKING: @@ -223,84 +222,22 @@ def build_kv_cache( sliding_window_size=sliding_window_size, ) - if effective_chunked_prefill_size is not None and disable_radix_cache: - if not is_hybrid_swa: - from sglang.srt.mem_cache.chunk_cache import ChunkCache - - tree_cache = ChunkCache(params) - else: - from sglang.srt.mem_cache.chunk_cache import SWAChunkCache - - tree_cache = SWAChunkCache(params) - else: - if envs.SGLANG_EXPERIMENTAL_CPP_RADIX_TREE.get(): - # lazy import to avoid JIT overhead - from sglang.srt.mem_cache.radix_cache_cpp import RadixCacheCpp - - logger.info("Using experimental C++ radix tree implementation.") - tree_cache = RadixCacheCpp(params=params, server_args=server_args) - elif envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get(): - from sglang.srt.mem_cache.unified_cache_components import ( - ComponentType, - ) - from sglang.srt.mem_cache.unified_radix_cache import ( - UnifiedRadixCache, - ) - - tree_components = [ComponentType.FULL] - if is_hybrid_swa or is_hybrid_ssm: - tree_components.append( - ComponentType.SWA if is_hybrid_swa else ComponentType.MAMBA - ) - params.tree_components = tuple(tree_components) - tree_cache = UnifiedRadixCache(params) - if enable_hierarchical_cache: - tree_cache.init_hicache(server_args, params) - tp_worker.register_hicache_layer_transfer_counter( - tree_cache.cache_controller.layer_done_counter - ) - elif enable_hierarchical_cache: - if is_hybrid_ssm: - from sglang.srt.mem_cache.hi_mamba_radix_cache import ( - HiMambaRadixCache, - ) - - tree_cache = HiMambaRadixCache(params=params, server_args=server_args) - else: - from sglang.srt.mem_cache.hiradix_cache import HiRadixCache - - tree_cache = HiRadixCache(params=params, server_args=server_args) - tp_worker.register_hicache_layer_transfer_counter( - tree_cache.cache_controller.layer_done_counter - ) - elif is_hybrid_swa: - from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache - - tree_cache = SWARadixCache(params=params) - elif is_hybrid_ssm: - from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache - - tree_cache = MambaRadixCache(params) - elif server_args.enable_lmcache: - from sglang.srt.mem_cache.storage.lmcache.lmc_radix_cache import ( - LMCRadixCache, - ) - - tree_cache = LMCRadixCache( - params=params, - model_config=model_config, - tp_size=ps.tp_size, - rank=ps.tp_rank, - tp_group=tp_group, - ) - else: - tree_cache = RadixCache(params) - - if ( - server_args.enable_streaming_session - and not tree_cache.supports_streaming_session() - ): - tree_cache = StreamingSession(tree_cache) + tree_cache = create_tree_cache( + TreeCacheBuildContext( + server_args=server_args, + params=params, + is_hybrid_swa=is_hybrid_swa, + is_hybrid_ssm=is_hybrid_ssm, + enable_hierarchical_cache=enable_hierarchical_cache, + disable_radix_cache=disable_radix_cache, + effective_chunked_prefill_size=effective_chunked_prefill_size, + tp_worker=tp_worker, + model_config=model_config, + tp_size=ps.tp_size, + tp_rank=ps.tp_rank, + tp_group=tp_group, + ) + ) embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get() init_mm_embedding_cache(embedding_cache_size * 1024 * 1024) diff --git a/python/sglang/srt/mem_cache/registry.py b/python/sglang/srt/mem_cache/registry.py new file mode 100644 index 000000000..c91aae91e --- /dev/null +++ b/python/sglang/srt/mem_cache/registry.py @@ -0,0 +1,194 @@ +"""Registry for pluggable RadixCache factories. + +If `--radix-cache-backend` is unset (by default), the built-in selection +chain is used to pick a cache implementation. + +To plug in a custom backend, register it under a string name via +`register_radix_cache_backend(name, factory)`, then select it with +`--radix-cache-backend ` (the flag accepts only registered names). +""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Callable, Optional + +from sglang.srt.environ import envs +from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache +from sglang.srt.mem_cache.cache_init_params import CacheInitParams + +if TYPE_CHECKING: + from sglang.srt.configs.model_config import ModelConfig + from sglang.srt.server_args import ServerArgs + +logger = logging.getLogger(__name__) + + +@dataclass +class TreeCacheBuildContext: + """Radix Cache construction arguments.""" + + server_args: ServerArgs + params: CacheInitParams + is_hybrid_swa: bool + is_hybrid_ssm: bool + enable_hierarchical_cache: bool + disable_radix_cache: bool + effective_chunked_prefill_size: Optional[int] + tp_worker: Any + model_config: ModelConfig + tp_size: int + tp_rank: int + tp_group: Any + + +RadixCacheFactory = Callable[[TreeCacheBuildContext], BasePrefixCache] + +_RADIX_CACHE_REGISTRY: dict[str, RadixCacheFactory] = {} + + +def register_radix_cache_backend(name: str, factory: RadixCacheFactory) -> None: + """Register a radix-cache factory under `name`. + + Raises ValueError if `name` is empty/whitespace-only or already + registered. + """ + if not name.strip(): + raise ValueError( + f"register_radix_cache_backend: name must be non-empty, got {name!r}" + ) + if name in _RADIX_CACHE_REGISTRY: + raise ValueError( + f"register_radix_cache_backend: {name!r} is already registered" + ) + _RADIX_CACHE_REGISTRY[name] = factory + + +def get_radix_cache_factory(name: str) -> Optional[RadixCacheFactory]: + return _RADIX_CACHE_REGISTRY.get(name) + + +def registered_radix_cache_backends() -> list[str]: + return list(_RADIX_CACHE_REGISTRY.keys()) + + +def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache: + """Built-in Radix Cache selection chain.""" + server_args = ctx.server_args + params = ctx.params + + if ctx.effective_chunked_prefill_size is not None and ctx.disable_radix_cache: + if not ctx.is_hybrid_swa: + from sglang.srt.mem_cache.chunk_cache import ChunkCache + + return ChunkCache(params) + from sglang.srt.mem_cache.chunk_cache import SWAChunkCache + + return SWAChunkCache(params) + + if envs.SGLANG_EXPERIMENTAL_CPP_RADIX_TREE.get(): + # lazy import to avoid JIT overhead + from sglang.srt.mem_cache.radix_cache_cpp import RadixCacheCpp + + logger.info("Using experimental C++ radix tree implementation.") + return RadixCacheCpp(params=params, server_args=server_args) + + if envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get(): + from sglang.srt.mem_cache.unified_cache_components import ComponentType + from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache + + tree_components = [ComponentType.FULL] + if ctx.is_hybrid_swa or ctx.is_hybrid_ssm: + tree_components.append( + ComponentType.SWA if ctx.is_hybrid_swa else ComponentType.MAMBA + ) + params.tree_components = tuple(tree_components) + cache = UnifiedRadixCache(params) + if ctx.enable_hierarchical_cache: + cache.init_hicache(server_args, params) + ctx.tp_worker.register_hicache_layer_transfer_counter( + cache.cache_controller.layer_done_counter + ) + return cache + + if ctx.enable_hierarchical_cache: + if ctx.is_hybrid_ssm: + from sglang.srt.mem_cache.hi_mamba_radix_cache import HiMambaRadixCache + + cache = HiMambaRadixCache(params=params, server_args=server_args) + else: + from sglang.srt.mem_cache.hiradix_cache import HiRadixCache + + cache = HiRadixCache(params=params, server_args=server_args) + ctx.tp_worker.register_hicache_layer_transfer_counter( + cache.cache_controller.layer_done_counter + ) + return cache + + if ctx.is_hybrid_swa: + from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache + + return SWARadixCache(params=params) + + if ctx.is_hybrid_ssm: + from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache + + return MambaRadixCache(params) + + if server_args.enable_lmcache: + from sglang.srt.mem_cache.storage.lmcache.lmc_radix_cache import ( + LMCRadixCache, + ) + + return LMCRadixCache( + params=params, + model_config=ctx.model_config, + tp_size=ctx.tp_size, + rank=ctx.tp_rank, + tp_group=ctx.tp_group, + ) + + from sglang.srt.mem_cache.radix_cache import RadixCache + + return RadixCache(params) + + +def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache: + """Route to the matching factory to construct Radix Cache.""" + name = ctx.server_args.radix_cache_backend + if name: + factory = get_radix_cache_factory(name) + if factory is None: + raise ValueError( + f"--radix-cache-backend={name!r} is not registered. " + f"Registered backends: {registered_radix_cache_backends()}. " + "External backends must call register_radix_cache_backend(...) at import time." + ) + cache = factory(ctx) + source = f"registered({name!r})" + else: + cache = default_radix_cache_factory(ctx) + source = "default" + + streaming_wrapped = False + if ( + ctx.server_args.enable_streaming_session + and not cache.supports_streaming_session() + ): + from sglang.srt.session.streaming_session import StreamingSession + + cache = StreamingSession(cache) + streaming_wrapped = True + + logger.info( + "Tree cache initialized: source=%s impl=%s hybrid_swa=%s hybrid_ssm=%s " + "hierarchical=%s streaming_wrapped=%s", + source, + type(cache).__name__, + ctx.is_hybrid_swa, + ctx.is_hybrid_ssm, + ctx.enable_hierarchical_cache, + streaming_wrapped, + ) + return cache diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index ae4cb7c12..d2aae837d 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -536,6 +536,10 @@ class ServerArgs: prefill_attention_backend: Optional[str] = None sampling_backend: Optional[str] = None grammar_backend: Optional[str] = None + # Name of a custom radix-cache factory registered via + # register_radix_cache_backend. Leave unset (by default) to use the + # built-in default cache selection chain. + radix_cache_backend: Optional[str] = None mm_attention_backend: Optional[str] = None fp8_gemm_runner_backend: str = "auto" fp4_gemm_runner_backend: str = "auto" @@ -5373,6 +5377,16 @@ class ServerArgs: default=ServerArgs.grammar_backend, help="Choose the backend for grammar-guided decoding.", ) + parser.add_argument( + "--radix-cache-backend", + type=str, + default=ServerArgs.radix_cache_backend, + help=( + "Name of a radix-cache backend previously registered via " + "register_radix_cache_backend. Omit this flag to use the " + "built-in default cache selection chain." + ), + ) parser.add_argument( "--mm-attention-backend", type=str, diff --git a/test/registered/unit/mem_cache/test_registry.py b/test/registered/unit/mem_cache/test_registry.py new file mode 100644 index 000000000..302d61f11 --- /dev/null +++ b/test/registered/unit/mem_cache/test_registry.py @@ -0,0 +1,291 @@ +"""Unit tests for the radix-cache registry, routing, and selection chain.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +import unittest +from unittest.mock import MagicMock, patch + +from sglang.srt.mem_cache.registry import ( + _RADIX_CACHE_REGISTRY, + TreeCacheBuildContext, + create_tree_cache, + default_radix_cache_factory, + get_radix_cache_factory, + register_radix_cache_backend, + registered_radix_cache_backends, +) +from sglang.test.test_utils import CustomTestCase + + +def _make_ctx( + *, + backend=None, + enable_streaming=False, + enable_lmcache=False, + is_hybrid_swa=False, + is_hybrid_ssm=False, + enable_hierarchical_cache=False, + disable_radix_cache=False, + effective_chunked_prefill_size=None, +): + server_args = MagicMock() + server_args.radix_cache_backend = backend + server_args.enable_streaming_session = enable_streaming + server_args.enable_lmcache = enable_lmcache + return TreeCacheBuildContext( + server_args=server_args, + params=MagicMock(), + is_hybrid_swa=is_hybrid_swa, + is_hybrid_ssm=is_hybrid_ssm, + enable_hierarchical_cache=enable_hierarchical_cache, + disable_radix_cache=disable_radix_cache, + effective_chunked_prefill_size=effective_chunked_prefill_size, + tp_worker=MagicMock(), + model_config=MagicMock(), + tp_size=1, + tp_rank=0, + tp_group=MagicMock(), + ) + + +class _RegistryIsolationMixin: + """Restore the global registry around each test so registrations + from one test don't leak into the next. + """ + + def setUp(self): + super().setUp() + self._registry_snapshot = dict(_RADIX_CACHE_REGISTRY) + + def tearDown(self): + _RADIX_CACHE_REGISTRY.clear() + _RADIX_CACHE_REGISTRY.update(self._registry_snapshot) + super().tearDown() + + +class TestRegisterRadixCacheBackend(_RegistryIsolationMixin, CustomTestCase): + def test_register_then_lookup(self): + factory = MagicMock() + register_radix_cache_backend("oss_unit_test", factory) + self.assertIs(get_radix_cache_factory("oss_unit_test"), factory) + self.assertIn("oss_unit_test", registered_radix_cache_backends()) + + def test_lookup_unknown_returns_none(self): + self.assertIsNone(get_radix_cache_factory("definitely_not_registered")) + + def test_empty_name_raises(self): + with self.assertRaises(ValueError): + register_radix_cache_backend("", MagicMock()) + + def test_whitespace_only_name_raises(self): + with self.assertRaises(ValueError): + register_radix_cache_backend(" ", MagicMock()) + + def test_duplicate_registration_raises(self): + register_radix_cache_backend("dupe", MagicMock()) + with self.assertRaises(ValueError): + register_radix_cache_backend("dupe", MagicMock()) + + +class TestCreateTreeCacheRouting(_RegistryIsolationMixin, CustomTestCase): + def test_dispatches_to_registered_factory(self): + cache = MagicMock() + cache.supports_streaming_session.return_value = True + factory = MagicMock(return_value=cache) + register_radix_cache_backend("custom", factory) + + result = create_tree_cache(_make_ctx(backend="custom")) + + factory.assert_called_once() + self.assertIs(result, cache) + + def test_unknown_backend_raises(self): + with self.assertRaises(ValueError): + create_tree_cache(_make_ctx(backend="not_a_real_backend")) + + @patch("sglang.srt.mem_cache.registry.default_radix_cache_factory") + def test_unset_backend_falls_back_to_default(self, default_factory): + cache = MagicMock() + cache.supports_streaming_session.return_value = True + default_factory.return_value = cache + + result = create_tree_cache(_make_ctx(backend=None)) + + default_factory.assert_called_once() + self.assertIs(result, cache) + + def test_streaming_wrap_when_cache_does_not_support_it(self): + inner = MagicMock() + inner.supports_streaming_session.return_value = False + register_radix_cache_backend("nonstreaming", MagicMock(return_value=inner)) + + with patch( + "sglang.srt.session.streaming_session.StreamingSession" + ) as session_cls: + session_cls.return_value = MagicMock(name="wrapped") + result = create_tree_cache( + _make_ctx(backend="nonstreaming", enable_streaming=True) + ) + + session_cls.assert_called_once_with(inner) + self.assertIs(result, session_cls.return_value) + + def test_no_streaming_wrap_when_cache_supports_it(self): + inner = MagicMock() + inner.supports_streaming_session.return_value = True + register_radix_cache_backend("streaming", MagicMock(return_value=inner)) + + result = create_tree_cache( + _make_ctx(backend="streaming", enable_streaming=True) + ) + + self.assertIs(result, inner) + + +class TestDefaultRadixCacheFactory(CustomTestCase): + """Branch coverage for the built-in radix cache selection chain. + + Each cache class is imported lazily inside the factory, so we patch + the class at its definition site to verify routing without depending + on each cache's real constructor or runtime state. + """ + + def test_chunk_cache_when_chunked_prefill_and_disable_radix(self): + ctx = _make_ctx(effective_chunked_prefill_size=512, disable_radix_cache=True) + with patch("sglang.srt.mem_cache.chunk_cache.ChunkCache") as ChunkCache: + ChunkCache.return_value = MagicMock() + result = default_radix_cache_factory(ctx) + ChunkCache.assert_called_once_with(ctx.params) + self.assertIs(result, ChunkCache.return_value) + + def test_swa_chunk_cache_when_chunked_prefill_disable_and_hybrid_swa(self): + ctx = _make_ctx( + effective_chunked_prefill_size=512, + disable_radix_cache=True, + is_hybrid_swa=True, + ) + with patch("sglang.srt.mem_cache.chunk_cache.SWAChunkCache") as SWAChunkCache: + SWAChunkCache.return_value = MagicMock() + result = default_radix_cache_factory(ctx) + SWAChunkCache.assert_called_once_with(ctx.params) + self.assertIs(result, SWAChunkCache.return_value) + + def test_cpp_radix_cache_when_env_flag_set(self): + ctx = _make_ctx() + # `radix_cache_cpp` requires ninja + C++ extension to import, so + # we inject a stand-in module rather than letting patch() trigger + # the real import. + fake_module = MagicMock() + with patch( + "sglang.srt.mem_cache.registry.envs.SGLANG_EXPERIMENTAL_CPP_RADIX_TREE.get", + return_value=True, + ), patch.dict( + "sys.modules", + {"sglang.srt.mem_cache.radix_cache_cpp": fake_module}, + ): + result = default_radix_cache_factory(ctx) + fake_module.RadixCacheCpp.assert_called_once_with( + params=ctx.params, server_args=ctx.server_args + ) + self.assertIs(result, fake_module.RadixCacheCpp.return_value) + + def test_unified_radix_cache_when_env_flag_set(self): + ctx = _make_ctx() + # Shim both factory imports — each transitively loads sgl_kernel. + fake_components = MagicMock() + fake_radix = MagicMock() + with patch( + "sglang.srt.mem_cache.registry.envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get", + return_value=True, + ), patch.dict( + "sys.modules", + { + "sglang.srt.mem_cache.unified_cache_components": fake_components, + "sglang.srt.mem_cache.unified_radix_cache": fake_radix, + }, + ): + result = default_radix_cache_factory(ctx) + fake_radix.UnifiedRadixCache.assert_called_once_with(ctx.params) + self.assertIs(result, fake_radix.UnifiedRadixCache.return_value) + + def test_hi_radix_cache_when_hierarchical(self): + ctx = _make_ctx(enable_hierarchical_cache=True) + # `hiradix_cache` imports `hicache_storage` and + # `memory_pool_host`, both of which transitively load + # `sgl_kernel`; inject a stand-in module. + fake_module = MagicMock() + with patch.dict( + "sys.modules", + {"sglang.srt.mem_cache.hiradix_cache": fake_module}, + ): + result = default_radix_cache_factory(ctx) + fake_module.HiRadixCache.assert_called_once_with( + params=ctx.params, server_args=ctx.server_args + ) + ctx.tp_worker.register_hicache_layer_transfer_counter.assert_called_once() + self.assertIs(result, fake_module.HiRadixCache.return_value) + + def test_hi_mamba_radix_cache_when_hierarchical_and_hybrid_ssm(self): + ctx = _make_ctx(enable_hierarchical_cache=True, is_hybrid_ssm=True) + # `hi_mamba_radix_cache` imports `hicache_storage`, which + # transitively loads `sgl_kernel`; inject a stand-in module. + fake_module = MagicMock() + with patch.dict( + "sys.modules", + {"sglang.srt.mem_cache.hi_mamba_radix_cache": fake_module}, + ): + result = default_radix_cache_factory(ctx) + fake_module.HiMambaRadixCache.assert_called_once_with( + params=ctx.params, server_args=ctx.server_args + ) + self.assertIs(result, fake_module.HiMambaRadixCache.return_value) + + def test_swa_radix_cache_when_hybrid_swa(self): + ctx = _make_ctx(is_hybrid_swa=True) + with patch("sglang.srt.mem_cache.swa_radix_cache.SWARadixCache") as SWA: + SWA.return_value = MagicMock() + result = default_radix_cache_factory(ctx) + SWA.assert_called_once_with(params=ctx.params) + self.assertIs(result, SWA.return_value) + + def test_mamba_radix_cache_when_hybrid_ssm(self): + ctx = _make_ctx(is_hybrid_ssm=True) + with patch("sglang.srt.mem_cache.mamba_radix_cache.MambaRadixCache") as Mamba: + Mamba.return_value = MagicMock() + result = default_radix_cache_factory(ctx) + Mamba.assert_called_once_with(ctx.params) + self.assertIs(result, Mamba.return_value) + + def test_lmc_radix_cache_when_enable_lmcache(self): + ctx = _make_ctx(enable_lmcache=True) + # The lmcache backend raises at import time when the `lmcache` + # package isn't installed, so inject a stand-in module instead + # of letting patch() trigger the real import. + fake_module = MagicMock() + with patch.dict( + "sys.modules", + {"sglang.srt.mem_cache.storage.lmcache.lmc_radix_cache": fake_module}, + ): + result = default_radix_cache_factory(ctx) + fake_module.LMCRadixCache.assert_called_once_with( + params=ctx.params, + model_config=ctx.model_config, + tp_size=ctx.tp_size, + rank=ctx.tp_rank, + tp_group=ctx.tp_group, + ) + self.assertIs(result, fake_module.LMCRadixCache.return_value) + + def test_fallback_to_radix_cache(self): + ctx = _make_ctx() + with patch("sglang.srt.mem_cache.radix_cache.RadixCache") as RadixCache: + RadixCache.return_value = MagicMock() + result = default_radix_cache_factory(ctx) + RadixCache.assert_called_once_with(ctx.params) + self.assertIs(result, RadixCache.return_value) + + +if __name__ == "__main__": + unittest.main()