[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,
|
||||
|
||||
Reference in New Issue
Block a user