[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:
co-authored by
Yuang Chen
Yuang Chen
parent
d96b1533ea
commit
14ffd447a4
@@ -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,
|
||||
|
||||
+7
-8
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user