[radix cache] pluggable RadixCache factory (--radix-cache-backend) (#25101)

This commit is contained in:
Jialin Ouyang
2026-05-20 10:05:04 -07:00
committed by GitHub
parent ccbbae00ea
commit 6e0b7f35ad
4 changed files with 516 additions and 80 deletions
+17 -80
View File
@@ -27,9 +27,8 @@ from sglang.srt.configs.model_config import ModelImpl
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.mm_utils import init_mm_embedding_cache 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.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.model_loader.utils import get_resolved_model_impl
from sglang.srt.session.streaming_session import StreamingSession
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -223,84 +222,22 @@ def build_kv_cache(
sliding_window_size=sliding_window_size, sliding_window_size=sliding_window_size,
) )
if effective_chunked_prefill_size is not None and disable_radix_cache: tree_cache = create_tree_cache(
if not is_hybrid_swa: TreeCacheBuildContext(
from sglang.srt.mem_cache.chunk_cache import ChunkCache server_args=server_args,
params=params,
tree_cache = ChunkCache(params) is_hybrid_swa=is_hybrid_swa,
else: is_hybrid_ssm=is_hybrid_ssm,
from sglang.srt.mem_cache.chunk_cache import SWAChunkCache enable_hierarchical_cache=enable_hierarchical_cache,
disable_radix_cache=disable_radix_cache,
tree_cache = SWAChunkCache(params) effective_chunked_prefill_size=effective_chunked_prefill_size,
else: tp_worker=tp_worker,
if envs.SGLANG_EXPERIMENTAL_CPP_RADIX_TREE.get(): model_config=model_config,
# lazy import to avoid JIT overhead tp_size=ps.tp_size,
from sglang.srt.mem_cache.radix_cache_cpp import RadixCacheCpp tp_rank=ps.tp_rank,
tp_group=tp_group,
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)
embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get() embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get()
init_mm_embedding_cache(embedding_cache_size * 1024 * 1024) init_mm_embedding_cache(embedding_cache_size * 1024 * 1024)
+194
View File
@@ -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 <name>` (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
+14
View File
@@ -536,6 +536,10 @@ class ServerArgs:
prefill_attention_backend: Optional[str] = None prefill_attention_backend: Optional[str] = None
sampling_backend: Optional[str] = None sampling_backend: Optional[str] = None
grammar_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 mm_attention_backend: Optional[str] = None
fp8_gemm_runner_backend: str = "auto" fp8_gemm_runner_backend: str = "auto"
fp4_gemm_runner_backend: str = "auto" fp4_gemm_runner_backend: str = "auto"
@@ -5373,6 +5377,16 @@ class ServerArgs:
default=ServerArgs.grammar_backend, default=ServerArgs.grammar_backend,
help="Choose the backend for grammar-guided decoding.", 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( parser.add_argument(
"--mm-attention-backend", "--mm-attention-backend",
type=str, type=str,
@@ -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()