Fix multimodal embedding cache retaining full batches through views (#39120)

This commit is contained in:
cctry
2026-09-12 21:37:14 -07:00
committed by GitHub
parent 7e3d18bbcc
commit 24b6c1c7f5
2 changed files with 11 additions and 10 deletions
@@ -1,6 +1,6 @@
import abc import abc
from collections import OrderedDict from collections import OrderedDict
from dataclasses import dataclass from dataclasses import dataclass, replace
from typing import List, Optional from typing import List, Optional
import torch import torch
@@ -65,7 +65,7 @@ class MultimodalCache(abc.ABC):
def _get_tensor_size(embedding: torch.Tensor): def _get_tensor_size(embedding: torch.Tensor):
return embedding.element_size() * embedding.numel() return embedding.untyped_storage().nbytes()
@dataclass(kw_only=True) @dataclass(kw_only=True)
@@ -109,13 +109,19 @@ class MultiModalStaticCache(MultimodalCache):
if mm_hash in self.mm_cache: if mm_hash in self.mm_cache:
self.mm_cache.move_to_end(mm_hash) self.mm_cache.move_to_end(mm_hash)
return True return True
data_size = _get_tensor_size(embedding.embedding) tensor = embedding.embedding
storage_size = _get_tensor_size(tensor)
data_size = min(storage_size, tensor.element_size() * tensor.numel())
while self.current_size + data_size > self.max_size: while self.current_size + data_size > self.max_size:
if not self.mm_cache: if not self.mm_cache:
return False return False
lru_hash, lru_embedding = self.mm_cache.popitem(last=False) lru_hash, lru_embedding = self.mm_cache.popitem(last=False)
self.current_size -= _get_tensor_size(lru_embedding.embedding) self.current_size -= _get_tensor_size(lru_embedding.embedding)
if storage_size > data_size:
# Clone admitted slices so one cached item cannot retain a whole batch.
embedding = replace(embedding, embedding=tensor.clone())
data_size = _get_tensor_size(embedding.embedding)
self.mm_cache[mm_hash] = embedding self.mm_cache[mm_hash] = embedding
self.current_size += data_size self.current_size += data_size
return True return True
@@ -177,20 +177,15 @@ def test_list_cache_entries_own_storage():
assert emb.untyped_storage().nbytes() == own_bytes assert emb.untyped_storage().nbytes() == own_bytes
def test_tensor_cache_entries_share_storage(): def test_tensor_cache_entries_own_storage():
# Documents the motivation for the per-item form: split views of the
# combined tensor keep the whole concatenated buffer alive.
mm_schedule.init_mm_embedding_cache(1 << 30) mm_schedule.init_mm_embedding_cache(1 << 30)
items = _make_items() items = _make_items()
mm_schedule._get_chunked_embedding_by_item( mm_schedule._get_chunked_embedding_by_item(
_encoder_tensor, items, ITEM_OFFSETS, 0, TOTAL_LEN, _CPU _encoder_tensor, items, ITEM_OFFSETS, 0, TOTAL_LEN, _CPU
) )
total_tokens = sum(_num_tokens(item) for item in items)
for item in items: for item in items:
emb = mm_schedule.embedding_cache.get_single(item.hash).embedding emb = mm_schedule.embedding_cache.get_single(item.hash).embedding
assert ( assert emb.untyped_storage().nbytes() == emb.numel() * emb.element_size()
emb.untyped_storage().nbytes() == total_tokens * HIDDEN * emb.element_size()
)
def test_by_item_mismatched_cache_entry_is_reencoded(): def test_by_item_mismatched_cache_entry_is_reencoded():