[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()
|
).element_size()
|
||||||
|
|
||||||
if get_mm().enable_mm_global_cache:
|
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,
|
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()
|
hidden_dims = self._infer_embedding_dims()
|
||||||
self.mm_global_cache = EmbeddingCacheController(
|
self.mm_global_cache = EmbeddingCacheController(
|
||||||
rank,
|
rank,
|
||||||
server_args.tp_size,
|
server_args.tp_size,
|
||||||
|
embedding_store=embedding_store,
|
||||||
hidden_dims=hidden_dims,
|
hidden_dims=hidden_dims,
|
||||||
tp_group=get_tp_group().cpu_group,
|
tp_group=get_tp_group().cpu_group,
|
||||||
all_rank_get=False,
|
all_rank_get=False,
|
||||||
|
|||||||
+7
-8
@@ -12,9 +12,7 @@ from typing import List, Optional, Tuple
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.managers.schedule_batch import Modality
|
from sglang.srt.managers.schedule_batch import Modality
|
||||||
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_embedding_store import (
|
from sglang.srt.mem_cache.embedding_store import EmbeddingStore
|
||||||
MooncakeEmbeddingStore,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -308,6 +306,7 @@ class EmbeddingCacheController:
|
|||||||
self,
|
self,
|
||||||
tp_rank,
|
tp_rank,
|
||||||
tp_size,
|
tp_size,
|
||||||
|
embedding_store: EmbeddingStore,
|
||||||
max_pool_size_gb=4.0,
|
max_pool_size_gb=4.0,
|
||||||
hidden_dims: dict = None,
|
hidden_dims: dict = None,
|
||||||
tp_group=None,
|
tp_group=None,
|
||||||
@@ -329,7 +328,7 @@ class EmbeddingCacheController:
|
|||||||
self.enable_eviction = enable_eviction
|
self.enable_eviction = enable_eviction
|
||||||
self.max_eviction_batch = max_eviction_batch
|
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.total_pool_size_bytes = int(max_pool_size_gb * 1024**3)
|
||||||
self.vision_pool, self.audio_pool = self._create_pools(pin_memory=True)
|
self.vision_pool, self.audio_pool = self._create_pools(pin_memory=True)
|
||||||
self.pools = {
|
self.pools = {
|
||||||
@@ -401,7 +400,7 @@ class EmbeddingCacheController:
|
|||||||
f"dim={pool.dim}, budget={pool.pool_size_bytes} bytes"
|
f"dim={pool.dim}, budget={pool.pool_size_bytes} bytes"
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
self.mooncake_store.register_buffer(pool.tensor)
|
self.embedding_store.register_buffer(pool.tensor)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"[Rank {self.tp_rank}] Registered {pool.modality} embedding pool: "
|
f"[Rank {self.tp_rank}] Registered {pool.modality} embedding pool: "
|
||||||
f"dim={pool.dim}, pages={pool.num_pages}, "
|
f"dim={pool.dim}, pages={pool.num_pages}, "
|
||||||
@@ -660,7 +659,7 @@ class EmbeddingCacheController:
|
|||||||
try:
|
try:
|
||||||
op = self.prefetch_queue.get_nowait()
|
op = self.prefetch_queue.get_nowait()
|
||||||
try:
|
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
|
op.keys, op.ptrs, op.sizes
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -680,7 +679,7 @@ class EmbeddingCacheController:
|
|||||||
try:
|
try:
|
||||||
op = self.insert_queue.get_nowait()
|
op = self.insert_queue.get_nowait()
|
||||||
try:
|
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
|
op.keys, op.ptrs, op.sizes
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -1008,7 +1007,7 @@ class EmbeddingCacheController:
|
|||||||
missing_hashes = [mm_hashes[i] for i in missing_indices]
|
missing_hashes = [mm_hashes[i] for i in missing_indices]
|
||||||
|
|
||||||
global_exists = await asyncio.to_thread(
|
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)
|
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
|
import logging
|
||||||
from typing import Any, List
|
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 (
|
from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import (
|
||||||
DEFAULT_TENANT_ID,
|
DEFAULT_TENANT_ID,
|
||||||
MooncakeBaseStore,
|
MooncakeBaseStore,
|
||||||
@@ -9,7 +10,7 @@ from sglang.srt.mem_cache.storage.mooncake_store.mooncake_store import (
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class MooncakeEmbeddingStore(MooncakeBaseStore):
|
class MooncakeEmbeddingStore(MooncakeBaseStore, EmbeddingStore):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
storage_config: Any = None,
|
storage_config: Any = None,
|
||||||
|
|||||||
@@ -2765,6 +2765,15 @@ class ServerArgs:
|
|||||||
"Enable global multimodal embedding cache to skip redundant ViT inference.",
|
"Enable global multimodal embedding cache to skip redundant ViT inference.",
|
||||||
NS("mm"),
|
NS("mm"),
|
||||||
] = False
|
] = 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[
|
disable_fast_image_processor: A[
|
||||||
bool, "Adopt base image processor instead of fast image processor.", NS("mm")
|
bool, "Adopt base image processor instead of fast image processor.", NS("mm")
|
||||||
] = False
|
] = False
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from unittest.mock import MagicMock
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.managers.schedule_batch import Modality
|
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,
|
EmbeddingCacheController,
|
||||||
EmbeddingCacheEntry,
|
EmbeddingCacheEntry,
|
||||||
EmbeddingPool,
|
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.element_size = torch.float32.itemsize
|
||||||
ctrl.enable_eviction = enable_eviction
|
ctrl.enable_eviction = enable_eviction
|
||||||
ctrl.max_eviction_batch = 10
|
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.total_pool_size_bytes = num_pages * page_size * dim * torch.float32.itemsize
|
||||||
ctrl.vision_pool = _make_pool(num_pages, dim, page_size)
|
ctrl.vision_pool = _make_pool(num_pages, dim, page_size)
|
||||||
ctrl.audio_pool = _make_pool(num_pages, dim, page_size, modality="audio")
|
ctrl.audio_pool = _make_pool(num_pages, dim, page_size, modality="audio")
|
||||||
|
|||||||
Reference in New Issue
Block a user