Coverage for transformer_lens/benchmarks/multimodal.py: 37%
106 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Multimodal benchmarks for TransformerBridge.
3Tests that multimodal models (LLaVA, Gemma3, etc.) correctly handle image inputs
4through forward(), generate(), and run_with_cache().
5"""
7import torch
9from transformer_lens.benchmarks.utils import (
10 BenchmarkResult,
11 BenchmarkSeverity,
12 is_tiny_test_model,
13)
14from transformer_lens.model_bridge import TransformerBridge
17def _create_test_image():
18 """Create a small synthetic test image using PIL.
20 Returns a 224x224 red image, or None if PIL is not available.
21 """
22 try:
23 from PIL import Image
25 return Image.new("RGB", (224, 224), color="red")
26 except ImportError:
27 return None
30def _prepare_test_inputs(bridge: TransformerBridge):
31 """Prepare multimodal test inputs using the bridge's processor.
33 Returns (input_ids, extra_kwargs, prompt) where extra_kwargs is a dict
34 containing pixel_values and any other processor outputs (e.g. image_sizes
35 for LlavaNext). Returns (None, None, None) on failure.
36 """
37 if bridge.processor is None:
38 return None, None, None
40 image = _create_test_image()
41 if image is None: 41 ↛ 42line 41 didn't jump to line 42 because the condition on line 41 was never true
42 return None, None, None
44 # Build a prompt with the model's image token placeholder. Families disagree
45 # on which token their processor accepts, and a rejected token raises:
46 # LLava: image_token = "<image>"
47 # Gemma3: only boi_token ("<start_of_image>") expands; image_token
48 # ("<image_soft_token>") is rejected
49 # Gemma4: only image_token ("<|image|>") expands; boi_token is a marker
50 if hasattr(bridge.processor, "post_process_generation"): 50 ↛ 53line 50 didn't jump to line 53 because the condition on line 50 was never true
51 # Task-prompt captioners (Florence-2 style) map task tokens to prompts
52 # themselves and ignore free-text instructions.
53 candidates = ["<CAPTION>"]
54 else:
55 tokens = [
56 getattr(bridge.processor, "image_token", None),
57 getattr(bridge.processor, "boi_token", None),
58 ]
59 seen = list(dict.fromkeys(t for t in tokens if t)) or ["<image>"]
60 candidates = [f"{t}\nDescribe this image." for t in seen]
62 # Some processors insert the image placeholders themselves; a manual
63 # placeholder would then double-count. Probe with a plain prompt first
64 # and prefer it if image tokens were auto-inserted. Idefics3-style
65 # processors RAISE on image inputs without image tokens in the text — a
66 # failed probe just means "not auto-inserting", so fall back below.
67 image_token_id = getattr(bridge.original_model.config, "image_token_id", None)
68 if image_token_id is not None:
69 try:
70 plain = "Describe this image."
71 probe = bridge.processor(text=plain, images=image, return_tensors="pt")
72 if (probe["input_ids"] == image_token_id).any(): 72 ↛ 77line 72 didn't jump to line 77 because the condition on line 72 was always true
73 candidates.insert(0, plain)
74 except Exception:
75 pass
77 for prompt in candidates:
78 try:
79 inputs = bridge.processor(text=prompt, images=image, return_tensors="pt")
80 input_ids = inputs["input_ids"].to(bridge.cfg.device)
82 # Collect all extra kwargs the model's forward() may need
83 # (pixel_values, image_sizes, pixel_attention_mask, etc.)
84 extra_kwargs = {}
85 for key, val in inputs.items():
86 if key == "input_ids":
87 continue
88 if hasattr(val, "to"): 88 ↛ 91line 88 didn't jump to line 91 because the condition on line 88 was always true
89 extra_kwargs[key] = val.to(bridge.cfg.device)
90 else:
91 extra_kwargs[key] = val
93 return input_ids, extra_kwargs, prompt
94 except Exception:
95 continue
97 return None, None, None
100def benchmark_multimodal_forward(
101 bridge: TransformerBridge,
102 test_text: str = "Describe this image.",
103 reference_model=None,
104) -> BenchmarkResult:
105 """Benchmark forward() with pixel_values for multimodal models.
107 Tests that passing pixel_values produces valid logits (non-NaN, correct shape).
109 Args:
110 bridge: TransformerBridge model to test.
111 test_text: Text prompt (used as fallback if processor unavailable).
112 reference_model: Not used, kept for API compatibility.
114 Returns:
115 BenchmarkResult with forward pass details.
116 """
117 if not getattr(bridge.cfg, "is_multimodal", False):
118 return BenchmarkResult(
119 name="multimodal_forward",
120 severity=BenchmarkSeverity.SKIPPED,
121 message="Skipped: model is not multimodal",
122 )
124 if is_tiny_test_model(getattr(bridge.cfg, "model_name", "") or ""):
125 return BenchmarkResult(
126 name="multimodal_forward",
127 severity=BenchmarkSeverity.INFO,
128 message="Skipped for tiny/test model",
129 )
131 input_ids, extra_kwargs, prompt = _prepare_test_inputs(bridge)
132 if input_ids is None:
133 return BenchmarkResult(
134 name="multimodal_forward",
135 severity=BenchmarkSeverity.SKIPPED,
136 message="Skipped: processor or PIL not available",
137 )
139 try:
140 with torch.no_grad():
141 logits = bridge.forward(input_ids, return_type="logits", **extra_kwargs)
143 if logits is None:
144 return BenchmarkResult(
145 name="multimodal_forward",
146 severity=BenchmarkSeverity.DANGER,
147 message="Forward pass returned None",
148 passed=False,
149 )
151 has_nan = torch.isnan(logits).any().item()
152 has_inf = torch.isinf(logits).any().item()
154 if has_nan or has_inf:
155 return BenchmarkResult(
156 name="multimodal_forward",
157 severity=BenchmarkSeverity.DANGER,
158 message=f"Logits contain NaN={has_nan}, Inf={has_inf}",
159 details={"shape": list(logits.shape)},
160 passed=False,
161 )
163 pixel_values = extra_kwargs.get("pixel_values")
164 return BenchmarkResult(
165 name="multimodal_forward",
166 severity=BenchmarkSeverity.INFO,
167 message=f"Multimodal forward pass successful, logits shape: {list(logits.shape)}",
168 details={
169 "logits_shape": list(logits.shape),
170 "input_ids_shape": list(input_ids.shape),
171 "pixel_values_shape": (
172 list(pixel_values.shape) if pixel_values is not None else None
173 ),
174 },
175 )
177 except Exception as e:
178 return BenchmarkResult(
179 name="multimodal_forward",
180 severity=BenchmarkSeverity.ERROR,
181 message=f"Multimodal forward pass failed: {str(e)}",
182 passed=False,
183 )
186def benchmark_multimodal_generation(
187 bridge: TransformerBridge,
188 test_text: str = "Describe this image.",
189 max_new_tokens: int = 10,
190 reference_model=None,
191) -> BenchmarkResult:
192 """Benchmark generate() with pixel_values for multimodal models.
194 Tests that generation with image input produces text output longer than input.
196 Args:
197 bridge: TransformerBridge model to test.
198 test_text: Text prompt.
199 max_new_tokens: Number of tokens to generate.
200 reference_model: Not used, kept for API compatibility.
202 Returns:
203 BenchmarkResult with generation details.
204 """
205 if not getattr(bridge.cfg, "is_multimodal", False):
206 return BenchmarkResult(
207 name="multimodal_generation",
208 severity=BenchmarkSeverity.SKIPPED,
209 message="Skipped: model is not multimodal",
210 )
212 if is_tiny_test_model(getattr(bridge.cfg, "model_name", "") or ""):
213 return BenchmarkResult(
214 name="multimodal_generation",
215 severity=BenchmarkSeverity.INFO,
216 message="Skipped for tiny/test model",
217 )
219 input_ids, extra_kwargs, prompt = _prepare_test_inputs(bridge)
220 if input_ids is None:
221 return BenchmarkResult(
222 name="multimodal_generation",
223 severity=BenchmarkSeverity.SKIPPED,
224 message="Skipped: processor or PIL not available",
225 )
227 try:
228 output = bridge.generate(
229 input_ids,
230 max_new_tokens=max_new_tokens,
231 return_type="tokens",
232 **extra_kwargs,
233 )
235 if not isinstance(output, torch.Tensor):
236 return BenchmarkResult(
237 name="multimodal_generation",
238 severity=BenchmarkSeverity.DANGER,
239 message="Generation did not return a tensor",
240 passed=False,
241 )
243 input_len = input_ids.shape[-1]
244 output_len = output.shape[-1]
246 # Encoder-decoder generate() returns decoder tokens only, so any
247 # token beyond the decoder start counts as new output.
248 if getattr(bridge.original_model.config, "is_encoder_decoder", False):
249 produced_new_tokens = output_len > 1
250 else:
251 produced_new_tokens = output_len > input_len
252 if not produced_new_tokens:
253 return BenchmarkResult(
254 name="multimodal_generation",
255 severity=BenchmarkSeverity.DANGER,
256 message="Generation produced no new tokens",
257 details={"input_tokens": input_len, "output_tokens": output_len},
258 passed=False,
259 )
261 generated_text = bridge.tokenizer.decode(output[0], skip_special_tokens=True)
263 return BenchmarkResult(
264 name="multimodal_generation",
265 severity=BenchmarkSeverity.INFO,
266 message=f"Multimodal generation successful: {input_len} -> {output_len} tokens",
267 details={
268 "input_tokens": input_len,
269 "output_tokens": output_len,
270 "max_new_tokens": max_new_tokens,
271 "generated_text": generated_text[:200],
272 },
273 )
275 except Exception as e:
276 return BenchmarkResult(
277 name="multimodal_generation",
278 severity=BenchmarkSeverity.ERROR,
279 message=f"Multimodal generation failed: {str(e)}",
280 passed=False,
281 )
284def benchmark_multimodal_cache(
285 bridge: TransformerBridge,
286 test_text: str = "Describe this image.",
287 reference_model=None,
288) -> BenchmarkResult:
289 """Benchmark run_with_cache() with pixel_values for multimodal models.
291 Tests that running with cache and image input populates the activation cache,
292 including vision encoder hooks if present.
294 Args:
295 bridge: TransformerBridge model to test.
296 test_text: Text prompt.
297 reference_model: Not used, kept for API compatibility.
299 Returns:
300 BenchmarkResult with cache details.
301 """
302 if not getattr(bridge.cfg, "is_multimodal", False):
303 return BenchmarkResult(
304 name="multimodal_cache",
305 severity=BenchmarkSeverity.SKIPPED,
306 message="Skipped: model is not multimodal",
307 )
309 if is_tiny_test_model(getattr(bridge.cfg, "model_name", "") or ""):
310 return BenchmarkResult(
311 name="multimodal_cache",
312 severity=BenchmarkSeverity.INFO,
313 message="Skipped for tiny/test model",
314 )
316 input_ids, extra_kwargs, prompt = _prepare_test_inputs(bridge)
317 if input_ids is None:
318 return BenchmarkResult(
319 name="multimodal_cache",
320 severity=BenchmarkSeverity.SKIPPED,
321 message="Skipped: processor or PIL not available",
322 )
324 try:
325 with torch.no_grad():
326 logits, cache = bridge.run_with_cache(input_ids, **extra_kwargs)
328 if cache is None or len(cache) == 0:
329 return BenchmarkResult(
330 name="multimodal_cache",
331 severity=BenchmarkSeverity.DANGER,
332 message="run_with_cache() returned empty cache",
333 passed=False,
334 )
336 cache_keys = list(cache.keys()) if hasattr(cache, "keys") else []
337 vision_keys = [k for k in cache_keys if "vision" in k.lower()]
339 return BenchmarkResult(
340 name="multimodal_cache",
341 severity=BenchmarkSeverity.INFO,
342 message=(
343 f"Multimodal cache populated: {len(cache_keys)} entries "
344 f"({len(vision_keys)} vision-related)"
345 ),
346 details={
347 "total_cache_entries": len(cache_keys),
348 "vision_cache_entries": len(vision_keys),
349 "vision_keys": vision_keys[:10],
350 "sample_keys": cache_keys[:10],
351 },
352 )
354 except Exception as e:
355 return BenchmarkResult(
356 name="multimodal_cache",
357 severity=BenchmarkSeverity.ERROR,
358 message=f"Multimodal cache test failed: {str(e)}",
359 passed=False,
360 )