[mm] Handle per-item embeddings in cache misses (#32498)
Co-authored-by: cctry <cctry@meta.com>
This commit is contained in:
@@ -671,10 +671,23 @@ def _batch_encode_per_image_misses(
|
|||||||
if not _can_skip_pre_embed_feature_move(data_embedding_func):
|
if not _can_skip_pre_embed_feature_move(data_embedding_func):
|
||||||
_move_items_to_device(miss_items, device)
|
_move_items_to_device(miss_items, device)
|
||||||
all_miss_embedding = data_embedding_func(miss_items)
|
all_miss_embedding = data_embedding_func(miss_items)
|
||||||
|
|
||||||
|
if isinstance(all_miss_embedding, list):
|
||||||
|
# Per-item embeddings: no split needed, and each cache entry owns
|
||||||
|
# its storage (a torch.split view would pin the whole concatenated
|
||||||
|
# buffer for as long as any single item stays cached). Mirrors
|
||||||
|
# _get_chunked_embedding_by_item.
|
||||||
|
assert len(all_miss_embedding) == len(miss_items), (
|
||||||
|
f"per-item embedding count {len(all_miss_embedding)} != "
|
||||||
|
f"cache-miss item count {len(miss_items)}"
|
||||||
|
)
|
||||||
|
split_embeddings = [
|
||||||
|
emb.reshape(-1, emb.shape[-1]) for emb in all_miss_embedding
|
||||||
|
]
|
||||||
|
else:
|
||||||
all_miss_embedding = all_miss_embedding.reshape(
|
all_miss_embedding = all_miss_embedding.reshape(
|
||||||
-1, all_miss_embedding.shape[-1]
|
-1, all_miss_embedding.shape[-1]
|
||||||
)
|
)
|
||||||
|
|
||||||
split_embeddings = torch.split(all_miss_embedding, token_counts, dim=0)
|
split_embeddings = torch.split(all_miss_embedding, token_counts, dim=0)
|
||||||
for h, emb in zip(ordered_hashes, split_embeddings):
|
for h, emb in zip(ordered_hashes, split_embeddings):
|
||||||
embedding_cache.set(h, EmbeddingResult(embedding=emb))
|
embedding_cache.set(h, EmbeddingResult(embedding=emb))
|
||||||
|
|||||||
Reference in New Issue
Block a user