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

1"""Multimodal benchmarks for TransformerBridge. 

2 

3Tests that multimodal models (LLaVA, Gemma3, etc.) correctly handle image inputs 

4through forward(), generate(), and run_with_cache(). 

5""" 

6 

7import torch 

8 

9from transformer_lens.benchmarks.utils import ( 

10 BenchmarkResult, 

11 BenchmarkSeverity, 

12 is_tiny_test_model, 

13) 

14from transformer_lens.model_bridge import TransformerBridge 

15 

16 

17def _create_test_image(): 

18 """Create a small synthetic test image using PIL. 

19 

20 Returns a 224x224 red image, or None if PIL is not available. 

21 """ 

22 try: 

23 from PIL import Image 

24 

25 return Image.new("RGB", (224, 224), color="red") 

26 except ImportError: 

27 return None 

28 

29 

30def _prepare_test_inputs(bridge: TransformerBridge): 

31 """Prepare multimodal test inputs using the bridge's processor. 

32 

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 

39 

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 

43 

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] 

61 

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 

76 

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) 

81 

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 

92 

93 return input_ids, extra_kwargs, prompt 

94 except Exception: 

95 continue 

96 

97 return None, None, None 

98 

99 

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. 

106 

107 Tests that passing pixel_values produces valid logits (non-NaN, correct shape). 

108 

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. 

113 

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 ) 

123 

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 ) 

130 

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 ) 

138 

139 try: 

140 with torch.no_grad(): 

141 logits = bridge.forward(input_ids, return_type="logits", **extra_kwargs) 

142 

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 ) 

150 

151 has_nan = torch.isnan(logits).any().item() 

152 has_inf = torch.isinf(logits).any().item() 

153 

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 ) 

162 

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 ) 

176 

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 ) 

184 

185 

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. 

193 

194 Tests that generation with image input produces text output longer than input. 

195 

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. 

201 

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 ) 

211 

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 ) 

218 

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 ) 

226 

227 try: 

228 output = bridge.generate( 

229 input_ids, 

230 max_new_tokens=max_new_tokens, 

231 return_type="tokens", 

232 **extra_kwargs, 

233 ) 

234 

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 ) 

242 

243 input_len = input_ids.shape[-1] 

244 output_len = output.shape[-1] 

245 

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 ) 

260 

261 generated_text = bridge.tokenizer.decode(output[0], skip_special_tokens=True) 

262 

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 ) 

274 

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 ) 

282 

283 

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. 

290 

291 Tests that running with cache and image input populates the activation cache, 

292 including vision encoder hooks if present. 

293 

294 Args: 

295 bridge: TransformerBridge model to test. 

296 test_text: Text prompt. 

297 reference_model: Not used, kept for API compatibility. 

298 

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 ) 

308 

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 ) 

315 

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 ) 

323 

324 try: 

325 with torch.no_grad(): 

326 logits, cache = bridge.run_with_cache(input_ids, **extra_kwargs) 

327 

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 ) 

335 

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()] 

338 

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 ) 

353 

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 )