vlm: cache kimi-k3 per-image processor artifacts (#34404)

This commit is contained in:
Mick
2026-08-16 19:51:13 +08:00
committed by GitHub
parent 968b355f12
commit 3d3194f6c3
20 changed files with 2617 additions and 276 deletions
@@ -33,6 +33,7 @@ from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReqInput
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import Modality, Req
from sglang.srt.multimodal.cache import media_preprocess_kwargs
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData
from sglang.srt.utils.common import safe_pickle_loads
@@ -725,6 +726,16 @@ def extract_original_req_id(part_req_id: str) -> str:
return part_req_id
def _encoder_media_item(mm_item: dict):
"""Keep per-media options aligned while preserving the legacy URL shape."""
item = {
key: value
for key, value in mm_item.items()
if key != "modality" and value is not None
}
return item["url"] if set(item) == {"url"} else item
def calculate_modality_num_parts(modalities, num_items_assigned):
"""
Calculate total number of parts and number of parts per modality.
@@ -1102,7 +1113,7 @@ class WaitingImageRDMARequest(WaitingImageRequest):
{
"encoder_idx": idx,
"mm_items": [
d["url"]
_encoder_media_item(d)
for d in mm_data_modality[
cum_num_items : cum_num_items + assigned_num
]
@@ -2171,7 +2182,7 @@ class MMReceiverBase(ABC):
return num_items_assigned
def _extract_url_data(self, request_obj) -> List[Dict]:
def _extract_url_data(self, request_obj: GenerateReqInput) -> List[Dict]:
def flatten_mm_items(items):
if not isinstance(items, list):
return [items]
@@ -2193,20 +2204,46 @@ class MMReceiverBase(ABC):
return mm_item
mm_data = []
for attr, modality in [
("image_data", Modality.IMAGE),
("video_data", Modality.VIDEO),
("audio_data", Modality.AUDIO),
image_hashes = request_obj.mm_content_hashes
image_index = 0
for mm_items, modality in [
(request_obj.image_data, Modality.IMAGE),
(request_obj.video_data, Modality.VIDEO),
(request_obj.audio_data, Modality.AUDIO),
]:
mm_items = getattr(request_obj, attr, None)
if mm_items:
mm_items = flatten_mm_items(mm_items)
for mm_item in mm_items:
mm_data.append(
{
entry = {
"url": to_raw_url(mm_item),
"modality": modality,
}
entry.update(
media_preprocess_kwargs(mm_item, defaults={"detail": "auto"})
)
if modality == Modality.IMAGE:
inline_hash = (
mm_item.content_hash
if isinstance(mm_item, ImageData)
else (
mm_item.get("content_hash")
if isinstance(mm_item, dict)
else None
)
)
explicit_hash = (
image_hashes[image_index]
if image_hashes is not None
and image_index < len(image_hashes)
else None
)
entry["content_hash"] = explicit_hash or inline_hash
image_index += 1
mm_data.append(entry)
if image_hashes is not None and image_index != len(image_hashes):
raise ValueError(
f"mm_content_hashes has {len(image_hashes)} entries for "
f"{image_index} images"
)
return mm_data
@@ -2329,7 +2366,7 @@ class MMReceiverHTTP(MMReceiverBase):
"encoder_idx": idx,
"encoder_url": effective_urls[idx],
"mm_items": [
mm_item.get("url")
_encoder_media_item(mm_item)
for mm_item in mm_data_modality[
cum_num_items : cum_num_items + assigned_num
]
@@ -58,10 +58,12 @@ from sglang.srt.model_executor.model_runner_components.load_model_utils import (
maybe_precompile_model_kernels_after_loading,
)
from sglang.srt.model_loader import get_model
from sglang.srt.multimodal.cache import parse_content_hash, snapshot_media
from sglang.srt.multimodal.encoder_preprocessing import (
EncoderPreprocessOutput,
get_encoder_preprocessed_items,
invoke_encoder_preprocessor,
resolve_encoder_media_processor_config,
)
from sglang.srt.multimodal.processors.qwen_vl import preprocess_video
from sglang.srt.observability.metrics_collector import EncoderMetricsCollector
@@ -369,6 +371,9 @@ class MMEncoder:
load_config=self.load_config,
device_config=self.device_config,
)
self.encoder_media_processor_config = resolve_encoder_media_processor_config(
self.model
)
maybe_precompile_model_kernels_after_loading(self.model, self.device)
self.context = zmq.asyncio.Context(2)
@@ -670,13 +675,27 @@ class MMEncoder:
Load a single multimodal data.
If data is precomputed, returns directly.
Static method that can be pickled for multiprocessing"""
media_metadata = {}
content_hash = None
if isinstance(data, dict):
if "url" not in data:
return data
media_metadata = {key: value for key, value in data.items() if key != "url"}
content_hash = parse_content_hash(data.get("content_hash"))
data = data["url"]
try:
if modality == Modality.IMAGE:
if content_hash is not None:
snapshot = snapshot_media(data)
if snapshot.content_digest != content_hash:
raise BadRequestError(
"Encoder media content hash mismatch: "
f"expected {content_hash}, got {snapshot.content_digest}"
)
data = snapshot.data
gpu_image_decode = (
"nvjpeg_fancy"
if self.use_image_processor_gpu and self.model_type == "kimi_k3"
self.encoder_media_processor_config.image_decode_mode
if self.use_image_processor_gpu
else False
)
img, _ = load_image(data, gpu_image_decode)
@@ -687,12 +706,23 @@ class MMEncoder:
):
# Needed only when `img` is a PIL image
img = img.convert("RGB")
if (
media_metadata
and self.encoder_media_processor_config.preserve_media_metadata
):
return {
"type": "image",
"image": img,
**media_metadata,
}
return img
elif modality == Modality.VIDEO:
return load_video(data, frame_count_limit)
elif modality == Modality.AUDIO:
return load_audio(data, self.model_audio_sr)
except MMError:
raise
except CLIENT_MEDIA_EXCEPTIONS as e:
# Not ValueError: the DP envelope classifies by `.code`, which only MMError carries.
raise BadRequestError(f"Error while loading data {data}: {e}") from e
+6 -14
View File
@@ -378,21 +378,13 @@ class MultimodalDataItem:
if self.pad_value is not None:
return
from sglang.srt.managers.mm_utils import hash_feature
from sglang.srt.multimodal.cache import resolve_multimodal_item_hash
if envs.SGLANG_MM_SKIP_COMPUTE_HASH.get():
import uuid
self.hash = uuid.uuid4().int
self.pad_value = _compute_pad_value(self.hash)
return
if self.hash is None:
if self.feature is not None:
hashed_feature = self.feature
else:
hashed_feature = self.precomputed_embeddings
self.hash = hash_feature(hashed_feature)
assert self.hash is not None
self.hash = resolve_multimodal_item_hash(
existing_hash=self.hash,
feature=self.feature,
precomputed_embeddings=self.precomputed_embeddings,
)
self.pad_value = _compute_pad_value(self.hash)
def is_modality(self, modality: Modality) -> bool:
+51 -16
View File
@@ -106,6 +106,7 @@ from sglang.srt.models.kimi_k3_vl import (
)
from sglang.srt.models.transformers import maybe_prefix
from sglang.srt.models.utils import WeightsMapper
from sglang.srt.multimodal.encoder_preprocessing import EncoderMediaProcessorConfig
from sglang.srt.multimodal.kimi_k3_image_processing import (
DEFERRED_PREPROCESSING_KEY,
fill_transparent_bg,
@@ -3075,6 +3076,10 @@ class KimiK3ForConditionalGeneration(nn.Module):
"""K3 multimodal wrapper: MoonViT3d tower + KimiK3LinearForCausalLM."""
supports_cuda_vmm_feature_transport = True
encoder_media_processor_config = EncoderMediaProcessorConfig(
image_decode_mode="nvjpeg_fancy",
preserve_media_metadata=True,
)
# Raw HF checkpoint prefixes, before hf_to_sglang_mapper is applied.
encoder_only_safetensors_weight_prefixes = (
@@ -3259,27 +3264,36 @@ class KimiK3ForConditionalGeneration(nn.Module):
for item in selected_items
]
if any(config is not None for config in deferred):
if not all(config is not None for config in deferred):
raise ValueError(
"Kimi-K3 cannot mix deferred and preprocessed image features"
)
first_config = deferred[0]
backend = first_config.backend
if any(config.backend != backend for config in deferred):
raise ValueError(
"Kimi-K3 cannot mix deferred preprocessing backends"
materialized = [None] * len(selected_items)
deferred_by_backend = {}
for index, (item, config) in enumerate(zip(selected_items, deferred)):
if config is None:
if not isinstance(item.feature, torch.Tensor):
raise TypeError(
"Kimi-K3 image feature must be a torch.Tensor, "
f"got {type(item.feature)}"
)
materialized[index] = item.feature
else:
deferred_by_backend.setdefault(config.backend, []).append(index)
for backend, indices in deferred_by_backend.items():
group_items = [selected_items[index] for index in indices]
group_configs = [deferred[index] for index in indices]
first_config = group_configs[0]
if backend == "gpu":
from sglang.srt.multimodal.processors.kimi_k25 import (
_gpu_preprocess_images,
)
image_scale, image_bias = normalization_tensors(
first_config.image_mean, first_config.image_std, device
first_config.image_mean,
first_config.image_std,
device,
)
pixel_values, _ = _gpu_preprocess_images(
[item.feature for item in selected_items],
[config.resize_config for config in deferred],
pixel_values, produced_grids = _gpu_preprocess_images(
[item.feature for item in group_items],
[config.resize_config for config in group_configs],
image_scale,
image_bias,
self.vision_tower.patch_size,
@@ -3288,20 +3302,41 @@ class KimiK3ForConditionalGeneration(nn.Module):
x, first_config.transparent_bg_config
),
)
expected_grids = grid_thws_host[indices]
if not torch.equal(produced_grids.cpu(), expected_grids):
raise ValueError(
"Kimi-K3 deferred GPU preprocessing produced wrong grids"
)
elif backend == "cpu":
from sglang.srt.multimodal.kimi_k3_image_processing import (
materialize_kimi_k3_cpu_features,
)
pixel_values = materialize_kimi_k3_cpu_features(
selected_items, self._encoder_image_processor
group_items, self._encoder_image_processor
)
pixel_values = pixel_values.to(device, non_blocking=True)
else:
raise ValueError(
f"Unsupported Kimi-K3 deferred preprocessing backend: {backend}"
)
return pixel_values.to(dtype=target_dtype)
patch_counts = [
int(grid_thws_host[index].prod().item()) for index in indices
]
if sum(patch_counts) != pixel_values.shape[0]:
raise ValueError(
"Kimi-K3 deferred feature length does not match image grids"
)
for index, feature in zip(
indices, pixel_values.split(patch_counts), strict=True
):
materialized[index] = feature
return materialize_multimodal_features(
materialized,
device=device,
dtype=target_dtype,
)
features = []
for item in selected_items:
+10 -2
View File
@@ -3,27 +3,35 @@
from sglang.srt.multimodal.cache.identity import (
CONTENT_HASH_PREFIX,
MediaSnapshot,
PreprocessFingerprintProvider,
build_artifact_key,
build_feature_hash,
build_processor_fingerprint,
media_preprocess_kwargs,
parse_content_hash,
resolve_multimodal_item_hash,
snapshot_media,
)
from sglang.srt.multimodal.cache.preprocess_cache import (
CacheLookup,
CacheMiss,
CacheSizeProvider,
MultimodalPreprocessCache,
estimate_cache_size_bytes,
)
__all__ = [
"CONTENT_HASH_PREFIX",
"CacheSizeProvider",
"CacheLookup",
"CacheMiss",
"MediaSnapshot",
"MultimodalPreprocessCache",
"PreprocessFingerprintProvider",
"build_artifact_key",
"build_feature_hash",
"build_processor_fingerprint",
"estimate_cache_size_bytes",
"media_preprocess_kwargs",
"parse_content_hash",
"resolve_multimodal_item_hash",
"snapshot_media",
]
+104 -28
View File
@@ -9,7 +9,7 @@ import struct
from dataclasses import dataclass
from enum import Enum
from pathlib import Path
from typing import Any, Mapping, Optional
from typing import TYPE_CHECKING, Any, Mapping, Optional, Protocol, runtime_checkable
from urllib.parse import unquote, urlparse
import numpy as np
@@ -17,8 +17,21 @@ import torch
import transformers
from PIL import Image
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
CONTENT_HASH_PREFIX = "sha256:"
_SHA256_HEX_LENGTH = 64
_MEDIA_ENVELOPE_FIELDS = frozenset(
{"type", "format", "url", "image", "video", "audio", "content_hash"}
)
@runtime_checkable
class PreprocessFingerprintProvider(Protocol):
"""Explicit source for settings that can change processor artifacts."""
def preprocess_fingerprint_payload(self) -> Any: ...
def parse_content_hash(value: Optional[str]) -> Optional[str]:
@@ -157,6 +170,43 @@ def snapshot_media(media: Any) -> MediaSnapshot:
raise TypeError(f"Unsupported media identity input: {type(media).__name__}")
def media_preprocess_kwargs(
source: Any, *, defaults: Optional[Mapping[str, Any]] = None
) -> dict[str, Any]:
"""Conservatively capture per-request options that can affect an artifact.
Unknown options are included instead of allow-listed. This may create a safe
false miss for a metadata-only option, but it prevents a new model option
from silently creating a false cache hit.
"""
defaults = defaults or {}
if dataclasses.is_dataclass(source):
values = {
field.name: value
for field, value in zip(
dataclasses.fields(source), dataclasses.astuple(source)
)
if field.name not in _MEDIA_ENVELOPE_FIELDS
}
elif isinstance(source, Mapping):
values = {
key: value
for key, value in source.items()
if key not in _MEDIA_ENVELOPE_FIELDS
}
else:
return {}
result = {}
for key, value in values.items():
if value is None or (isinstance(value, Mapping) and not value):
continue
if key in defaults and _canonicalize(value) == _canonicalize(defaults[key]):
continue
result[key] = value
return result
def _qualified_type_name(value: Any) -> str:
value_type = type(value)
return f"{value_type.__module__}.{value_type.__qualname__}"
@@ -173,8 +223,10 @@ def _canonicalize(value: Any) -> Any:
"type": "dataclass",
"class": _qualified_type_name(value),
"fields": [
[field.name, _canonicalize(getattr(value, field.name))]
for field in dataclasses.fields(value)
[field.name, _canonicalize(field_value)]
for field, field_value in zip(
dataclasses.fields(value), dataclasses.astuple(value)
)
],
}
if isinstance(value, Enum):
@@ -273,24 +325,49 @@ def build_artifact_key(
return _digest_bytes(_canonical_json(payload))
def build_feature_hash(artifact_key: str, processor_output_hash: int) -> int:
"""Namespace a processor-output hash by its complete artifact identity."""
artifact_key = parse_content_hash(artifact_key)
if (
isinstance(processor_output_hash, bool)
or not isinstance(processor_output_hash, int)
or processor_output_hash < 0
):
raise ValueError("processor_output_hash must be a non-negative integer")
output_hash_bytes = processor_output_hash.to_bytes(
max(1, (processor_output_hash.bit_length() + 7) // 8),
byteorder="big",
signed=False,
def resolve_multimodal_item_hash(
*,
existing_hash: Optional[int] = None,
feature: Any = None,
precomputed_embeddings: Any = None,
namespace: Optional[str] = None,
) -> int:
"""Unified helper for resolving a hash for MultimodalDataItem cache, optionally scoped to an artifact identity.
Args:
namespace: Optional SHA-256 identity covering every input that can change the preprocessing result.
It scopes the feature hash so downstream caches cannot reuse embeddings across different preprocessing settings.
"""
from sglang.srt.environ import envs
if envs.SGLANG_MM_SKIP_COMPUTE_HASH.get():
import uuid
item_hash = uuid.uuid4().int
elif existing_hash is not None:
# if exists, reuse
item_hash = existing_hash
else:
# hash from feature
from sglang.srt.managers.mm_utils import hash_feature
value = feature if feature is not None else precomputed_embeddings
item_hash = hash_feature(value)
if namespace is None:
return item_hash
if isinstance(item_hash, bool) or not isinstance(item_hash, int) or item_hash < 0:
raise ValueError("item hash must be a non-negative integer")
namespace = parse_content_hash(namespace)
assert namespace is not None
hash_bytes = item_hash.to_bytes(
max(1, (item_hash.bit_length() + 7) // 8), byteorder="big", signed=False
)
digest = _hash_parts(
b"multimodal-feature-v1",
bytes.fromhex(artifact_key[len(CONTENT_HASH_PREFIX) :]),
output_hash_bytes,
bytes.fromhex(namespace[len(CONTENT_HASH_PREFIX) :]),
hash_bytes,
)
return int.from_bytes(
bytes.fromhex(digest[len(CONTENT_HASH_PREFIX) :])[:8],
@@ -302,27 +379,26 @@ def build_feature_hash(artifact_key: str, processor_output_hash: int) -> int:
def build_processor_fingerprint(
processor: Any,
hf_config: Any,
server_args: Any,
server_args: ServerArgs,
*,
extra: Optional[Mapping[str, Any]] = None,
) -> str:
"""Fingerprint preprocessing choices that can change processor output."""
processor_payload = (
processor.preprocess_fingerprint_payload()
if hasattr(processor, "preprocess_fingerprint_payload")
if isinstance(processor, PreprocessFingerprintProvider)
else {}
)
hf_payload = hf_config.to_dict()
payload = {
"transformers": transformers.__version__,
"processor_class": f"{type(processor).__module__}.{type(processor).__qualname__}",
"model_type": getattr(hf_config, "model_type", None),
"architectures": getattr(hf_config, "architectures", None),
"model_revision": getattr(server_args, "revision", None),
"tokenizer_revision": getattr(server_args, "tokenizer_revision", None),
"disable_fast_image_processor": getattr(
server_args, "disable_fast_image_processor", False
),
"mm_process_config": getattr(server_args, "mm_process_config", None) or {},
"model_type": hf_payload.get("model_type"),
"architectures": hf_payload.get("architectures"),
"model_revision": server_args.revision,
"processor_revision": server_args.revision,
"disable_fast_image_processor": server_args.disable_fast_image_processor,
"mm_process_config": server_args.mm_process_config or {},
"processor": processor_payload,
"extra": extra or {},
}
+202 -14
View File
@@ -1,15 +1,29 @@
"""Bounded CPU cache and single-flight coordination for MM preprocessing."""
"""Bounded CPU storage and single-flight coordination for MM preprocessing.
Model processors store prompt-independent ``MediaArtifact`` values here. This
module knows nothing about a model or media format: it provides byte-accounted
LRU storage and ensures concurrent misses for one key share one computation.
"""
from __future__ import annotations
import asyncio
import concurrent.futures
import dataclasses
import sys
import threading
from collections import OrderedDict
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Awaitable, Callable, Generic, Optional, TypeVar
from typing import (
Any,
Awaitable,
Callable,
Generic,
Optional,
Protocol,
TypeVar,
runtime_checkable,
)
import numpy as np
import torch
@@ -17,21 +31,45 @@ from PIL import Image
K = TypeVar("K")
V = TypeVar("V")
_USE_RESULT = object()
@dataclass(frozen=True)
class CacheLookup(Generic[V]):
"""A resolved value returned immediately or after shared computation."""
value: V
hit: bool
joined: bool = False
@dataclass(frozen=True)
class CacheMiss(Generic[K, V]):
"""Handle for one in-flight cache miss.
Exactly one handle has ``should_compute=True`` and must publish the result.
Other handles for the same key wait on the shared ``future``.
"""
key: K
future: concurrent.futures.Future[V]
generation: int
should_compute: bool
@dataclass
class _Entry(Generic[V]):
value: V
size_bytes: int
@runtime_checkable
class CacheSizeProvider(Protocol):
"""Explicitly expose the owned values that count against a cache budget."""
def cache_size_items(self) -> Sequence[Any]: ...
def estimate_cache_size_bytes(value: Any) -> Optional[int]:
"""Estimate owned CPU bytes, returning None for GPU-backed artifacts."""
seen: set[int] = set()
@@ -56,9 +94,9 @@ def estimate_cache_size_bytes(value: Any) -> Optional[int]:
return len(item)
if isinstance(item, str):
return len(item.encode())
if dataclasses.is_dataclass(item):
return visit(dataclasses.asdict(item))
if isinstance(item, dict):
if isinstance(item, CacheSizeProvider):
return visit(item.cache_size_items())
if isinstance(item, Mapping):
total = 0
for key, child in item.items():
key_size = visit(key)
@@ -81,7 +119,13 @@ def estimate_cache_size_bytes(value: Any) -> Optional[int]:
class MultimodalPreprocessCache(Generic[K, V]):
"""Thread-safe byte-accounted LRU with per-key async single-flight."""
"""Thread-safe CPU LRU with per-key async single-flight.
``max_size_bytes`` and ``max_entries`` bound retained values. In-flight
computations are tracked separately and are not part of the LRU budget.
``clear()`` invalidates cache writes from old computations so a flush cannot
be undone by work that started before it.
"""
def __init__(self, max_size_bytes: int, max_entries: int = 8192):
if max_size_bytes < 0:
@@ -103,6 +147,7 @@ class MultimodalPreprocessCache(Generic[K, V]):
@property
def enabled(self) -> bool:
"""Whether values can be retained; zero bytes is the cache kill switch."""
return self.max_size_bytes > 0
def __len__(self) -> int:
@@ -114,6 +159,7 @@ class MultimodalPreprocessCache(Generic[K, V]):
return key in self._entries
def get(self, key: K) -> Optional[V]:
"""Read and touch an LRU entry, recording a hit or miss."""
with self._lock:
entry = self._entries.get(key)
if entry is None:
@@ -123,6 +169,31 @@ class MultimodalPreprocessCache(Generic[K, V]):
self.hits += 1
return entry.value
def get_if_present(
self,
key: K,
predicate: Callable[[V], bool],
*,
evict_on_reject: bool = False,
) -> Optional[V]:
"""Use a compatible entry without recording an absent speculative miss.
The predicate runs while holding the cache lock, so a caller cannot use
an entry that another thread replaces between validation and lookup.
"""
with self._lock:
entry = self._entries.get(key)
if entry is None:
return None
if not predicate(entry.value):
if evict_on_reject:
self._entries.pop(key)
self.current_size_bytes -= entry.size_bytes
return None
self._entries.move_to_end(key)
self.hits += 1
return entry.value
def put(
self,
key: K,
@@ -131,6 +202,12 @@ class MultimodalPreprocessCache(Generic[K, V]):
*,
_generation: Optional[int] = None,
) -> bool:
"""Insert a value if it is CPU-sizeable and fits the configured budget.
With automatic sizing, returns ``False`` when caching is disabled, the
value contains a GPU tensor, the value is too large, or its generation
predates ``clear()``.
"""
if not self.enabled:
return False
if size_bytes is None:
@@ -139,6 +216,8 @@ class MultimodalPreprocessCache(Generic[K, V]):
return False
with self._lock:
# A pre-flush computation may finish, but it must not repopulate the
# new cache generation.
if _generation is not None and _generation != self._generation:
return False
old = self._entries.pop(key, None)
@@ -156,6 +235,7 @@ class MultimodalPreprocessCache(Generic[K, V]):
return True
def pop(self, key: K) -> Optional[V]:
"""Remove and return one entry without changing hit/miss counters."""
with self._lock:
entry = self._entries.pop(key, None)
if entry is None:
@@ -164,6 +244,7 @@ class MultimodalPreprocessCache(Generic[K, V]):
return entry.value
def clear(self) -> None:
"""Drop values and prevent older in-flight work from repopulating them."""
with self._lock:
self._entries.clear()
self.current_size_bytes = 0
@@ -178,6 +259,14 @@ class MultimodalPreprocessCache(Generic[K, V]):
*,
size_bytes: Optional[Callable[[V], Optional[int]]] = None,
) -> CacheLookup[V]:
"""Return a cached value or share one async computation for ``key``.
Cancellation affects only the caller that is awaiting the result. The
shared computation remains alive for other callers.
"""
if not self.enabled:
return CacheLookup(await compute(), hit=False)
cached = self.get(key)
if cached is not None:
return CacheLookup(cached, hit=True)
@@ -188,25 +277,25 @@ class MultimodalPreprocessCache(Generic[K, V]):
future = concurrent.futures.Future()
generation = self._generation
self._inflight[key] = (future, generation)
owner = True
should_compute = True
else:
future, generation = inflight
self.singleflight_joins += 1
owner = False
should_compute = False
if owner:
if should_compute:
self.create_background_task(
self._compute_owned_value(
self._compute_shared_value(
key, future, generation, compute, size_bytes=size_bytes
)
)
# The cache owns the shared computation. Cancelling either its first
# caller or a later joiner ends only that caller's local await.
# computing caller or a later waiter ends only that caller's local await.
value = await asyncio.shield(asyncio.wrap_future(future))
return CacheLookup(value, hit=False, joined=not owner)
return CacheLookup(value, hit=False, joined=not should_compute)
async def _compute_owned_value(
async def _compute_shared_value(
self,
key: K,
future: concurrent.futures.Future[V],
@@ -215,6 +304,7 @@ class MultimodalPreprocessCache(Generic[K, V]):
*,
size_bytes: Optional[Callable[[V], Optional[int]]],
) -> None:
"""Compute once, cache the result, and wake every caller for this key."""
try:
value = await compute()
measured = size_bytes(value) if size_bytes is not None else None
@@ -246,7 +336,105 @@ class MultimodalPreprocessCache(Generic[K, V]):
task.add_done_callback(self._background_task_done)
return task
def lookup_or_claim_many(
self,
keys: list[K],
*,
predicate: Optional[Callable[[K, V], bool]] = None,
) -> list[CacheLookup[V] | CacheMiss[K, V]]:
"""Return hits and single-flight miss handles in input-key order.
For each missing key, one result has ``should_compute=True``. Repeated
keys or concurrent callers receive handles with ``should_compute=False``
and should call ``wait_for_miss`` instead of recomputing the value.
"""
results: list[CacheLookup[V] | CacheMiss[K, V]] = []
with self._lock:
for key in keys:
if not self.enabled:
future: concurrent.futures.Future[V] = concurrent.futures.Future()
results.append(
CacheMiss(key, future, self._generation, should_compute=True)
)
continue
entry = self._entries.get(key)
if entry is not None and (
predicate is None or predicate(key, entry.value)
):
self._entries.move_to_end(key)
self.hits += 1
results.append(CacheLookup(entry.value, hit=True))
continue
if entry is not None:
self._entries.pop(key)
self.current_size_bytes -= entry.size_bytes
self.misses += 1
inflight = self._inflight.get(key)
if inflight is None or inflight[1] != self._generation:
future: concurrent.futures.Future[V] = concurrent.futures.Future()
generation = self._generation
self._inflight[key] = (future, generation)
results.append(
CacheMiss(key, future, generation, should_compute=True)
)
else:
future, generation = inflight
self.singleflight_joins += 1
results.append(
CacheMiss(key, future, generation, should_compute=False)
)
return results
def complete_miss(
self,
miss: CacheMiss[K, V],
value: V,
*,
cache_value: V | object = _USE_RESULT,
size_bytes: Optional[int] = None,
) -> None:
"""Publish a computed miss to waiters and optionally retain a copy.
``value`` is returned to current waiters. ``cache_value`` may be a
smaller representation retained for future requests.
"""
if not miss.should_compute:
raise ValueError("Only the caller computing a cache miss can complete it")
self.put(
miss.key,
value if cache_value is _USE_RESULT else cache_value,
size_bytes,
_generation=miss.generation,
)
miss.future.set_result(value)
with self._lock:
if self._inflight.get(miss.key) == (
miss.future,
miss.generation,
):
self._inflight.pop(miss.key, None)
def fail_miss(self, miss: CacheMiss[K, V], error: BaseException) -> None:
"""Publish a computation failure to every waiter for this miss."""
if not miss.should_compute:
raise ValueError("Only the caller computing a cache miss can fail it")
miss.future.set_exception(error)
miss.future.exception()
with self._lock:
if self._inflight.get(miss.key) == (
miss.future,
miss.generation,
):
self._inflight.pop(miss.key, None)
async def wait_for_miss(self, miss: CacheMiss[K, V]) -> V:
"""Wait for another caller's computation without cancelling it."""
return await asyncio.shield(asyncio.wrap_future(miss.future))
def stats(self) -> dict[str, int]:
"""Return a lock-consistent snapshot of cache and single-flight state."""
with self._lock:
return {
"entries": len(self._entries),
@@ -1,7 +1,8 @@
import hashlib
import inspect
from collections.abc import Mapping
from typing import Any, Callable, Iterable, Sequence
from dataclasses import dataclass
from typing import Any, Callable, Iterable, Protocol, Sequence, runtime_checkable
import numpy as np
import torch
@@ -11,6 +12,30 @@ from sglang.srt.managers.schedule_batch import MultimodalDataItem
LOCAL_PREPROCESSED_KEY = "encoder_local_preprocessed"
@dataclass(frozen=True)
class EncoderMediaProcessorConfig:
"""Optional model-declared media loading behavior for encoder mode."""
image_decode_mode: bool | str = False
preserve_media_metadata: bool = False
@runtime_checkable
class EncoderMediaProcessorConfigProvider(Protocol):
"""Model contract for optional encoder-side media preprocessing."""
encoder_media_processor_config: EncoderMediaProcessorConfig
def resolve_encoder_media_processor_config(
model: object,
) -> EncoderMediaProcessorConfig:
"""Resolve a model-declared capability without model-name dispatch."""
if isinstance(model, EncoderMediaProcessorConfigProvider):
return model.encoder_media_processor_config
return EncoderMediaProcessorConfig()
def hash_raw_encoder_item(value: Any) -> int:
"""Hash raw CPU media including layout metadata, before owner materialization."""
if isinstance(value, torch.Tensor):
@@ -46,7 +46,13 @@ def prepare_kimi_k3_encoder_inputs(
navit_resize_config,
)
media_proc_cfg = getattr(image_processor, "media_proc_cfg", None)
try:
media_proc_cfg = image_processor.media_proc_cfg
except AttributeError as exc:
raise ValueError(
"Kimi-K3 EPD owner-side preprocessing requires "
"image_processor.media_proc_cfg"
) from exc
if not isinstance(media_proc_cfg, dict):
raise ValueError(
"Kimi-K3 EPD owner-side preprocessing requires "
@@ -69,12 +75,16 @@ def prepare_kimi_k3_encoder_inputs(
)
concrete_images = []
content_digests = []
for image in images:
content_digest = None
if isinstance(image, dict):
if image.get("type") != "image" or "image" not in image:
raise ValueError(f"Unsupported Kimi-K3 encoder media item: {image}")
content_digest = image.get("content_hash")
image = image["image"]
concrete_images.append(image)
content_digests.append(content_digest)
patch_size = int(media_proc_cfg["patch_size"])
merge_kernel_size = int(media_proc_cfg["merge_kernel_size"])
@@ -89,7 +99,7 @@ def prepare_kimi_k3_encoder_inputs(
items = []
grids = []
original_image_sizes = []
for image in concrete_images:
for image, content_digest in zip(concrete_images, content_digests):
width, height = (
(int(image.shape[-1]), int(image.shape[-2]))
if isinstance(image, torch.Tensor)
@@ -106,15 +116,18 @@ def prepare_kimi_k3_encoder_inputs(
)
grid_thw = _grid_thw_from_resize_config(resize_config, patch_size)
grid_tensor = torch.tensor([grid_thw], dtype=torch.int64)
item = MultimodalDataItem(
modality=Modality.IMAGE,
feature=to_chw_uint8(image) if use_gpu_preprocessing else image,
model_specific_data={
model_specific_data = {
"grid_thws": grid_tensor,
DEFERRED_PREPROCESSING_KEY: deferred_preprocessing(
resize_config=resize_config
),
},
}
if content_digest is not None:
model_specific_data["content_digest"] = content_digest
item = MultimodalDataItem(
modality=Modality.IMAGE,
feature=to_chw_uint8(image) if use_gpu_preprocessing else image,
model_specific_data=model_specific_data,
)
if not use_gpu_preprocessing:
item.set_hash(hash_raw_encoder_item(image))
@@ -0,0 +1,31 @@
"""Bridge multimodal processors and the shared preprocess cache.
A model processor turns raw media into model-specific ``MediaArtifact``
objects and later combines those artifacts with the current prompt. The
``MultimodalPreprocessCache`` stores cache-safe copies of the artifacts so a
later request can skip model preprocessing. With a trusted caller-provided
content hash, a hot hit can skip loading the media source as well.
This package defines that boundary:
* ``MediaArtifact`` is the common contract for model-specific cache items.
* ``MediaArtifactInput`` carries a decoded cache miss into a model processor.
* ``MediaArtifactCacheMixin`` owns lookup, single-flight miss handling, and
result ordering around the model's ``prepare_artifact_batch`` method.
The package does not implement model preprocessing or a second cache. Model
logic stays in ``multimodal.processors``; storage and concurrency stay in
``multimodal.cache``.
"""
from sglang.srt.multimodal.media_artifacts.base import (
MediaArtifact,
MediaArtifactCacheMixin,
MediaArtifactInput,
)
__all__ = [
"MediaArtifact",
"MediaArtifactCacheMixin",
"MediaArtifactInput",
]
@@ -0,0 +1,485 @@
"""Shared contracts and coordination for reusable multimodal artifacts.
A media artifact is the model-specific, prompt-independent state produced from
one media input. It keeps the metadata needed to rebuild a request (for example,
image size, token count, and encoder grid) and, when cacheable on CPU, the
processor feature itself. Prompt tokens and offsets are deliberately excluded.
The artifact connects the media preprocessor to request composition::
raw media -> identity/cache lookup -> MediaArtifact
cache miss: MediaArtifactInput -> prepare_artifact_batch()
cache hit: reuse the stored artifact
MediaArtifact + current prompt -> MultimodalDataItem -> encoder/ViT
The current request uses the full artifact returned by preprocessing. The
preprocess cache stores ``artifact.cache_value()``, which may omit a CUDA feature
and retain only reusable metadata. Such a featureless artifact is usable only
when the downstream embedding cache already contains the encoded feature.
An artifact is therefore the logical preprocess-cache item. It is not the raw
media, a prompt-specific ``MultimodalDataItem``, or a ViT embedding-cache entry.
"""
from __future__ import annotations
import asyncio
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Optional, Protocol, runtime_checkable
import numpy as np
import torch
from PIL import Image
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.multimodal.cache import (
CacheLookup,
CacheMiss,
MediaSnapshot,
build_artifact_key,
media_preprocess_kwargs,
parse_content_hash,
snapshot_media,
)
from sglang.srt.utils import load_image
@runtime_checkable
class MediaArtifact(Protocol):
"""Common contract implemented by each model's preprocess artifact.
``content_digest`` identifies the media contents. ``artifact_key`` also
includes every preprocessing choice that can change the artifact.
``feature_hash`` becomes ``MultimodalDataItem.hash`` and identifies the
corresponding encoder embedding.
"""
content_digest: str
artifact_key: str
feature_hash: int
@property
def has_feature(self) -> bool: ...
def cache_value(self) -> MediaArtifact:
"""Return the cache-safe representation, possibly without a feature."""
...
def cache_size_items(self) -> Sequence[Any]:
"""Return owned values counted against the preprocess-cache budget."""
...
@dataclass(frozen=True)
class MediaArtifactInput:
"""One decoded raw multimodal input that missed the preprocess cache.
The shared cache layer has already validated its content digest, derived
its artifact key, and claimed the cache miss. The model-specific artifact
builder preprocesses ``media`` into a reusable ``MediaArtifact`` without
loading or hashing the source again. This object is a transient handoff;
it is not itself stored in the cache.
"""
# hash of the media content
content_digest: str
# cache key containing the processor fingerprint, preprocess kwargs, media content
artifact_key: str
modality: Modality
# the original media input that has been loaded and decoded (e.g., PIL.Image)
media: Any
class MediaArtifactCacheMixin:
"""Turn media inputs into ordered artifacts, reusing cached work per item.
The shared layer owns identity, cache lookup, single-flight, partial hits,
and result ordering. A model adapter implements ``prepare_artifact_batch``
and later composes the returned artifacts with the current prompt. Models
can override snapshot/decode/key hooks for each modality without copying the
cache algorithm. Every artifact-producing setting must be exposed through
``preprocess_fingerprint_payload``.
"""
artifact_modality: Optional[Modality] = None
artifact_option_defaults: Mapping[str, Any] = {"detail": "auto"}
def artifact_preprocess_kwargs(
self, source: Any, modality: Modality
) -> Mapping[str, Any]:
"""Return request options that can change this media's artifact.
These options become part of the artifact key. Model adapters can
override this hook when their request schema has additional knobs.
"""
return media_preprocess_kwargs(source, defaults=self.artifact_option_defaults)
def _resolve_artifact_modality(self, modality: Optional[Modality]) -> Modality:
modality = modality or self.artifact_modality
if modality is None:
raise ValueError("A modality is required for artifact caching")
return modality
def _artifact_key(
self,
content_digest: str,
source: Any,
*,
modality: Optional[Modality] = None,
) -> str:
"""Identify one artifact by media content and all preprocess choices."""
if self.processor_fingerprint is None:
raise RuntimeError("Artifact caching requires a processor fingerprint")
modality = self._resolve_artifact_modality(modality)
return build_artifact_key(
content_digest,
modality=modality.name.lower(),
processor_fingerprint=self.processor_fingerprint,
preprocess_kwargs=self.artifact_preprocess_kwargs(source, modality),
)
def decode_media_snapshot(self, snapshot: MediaSnapshot, modality: Modality) -> Any:
"""Decode an immutable snapshot for the model adapter.
Image is the shared default. Future video/audio adapters must make their
decode/sampling contract explicit by overriding this hook.
"""
if modality != Modality.IMAGE:
raise NotImplementedError(
f"{modality.name.lower()} artifact decoding " "requires a model adapter"
)
data = snapshot.data
if isinstance(data, torch.Tensor):
return data
if isinstance(data, np.ndarray):
return torch.from_numpy(data)
if isinstance(data, Image.Image):
data.load()
return data
image, _ = load_image(data, self.gpu_image_decode)
if isinstance(image, Image.Image):
image.load()
return image
def snapshot_media_source(self, source: Any, modality: Modality) -> MediaSnapshot:
"""Capture immutable media content before decode and preprocessing.
The shared implementation covers images. Video/audio adapters can add
streaming or frame-sampling identities here without changing the cache
coordinator.
"""
if modality != Modality.IMAGE:
raise NotImplementedError(
f"{modality.name.lower()} artifact identity " "requires a model adapter"
)
return snapshot_media(source)
def prepare_artifact_batch(
self, entries: Sequence[MediaArtifactInput]
) -> list[MediaArtifact]:
"""Preprocess raw multimodal inputs that missed the preprocess cache.
Each entry is one unique, decoded cache miss. Implementations must
return one reusable artifact (the preprocess-cache item) per entry, in
the same order, while preserving its content digest and artifact key.
The shared layer uses the artifact for the current request and stores
``artifact.cache_value()`` for reuse.
"""
raise NotImplementedError
def artifact_usable(
self, artifact: MediaArtifact, *, allow_featureless: bool
) -> bool:
"""Whether this request can use an artifact that may omit its feature.
A metadata-only artifact is valid only after the scheduler has confirmed
that the corresponding encoder embedding is already cached.
"""
return artifact.has_feature or allow_featureless
@staticmethod
def validate_artifact(artifact: MediaArtifact, entry: MediaArtifactInput) -> None:
"""Enforce identity invariants shared by every model adapter."""
if artifact.content_digest != entry.content_digest:
raise ValueError("prepare_artifact_batch changed the media content digest")
if artifact.artifact_key != entry.artifact_key:
raise ValueError("prepare_artifact_batch changed the media artifact key")
if (
isinstance(artifact.feature_hash, bool)
or not isinstance(artifact.feature_hash, int)
or artifact.feature_hash < 0
):
raise ValueError(
"Media artifact feature_hash must be a non-negative integer"
)
async def _run_preprocess_and_build_artifact_batch(
self, entries: Sequence[MediaArtifactInput]
) -> list[MediaArtifact]:
"""Run model preprocessing locally or on the processor worker pool, return the artifact"""
if self.mm_processor_executor is None:
return self.prepare_artifact_batch(entries)
return await self.mm_processor_executor.run(
self.prepare_artifact_batch, entries
)
def _get_cached_artifact(
self,
key: str,
content_digest: str,
modality: Modality,
*,
allow_featureless: bool,
) -> Optional[MediaArtifact]:
"""Return a compatible cached artifact without recording a cold miss.
Identity mismatches are corrupt entries and are evicted. A featureless
entry that this request cannot use is left for the miss path, which
removes it temporarily and verifies the recomputed feature hash.
"""
artifact = self.mm_preprocess_cache.get_if_present(
key,
lambda value: (
isinstance(value, MediaArtifact)
and value.artifact_key == key
and value.content_digest == content_digest
),
evict_on_reject=True,
)
if artifact is not None:
self.validate_artifact(
artifact,
MediaArtifactInput(content_digest, key, modality, None),
)
if not self.artifact_usable(artifact, allow_featureless=allow_featureless):
return None
return artifact
async def prepare_media_artifacts(
self,
media_data: Sequence[Any],
*,
content_hashes: Optional[Sequence[Optional[str]]] = None,
featureless_hit_mask: Optional[Sequence[bool]] = None,
modality: Optional[Modality] = None,
) -> list[MediaArtifact]:
"""Try resolving one preprocess-cache artifact for each processor input.
Each media input is looked up independently, and results preserve the
input order. A cache hit returns the stored artifact (the cache item).
A miss snapshots and decodes the raw input, runs
``prepare_artifact_batch``, stores its cache-safe artifact, and returns
the prepared artifact to the current request. Duplicate and concurrent
misses share the same preprocessing work.
This stage is prompt-independent. It does not create prompt tokens,
offsets, or ``MultimodalDataItem`` objects; the model processor uses the
returned artifacts to compose those request-specific values afterward.
"""
modality = self._resolve_artifact_modality(modality)
media_count = len(media_data)
if content_hashes is None:
content_hashes = [None] * media_count
if len(content_hashes) != media_count:
raise ValueError(
f"mm_content_hashes has {len(content_hashes)} entries for "
f"{media_count} {modality.name.lower()} items"
)
content_hashes = [parse_content_hash(value) for value in content_hashes]
if featureless_hit_mask is None:
featureless_hit_mask = [False] * media_count
if len(featureless_hit_mask) != media_count:
raise ValueError("featureless_hit_mask must align with media_data")
# keep per-input state aligned for duplicates and partial hits
artifacts: list[Optional[MediaArtifact]] = [None] * media_count
snapshots: list[Optional[MediaSnapshot]] = [None] * media_count
keys: list[Optional[str]] = [None] * media_count
# 1. fast path: resolve trusted provided hash hits without reading media
# e.g., an image could be submitted with a provided hash:
# "image_url": {
# "url": "https://example.com/image.jpg",
# "content_hash": "sha256:<64-hex>"
# }
load_indices = []
for index, (source, caller_hash, allow_featureless) in enumerate(
zip(media_data, content_hashes, featureless_hit_mask)
):
if self.trust_mm_content_hashes and caller_hash is not None:
key = self._artifact_key(caller_hash, source, modality=modality)
keys[index] = key
artifact = self._get_cached_artifact(
key,
caller_hash,
modality,
allow_featureless=allow_featureless,
)
if artifact is not None:
artifacts[index] = artifact
continue
load_indices.append(index)
# 2. read cache: build artifact key from media snapshot then try reading cache
snapshot_futures = {
index: self.io_executor.submit(
self.snapshot_media_source, media_data[index], modality
)
for index in load_indices
}
for index, future in snapshot_futures.items():
snapshot = await asyncio.wrap_future(future)
caller_hash = content_hashes[index]
if caller_hash is not None and caller_hash != snapshot.content_digest:
raise ValueError(
f"content hash mismatch for media_data[{index}]: "
f"expected {caller_hash}, got {snapshot.content_digest}"
)
snapshots[index] = snapshot
key = self._artifact_key(
snapshot.content_digest, media_data[index], modality=modality
)
keys[index] = key
artifacts[index] = self._get_cached_artifact(
key,
snapshot.content_digest,
modality,
allow_featureless=featureless_hit_mask[index],
)
# 3. deduplicate misses before decode
first_index_by_key: dict[str, int] = {}
previous_metadata: dict[str, MediaArtifact] = {}
for index in load_indices:
if artifacts[index] is not None:
continue
key = keys[index]
assert key is not None
if key not in first_index_by_key:
first_index_by_key[key] = index
previous = self.mm_preprocess_cache.pop(key)
if previous is not None:
previous_metadata[key] = previous
unique_keys = list(first_index_by_key)
# 4. submit one computation (preprocess) for each unique miss
cache_results = self.mm_preprocess_cache.lookup_or_claim_many(
unique_keys,
predicate=lambda key, artifact: self.artifact_usable(
artifact,
allow_featureless=featureless_hit_mask[first_index_by_key[key]],
),
)
resolved_by_key: dict[str, MediaArtifact] = {}
misses_to_compute: list[CacheMiss[str, MediaArtifact]] = []
for key, result in zip(unique_keys, cache_results):
if isinstance(result, CacheLookup):
resolved_by_key[key] = result.value
elif result.should_compute:
misses_to_compute.append(result)
if misses_to_compute:
missed_task = self.mm_preprocess_cache.create_background_task(
self._compute_cache_misses(
misses_to_compute,
first_index_by_key,
snapshots,
previous_metadata,
resolved_by_key,
modality,
)
)
# shared work outlives cancellation of this request
await asyncio.shield(missed_task)
# 5. wait for misses already claimed by another request
for key, result in zip(unique_keys, cache_results):
if isinstance(result, CacheMiss) and not result.should_compute:
resolved_by_key[key] = await self.mm_preprocess_cache.wait_for_miss(
result
)
# 6. restore the original processor-input order
for index, artifact in enumerate(artifacts):
if artifact is None:
key = keys[index]
assert key is not None
artifacts[index] = resolved_by_key[key]
if any(artifact is None for artifact in artifacts):
raise RuntimeError("Artifact cache did not resolve every media item")
return [artifact for artifact in artifacts if artifact is not None]
async def _compute_cache_misses(
self,
misses_to_compute: Sequence[CacheMiss[str, MediaArtifact]],
first_index_by_key: Mapping[str, int],
snapshots: Sequence[Optional[MediaSnapshot]],
previous_metadata: Mapping[str, MediaArtifact],
resolved_by_key: dict[str, MediaArtifact],
modality: Modality,
) -> None:
"""Decode and preprocess claimed misses, then wake concurrent waiters.
The full artifact is returned to requests waiting on the miss. A
possibly smaller ``artifact.cache_value()`` is retained in the bounded
CPU cache. The two values differ when a CUDA feature must not be cached.
"""
try:
# 1. decode (load media) each unique miss
missed_media = []
for missed in misses_to_compute:
index = first_index_by_key[missed.key]
snapshot = snapshots[index]
assert snapshot is not None
media = await asyncio.wrap_future(
self.io_executor.submit(
self.decode_media_snapshot, snapshot, modality
)
)
missed_media.append(
MediaArtifactInput(
content_digest=snapshot.content_digest,
artifact_key=missed.key,
modality=modality,
media=media,
)
)
# 2. preprocess all decoded misses as one model batch
missed_artifacts = await self._run_preprocess_and_build_artifact_batch(
missed_media
)
if len(missed_artifacts) != len(misses_to_compute):
raise ValueError(
"prepare_artifact_batch must return one artifact per cache miss"
)
for missed, entry, artifact in zip(
misses_to_compute, missed_media, missed_artifacts
):
self.validate_artifact(artifact, entry)
previous = previous_metadata.get(missed.key)
if (
previous is not None
and previous.feature_hash != artifact.feature_hash
):
raise ValueError(
"Cached media artifact feature hash changed for identical "
f"identity {missed.key}"
)
cache_value = artifact.cache_value()
self.validate_artifact(cache_value, entry)
# 3. return full artifacts and retain cache-safe copies
self.mm_preprocess_cache.complete_miss(
missed,
artifact,
cache_value=cache_value,
)
resolved_by_key[missed.key] = artifact
except BaseException as error:
for missed in misses_to_compute:
if not missed.future.done():
self.mm_preprocess_cache.fail_miss(missed, error)
raise
@@ -0,0 +1,140 @@
"""Prompt-independent Kimi-K3 image preprocessing artifacts."""
from __future__ import annotations
from collections.abc import Mapping
from copy import deepcopy
from dataclasses import dataclass, replace
from typing import Any, Optional, Protocol
import torch
from sglang.srt.multimodal.kimi_k3_image_processing import (
KimiK3DeferredPreprocessing,
)
class KimiK3MediaProcessorConfigProvider(Protocol):
"""Typed view of the HF media processor state consumed by this adapter."""
media_proc_cfg: Mapping[str, Any]
@dataclass(frozen=True)
class KimiK3PreprocessConfig:
"""The single source of truth for K3 artifact-producing choices."""
patch_size: int
merge_kernel_size: int
in_patch_limit: int
patch_limit_on_one_side: int
fixed_output_tokens: Optional[int]
image_mean: tuple[float, ...]
image_std: tuple[float, ...]
transparent_bg_config: Optional[dict]
@classmethod
def from_media_processor(
cls, media_processor: KimiK3MediaProcessorConfigProvider
) -> KimiK3PreprocessConfig:
config = media_processor.media_proc_cfg
return cls(
patch_size=int(config["patch_size"]),
merge_kernel_size=int(config["merge_kernel_size"]),
in_patch_limit=int(config["in_patch_limit"]),
patch_limit_on_one_side=int(config["patch_limit_on_one_side"]),
fixed_output_tokens=(
None
if config.get("fixed_output_tokens") is None
else int(config["fixed_output_tokens"])
),
image_mean=tuple(float(value) for value in config["image_mean"]),
image_std=tuple(float(value) for value in config["image_std"]),
transparent_bg_config=deepcopy(config.get("transparent_bg_config")),
)
@dataclass(frozen=True)
class KimiK3ResizeConfig:
num_tokens: int
new_width: int
new_height: int
pad_width: int
pad_height: int
@classmethod
def from_dict(cls, value: dict) -> KimiK3ResizeConfig:
return cls(
num_tokens=int(value["num_tokens"]),
new_width=int(value["new_width"]),
new_height=int(value["new_height"]),
pad_width=int(value["pad_width"]),
pad_height=int(value["pad_height"]),
)
def as_dict(self) -> dict[str, int]:
return {
"num_tokens": self.num_tokens,
"new_width": self.new_width,
"new_height": self.new_height,
"pad_width": self.pad_width,
"pad_height": self.pad_height,
}
@dataclass(frozen=True)
class KimiK3ImagePreprocessArtifact:
"""K3's prompt-independent preprocess result for one image, containing the feature and everything
``original_size`` and ``resize_config`` rebuild the K3 image tokens for
each prompt; ``grid_thw`` becomes encoder metadata; ``feature`` is either
the prepared encoder input or a raw tensor paired with deferred GPU
preprocessing. ``feature_hash`` links the artifact to the embedding cache.
"""
content_digest: str
artifact_key: str
feature_hash: int
original_size: tuple[int, int]
resize_config: KimiK3ResizeConfig
grid_thw: tuple[int, int, int]
feature: Optional[torch.Tensor]
deferred: Optional[KimiK3DeferredPreprocessing] = None
@property
def has_feature(self) -> bool:
return self.feature is not None
def cache_value(self) -> KimiK3ImagePreprocessArtifact:
"""Return the CPU-cacheable copy; never retain a CUDA tensor."""
if self.feature is None or self.feature.device.type == "cpu":
return self
return replace(self, feature=None)
def cache_size_items(self) -> tuple:
"""Return every owned value that contributes to the CPU cache budget."""
deferred = None
if self.deferred is not None:
deferred = (
self.deferred.backend,
self.deferred.image_mean,
self.deferred.image_std,
self.deferred.transparent_bg_config,
self.deferred.resize_config,
)
return (
self.content_digest,
self.artifact_key,
self.feature_hash,
self.original_size,
(
self.resize_config.num_tokens,
self.resize_config.new_width,
self.resize_config.new_height,
self.resize_config.pad_width,
self.resize_config.pad_height,
),
self.grid_thw,
self.feature,
deferred,
)
@@ -6,7 +6,15 @@ import multiprocessing as mp
import os
import re
from abc import ABC, abstractmethod
from typing import Any, Dict, Iterator, List, Optional, Tuple, Union
from typing import (
Any,
Dict,
Iterator,
List,
Optional,
Tuple,
Union,
)
import numpy as np
import torch
@@ -21,6 +29,7 @@ from sglang.srt.managers.schedule_batch import (
)
from sglang.srt.multimodal.cache import (
MultimodalPreprocessCache,
PreprocessFingerprintProvider,
build_processor_fingerprint,
)
from sglang.srt.multimodal.processors.executor import MultimodalProcessorExecutor
@@ -190,6 +199,8 @@ class BaseMultimodalProcessor(ABC):
preserve_processor_input_ids = False
auto_mm_processor_worker_num = 1
auto_mm_io_worker_num = 4
# Models opt in by assigning a non-zero default. A user-provided server
# argument overrides this value; zero disables storage and cache-key work.
auto_mm_preprocess_cache_size_mb = 0
supports_mm_processor_concurrency = False
@@ -204,9 +215,7 @@ class BaseMultimodalProcessor(ABC):
server_args.allowed_media_domains,
server_args.media_url_max_file_size_mb,
)
configured_mm_feature_transport = getattr(
server_args, "mm_feature_transport", "cpu"
)
configured_mm_feature_transport = server_args.mm_feature_transport
self.mm_feature_transport = (
configured_mm_feature_transport
if configured_mm_feature_transport in ("cpu", "cuda_ipc", "cuda_vmm")
@@ -216,10 +225,8 @@ class BaseMultimodalProcessor(ABC):
self.use_ipc_pool_handle_cache = (
self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
)
self.image_processor_backend = getattr(
server_args, "image_processor_backend", "auto"
)
if getattr(server_args, "disable_fast_image_processor", False):
self.image_processor_backend = server_args.image_processor_backend
if server_args.disable_fast_image_processor:
self.image_processor_backend = "pil"
self.disable_fast_image_processor = self.image_processor_backend == "pil"
self.skip_tokenizer_init = server_args.skip_tokenizer_init
@@ -229,25 +236,24 @@ class BaseMultimodalProcessor(ABC):
self.video_config = mm_process_config.get("video", {})
self.audio_config = mm_process_config.get("audio", {})
requested_cache_mb = getattr(
self.server_args, "mm_preprocess_cache_size_mb", None
)
# Each tokenizer worker is a separate process with its own CPU cache.
# Split the requested service-wide budget so increasing worker count
# does not silently multiply host-memory usage.
requested_cache_mb = self.server_args.mm_preprocess_cache_size_mb
total_cache_mb = (
self.auto_mm_preprocess_cache_size_mb
if requested_cache_mb is None
else requested_cache_mb
)
tokenizer_worker_num = max(
int(getattr(self.server_args, "tokenizer_worker_num", 1)), 1
)
tokenizer_worker_num = max(int(self.server_args.tokenizer_worker_num), 1)
worker_cache_bytes = total_cache_mb * 1024 * 1024 // tokenizer_worker_num
self.mm_preprocess_cache = MultimodalPreprocessCache(
max_size_bytes=worker_cache_bytes,
max_entries=8192,
)
self.trust_mm_content_hashes = bool(
getattr(self.server_args, "trust_mm_content_hashes", False)
)
self.trust_mm_content_hashes = bool(self.server_args.trust_mm_content_hashes)
# The fingerprint is needed only to build artifact keys. Avoid inspecting
# processor state when this processor will never retain artifacts.
self.processor_fingerprint = (
build_processor_fingerprint(self, hf_config, server_args)
if self.mm_preprocess_cache.enabled
@@ -418,26 +424,46 @@ class BaseMultimodalProcessor(ABC):
@property
def keep_mm_features_on_device(self) -> bool:
"""Whether feature transport expects processor outputs to stay on GPU."""
return self.mm_feature_transport in ("cuda_ipc", "cuda_vmm")
def preprocess_fingerprint_payload(self) -> dict[str, Any]:
"""Stable processor choices that may change per-media artifacts."""
"""Return every stable setting that can change a media artifact.
The payload is hashed once at startup and becomes part of every
artifact key. Model processors must extend this method when they add an
output-affecting option. The wrapped HF processor can expose its own
typed payload through ``PreprocessFingerprintProvider``.
"""
wrapped_processor = (
self._processor.preprocess_fingerprint_payload()
if isinstance(self._processor, PreprocessFingerprintProvider)
else None
)
return {
"wrapper_class": (
f"{type(self._processor).__module__}."
f"{type(self._processor).__qualname__}"
),
"gpu_image_decode": self.gpu_image_decode,
"image_processor_backend": self.image_processor_backend,
"feature_transport": self.mm_feature_transport,
"image_config": self.image_config,
"video_config": self.video_config,
"audio_config": self.audio_config,
"wrapped_processor": wrapped_processor,
}
def clear_preprocess_cache(self) -> None:
"""Drop artifacts and reject cache writes from pre-flush work.
Active requests continue and still receive their preprocessing result;
they simply cannot repopulate the freshly cleared cache.
"""
self.mm_preprocess_cache.clear()
def shutdown(self) -> None:
"""Release executor resources and cached CPU artifacts."""
"""Drop cached artifacts and stop every processor-side executor."""
self.clear_preprocess_cache()
self.io_executor.shutdown(wait=False, cancel_futures=True)
self.cpu_executor.shutdown(wait=False, cancel_futures=True)
@@ -9,8 +9,9 @@ at load time.
"""
import functools
import math
import re
from typing import Dict, List, Union
from typing import Dict, List, Optional, Union
import numpy as np
import torch
@@ -22,6 +23,7 @@ from sglang.srt.managers.schedule_batch import (
MultimodalProcessorOutput,
)
from sglang.srt.models.kimi_k3 import KimiK3ForConditionalGeneration
from sglang.srt.multimodal.cache import resolve_multimodal_item_hash
from sglang.srt.multimodal.kimi_k3_image_processing import (
DEFERRED_PREPROCESSING_KEY,
KimiK3DeferredPreprocessing,
@@ -32,6 +34,15 @@ from sglang.srt.multimodal.kimi_k3_image_processing import (
from sglang.srt.multimodal.kimi_k3_image_processing import (
to_chw_uint8,
)
from sglang.srt.multimodal.media_artifacts import (
MediaArtifactCacheMixin,
MediaArtifactInput,
)
from sglang.srt.multimodal.media_artifacts.kimi_k3 import (
KimiK3ImagePreprocessArtifact,
KimiK3PreprocessConfig,
KimiK3ResizeConfig,
)
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor as SGLangBaseProcessor,
)
@@ -152,9 +163,24 @@ def _k3_to_cuda_chw(image: Union[torch.Tensor, Image.Image]) -> torch.Tensor:
class KimiK3GPUProcessorWrapper(KimiGPUProcessorWrapper):
def __init__(self, *args, transparent_bg_config=None, **kwargs):
super().__init__(*args, **kwargs)
self._transparent_bg_config = transparent_bg_config
def __init__(self, hf_processor, image_token, image_token_id, config):
self.preprocess_config = config
super().__init__(
hf_processor,
image_token=image_token,
image_token_id=image_token_id,
patch_size=config.patch_size,
merge_kernel_size=config.merge_kernel_size,
in_patch_limit=config.in_patch_limit,
patch_limit_on_one_side=config.patch_limit_on_one_side,
fixed_output_tokens=config.fixed_output_tokens,
image_mean=config.image_mean,
image_std=config.image_std,
)
self._transparent_bg_config = config.transparent_bg_config
def preprocess_fingerprint_payload(self):
return self.preprocess_config
def _prepare_input_ids(
self, input_text, resize_configs, original_input_ids, image_sizes
@@ -285,9 +311,65 @@ class KimiK3GPUProcessorWrapper(KimiGPUProcessorWrapper):
)
return input_ids, resize_configs, deferred_preprocessing
def prepare_image_features(self, images):
"""Prepare prompt-independent, per-image features in one processor call."""
image_sizes = [_get_image_dimensions(image) for image in images]
resize_configs = [
navit_resize_config(
width,
height,
self._patch_size,
self._merge_kernel_size,
self._in_patch_limit,
self._patch_limit_on_one_side,
self._fixed_output_tokens,
)
for width, height in image_sizes
]
class KimiK3ImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
if images and torch.cuda.is_available():
image_scale, image_bias = self._get_gpu_norm_tensors()
pixel_values, grid_thws = _gpu_preprocess_images(
images,
resize_configs,
image_scale,
image_bias,
self._patch_size,
to_chw=_k3_to_cuda_chw,
post_resize=lambda x: _fill_transparent_bg(
x, self._transparent_bg_config
),
)
else:
# The checkpoint CPU processor couples prompt composition with media
# preprocessing. A synthetic prompt keeps that API but is discarded;
# image features and grids are independent of its text.
output = self._cpu_call(self._image_token * len(images), images)
pixel_values = output["pixel_values"]
grid_thws = output["image_grid_thw"]
grids = [tuple(int(value) for value in grid) for grid in grid_thws.tolist()]
patch_counts = [math.prod(grid) for grid in grids]
if sum(patch_counts) != pixel_values.shape[0]:
raise ValueError(
"Kimi-K3 processor feature length does not match image grids: "
f"{pixel_values.shape[0]} != {sum(patch_counts)}"
)
return (
list(pixel_values.split(patch_counts)),
image_sizes,
resize_configs,
grids,
)
class KimiK3ImageProcessor(
KimiGridMMDataMixin,
MediaArtifactCacheMixin,
SGLangBaseProcessor,
):
models = [KimiK3ForConditionalGeneration]
artifact_modality = Modality.IMAGE
# K3 accuracy is sensitive to the chroma upsampling used for common 4:2:0
# JPEG inputs. This mode uses interpolated nvJPEG upsampling when the K3
# image dependency is installed and otherwise falls back to PIL.
@@ -307,25 +389,23 @@ class KimiK3ImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
image_token_regex=re.compile(r"(?:<\|media_pad\|>)+"),
).build(_processor)
media_proc_cfg = _processor.media_processor.media_proc_cfg
preprocess_config = KimiK3PreprocessConfig.from_media_processor(
_processor.media_processor
)
processor = KimiK3GPUProcessorWrapper(
_processor,
image_token=mm_tokens.image_token,
image_token_id=mm_tokens.image_token_id,
patch_size=media_proc_cfg["patch_size"],
merge_kernel_size=media_proc_cfg["merge_kernel_size"],
in_patch_limit=media_proc_cfg["in_patch_limit"],
patch_limit_on_one_side=media_proc_cfg["patch_limit_on_one_side"],
fixed_output_tokens=media_proc_cfg.get("fixed_output_tokens"),
image_mean=media_proc_cfg["image_mean"],
image_std=media_proc_cfg["image_std"],
transparent_bg_config=media_proc_cfg.get("transparent_bg_config"),
config=preprocess_config,
)
super().__init__(hf_config, server_args, processor, *args, **kwargs)
self.mm_tokens = mm_tokens
def _should_defer_gpu_preprocessing(self, images) -> bool:
"""
when raw_bytes <= processed_bytes, preprocess first would introduce larger payload, so deferring gpu preprocessing would benefit
"""
if (
not images
or self.mm_feature_transport != "cpu"
@@ -340,17 +420,18 @@ class KimiK3ImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
raw_bytes = 0
processed_bytes = 0
patch_size = self._processor._patch_size
config = self._processor.preprocess_config
patch_size = config.patch_size
for image in images:
width, height = _get_image_dimensions(image)
resize_config = navit_resize_config(
width,
height,
patch_size,
self._processor._merge_kernel_size,
self._processor._in_patch_limit,
self._processor._patch_limit_on_one_side,
self._processor._fixed_output_tokens,
config.merge_kernel_size,
config.in_patch_limit,
config.patch_limit_on_one_side,
config.fixed_output_tokens,
)
if isinstance(image, torch.Tensor):
channels = (
@@ -391,7 +472,7 @@ class KimiK3ImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
base_output.images, resize_configs, offsets
):
grid_thw = _grid_thw_from_resize_config(
resize_config, self._processor._patch_size
resize_config, self._processor.preprocess_config.patch_size
)
item = MultimodalDataItem(
modality=Modality.IMAGE,
@@ -413,6 +494,217 @@ class KimiK3ImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
im_token_id=self.mm_tokens.image_token_id,
)
def _make_artifact(
self,
*,
content_digest: str,
artifact_key: str,
original_size: tuple[int, int],
resize_config: dict,
grid_thw: tuple[int, int, int],
feature: torch.Tensor,
deferred: Optional[KimiK3DeferredPreprocessing] = None,
) -> KimiK3ImagePreprocessArtifact:
"""Freeze one image's prompt-independent preprocessing result."""
# Use the same feature-hash contract as MultimodalDataItem.
feature_hash = resolve_multimodal_item_hash(
feature=feature, namespace=artifact_key
)
if not self.keep_mm_features_on_device and feature.device.type != "cpu":
feature = feature.cpu()
return KimiK3ImagePreprocessArtifact(
content_digest=content_digest,
artifact_key=artifact_key,
feature_hash=feature_hash,
original_size=original_size,
resize_config=KimiK3ResizeConfig.from_dict(resize_config),
grid_thw=grid_thw,
feature=feature,
deferred=deferred,
)
def prepare_artifact_batch(
self,
entries: list[MediaArtifactInput],
*,
processor=None,
) -> list[KimiK3ImagePreprocessArtifact]:
"""Preprocess raw cache misses into reusable per-image cache items.
Each entry is a confirmed cache miss. It is either processed now or
stored with the metadata needed for deferred GPU preprocessing.
"""
processor = processor or self._processor
artifacts: list[Optional[KimiK3ImagePreprocessArtifact]] = [None] * len(entries)
# 1. collect inputs that must be preprocessed now instead of deferred
eager_entry_indices = []
eager_images = []
config = processor.preprocess_config
for index, entry in enumerate(entries):
image = entry.media
if not self._should_defer_gpu_preprocessing([image]):
eager_entry_indices.append(index)
eager_images.append(image)
continue
width, height = _get_image_dimensions(image)
resize_config = navit_resize_config(
width,
height,
config.patch_size,
config.merge_kernel_size,
config.in_patch_limit,
config.patch_limit_on_one_side,
config.fixed_output_tokens,
)
grid_thw = _grid_thw_from_resize_config(resize_config, config.patch_size)
feature = to_chw_uint8(image).cpu().contiguous()
artifacts[index] = self._make_artifact(
content_digest=entry.content_digest,
artifact_key=entry.artifact_key,
original_size=(width, height),
resize_config=resize_config,
grid_thw=grid_thw,
feature=feature,
deferred=KimiK3DeferredPreprocessing(
backend="gpu",
image_mean=list(config.image_mean),
image_std=list(config.image_std),
transparent_bg_config=config.transparent_bg_config,
resize_config=resize_config,
),
)
# 2. preprocess CPU eager inputs as one batch
if eager_images:
features, sizes, configs, grids = processor.prepare_image_features(
eager_images
)
for index, feature, size, resize_config, grid in zip(
eager_entry_indices, features, sizes, configs, grids
):
entry = entries[index]
artifacts[index] = self._make_artifact(
content_digest=entry.content_digest,
artifact_key=entry.artifact_key,
original_size=size,
resize_config=resize_config,
grid_thw=grid,
feature=feature,
)
# 3. return artifacts in the original processor-input order
if any(artifact is None for artifact in artifacts):
raise RuntimeError("Kimi-K3 artifact batch did not produce every image")
return [artifact for artifact in artifacts if artifact is not None]
def compose_request(
self,
input_text,
artifacts: list[KimiK3ImagePreprocessArtifact],
) -> MultimodalProcessorOutput:
"""Compose the current request from its prompt and ordered artifacts.
``prepare_media_artifacts`` has already returned one artifact for each
processor input, either from the preprocess cache or from fresh
preprocessing. This method expands the current prompt's image tokens
and converts each artifact into its request-specific
``MultimodalDataItem`` with offsets, grid metadata, feature, and feature
hash. It does not read raw media or access the preprocess cache.
"""
# 1. rebuild prompt-specific tokens and offsets
original_ids = (
input_text
if isinstance(input_text, (list, torch.Tensor))
else _encode_k3_special_tokens(self._tokenizer, input_text)
)
input_ids = _expand_k3_image_prompt_token_ids(
original_ids,
self.mm_tokens.image_token_id,
[artifact.resize_config.num_tokens for artifact in artifacts],
[artifact.original_size for artifact in artifacts],
self._tokenizer,
).flatten()
offsets = self.get_mm_items_offset(input_ids, self.mm_tokens.image_token_id)
if len(offsets) != len(artifacts):
raise ValueError("Expected one Kimi-K3 image span for each image")
# 2. build request-owned items from prompt-independent artifacts
items = []
for artifact, offset in zip(artifacts, offsets):
model_specific_data = {
"image_grid_thw": torch.tensor([artifact.grid_thw], dtype=torch.int64)
}
if artifact.deferred is not None:
model_specific_data[DEFERRED_PREPROCESSING_KEY] = artifact.deferred
item = MultimodalDataItem(
modality=Modality.IMAGE,
feature=artifact.feature,
offsets=[offset],
model_specific_data=model_specific_data,
)
item.set_hash(artifact.feature_hash)
if self.use_cuda_ipc and isinstance(item.feature, torch.Tensor):
item.feature = self._wrap_tensor_for_cuda_ipc(item.feature)
if self.keep_mm_features_on_device and item.feature is not None:
item.model_specific_data[DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY] = (
True
)
items.append(item)
return MultimodalProcessorOutput(
input_ids=input_ids.tolist(),
mm_items=items,
im_token_id=self.mm_tokens.image_token_id,
)
async def _process_mm_data_uncached(
self, image_data, input_text, request_obj, **kwargs
):
"""Compatibility path for precomputed inputs and lightweight test stubs."""
expected_image_count = len(image_data or [])
placeholder_count = self.count_image_placeholders(
input_text, self.mm_tokens.image_token_id
)
if placeholder_count is not None:
base_output = await self.fast_load_mm_data(
prompt=input_text,
image_data=image_data,
multimodal_tokens=self.mm_tokens,
discard_alpha_channel=False,
input_ids=input_text,
)
else:
base_output = await self.load_mm_data(
prompt=input_text,
image_data=image_data,
multimodal_tokens=self.mm_tokens,
discard_alpha_channel=False,
)
if len(base_output.images) != expected_image_count:
raise ValueError(
"Kimi image placeholders must map one-to-one to image data: "
f"expected {expected_image_count}, loaded {len(base_output.images)}"
)
if self._should_defer_gpu_preprocessing(base_output.images):
return self._build_deferred_output(base_output)
mm_items, input_ids, _ = await self.process_and_combine_mm_data_async(
base_output,
self.mm_tokens,
sglang_original_input_ids=base_output.input_ids,
)
if self.keep_mm_features_on_device:
for item in mm_items:
item.model_specific_data[DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY] = (
True
)
return MultimodalProcessorOutput(
input_ids=input_ids.tolist(),
mm_items=mm_items,
im_token_id=self.mm_tokens.image_token_id,
)
async def process_mm_data_async(
self,
image_data: List[Union[str, bytes, Dict]],
@@ -421,7 +713,7 @@ class KimiK3ImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
*args,
**kwargs,
):
if getattr(request_obj, "video_data", None) or kwargs.get("audio_data"):
if request_obj.video_data or kwargs.get("audio_data"):
raise ValueError("Kimi-K3 supports image input only")
expected_image_count = len(image_data or [])
@@ -434,60 +726,21 @@ class KimiK3ImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
"Kimi image placeholders must map one-to-one to image data: "
f"expected {expected_image_count}, found {placeholder_count} token(s)"
)
# Keep structural media tokens distinct from user text that happens to
# spell ``<|media_pad|>``. Decoding the whole prompt and matching the
# resulting string would lose that distinction and could bind an image
# to user-provided text instead of the renderer-inserted token.
base_output = await self.fast_load_mm_data(
prompt=input_text,
image_data=image_data,
multimodal_tokens=self.mm_tokens,
discard_alpha_channel=False,
# Unlike load_mm_data, fast_load_mm_data does not derive
# input_ids from the prompt. Without this the wrapper falls
# back to re-encoding the decoded string, which is the loss of
# the structural/user distinction described above.
input_ids=input_text,
)
else:
base_output = await self.load_mm_data(
prompt=input_text,
image_data=image_data,
multimodal_tokens=self.mm_tokens,
discard_alpha_channel=False,
if (
any(self._is_preprocessed_input(item) for item in image_data)
or not self.mm_preprocess_cache.enabled
):
# 1. keep preprocessed inputs and cache-off requests on the legacy path
return await self._process_mm_data_uncached(
image_data, input_text, request_obj, **kwargs
)
if len(base_output.images) != expected_image_count:
raise ValueError(
"Kimi image placeholders must map one-to-one to image data: "
f"expected {expected_image_count}, loaded {len(base_output.images)}"
)
if self._should_defer_gpu_preprocessing(base_output.images):
return self._build_deferred_output(base_output)
mm_items, input_ids, _ = await self.process_and_combine_mm_data_async(
base_output,
self.mm_tokens,
sglang_original_input_ids=base_output.input_ids,
)
# K3's tower is unconditionally image-wise data-parallel (each image
# is consumed by exactly one TP rank), so keep IPC proxies lazy until
# that assignment is known: one tokenizer/scheduler crossing per
# image instead of one per rank. K2.5 gates this on
# --mm-enable-dp-encoder; K3 needs no flag.
if self.keep_mm_features_on_device:
for item in mm_items:
item.model_specific_data[DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY] = (
True
)
return MultimodalProcessorOutput(
input_ids=input_ids.tolist(),
mm_items=mm_items,
im_token_id=self.mm_tokens.image_token_id,
# 2. resolve per-image artifacts before composing the current prompt
artifacts = await self.prepare_media_artifacts(
image_data,
content_hashes=request_obj.mm_content_hashes,
)
return self.compose_request(input_text, artifacts)
def get_mm_data(self, prompt, embeddings, **kwargs):
img_grid_thw = kwargs.get("img_grid_thw", None)
@@ -19,6 +19,7 @@ from sglang.srt.disaggregation.encode_receiver import (
EmbeddingData,
MMReceiverHTTP,
MultiModalEmbeddingData,
_encoder_media_item,
_select_mm_processor_prompt,
)
from sglang.srt.disaggregation.encode_server import MMEncoder, _get_mm_grid_dim
@@ -27,8 +28,10 @@ from sglang.srt.managers.tokenizer_manager import (
_reject_missing_dispatched_encoder_embedding,
)
from sglang.srt.models.kimi_k3 import KimiK3ForConditionalGeneration
from sglang.srt.multimodal.cache import snapshot_media
from sglang.srt.multimodal.encoder_preprocessing import (
LOCAL_PREPROCESSED_KEY,
EncoderMediaProcessorConfig,
EncoderPreprocessOutput,
get_encoder_preprocessed_items,
hash_raw_encoder_item,
@@ -41,6 +44,7 @@ from sglang.srt.multimodal.kimi_k3_image_processing import (
)
from sglang.srt.runtime_context import get_context
from sglang.srt.server_args import resolve_encoder_transfer_backend
from sglang.srt.utils import ImageData
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@@ -208,6 +212,11 @@ def _encoder(model_type="kimi_k3"):
vision_config=SimpleNamespace(merge_kernel_size=(2, 2))
)
)
encoder.encoder_media_processor_config = (
KimiK3ForConditionalGeneration.encoder_media_processor_config
if model_type == "kimi_k3"
else EncoderMediaProcessorConfig()
)
return encoder
@@ -296,6 +305,19 @@ def test_kimi_k3_epd_preprocess_preserves_raw_per_image_items():
assert deferred.image_std == [0.5, 0.5, 0.5]
def test_kimi_k3_epd_preserves_verified_content_identity():
image = Image.new("RGB", (8, 6), color=(1, 2, 3))
digest = "sha256:" + "ab" * 32
output = prepare_kimi_k3_encoder_inputs(
[{"type": "image", "image": image, "content_hash": digest}],
_kimi_k3_image_processor(),
)
item = get_encoder_preprocessed_items(output)[0]
assert item.model_specific_data["content_digest"] == digest
def test_kimi_k3_epd_model_preprocessor_receives_image_processor():
image = Image.new("RGB", (8, 6), color=(1, 2, 3))
image_processor = _kimi_k3_image_processor()
@@ -458,6 +480,67 @@ def test_kimi_k3_epd_selects_matching_jpeg_decode_mode(
load.assert_called_once_with(b"jpeg", expected_decode_mode)
def test_kimi_k3_epd_verifies_content_hash_before_decode():
payload = b"jpeg"
digest = snapshot_media(payload).content_digest
expected = torch.zeros((3, 2, 3), dtype=torch.uint8)
encoder = _encoder()
encoder.use_image_processor_gpu = False
with patch(
"sglang.srt.disaggregation.encode_server.load_image",
return_value=(expected, None),
) as load:
output = encoder._load_single_item(
{"url": payload, "content_hash": digest}, Modality.IMAGE
)
assert output == {
"type": "image",
"image": expected,
"content_hash": digest,
}
load.assert_called_once_with(payload, False)
def test_epd_receiver_keeps_content_hash_aligned_with_image():
digest = "sha256:" + "cd" * 32
receiver = MMReceiverHTTP.__new__(MMReceiverHTTP)
request = SimpleNamespace(
image_data=[
ImageData(
url="image",
detail="high",
max_dynamic_patch=12,
preprocess_kwargs={"crop": False},
content_hash=digest,
)
],
video_data=None,
audio_data=None,
mm_content_hashes=[digest],
)
assert receiver._extract_url_data(request) == [
{
"url": "image",
"modality": Modality.IMAGE,
"detail": "high",
"max_dynamic_patch": 12,
"preprocess_kwargs": {"crop": False},
"content_hash": digest,
}
]
assert _encoder_media_item(receiver._extract_url_data(request)[0]) == {
"url": "image",
"detail": "high",
"max_dynamic_patch": 12,
"preprocess_kwargs": {"crop": False},
"content_hash": digest,
}
def test_kimi_k3_epd_aggregates_original_image_sizes_in_part_order():
first = EmbeddingData(
req_id="request",
@@ -127,12 +127,14 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
with patch.dict(os.environ, {}, clear=False):
os.environ.pop("SGLANG_IO_WORKERS", None)
with patch.object(
with (
patch.object(
BaseMultimodalProcessor, "auto_mm_processor_worker_num", 4
), patch.object(
BaseMultimodalProcessor, "auto_mm_io_worker_num", 16
), patch.object(
),
patch.object(BaseMultimodalProcessor, "auto_mm_io_worker_num", 16),
patch.object(
BaseMultimodalProcessor, "supports_mm_processor_concurrency", True
),
):
proc = self._make_processor({})
try:
@@ -172,14 +174,18 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
def _server_args(mm_feature_transport):
return SimpleNamespace(
mm_feature_transport=mm_feature_transport,
image_processor_backend="auto",
disable_fast_image_processor=False,
skip_tokenizer_init=False,
mm_process_config={},
mm_preprocess_cache_size_mb=0,
trust_mm_content_hashes=False,
mm_processor_worker_num=0,
mm_io_worker_num=0,
tokenizer_worker_num=1,
base_gpu_id=2,
tp_size=8,
rl_on_policy_target=None,
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
@@ -195,9 +201,13 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
# transport policy must still resolve from the instance's ServerArgs.
from sglang.srt.multimodal.processors import base_processor
with envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(True), patch.object(
with (
envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(True),
patch.object(
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
),
patch.object(base_processor, "MmItemMemoryPool") as memory_pool,
):
processor = base_processor.BaseMultimodalProcessor(
hf_config=MagicMock(),
server_args=self._server_args("cuda_ipc"),
@@ -213,9 +223,13 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
def test_cuda_ipc_pool_handle_cache_can_be_disabled(self):
from sglang.srt.multimodal.processors import base_processor
with envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(False), patch.object(
with (
envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(False),
patch.object(
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
),
patch.object(base_processor, "MmItemMemoryPool") as memory_pool,
):
processor = base_processor.BaseMultimodalProcessor(
hf_config=MagicMock(),
server_args=self._server_args("cuda_ipc"),
@@ -230,9 +244,13 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
def test_cpu_transport_does_not_allocate_ipc_pool(self):
from sglang.srt.multimodal.processors import base_processor
with envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(True), patch.object(
with (
envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.override(True),
patch.object(
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
),
patch.object(base_processor, "MmItemMemoryPool") as memory_pool,
):
processor = base_processor.BaseMultimodalProcessor(
hf_config=MagicMock(),
server_args=self._server_args("cpu"),
@@ -251,9 +269,12 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
hf_processor = self._processor()
feature = torch.empty(1, device="meta")
hf_processor.return_value = {"pixel_values": feature}
with patch.object(
with (
patch.object(
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
),
patch.object(base_processor, "MmItemMemoryPool") as memory_pool,
):
processor = base_processor.BaseMultimodalProcessor(
hf_config=MagicMock(),
server_args=self._server_args("cuda_vmm"),
@@ -365,9 +386,10 @@ class TestPrecomputeHashBeforeCpuTransfer(CustomTestCase):
BaseMultimodalProcessor,
)
with patch.object(
BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(BaseMultimodalProcessor, "__init__", lambda self: None):
with (
patch.object(BaseMultimodalProcessor, "__abstractmethods__", set()),
patch.object(BaseMultimodalProcessor, "__init__", lambda self: None),
):
processor = BaseMultimodalProcessor()
processor.precompute_hash_before_cpu_transfer = enabled
processor.use_cuda_ipc = False
@@ -409,9 +431,10 @@ class TestMultimodalProcessorConcurrency(unittest.IsolatedAsyncioTestCase):
MultimodalProcessorExecutor,
)
with patch.object(
BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(BaseMultimodalProcessor, "__init__", lambda self: None):
with (
patch.object(BaseMultimodalProcessor, "__abstractmethods__", set()),
patch.object(BaseMultimodalProcessor, "__init__", lambda self: None),
):
processor = BaseMultimodalProcessor()
processor.mm_processor_executor = MultimodalProcessorExecutor(
@@ -438,9 +461,10 @@ class TestMultimodalProcessorConcurrency(unittest.IsolatedAsyncioTestCase):
BaseMultimodalProcessor,
)
with patch.object(
BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(BaseMultimodalProcessor, "__init__", lambda self: None):
with (
patch.object(BaseMultimodalProcessor, "__abstractmethods__", set()),
patch.object(BaseMultimodalProcessor, "__init__", lambda self: None),
):
processor = BaseMultimodalProcessor()
processor.mm_processor_executor = None
+558 -19
View File
@@ -1,7 +1,14 @@
"""CPU coverage for Kimi-K2.5/K2.7 encoder-DP wiring."""
import asyncio
import base64
import functools
import io
import pickle
import tempfile
from concurrent.futures import ThreadPoolExecutor
from dataclasses import replace
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock, patch
@@ -18,11 +25,26 @@ from sglang.srt.managers.schedule_batch import (
MultimodalInputs,
MultimodalProcessorOutput,
)
from sglang.srt.models.kimi_k3 import KimiK3ForConditionalGeneration
from sglang.srt.models.kimi_k25 import (
KimiK25ForConditionalGeneration,
mm_projection_auto,
)
from sglang.srt.models.kimi_vl_moonvit import tpool_patch_merger
from sglang.srt.multimodal.cache import (
MultimodalPreprocessCache,
resolve_multimodal_item_hash,
snapshot_media,
)
from sglang.srt.multimodal.kimi_k3_image_processing import (
DEFERRED_PREPROCESSING_KEY,
KimiK3DeferredPreprocessing,
)
from sglang.srt.multimodal.media_artifacts.kimi_k3 import (
KimiK3ImagePreprocessArtifact,
KimiK3PreprocessConfig,
KimiK3ResizeConfig,
)
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
from sglang.srt.multimodal.processors.kimi_common import KimiGridMMDataMixin
@@ -45,6 +67,8 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
CudaIpcTensorTransportProxy,
)
from sglang.srt.runtime_context import get_context, get_parallel
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
@@ -320,8 +344,9 @@ def test_dp_helper_supports_moonvit3d_packed_embeddings_on_tp1():
# The IPC consumer count asks for the *configured* TP size (matching
# MmItemMemoryPool.try_to_recycle), so the double publishes one too.
with get_context().override_server_args(tp_size=1), get_parallel().override(
tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0
with (
get_context().override_server_args(tp_size=1),
get_parallel().override(tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0),
):
output = run_dp_sharded_mrope_vision_model(
tower, pixel_values, [[1, 2, 2]], rope_type="rope_2d_packed"
@@ -338,8 +363,9 @@ def test_dp_helper_can_lazily_load_kimi_features_on_tp1():
# The IPC consumer count asks for the *configured* TP size (matching
# MmItemMemoryPool.try_to_recycle), so the double publishes one too.
with get_context().override_server_args(tp_size=1), get_parallel().override(
tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0
with (
get_context().override_server_args(tp_size=1),
get_parallel().override(tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0),
):
output = run_dp_sharded_mrope_vision_model(
tower,
@@ -490,8 +516,9 @@ def test_kimi_non_dp_keeps_grid_thws_on_the_host():
# The IPC consumer count asks for the *configured* TP size (matching
# MmItemMemoryPool.try_to_recycle), so the double publishes one too.
with get_context().override_server_args(tp_size=1), get_parallel().override(
tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0
with (
get_context().override_server_args(tp_size=1),
get_parallel().override(tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0),
):
model.get_image_feature(items)
@@ -582,6 +609,30 @@ class _HFProcessor:
)
class _AnySizeTokenizer:
def encode(self, text, allowed_special=None):
if text.startswith("<|media_begin|>image "):
return [10, 11]
if text == "<|media_end|>":
return [14]
return []
def _k3_preprocess_config(
*, patch_size=14, in_patch_limit=16384
) -> KimiK3PreprocessConfig:
return KimiK3PreprocessConfig(
patch_size=patch_size,
merge_kernel_size=2,
in_patch_limit=in_patch_limit,
patch_limit_on_one_side=512,
fixed_output_tokens=None,
image_mean=(0.5, 0.5, 0.5),
image_std=(0.5, 0.5, 0.5),
transparent_bg_config=None,
)
@pytest.mark.parametrize(
("processor_cls", "wrapper_cls"),
[
@@ -592,13 +643,17 @@ class _HFProcessor:
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
server_args = SimpleNamespace(
mm_feature_transport="cpu",
image_processor_backend="auto",
disable_fast_image_processor=False,
skip_tokenizer_init=False,
mm_process_config={},
mm_io_worker_num=0,
mm_processor_worker_num=0,
tokenizer_worker_num=1,
mm_preprocess_cache_size_mb=0,
trust_mm_content_hashes=False,
base_gpu_id=0,
rl_on_policy_target=None,
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
@@ -615,6 +670,12 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls
assert isinstance(processor._processor, wrapper_cls)
assert isinstance(worker_processor, wrapper_cls)
assert worker_processor is not processor._processor
if processor_cls is KimiK3ImageProcessor:
fingerprint_config = processor.preprocess_fingerprint_payload()[
"wrapped_processor"
]
assert isinstance(fingerprint_config, KimiK3PreprocessConfig)
assert fingerprint_config.patch_size == 14
finally:
processor.mm_processor_executor.shutdown()
processor.io_executor.shutdown()
@@ -690,18 +751,498 @@ def test_kimi_k3_epd_rebuild_uses_the_same_media_contract():
)
def test_kimi_k3_cpu_transport_defers_gpu_preprocessing():
from sglang.srt.multimodal.kimi_k3_image_processing import (
DEFERRED_PREPROCESSING_KEY,
KimiK3DeferredPreprocessing,
def _cached_k3_artifact(content_digest, artifact_key, value=1):
return KimiK3ImagePreprocessArtifact(
content_digest=content_digest,
artifact_key=artifact_key,
feature_hash=123,
original_size=(1536, 1024),
resize_config=KimiK3ResizeConfig(
num_tokens=3,
new_width=6,
new_height=2,
pad_width=0,
pad_height=0,
),
grid_thw=(1, 2, 6),
feature=torch.full((12, 2), value, dtype=torch.float32),
)
def test_kimi_k3_cached_artifact_is_composed_per_prompt():
processor = object.__new__(KimiK3ImageProcessor)
processor.mm_tokens = SimpleNamespace(image_token_id=99)
processor._tokenizer = _Tokenizer()
processor.mm_feature_transport = "cpu"
processor.use_cuda_ipc = False
artifact = _cached_k3_artifact("sha256:" + "ab" * 32, "artifact")
first = processor.compose_request([1, 99, 2], [artifact])
second = processor.compose_request([3, 4, 99, 5], [artifact])
assert first.input_ids != second.input_ids
assert first.mm_items[0].offsets == [(3, 5)]
assert second.mm_items[0].offsets == [(4, 6)]
assert first.mm_items[0].hash == second.mm_items[0].hash == 123
torch.testing.assert_close(first.mm_items[0].feature, second.mm_items[0].feature)
def test_kimi_k3_cached_deferred_artifact_has_model_contract():
processor = object.__new__(KimiK3ImageProcessor)
processor.mm_feature_transport = "cpu"
feature = torch.zeros((3, 2, 2), dtype=torch.uint8)
artifact = processor._make_artifact(
content_digest="sha256:" + "ab" * 32,
artifact_key="sha256:" + "cd" * 32,
original_size=(2, 2),
resize_config={
"num_tokens": 1,
"new_width": 2,
"new_height": 2,
"pad_width": 0,
"pad_height": 0,
},
grid_thw=(1, 1, 1),
feature=feature,
deferred=KimiK3DeferredPreprocessing(
backend="gpu",
image_mean=[0.5, 0.5, 0.5],
image_std=[0.5, 0.5, 0.5],
transparent_bg_config=None,
resize_config={
"num_tokens": 1,
"new_width": 2,
"new_height": 2,
"pad_width": 0,
"pad_height": 0,
},
),
)
config = artifact.deferred
assert config.backend == "gpu"
assert config.resize_config["new_width"] == 2
def test_kimi_k3_normal_cache_path_connects_real_producer_to_model_consumer():
hf_processor = _HFProcessor()
hf_processor.tokenizer = _AnySizeTokenizer()
hf_config = SimpleNamespace(
media_placeholder_token_id=42,
to_dict=lambda: {
"model_type": "kimi_k3",
"architectures": ["KimiK3ForConditionalGeneration"],
},
)
server_args = ServerArgs(
model_path="dummy",
mm_feature_transport="cpu",
mm_process_config={},
mm_io_worker_num=1,
mm_processor_worker_num=0,
tokenizer_worker_num=1,
mm_preprocess_cache_size_mb=1,
)
processor = KimiK3ImageProcessor(
hf_config=hf_config,
server_args=server_args,
_processor=hf_processor,
transport_mode=None,
)
image = Image.new("RGB", (28, 28), color=(1, 2, 3))
encoded_image = io.BytesIO()
image.save(encoded_image, format="PNG")
image_data = ImageData(
url="data:image/png;base64,"
+ base64.b64encode(encoded_image.getvalue()).decode()
)
request = SimpleNamespace(video_data=None, mm_content_hashes=None)
class _Tower(nn.Module):
device = torch.device("cpu")
patch_size = 14
def __init__(self):
super().__init__()
self.patch_embed = SimpleNamespace(
proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.float32))
)
def forward(self, pixel_values, _grid_thws):
return pixel_values
model = KimiK3ForConditionalGeneration.__new__(KimiK3ForConditionalGeneration)
nn.Module.__init__(model)
model.use_data_parallel = False
model.vision_tower = _Tower()
model.mm_projector = _Projector()
try:
with (
patch(
"sglang.srt.multimodal.processors.kimi_k3.is_cuda", return_value=True
),
patch.object(
processor,
"prepare_artifact_batch",
wraps=processor.prepare_artifact_batch,
) as prepare_artifacts,
):
cold = asyncio.run(
processor.process_mm_data_async([image_data], [1, 42, 2], request)
)
hot = asyncio.run(
processor.process_mm_data_async([image_data], [3, 42, 4], request)
)
cold_items = pickle.loads(pickle.dumps(cold.mm_items))
hot_items = pickle.loads(pickle.dumps(hot.mm_items))
with (
patch("sglang.srt.models.kimi_k3.configured_tp_size", return_value=1),
patch(
"sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images",
return_value=(
torch.ones((4, 3), dtype=torch.float32),
torch.tensor([[1, 2, 2]], dtype=torch.int64),
),
),
):
cold_features = model.get_image_feature(cold_items)
hot_features = model.get_image_feature(hot_items)
finally:
processor.shutdown()
assert prepare_artifacts.call_count == 1
assert cold.mm_items[0].hash == hot.mm_items[0].hash
assert cold.mm_items[0].offsets == hot.mm_items[0].offsets == [(3, 3)]
assert (
cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend == "gpu"
)
torch.testing.assert_close(cold_features, hot_features)
def test_kimi_k3_model_accepts_mixed_cached_eager_and_deferred_artifacts():
class _Tower(nn.Module):
device = torch.device("cpu")
patch_size = 2
def __init__(self):
super().__init__()
self.patch_embed = SimpleNamespace(
proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.float32))
)
def forward(self, pixel_values, _grid_thws):
return pixel_values
model = KimiK3ForConditionalGeneration.__new__(KimiK3ForConditionalGeneration)
nn.Module.__init__(model)
model.use_data_parallel = False
model.vision_tower = _Tower()
model.mm_projector = _Projector()
eager = _image_item(torch.ones((1, 3)), [[1, 1, 1]])
deferred = _image_item(torch.zeros((3, 2, 2), dtype=torch.uint8), [[1, 1, 1]])
deferred.model_specific_data[DEFERRED_PREPROCESSING_KEY] = (
KimiK3DeferredPreprocessing(
backend="gpu",
image_mean=[0.5, 0.5, 0.5],
image_std=[0.5, 0.5, 0.5],
transparent_bg_config=None,
resize_config={
"num_tokens": 1,
"new_width": 2,
"new_height": 2,
"pad_width": 0,
"pad_height": 0,
},
)
)
with (
patch("sglang.srt.models.kimi_k3.configured_tp_size", return_value=1),
patch(
"sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images",
return_value=(torch.full((1, 3), 2.0), torch.tensor([[1, 1, 1]])),
),
):
output = model.get_image_feature([eager, deferred])
torch.testing.assert_close(output, torch.tensor([[1.0, 1.0, 1.0], [2.0, 2.0, 2.0]]))
def test_kimi_k3_trusted_hot_hit_skips_media_read():
processor = object.__new__(KimiK3ImageProcessor)
processor.processor_fingerprint = "processor"
processor.trust_mm_content_hashes = True
processor.mm_preprocess_cache = MultimodalPreprocessCache(1024 * 1024)
processor.io_executor = ThreadPoolExecutor(max_workers=1)
digest = "sha256:" + "ab" * 32
key = processor._artifact_key(digest, "unread-source")
artifact = _cached_k3_artifact(digest, key)
processor.mm_preprocess_cache.put(key, artifact)
try:
with patch(
"sglang.srt.multimodal.media_artifacts.base.snapshot_media",
side_effect=AssertionError("trusted cache hit must not read media"),
):
result = asyncio.run(
processor.prepare_media_artifacts(
["unread-source"],
content_hashes=[digest],
)
)
finally:
processor.io_executor.shutdown()
assert result == [artifact]
def test_kimi_k3_default_media_options_share_one_artifact_key():
processor = object.__new__(KimiK3ImageProcessor)
processor.processor_fingerprint = "processor"
digest = "sha256:" + "ab" * 32
keys = {
processor._artifact_key(digest, "image.png"),
processor._artifact_key(digest, ImageData(url="image.png")),
processor._artifact_key(digest, {"url": "image.png", "detail": "auto"}),
}
assert len(keys) == 1
def test_kimi_k3_output_affecting_media_options_do_not_share_artifacts():
processor = object.__new__(KimiK3ImageProcessor)
processor.processor_fingerprint = "processor"
digest = "sha256:" + "ab" * 32
base = processor._artifact_key(digest, ImageData(url="image.png"))
assert base != processor._artifact_key(
digest, ImageData(url="image.png", detail="low")
)
assert base != processor._artifact_key(
digest, ImageData(url="image.png", max_dynamic_patch=4)
)
assert base != processor._artifact_key(
digest,
ImageData(url="image.png", preprocess_kwargs={"max_pixels": 1024}),
)
assert base != processor._artifact_key(
digest,
{"url": "image.png", "future_model_option": "new-behavior"},
)
def test_kimi_k3_rejects_changed_feature_hash_for_same_artifact():
processor = object.__new__(KimiK3ImageProcessor)
processor.processor_fingerprint = "processor"
processor.trust_mm_content_hashes = False
processor.mm_preprocess_cache = MultimodalPreprocessCache(1024 * 1024)
processor.mm_processor_executor = None
processor.io_executor = ThreadPoolExecutor(max_workers=2)
image = Image.new("RGB", (2, 2), color=(1, 2, 3))
digest = snapshot_media(image).content_digest
key = processor._artifact_key(digest, image)
old = replace(_cached_k3_artifact(digest, key), feature=None)
new = replace(_cached_k3_artifact(digest, key), feature_hash=old.feature_hash + 1)
processor.mm_preprocess_cache.put(key, old)
async def prepare(_entries):
return [new]
processor._run_preprocess_and_build_artifact_batch = prepare
try:
with pytest.raises(ValueError, match="feature hash changed"):
asyncio.run(processor.prepare_media_artifacts([image]))
finally:
processor.io_executor.shutdown()
assert key not in processor.mm_preprocess_cache
def test_kimi_k3_artifact_and_data_item_share_hash_resolution():
processor = object.__new__(KimiK3ImageProcessor)
processor.mm_feature_transport = "cpu"
processor.mm_tokens = SimpleNamespace(image_token_id=99)
processor._tokenizer = _Tokenizer()
processor.use_cuda_ipc = False
feature = torch.zeros((4, 3), dtype=torch.float32)
digest = "sha256:" + "ab" * 32
artifact = processor._make_artifact(
content_digest=digest,
artifact_key="sha256:" + "01" * 32,
original_size=(2, 2),
resize_config={
"num_tokens": 1,
"new_width": 2,
"new_height": 2,
"pad_width": 0,
"pad_height": 0,
},
grid_thw=(1, 1, 1),
feature=feature,
)
direct_item = MultimodalDataItem(modality=Modality.IMAGE, feature=feature)
direct_item.set_pad_value()
composed_item = processor.compose_request([1, 99, 2], [artifact]).mm_items[0]
expected_hash = resolve_multimodal_item_hash(
existing_hash=direct_item.hash,
namespace=artifact.artifact_key,
)
expected_item = MultimodalDataItem(modality=Modality.IMAGE, hash=expected_hash)
expected_item.set_pad_value()
assert artifact.feature_hash == composed_item.hash == expected_hash
assert composed_item.pad_value == expected_item.pad_value
def test_kimi_k3_untrusted_path_change_is_a_cache_miss():
processor = object.__new__(KimiK3ImageProcessor)
processor.processor_fingerprint = "processor"
processor.trust_mm_content_hashes = False
processor.mm_preprocess_cache = MultimodalPreprocessCache(1024 * 1024)
processor.mm_processor_executor = None
processor.io_executor = ThreadPoolExecutor(max_workers=2)
async def prepare(entries):
return [
_cached_k3_artifact(
entry.content_digest,
entry.artifact_key,
entry.media.getpixel((0, 0))[0],
)
for entry in entries
]
processor._run_preprocess_and_build_artifact_batch = prepare
try:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "mutable.png"
Image.new("RGB", (2, 2), color=(1, 0, 0)).save(path)
first = asyncio.run(processor.prepare_media_artifacts([str(path)]))[0]
Image.new("RGB", (2, 2), color=(2, 0, 0)).save(path)
second = asyncio.run(processor.prepare_media_artifacts([str(path)]))[0]
finally:
processor.io_executor.shutdown()
assert first.content_digest != second.content_digest
assert first.artifact_key != second.artifact_key
assert first.feature[0, 0].item() == 1
assert second.feature[0, 0].item() == 2
def test_kimi_k3_partial_hits_deduplicate_misses_and_preserve_order():
processor = object.__new__(KimiK3ImageProcessor)
processor.processor_fingerprint = "processor"
processor.trust_mm_content_hashes = False
processor.mm_preprocess_cache = MultimodalPreprocessCache(1024 * 1024)
processor.mm_processor_executor = None
processor.io_executor = ThreadPoolExecutor(max_workers=4)
cached_image = Image.new("RGB", (2, 2), color=(1, 0, 0))
missed_image = Image.new("RGB", (2, 2), color=(2, 0, 0))
cached_digest = snapshot_media(cached_image).content_digest
missed_digest = snapshot_media(missed_image).content_digest
cached_key = processor._artifact_key(cached_digest, cached_image)
cached = _cached_k3_artifact(cached_digest, cached_key, value=1)
processor.mm_preprocess_cache.put(cached_key, cached)
batches = []
async def prepare(entries):
batches.append(entries)
return [
replace(
_cached_k3_artifact(
entry.content_digest,
entry.artifact_key,
entry.media.getpixel((0, 0))[0],
),
feature_hash=456,
)
for entry in entries
]
processor._run_preprocess_and_build_artifact_batch = prepare
try:
artifacts = asyncio.run(
processor.prepare_media_artifacts(
[cached_image, missed_image, missed_image, cached_image],
)
)
finally:
processor.io_executor.shutdown()
assert len(batches) == 1
assert len(batches[0]) == 1
assert batches[0][0].content_digest == missed_digest
assert batches[0][0].artifact_key == processor._artifact_key(
missed_digest, missed_image
)
assert [artifact.content_digest for artifact in artifacts] == [
cached_digest,
missed_digest,
missed_digest,
cached_digest,
]
assert artifacts[0] is artifacts[3]
assert artifacts[1] is artifacts[2]
def test_kimi_k3_cancelled_artifact_owner_does_not_fail_joiner():
processor = object.__new__(KimiK3ImageProcessor)
processor.processor_fingerprint = "processor"
processor.trust_mm_content_hashes = False
processor.mm_preprocess_cache = MultimodalPreprocessCache(1024 * 1024)
processor.mm_processor_executor = None
processor.io_executor = ThreadPoolExecutor(max_workers=2)
processor._preprocess_metrics_callback = None
image = Image.new("RGB", (2, 2), color=(1, 2, 3))
started = asyncio.Event()
release = asyncio.Event()
async def prepare(entries):
started.set()
await release.wait()
return [
_cached_k3_artifact(
entry.content_digest,
entry.artifact_key,
entry.media.getpixel((0, 0))[0],
)
for entry in entries
]
processor._run_preprocess_and_build_artifact_batch = prepare
async def run():
owner = asyncio.create_task(processor.prepare_media_artifacts([image]))
await started.wait()
joiner = asyncio.create_task(processor.prepare_media_artifacts([image]))
await asyncio.sleep(0)
owner.cancel()
with pytest.raises(asyncio.CancelledError):
await owner
release.set()
artifacts = await joiner
assert len(artifacts) == 1
assert artifacts[0].feature[0, 0].item() == 1
try:
asyncio.run(run())
finally:
processor.io_executor.shutdown()
def test_kimi_k3_cpu_transport_defers_gpu_preprocessing():
processor = object.__new__(KimiK3ImageProcessor)
processor.mm_tokens = SimpleNamespace(image_token_id=99)
processor.mm_feature_transport = "cpu"
processor.use_cuda_ipc = False
processor._processor = SimpleNamespace(
_patch_size=2,
preprocess_config=_k3_preprocess_config(patch_size=2),
prepare_deferred=Mock(
return_value=(
torch.tensor([[1, 99, 99, 2, 99, 3]]),
@@ -776,11 +1317,7 @@ def test_kimi_k3_defers_only_when_raw_transport_is_smaller(
processor = object.__new__(KimiK3ImageProcessor)
processor.mm_feature_transport = "cpu"
processor._processor = SimpleNamespace(
_patch_size=14,
_merge_kernel_size=2,
_in_patch_limit=in_patch_limit,
_patch_limit_on_one_side=512,
_fixed_output_tokens=None,
preprocess_config=_k3_preprocess_config(in_patch_limit=in_patch_limit),
)
image = torch.zeros(image_shape, dtype=torch.uint8)
@@ -817,7 +1354,7 @@ def test_kimi_k3_eager_preprocessing_preserves_float_tensor_support():
assert output.shape == (3, 4, 4)
@pytest.mark.parametrize("transport", ["cuda_ipc", "fabric"])
@pytest.mark.parametrize("transport", ["cuda_ipc", "cuda_vmm"])
def test_kimi_k3_keeps_gpu_transport_preprocessing_eager(transport):
processor = object.__new__(KimiK3ImageProcessor)
processor.mm_feature_transport = transport
@@ -831,6 +1368,7 @@ def test_kimi_k3_keeps_gpu_transport_preprocessing_eager(transport):
def test_kimi_k3_rejects_silently_dropped_images():
processor = object.__new__(KimiK3ImageProcessor)
processor.mm_tokens = Mock()
processor.mm_preprocess_cache = MultimodalPreprocessCache(0)
processor.load_mm_data = AsyncMock(return_value=SimpleNamespace(images=[object()]))
with pytest.raises(ValueError, match="expected 2, loaded 1"):
@@ -838,13 +1376,14 @@ def test_kimi_k3_rejects_silently_dropped_images():
processor.process_mm_data_async(
image_data=["image-1", "image-2"],
input_text="<|media_pad|><|media_pad|>",
request_obj=SimpleNamespace(video_data=None),
request_obj=SimpleNamespace(video_data=None, mm_content_hashes=None),
)
)
def test_kimi_k3_uses_token_ids_to_preserve_media_boundaries():
processor = object.__new__(KimiK3ImageProcessor)
processor.mm_preprocess_cache = MultimodalPreprocessCache(0)
processor.mm_feature_transport = "cpu"
processor.mm_tokens = SimpleNamespace(image_token_id=99)
processor.mm_feature_transport = "cuda_ipc"
@@ -863,7 +1402,7 @@ def test_kimi_k3_uses_token_ids_to_preserve_media_boundaries():
processor.process_mm_data_async(
image_data=["image-1", "image-2"],
input_text=[1, 99, 2, 99, 3],
request_obj=SimpleNamespace(video_data=None),
request_obj=SimpleNamespace(video_data=None, mm_content_hashes=None),
)
)
@@ -69,8 +69,11 @@ def make_processor(config, image_processor_cls=None):
model_impl="sglang",
keep_mm_feature_on_device=False,
mm_feature_transport="cpu",
image_processor_backend="auto",
disable_fast_image_processor=True,
skip_tokenizer_init=False,
mm_preprocess_cache_size_mb=0,
trust_mm_content_hashes=False,
# Read by NativeMmHost._use_feature_shm (single-rank fixture → the
# inline zero-copy transport, like the 1-GPU e2e).
tp_size=1,
@@ -80,6 +83,7 @@ def make_processor(config, image_processor_cls=None):
mm_processor_worker_num=1,
tokenizer_worker_num=1,
base_gpu_id=0,
rl_on_policy_target=None,
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
@@ -0,0 +1,238 @@
import asyncio
import unittest
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, replace
from typing import Optional
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.multimodal.cache import MultimodalPreprocessCache, snapshot_media
from sglang.srt.multimodal.media_artifacts import (
MediaArtifactCacheMixin,
MediaArtifactInput,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@dataclass(frozen=True)
class _Artifact:
content_digest: str
artifact_key: str
feature_hash: int
feature: Optional[bytes]
@property
def has_feature(self) -> bool:
return self.feature is not None
def cache_value(self):
return self
def cache_size_items(self):
return (
self.content_digest,
self.artifact_key,
self.feature_hash,
self.feature,
)
@dataclass(frozen=True)
class _FutureMediaInput:
url: str
content_hash: Optional[str] = None
frame_sampling: int = 2
class _Processor(MediaArtifactCacheMixin):
artifact_modality = Modality.IMAGE
artifact_option_defaults = {"detail": "auto", "frame_sampling": 2}
def __init__(self):
self.processor_fingerprint = "processor"
self.trust_mm_content_hashes = False
self.mm_preprocess_cache = MultimodalPreprocessCache(1024 * 1024)
self.mm_processor_executor = None
self.io_executor = ThreadPoolExecutor(max_workers=4)
self.batches = []
def decode_media_snapshot(self, snapshot, modality):
self.assert_artifact_modality(modality)
return snapshot.data
@staticmethod
def assert_artifact_modality(modality):
if modality != Modality.IMAGE:
raise AssertionError(f"unexpected modality: {modality}")
def prepare_artifact_batch(
self, entries: list[MediaArtifactInput]
) -> list[_Artifact]:
self.batches.append(entries)
return [
_Artifact(
content_digest=entry.content_digest,
artifact_key=entry.artifact_key,
feature_hash=int(entry.content_digest[-16:], 16),
feature=entry.media,
)
for entry in entries
]
def close(self):
self.io_executor.shutdown()
class TestMediaArtifactProcessor(unittest.TestCase):
def test_unknown_model_option_is_part_of_artifact_identity(self):
processor = _Processor()
digest = snapshot_media(b"image").content_digest
try:
base = processor._artifact_key(digest, _FutureMediaInput(url="image.png"))
self.assertEqual(
base,
processor._artifact_key(
digest,
{
"url": "image.png",
"content_hash": digest,
"frame_sampling": 2,
},
),
)
self.assertNotEqual(
base,
processor._artifact_key(
digest,
{"url": "image.png", "future_model_knob": "different"},
),
)
self.assertNotEqual(
base,
processor._artifact_key(
digest,
_FutureMediaInput(url="image.png"),
modality=Modality.VIDEO,
),
)
finally:
processor.close()
def test_non_image_models_can_override_identity_and_decode_hooks(self):
class _VideoProcessor(_Processor):
artifact_modality = Modality.VIDEO
def snapshot_media_source(self, source, modality):
self.assert_video_modality(modality)
return snapshot_media(source.encode())
def decode_media_snapshot(self, snapshot, modality):
self.assert_video_modality(modality)
return snapshot.data
@staticmethod
def assert_video_modality(modality):
if modality != Modality.VIDEO:
raise AssertionError(f"unexpected modality: {modality}")
processor = _VideoProcessor()
try:
artifacts = asyncio.run(processor.prepare_media_artifacts(["clip.mp4"]))
finally:
processor.close()
self.assertEqual(len(artifacts), 1)
self.assertEqual(
artifacts[0].content_digest, snapshot_media(b"clip.mp4").content_digest
)
def test_partial_hits_and_duplicate_misses_are_shared_by_contract(self):
processor = _Processor()
first_digest = snapshot_media(b"first").content_digest
first_key = processor._artifact_key(first_digest, b"first")
first = _Artifact(first_digest, first_key, 1, b"first")
processor.mm_preprocess_cache.put(first_key, first)
try:
artifacts = asyncio.run(
processor.prepare_media_artifacts(
[b"first", b"second", b"second", b"first"]
)
)
finally:
processor.close()
self.assertEqual(len(processor.batches), 1)
self.assertEqual(len(processor.batches[0]), 1)
self.assertEqual(
[artifact.content_digest for artifact in artifacts],
[
first_digest,
snapshot_media(b"second").content_digest,
snapshot_media(b"second").content_digest,
first_digest,
],
)
self.assertIs(artifacts[0], artifacts[3])
self.assertIs(artifacts[1], artifacts[2])
def test_adapter_cannot_change_validated_artifact_identity(self):
processor = _Processor()
async def wrong_identity(entries):
artifact = processor.prepare_artifact_batch(entries)[0]
return [replace(artifact, artifact_key="sha256:" + "0" * 64)]
processor._run_preprocess_and_build_artifact_batch = wrong_identity
try:
with self.assertRaisesRegex(ValueError, "changed the media artifact key"):
asyncio.run(processor.prepare_media_artifacts([b"image"]))
finally:
processor.close()
def test_trusted_hit_uses_identity_without_reading_source(self):
processor = _Processor()
processor.trust_mm_content_hashes = True
digest = snapshot_media(b"cached").content_digest
key = processor._artifact_key(digest, "unread-source")
artifact = _Artifact(digest, key, 1, b"cached")
processor.mm_preprocess_cache.put(key, artifact)
try:
artifacts = asyncio.run(
processor.prepare_media_artifacts(
["unread-source"], content_hashes=[digest]
)
)
finally:
processor.close()
self.assertEqual(artifacts, [artifact])
self.assertEqual(processor.batches, [])
def test_cached_artifact_must_match_content_identity(self):
processor = _Processor()
digest = snapshot_media(b"fresh").content_digest
key = processor._artifact_key(digest, b"fresh")
processor.mm_preprocess_cache.put(
key,
_Artifact(
snapshot_media(b"stale").content_digest,
key,
1,
b"stale",
),
)
try:
artifacts = asyncio.run(processor.prepare_media_artifacts([b"fresh"]))
finally:
processor.close()
self.assertEqual(artifacts[0].content_digest, digest)
self.assertEqual(artifacts[0].feature, b"fresh")
self.assertEqual(len(processor.batches), 1)
if __name__ == "__main__":
unittest.main()
@@ -4,22 +4,24 @@ import os
import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch
import numpy as np
import torch
from PIL import Image
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.multimodal.cache import (
CacheMiss,
MultimodalPreprocessCache,
build_artifact_key,
build_feature_hash,
build_processor_fingerprint,
estimate_cache_size_bytes,
parse_content_hash,
resolve_multimodal_item_hash,
snapshot_media,
)
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
@@ -209,21 +211,25 @@ class TestMediaIdentity(unittest.TestCase):
def preprocess_fingerprint_payload(self):
return {"backend": self.backend, "antialias": True}
config = SimpleNamespace(model_type="vlm", architectures=["VLM"])
args = SimpleNamespace(
class Config:
def to_dict(self):
return {"model_type": "vlm", "architectures": ["VLM"]}
config = Config()
args = ServerArgs(
model_path="dummy",
revision="model-revision",
tokenizer_revision="tokenizer-revision",
disable_fast_image_processor=False,
mm_process_config={"image": {"max_pixels": 1024}},
)
base = build_processor_fingerprint(Processor("gpu"), config, args)
changed_backend = build_processor_fingerprint(Processor("cpu"), config, args)
changed_args = SimpleNamespace(
**{
**vars(args),
"mm_process_config": {"image": {"max_pixels": 2048}},
}
changed_args = ServerArgs(
model_path="dummy",
revision="model-revision",
disable_fast_image_processor=False,
mm_process_config={"image": {"max_pixels": 2048}},
)
changed_config = build_processor_fingerprint(
Processor("gpu"), config, changed_args
@@ -231,7 +237,7 @@ class TestMediaIdentity(unittest.TestCase):
self.assertNotEqual(base, changed_backend)
self.assertNotEqual(base, changed_config)
def test_feature_hash_includes_artifact_and_processor_output(self):
def test_item_hash_namespace_covers_identity_and_processor_output(self):
digest = snapshot_media(b"image").content_digest
first = build_artifact_key(
digest,
@@ -243,11 +249,25 @@ class TestMediaIdentity(unittest.TestCase):
modality="image",
processor_fingerprint="processor-b",
)
self.assertNotEqual(build_feature_hash(first, 1), build_feature_hash(second, 1))
self.assertNotEqual(build_feature_hash(first, 1), build_feature_hash(first, 2))
self.assertIsInstance(build_feature_hash(first, 1 << 128), int)
self.assertNotEqual(
resolve_multimodal_item_hash(existing_hash=1, namespace=first),
resolve_multimodal_item_hash(existing_hash=1, namespace=second),
)
self.assertNotEqual(
resolve_multimodal_item_hash(existing_hash=1, namespace=first),
resolve_multimodal_item_hash(existing_hash=2, namespace=first),
)
with self.assertRaises(ValueError):
build_feature_hash(first, -1)
resolve_multimodal_item_hash(existing_hash=-1, namespace=first)
def test_multimodal_data_item_uses_shared_feature_hash(self):
feature = torch.arange(12, dtype=torch.float32).reshape(4, 3)
expected = resolve_multimodal_item_hash(feature=feature)
item = MultimodalDataItem(modality=Modality.IMAGE, feature=feature)
item.set_pad_value()
self.assertEqual(item.hash, expected)
class TestMultimodalPreprocessCache(unittest.TestCase):
@@ -262,6 +282,31 @@ class TestMultimodalPreprocessCache(unittest.TestCase):
self.assertIn("c", cache)
self.assertEqual(cache.current_size_bytes, 6)
def test_compatible_lookup_is_atomic_and_does_not_count_bypass_as_miss(self):
cache = MultimodalPreprocessCache[str, bytes](max_size_bytes=1024)
cache.put("key", b"metadata-only")
self.assertIsNone(cache.get_if_present("key", lambda value: False))
self.assertEqual((cache.hits, cache.misses), (0, 0))
self.assertEqual(
cache.get_if_present("key", lambda value: value.startswith(b"metadata")),
b"metadata-only",
)
self.assertEqual((cache.hits, cache.misses), (1, 0))
def test_claimed_miss_rejects_an_incompatible_racing_entry(self):
cache = MultimodalPreprocessCache[str, bytes](max_size_bytes=1024)
cache.put("key", b"metadata-only")
miss = cache.lookup_or_claim_many(
["key"], predicate=lambda key, value: value == b"full-feature"
)[0]
self.assertIsInstance(miss, CacheMiss)
self.assertTrue(miss.should_compute)
self.assertNotIn("key", cache)
self.assertEqual(cache.current_size_bytes, 0)
def test_gpu_backed_values_are_not_implicitly_copied(self):
if not torch.cuda.is_available():
self.skipTest("CUDA is not available")
@@ -370,6 +415,61 @@ class TestMultimodalPreprocessCache(unittest.TestCase):
asyncio.run(run())
def test_lookup_or_claim_many_batches_owned_and_joined_misses(self):
async def run():
cache = MultimodalPreprocessCache[str, bytes](max_size_bytes=1024)
results = cache.lookup_or_claim_many(["a", "b", "a"])
misses_to_compute = [
item
for item in results
if isinstance(item, CacheMiss) and item.should_compute
]
self.assertEqual([item.key for item in misses_to_compute], ["a", "b"])
cache.complete_miss(misses_to_compute[0], b"value-a")
cache.complete_miss(misses_to_compute[1], b"value-b")
self.assertEqual(await cache.wait_for_miss(results[2]), b"value-a")
self.assertEqual(cache.get("b"), b"value-b")
asyncio.run(run())
def test_cancelled_miss_waiter_does_not_cancel_computing_caller(self):
async def run():
cache = MultimodalPreprocessCache[str, bytes](max_size_bytes=1024)
computing_miss = cache.lookup_or_claim_many(["key"])[0]
waiting_miss = cache.lookup_or_claim_many(["key"])[0]
self.assertTrue(computing_miss.should_compute)
self.assertFalse(waiting_miss.should_compute)
waiter = asyncio.create_task(cache.wait_for_miss(waiting_miss))
await asyncio.sleep(0)
waiter.cancel()
with self.assertRaises(asyncio.CancelledError):
await waiter
cache.complete_miss(computing_miss, b"artifact")
self.assertEqual(computing_miss.future.result(), b"artifact")
self.assertEqual(cache.get("key"), b"artifact")
asyncio.run(run())
def test_disabled_cache_does_not_join_or_retain(self):
async def run():
cache = MultimodalPreprocessCache[str, bytes](max_size_bytes=0)
misses = cache.lookup_or_claim_many(["a", "a"])
self.assertTrue(
all(
isinstance(item, CacheMiss) and item.should_compute
for item in misses
)
)
for item in misses:
cache.complete_miss(item, b"value")
self.assertEqual(len(cache), 0)
self.assertEqual(cache.stats()["singleflight_joins"], 0)
asyncio.run(run())
def test_clear_starts_a_new_singleflight_generation(self):
async def run():
cache = MultimodalPreprocessCache[str, bytes](max_size_bytes=1024)
@@ -397,6 +497,20 @@ class TestMultimodalPreprocessCache(unittest.TestCase):
asyncio.run(run())
def test_clear_starts_a_new_cache_miss_generation(self):
cache = MultimodalPreprocessCache[str, bytes](max_size_bytes=1024)
old = cache.lookup_or_claim_many(["key"])[0]
cache.clear()
new = cache.lookup_or_claim_many(["key"])[0]
self.assertTrue(old.should_compute)
self.assertTrue(new.should_compute)
self.assertIsNot(old.future, new.future)
cache.complete_miss(old, b"old")
self.assertNotIn("key", cache)
cache.complete_miss(new, b"new")
self.assertEqual(cache.get("key"), b"new")
if __name__ == "__main__":
unittest.main()