[Fix] Re-encode multimodal embeddings after cache mismatch (#36595)

Co-authored-by: john.xh <dengxuhuijohn@gmail.com>
This commit is contained in:
Mick
2026-08-27 23:24:30 +08:00
committed by GitHub
co-authored by john.xh
parent 024a7a1031
commit 20a491d1d3
2 changed files with 182 additions and 22 deletions
+82 -22
View File
@@ -160,6 +160,29 @@ def _flatten_embedding_result(
return embedding
def _embedding_token_count(embedding: torch.Tensor) -> int:
"""Return the number of multimodal tokens represented by an embedding."""
# Vision encoders may return [tokens, hidden] or a higher-rank tensor. The
# scheduler always consumes the flattened token dimension.
return embedding.reshape(-1, embedding.shape[-1]).shape[0]
def _discard_mismatched_cached_embedding(
cache_key: Optional[int],
expected_token_count: int,
cached_token_count: int,
) -> None:
"""Log and remove a cache entry that cannot serve the current item."""
logger.warning(
"Discarding cached multimodal embedding due to a token-count mismatch: "
"cache_key=%s, expected_tokens=%d, cached_tokens=%d. Recomputing embedding.",
cache_key,
expected_token_count,
cached_token_count,
)
embedding_cache.free(cache_key, None)
def _can_skip_pre_embed_feature_move(data_embedding_func: DataEmbeddingFunc) -> bool:
"""Models that materialize and batch visual features inside their encoder.
@@ -230,6 +253,21 @@ def _get_chunked_embedding_full(
embedding_items_hash = MultiModalStaticCache.combine_hashes(item_hashes)
embedding_per_req = embedding_cache.get(item_hashes)
# A compact feature hash can collide for inputs with different token
# counts. Never feed a stale cache entry into the scheduler: the length
# mismatch would otherwise surface much later in _adjust_embedding_length
# as an unrecoverable prefill crash.
if embedding_per_req is not None and not isinstance(
embedding_per_req, EVSEmbeddingResult
):
expected_token_count = sum(end - start + 1 for start, end in items_offset)
cached_token_count = _embedding_token_count(embedding_per_req.embedding)
if cached_token_count != expected_token_count:
_discard_mismatched_cached_embedding(
embedding_items_hash, expected_token_count, cached_token_count
)
embedding_per_req = None
if embedding_per_req is None:
if not _can_skip_pre_embed_feature_move(data_embedding_func):
_move_items_to_device(embedding_items_per_req, device)
@@ -286,16 +324,19 @@ def _batch_encode_per_image_misses(
data_embedding_func: DataEmbeddingFunc,
per_image_requests: List[PerImageRequestInfo],
device: torch.device,
) -> Dict[int, torch.Tensor]:
) -> Dict[Tuple[Optional[int], int], torch.Tensor]:
"""
Collect cache misses across ALL per-image requests, deduplicate by hash,
encode in a single ViT call, and populate the cache.
Collect cache misses across ALL per-image requests, deduplicate by hash and
expected token count, encode in a single ViT call, and populate the cache.
Returns:
hash_to_embedding: mapping from item.hash to its full embedding tensor.
hash_to_embedding: mapping from (item.hash, token_count) to its full
embedding tensor. Including the token count prevents two
colliding hashes with different placeholder spans from being
deduplicated within the same batch.
"""
unique_misses: Dict[int, Tuple[MultimodalDataItem, int]] = {}
hash_to_embedding: Dict[int, torch.Tensor] = {}
unique_misses: Dict[Tuple[Optional[int], int], Tuple[MultimodalDataItem, int]] = {}
hash_to_embedding: Dict[Tuple[Optional[int], int], torch.Tensor] = {}
# Phase 1a: find overlapping items per request and collect cache misses
for req_info in per_image_requests:
@@ -311,20 +352,29 @@ def _batch_encode_per_image_misses(
req_info.overlapping = overlapping
for _idx, item, start, end in overlapping:
if item.hash in hash_to_embedding:
expected_token_count = end - start + 1
cache_key = (item.hash, expected_token_count)
if cache_key in hash_to_embedding:
continue
cached = embedding_cache.get_single(item.hash)
if cached is not None:
hash_to_embedding[item.hash] = cached.embedding
elif item.hash not in unique_misses:
token_count = end - start + 1
unique_misses[item.hash] = (item, token_count)
cached_embedding = cached.embedding
cached_token_count = _embedding_token_count(cached_embedding)
if cached_token_count == expected_token_count:
hash_to_embedding[cache_key] = cached_embedding
else:
_discard_mismatched_cached_embedding(
item.hash, expected_token_count, cached_token_count
)
unique_misses[cache_key] = (item, expected_token_count)
elif cache_key not in unique_misses:
unique_misses[cache_key] = (item, expected_token_count)
# Phase 1b: single ViT call for all unique cache misses
if unique_misses:
ordered_hashes = list(unique_misses.keys())
miss_items = [unique_misses[h][0] for h in ordered_hashes]
token_counts = [unique_misses[h][1] for h in ordered_hashes]
ordered_cache_keys = list(unique_misses.keys())
miss_items = [unique_misses[key][0] for key in ordered_cache_keys]
token_counts = [unique_misses[key][1] for key in ordered_cache_keys]
if not _can_skip_pre_embed_feature_move(data_embedding_func):
_move_items_to_device(miss_items, device)
@@ -347,10 +397,10 @@ def _batch_encode_per_image_misses(
-1, all_miss_embedding.shape[-1]
)
split_embeddings = torch.split(all_miss_embedding, token_counts, dim=0)
for h, emb in zip(ordered_hashes, split_embeddings):
embedding_cache.set(h, EmbeddingResult(embedding=emb))
for cache_key, emb in zip(ordered_cache_keys, split_embeddings):
embedding_cache.set(cache_key[0], EmbeddingResult(embedding=emb))
# Keep a local ref (no extra GPU memory) so assembly never fails due to LRU eviction.
hash_to_embedding[h] = emb
hash_to_embedding[cache_key] = emb
return hash_to_embedding
@@ -386,10 +436,19 @@ def _get_chunked_embedding_by_item(
cached_embeddings = {}
miss_items = []
for idx, item, start, end in overlapping:
expected_token_count = end - start + 1
cached = embedding_cache.get_single(item.hash)
if cached is not None:
cached_embeddings[idx] = cached.embedding
_acknowledge_deferred_cuda_ipc_cache_hits([item])
cached_embedding = cached.embedding
cached_token_count = _embedding_token_count(cached_embedding)
if cached_token_count == expected_token_count:
cached_embeddings[idx] = cached_embedding
_acknowledge_deferred_cuda_ipc_cache_hits([item])
else:
_discard_mismatched_cached_embedding(
item.hash, expected_token_count, cached_token_count
)
miss_items.append((idx, item, start, end))
else:
miss_items.append((idx, item, start, end))
@@ -436,7 +495,7 @@ def _get_chunked_embedding_by_item(
def _assemble_per_image_chunk(
overlapping: List[Tuple[int, MultimodalDataItem, int, int]],
hash_to_embedding: Dict[int, torch.Tensor],
hash_to_embedding: Dict[Tuple[Optional[int], int], torch.Tensor],
extend_prefix_len: int,
extend_seq_len: int,
) -> Optional[torch.Tensor]:
@@ -452,7 +511,8 @@ def _assemble_per_image_chunk(
chunk_slices = []
for _idx, item, start, end in overlapping:
emb = hash_to_embedding[item.hash] # shape: (end - start + 1, hidden)
cache_key = (item.hash, end - start + 1)
emb = hash_to_embedding[cache_key] # shape: (end - start + 1, hidden)
overlap_start = max(start, chunk_start)
overlap_end = min(end, chunk_end - 1) # inclusive
local_start = overlap_start - start
@@ -530,7 +590,7 @@ def _get_chunked_prefill_embedding(
full_path_requests.append(req_info)
# Phase 1: batch encode all per-image cache misses in ONE ViT call
hash_to_embedding: Dict[int, torch.Tensor] = {}
hash_to_embedding: Dict[Tuple[Optional[int], int], torch.Tensor] = {}
if per_image_requests:
hash_to_embedding = _batch_encode_per_image_misses(
data_embedding_func, per_image_requests, device
@@ -9,6 +9,9 @@ of the combined tensor pins the whole concatenated buffer).
CPU-only: exercises mm_schedule internals directly, no engine or GPU.
"""
import logging
from unittest.mock import Mock
import pytest
import torch
@@ -185,6 +188,103 @@ def test_tensor_cache_entries_share_storage():
)
def test_by_item_mismatched_cache_entry_is_reencoded():
mm_schedule.init_mm_embedding_cache(1 << 30)
items = _make_items()
first_item = items[0]
mm_schedule.embedding_cache.set(
first_item.hash,
mm_schedule.EmbeddingResult(embedding=torch.zeros(1, HIDDEN)),
)
encoder = Mock(side_effect=_encoder_list)
chunk = mm_schedule._get_chunked_embedding_by_item(
encoder, items, ITEM_OFFSETS, 0, TOTAL_LEN, _CPU
)
assert chunk.shape == (sum(_num_tokens(item) for item in items), HIDDEN)
encoder.assert_called_once()
assert mm_schedule.embedding_cache.get_single(first_item.hash).embedding.shape[
0
] == _num_tokens(first_item)
def test_batched_mismatched_cache_entry_is_reencoded():
mm_schedule.init_mm_embedding_cache(1 << 30)
items = _make_items()
first_item = items[0]
mm_schedule.embedding_cache.set(
first_item.hash,
mm_schedule.EmbeddingResult(embedding=torch.zeros(1, HIDDEN)),
)
request = mm_schedule.PerImageRequestInfo(
req_idx=0,
items=items,
items_offset=ITEM_OFFSETS,
extend_prefix_len=0,
extend_seq_len=TOTAL_LEN,
)
encoder = Mock(side_effect=_encoder_list)
embeddings = mm_schedule._batch_encode_per_image_misses(encoder, [request], _CPU)
assert embeddings[(first_item.hash, _num_tokens(first_item))].shape == (
_num_tokens(first_item),
HIDDEN,
)
encoder.assert_called_once()
def test_batched_colliding_hashes_with_different_lengths_are_not_deduplicated():
mm_schedule.init_mm_embedding_cache(1 << 30)
items = _make_items()
# Simulate the compact-hash collision that motivated the cache guard.
items[1].hash = items[0].hash
requests = [
mm_schedule.PerImageRequestInfo(
req_idx=0,
items=items[:2],
items_offset=ITEM_OFFSETS[:2],
extend_prefix_len=0,
extend_seq_len=TOTAL_LEN,
)
]
encoder = Mock(side_effect=_encoder_list)
embeddings = mm_schedule._batch_encode_per_image_misses(encoder, requests, _CPU)
first_key = (items[0].hash, _num_tokens(items[0]))
second_key = (items[1].hash, _num_tokens(items[1]))
assert embeddings[first_key].shape == (_num_tokens(items[0]), HIDDEN)
assert embeddings[second_key].shape == (_num_tokens(items[1]), HIDDEN)
encoder.assert_called_once()
def test_full_mismatched_cache_entry_is_reencoded(caplog):
mm_schedule.init_mm_embedding_cache(1 << 30)
items = _make_items()
combined_hash = mm_schedule.MultiModalStaticCache.combine_hashes(
[item.hash for item in items]
)
mm_schedule.embedding_cache.set(
combined_hash,
mm_schedule.EmbeddingResult(embedding=torch.zeros(1, HIDDEN)),
)
input_ids = torch.zeros(TOTAL_LEN, dtype=torch.long)
encoder = Mock(side_effect=_encoder_tensor)
with caplog.at_level(logging.WARNING, logger=mm_schedule.logger.name):
chunk, _ = mm_schedule._get_chunked_embedding_full(
encoder, items, ITEM_OFFSETS, 0, TOTAL_LEN, input_ids, _CPU
)
assert chunk.shape == (sum(_num_tokens(item) for item in items), HIDDEN)
encoder.assert_called_once()
assert "Discarding cached multimodal embedding" in caplog.text
assert "expected_tokens=15" in caplog.text
assert "cached_tokens=1" in caplog.text
if __name__ == "__main__":
import sys