[FEAT] Decouple multimodal global cache from Mooncake (#30392)

Co-authored-by: Yuang Chen <cya539102@antgroup.com>
Co-authored-by: Yuang Chen <1131578721@qq.com>
This commit is contained in:
siyu
2026-08-10 19:23:19 +08:00
committed by GitHub
co-authored by Yuang Chen Yuang Chen
parent d96b1533ea
commit 14ffd447a4
6 changed files with 153 additions and 12 deletions
@@ -387,14 +387,19 @@ class MMEncoder:
).element_size()
if get_mm().enable_mm_global_cache:
from sglang.srt.mem_cache.storage.mooncake_store.embedding_cache_controller import (
from sglang.srt.mem_cache.embedding_cache_controller import (
EmbeddingCacheController,
)
from sglang.srt.mem_cache.embedding_store import EmbeddingStoreFactory
embedding_store = EmbeddingStoreFactory.create_backend(
get_mm().mm_global_cache_backend,
)
hidden_dims = self._infer_embedding_dims()
self.mm_global_cache = EmbeddingCacheController(
rank,
server_args.tp_size,
embedding_store=embedding_store,
hidden_dims=hidden_dims,
tp_group=get_tp_group().cpu_group,
all_rank_get=False,
@@ -12,9 +12,7 @@ from typing import List, Optional, Tuple
import torch
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_embedding_store import (
MooncakeEmbeddingStore,
)
from sglang.srt.mem_cache.embedding_store import EmbeddingStore
logger = logging.getLogger(__name__)
@@ -308,6 +306,7 @@ class EmbeddingCacheController:
self,
tp_rank,
tp_size,
embedding_store: EmbeddingStore,
max_pool_size_gb=4.0,
hidden_dims: dict = None,
tp_group=None,
@@ -329,7 +328,7 @@ class EmbeddingCacheController:
self.enable_eviction = enable_eviction
self.max_eviction_batch = max_eviction_batch
self.mooncake_store = MooncakeEmbeddingStore()
self.embedding_store = embedding_store
self.total_pool_size_bytes = int(max_pool_size_gb * 1024**3)
self.vision_pool, self.audio_pool = self._create_pools(pin_memory=True)
self.pools = {
@@ -401,7 +400,7 @@ class EmbeddingCacheController:
f"dim={pool.dim}, budget={pool.pool_size_bytes} bytes"
)
return
self.mooncake_store.register_buffer(pool.tensor)
self.embedding_store.register_buffer(pool.tensor)
logger.info(
f"[Rank {self.tp_rank}] Registered {pool.modality} embedding pool: "
f"dim={pool.dim}, pages={pool.num_pages}, "
@@ -660,7 +659,7 @@ class EmbeddingCacheController:
try:
op = self.prefetch_queue.get_nowait()
try:
results = self.mooncake_store.batch_get_into_multi_buffers(
results = self.embedding_store.batch_get_into_multi_buffers(
op.keys, op.ptrs, op.sizes
)
except Exception:
@@ -680,7 +679,7 @@ class EmbeddingCacheController:
try:
op = self.insert_queue.get_nowait()
try:
results = self.mooncake_store.batch_put_from_multi_buffers(
results = self.embedding_store.batch_put_from_multi_buffers(
op.keys, op.ptrs, op.sizes
)
except Exception:
@@ -1008,7 +1007,7 @@ class EmbeddingCacheController:
missing_hashes = [mm_hashes[i] for i in missing_indices]
global_exists = await asyncio.to_thread(
self.mooncake_store.batch_is_exist, missing_hashes
self.embedding_store.batch_is_exist, missing_hashes
)
global_hit_count = sum(global_exists)
@@ -0,0 +1,127 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to SGLang project
import abc
import importlib
import logging
from typing import Any, Dict, List
import torch
logger = logging.getLogger(__name__)
class EmbeddingStore(abc.ABC):
"""Abstract base class for multimodal embedding storage backends.
Stores pre-computed vision/audio embeddings by content hash so they
can be shared across nodes without re-running the encoder.
"""
@abc.abstractmethod
def batch_get(
self, hashes: List[str], ptrs: List[int], sizes: List[int]
) -> List[bool]:
raise NotImplementedError
@abc.abstractmethod
def batch_put(
self, hashes: List[str], ptrs: List[int], sizes: List[int]
) -> List[bool]:
raise NotImplementedError
@abc.abstractmethod
def batch_get_into_multi_buffers(
self,
hashes: List[str],
ptrs: List[List[int]],
sizes: List[List[int]],
) -> List[bool]:
raise NotImplementedError
@abc.abstractmethod
def batch_put_from_multi_buffers(
self,
hashes: List[str],
ptrs: List[List[int]],
sizes: List[List[int]],
) -> List[bool]:
raise NotImplementedError
@abc.abstractmethod
def batch_is_exist(self, hashes: List[str]) -> List[bool]:
raise NotImplementedError
def register_buffer(self, tensor: torch.Tensor) -> None:
pass
def get_key(self, mm_hash: str) -> str:
return f"emb_{mm_hash}"
class EmbeddingStoreFactory:
"""Factory for creating embedding store backend instances."""
_registry: Dict[str, Dict[str, Any]] = {}
@staticmethod
def _load_backend_class(
module_path: str, class_name: str, backend_name: str
) -> type:
try:
module = importlib.import_module(module_path)
backend_class = getattr(module, class_name)
if not issubclass(backend_class, EmbeddingStore):
raise TypeError(
f"Backend class {class_name} must inherit from EmbeddingStore"
)
return backend_class
except ImportError as e:
raise ImportError(
f"Failed to import embedding store backend '{backend_name}' "
f"from '{module_path}': {e}"
) from e
except AttributeError as e:
raise AttributeError(
f"Class '{class_name}' not found in module '{module_path}': {e}"
) from e
@classmethod
def register_backend(cls, name: str, module_path: str, class_name: str) -> None:
if name in cls._registry:
logger.warning(
f"Embedding store backend '{name}' is already registered, overwriting"
)
def loader() -> type:
return cls._load_backend_class(module_path, class_name, name)
cls._registry[name] = {
"loader": loader,
"module_path": module_path,
"class_name": class_name,
}
@classmethod
def create_backend(cls, backend_name: str, **kwargs) -> EmbeddingStore:
if backend_name not in cls._registry:
available = list(cls._registry.keys())
raise ValueError(
f"Unknown embedding store backend '{backend_name}'. "
f"Registered backends: {available}."
)
entry = cls._registry[backend_name]
backend_class = entry["loader"]()
logger.info(
f"Creating embedding store backend '{backend_name}' "
f"({entry['module_path']}.{entry['class_name']})"
)
return backend_class(**kwargs)
EmbeddingStoreFactory.register_backend(
"mooncake",
"sglang.srt.mem_cache.storage.mooncake_store.mooncake_embedding_store",
"MooncakeEmbeddingStore",
)
@@ -1,6 +1,7 @@
import logging
from typing import Any, List
from sglang.srt.mem_cache.embedding_store import EmbeddingStore
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import (
DEFAULT_TENANT_ID,
MooncakeBaseStore,
@@ -9,7 +10,7 @@ from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import (
logger = logging.getLogger(__name__)
class MooncakeEmbeddingStore(MooncakeBaseStore):
class MooncakeEmbeddingStore(MooncakeBaseStore, EmbeddingStore):
def __init__(
self,
storage_config: Any = None,
+9
View File
@@ -2765,6 +2765,15 @@ class ServerArgs:
"Enable global multimodal embedding cache to skip redundant ViT inference.",
NS("mm"),
] = False
mm_global_cache_backend: A[
str,
Arg(
help="Storage backend for the multimodal global embedding cache. "
"Used when --enable-mm-global-cache is set.",
choices=["mooncake"],
),
NS("mm"),
] = "mooncake"
disable_fast_image_processor: A[
bool, "Adopt base image processor instead of fast image processor.", NS("mm")
] = False