From 20a491d1d311553bbab3f22e19bbafb86ef3c0cc Mon Sep 17 00:00:00 2001 From: Mick Date: Thu, 27 Aug 2026 23:24:30 +0800 Subject: [PATCH] [Fix] Re-encode multimodal embeddings after cache mismatch (#36595) Co-authored-by: john.xh --- python/sglang/srt/managers/mm_schedule.py | 104 ++++++++++++++---- .../test_mm_chunked_embedding_unit.py | 100 +++++++++++++++++ 2 files changed, 182 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/managers/mm_schedule.py b/python/sglang/srt/managers/mm_schedule.py index 45dd91bbe..dbf6d2f21 100644 --- a/python/sglang/srt/managers/mm_schedule.py +++ b/python/sglang/srt/managers/mm_schedule.py @@ -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 diff --git a/test/registered/chunked_prefill/test_mm_chunked_embedding_unit.py b/test/registered/chunked_prefill/test_mm_chunked_embedding_unit.py index e286ad565..b7245479e 100644 --- a/test/registered/chunked_prefill/test_mm_chunked_embedding_unit.py +++ b/test/registered/chunked_prefill/test_mm_chunked_embedding_unit.py @@ -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