[radix cache] pluggable RadixCache factory (--radix-cache-backend) (#25101)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user