[Fix] Re-encode multimodal embeddings after cache mismatch (#36595)
Co-authored-by: john.xh <dengxuhuijohn@gmail.com>
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user