diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 03e03bf4f..1e913091b 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -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, diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/embedding_cache_controller.py b/python/sglang/srt/mem_cache/embedding_cache_controller.py similarity index 98% rename from python/sglang/srt/mem_cache/storage/mooncake_store/embedding_cache_controller.py rename to python/sglang/srt/mem_cache/embedding_cache_controller.py index 938b092b2..88abf1437 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/embedding_cache_controller.py +++ b/python/sglang/srt/mem_cache/embedding_cache_controller.py @@ -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) diff --git a/python/sglang/srt/mem_cache/embedding_store.py b/python/sglang/srt/mem_cache/embedding_store.py new file mode 100644 index 000000000..363692fc1 --- /dev/null +++ b/python/sglang/srt/mem_cache/embedding_store.py @@ -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", +) diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_embedding_store.py b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_embedding_store.py index 79d007760..b68ff7b1b 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_embedding_store.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_embedding_store.py @@ -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, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 3cd4f8cec..27793f396 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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 diff --git a/test/registered/unit/mem_cache/test_embedding_cache_controller.py b/test/registered/unit/mem_cache/test_embedding_cache_controller.py index 77926b5cc..557c31af1 100644 --- a/test/registered/unit/mem_cache/test_embedding_cache_controller.py +++ b/test/registered/unit/mem_cache/test_embedding_cache_controller.py @@ -8,7 +8,7 @@ from unittest.mock import MagicMock import torch from sglang.srt.managers.schedule_batch import Modality -from sglang.srt.mem_cache.storage.mooncake_store.embedding_cache_controller import ( +from sglang.srt.mem_cache.embedding_cache_controller import ( EmbeddingCacheController, EmbeddingCacheEntry, EmbeddingPool, @@ -55,7 +55,7 @@ def _make_controller(num_pages=16, dim=4, page_size=2, enable_eviction=True): ctrl.element_size = torch.float32.itemsize ctrl.enable_eviction = enable_eviction ctrl.max_eviction_batch = 10 - ctrl.mooncake_store = MagicMock() + ctrl.embedding_store = MagicMock() ctrl.total_pool_size_bytes = num_pages * page_size * dim * torch.float32.itemsize ctrl.vision_pool = _make_pool(num_pages, dim, page_size) ctrl.audio_pool = _make_pool(num_pages, dim, page_size, modality="audio")