Skip to content
Draft
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
250 changes: 217 additions & 33 deletions tensorrt_llm/_torch/models/modeling_multimodal_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
# SPDX-License-Identifier: Apache-2.0

import contextlib
import copy
from dataclasses import dataclass
from typing import (
TYPE_CHECKING,
Expand Down Expand Up @@ -115,6 +116,28 @@ class PreparedLlmInputs:
extra_embeds: Sequence[torch.Tensor] = ()


@dataclass(frozen=True)
class EncoderCachePartition:
"""Per-item cache partition for a single `MultimodalParams`.

`hits` maps item index to its cached embedding row-block; `miss_indices` lists item
indices that still require encoder work; `keys` is aligned to item order so miss
embeddings can be written back after they are computed.
"""

hits: Dict[int, torch.Tensor]
miss_indices: list[int]
keys: list[Hashable]

@property
def is_full_hit(self) -> bool:
return bool(self.keys) and not self.miss_indices

@property
def is_full_miss(self) -> bool:
return bool(self.keys) and not self.hits


class MultimodalModelMixin:
"""Template-method mixin for PyTorch multimodal causal LM models.

Expand All @@ -124,10 +147,11 @@ class MultimodalModelMixin:
Current limitations:

* For the time being, the persistent multimodal encoder cache stores per-item embeddings for
single-modality `MultimodalParams` objects.
* Cache reuse is all-or-nothing for each `MultimodalParams` object: every item in that object
hit the cache before cached embeddings are reused. Mixed-modality `MultimodalParams` objects
bypass the persistent cache.
single-modality `MultimodalParams` objects. Mixed-modality objects bypass the cache.
* A partially cached `MultimodalParams` is handled by encoding only its miss items
and interleaving cached items back in original per-item order. The default
`build_multimodal_encoder_input` handles stacked-on-dim-0 and packed-with-grid-thw
layouts; models with other layouts override that method.
"""

supports_encoder_cache: ClassVar[bool] = False
Expand Down Expand Up @@ -240,6 +264,81 @@ def after_active_multimodal_embeddings(
# them as extra embeds without changing the base flow.
return active_embeddings, ()

def build_multimodal_encoder_input(
self,
param: MultimodalParams,
item_indices: Sequence[int],
) -> MultimodalParams:
"""Return a `MultimodalParams` whose raw modality inputs contain only
`item_indices` from `param`, in that order.

Default handles two common single-modality layouts:

- Stacked on dim 0 (Mistral 3 / Pixtral / LLaVA-family): `pixel_values`
`[B, C, H, W]` with a parallel `image_sizes` list; both sliced by item.
- Packed with `*_grid_thw` offsets (Qwen2-VL family): `pixel_values`
`[total_patches, feat]` + `image_grid_thw` `[B, 3]`; prefix-summed patch
counts locate each item's slice, and `image_grid_thw` is sliced in parallel.

Models with a different layout (e.g. mixed-modality per param, custom packed
formats) override this method. The parallel per-item metadata
(`multimodal_embedding_lengths`, `multimodal_hashes`) is re-sliced by the mixin
after this returns, so overrides need only handle modality data.
"""
modality = self._encoder_cache_modality(param)
if modality is None:
raise NotImplementedError(
"Default `build_multimodal_encoder_input` only supports single-modality "
"params. Override for other layouts."
)
modality_data = param.multimodal_data[modality]
if not isinstance(modality_data, dict):
raise TypeError(
f"multimodal_data[{modality!r}] must be a dict, got {type(modality_data).__name__}"
)

indices = list(item_indices)
grid_key = {"image": "image_grid_thw", "video": "video_grid_thw"}.get(modality)
pixel_key = {"image": "pixel_values", "video": "pixel_values_videos"}.get(modality)

if grid_key and pixel_key and grid_key in modality_data and pixel_key in modality_data:
# Packed layout: prefix-sum patch counts to locate each item's slab, then
# concat the requested subset in item-index order.
grids = modality_data[grid_key]
patch_counts = [int(c) for c in torch.prod(grids, dim=1).tolist()]
per_item = torch.split(modality_data[pixel_key], patch_counts, dim=0)
sliced = {
**modality_data,
pixel_key: torch.cat([per_item[i] for i in indices], dim=0),
grid_key: grids[indices],
}
elif (
modality == "image"
and "pixel_values" in modality_data
and "image_sizes" in modality_data
):
# Stacked layout: dim-0 select from `pixel_values` and list-index `image_sizes`.
sliced = {
**modality_data,
"pixel_values": modality_data["pixel_values"][indices],
"image_sizes": [modality_data["image_sizes"][i] for i in indices],
}
else:
raise NotImplementedError(
f"Default `build_multimodal_encoder_input` cannot slice {modality} layout "
f"with fields {sorted(modality_data)}; override this method."
)

# Shallow-copy `multimodal_input` so `_apply_metadata_slice` can rewrite
# `multimodal_hashes` on the residual without mutating the source.
residual_input = (
copy.copy(param.multimodal_input) if param.multimodal_input is not None else None
)
return MultimodalParams(
multimodal_data={**param.multimodal_data, modality: sliced},
multimodal_input=residual_input,
)

# A future optional mixin-owned forward can build on the same template method.
def prepare_multimodal_inputs(
self,
Expand Down Expand Up @@ -312,15 +411,29 @@ def _get_or_encode_multimodal_embeddings(
the single tensor contract for both encoded and cached-only paths.
"""
encoder_cache = self._get_multimodal_encoder_cache()
cache_misses = []
cache_misses: list[MultimodalParams] = []
partial_hits: list[tuple[MultimodalParams, EncoderCachePartition]] = []
if encoder_cache is not None:
for param in multimodal_params:
if param.multimodal_data.get("multimodal_embedding") is not None:
# The forward that attached this request-local embedding already populated the
# persistent cache.
continue
if not self._attach_encoder_cache_hit(param, encoder_cache):
partition = self.partition_encoder_cache(param, encoder_cache)
if partition is None or partition.is_full_miss:
cache_misses.append(param)
continue
if partition.is_full_hit:
param.multimodal_data["multimodal_embedding"] = self.assemble_full_embedding(
partition.hits, len(partition.keys)
)
continue
partial_hits.append((param, partition))

if partial_hits:
# `encoder_cache` is non-None here because partitions are only produced when the cache
# exists.
self._encode_with_partial_cache(partial_hits, encoder_cache)

embeddings = get_multimodal_embeddings(
encoder_forward_fn=self.encode_multimodal_inputs,
Expand Down Expand Up @@ -476,47 +589,118 @@ def _encoder_cache_keys(
]

@classmethod
def _attach_encoder_cache_hit(
def partition_encoder_cache(
cls,
param: MultimodalParams,
encoder_cache: TensorLRUCache,
) -> bool:
"""Attach a full persistent-cache hit and report whether one was found."""
) -> Optional[EncoderCachePartition]:
"""Look up every item of `param` and return the per-item partition.

Returns `None` when the param is not cacheable (mixed modality, missing metadata,
request-local embedding already attached); the caller treats that as a full miss.
"""
if param.multimodal_data.get("multimodal_embedding") is not None:
logger.debug(
f"{_MM_ENCODER_CACHE_LOG_NAME}: request-local multimodal embedding present; "
"skipping persistent cache lookup"
)
return False
return None

keys = cls._encoder_cache_keys(param)
if not keys:
return False

cached_embeddings = []
for key in keys:
cached_embedding = encoder_cache.get(key)
if cached_embedding is None:
# TODO(TRTLLM-13996): allow re-computing only the uncached items.
# `get_multimodal_embeddings` treats a param as either fully cached or uncached.
# Attaching partial hits would make the later concatenated tensor ambiguous because
# there is no placeholder for missing item rows inside `multimodal_embedding`.
logger.debug(
f"{_MM_ENCODER_CACHE_LOG_NAME}: cache miss; hit_items={len(cached_embeddings)},"
f" total_items={len(keys)}."
)
return False
cached_embeddings.append(cached_embedding)
return None

hits: Dict[int, torch.Tensor] = {}
miss_indices: list[int] = []
for i, key in enumerate(keys):
cached = encoder_cache.get(key)
if cached is None:
miss_indices.append(i)
else:
hits[i] = cached

if len(cached_embeddings) == 1:
param.multimodal_data["multimodal_embedding"] = cached_embeddings[0]
else:
param.multimodal_data["multimodal_embedding"] = torch.cat(cached_embeddings, dim=0)
logger.debug(
f"{_MM_ENCODER_CACHE_LOG_NAME}: full cache hit for {len(keys)} item entries, "
f"rows={param.multimodal_data['multimodal_embedding'].shape[0]}."
f"{_MM_ENCODER_CACHE_LOG_NAME}: partition hit_items={len(hits)}, "
f"miss_items={len(miss_indices)}, total_items={len(keys)}."
)
return True
return EncoderCachePartition(hits=hits, miss_indices=miss_indices, keys=keys)

@staticmethod
def assemble_full_embedding(
item_tensors: Dict[int, torch.Tensor],
total_items: int,
) -> torch.Tensor:
"""Concatenate per-item embedding tensors in item-index order.

`item_tensors` must contain every index in `[0, total_items)`.
"""
if total_items == 1:
return item_tensors[0]
return torch.cat([item_tensors[i] for i in range(total_items)], dim=0)

@staticmethod
def _apply_metadata_slice(
residual: MultimodalParams,
source: MultimodalParams,
item_indices: Sequence[int],
) -> None:
"""Overwrite `residual`'s per-item metadata to match the sliced items.

Models slice raw modality tensors in `build_multimodal_encoder_input`; the mixin owns
the parallel per-item metadata slice so every model gets it identically.
"""
source_lengths = source.multimodal_data["multimodal_embedding_lengths"]
residual.multimodal_data["multimodal_embedding_lengths"] = [
source_lengths[i] for i in item_indices
]
if residual.multimodal_input is not None and source.multimodal_input is not None:
source_hashes = source.multimodal_input.multimodal_hashes
residual.multimodal_input.multimodal_hashes = [source_hashes[i] for i in item_indices]

def _encode_with_partial_cache(
self,
partials: Sequence[tuple[MultimodalParams, EncoderCachePartition]],
encoder_cache: TensorLRUCache,
) -> None:
"""Encode only the miss items of each partial-hit param and stitch results.

After this returns, each param's `multimodal_embedding` has the same shape as a
full encoder run so downstream `get_multimodal_embeddings` treats it as fully
cached.
"""
for param, partition in partials:
residual = self.build_multimodal_encoder_input(param, partition.miss_indices)
self._apply_metadata_slice(residual, param, partition.miss_indices)

residual_output = self.encode_multimodal_inputs([residual])
miss_lengths = [
param.multimodal_data["multimodal_embedding_lengths"][i]
for i in partition.miss_indices
]
miss_tensors = torch.split(residual_output, miss_lengths, dim=0)

by_item: Dict[int, torch.Tensor] = dict(partition.hits)
for miss_idx, tensor in zip(partition.miss_indices, miss_tensors, strict=True):
by_item[miss_idx] = tensor
param.multimodal_data["multimodal_embedding"] = self.assemble_full_embedding(
by_item, len(partition.keys)
)

inserted = 0
rejected = 0
for miss_idx, tensor in zip(partition.miss_indices, miss_tensors, strict=True):
if encoder_cache.put(partition.keys[miss_idx], tensor):
inserted += 1
else:
rejected += 1
logger.debug(
f"{_MM_ENCODER_CACHE_LOG_NAME}: partial-hit encode "
f"total_items={len(partition.keys)} "
f"hit_items={len(partition.hits)} "
f"encoded_items={len(partition.miss_indices)} "
f"cache_writes_inserted={inserted} "
f"cache_writes_rejected={rejected}"
)

@classmethod
def _write_encoder_cache_entries(
Expand Down
Loading