Coverage for transformer_lens/benchmarks/forward_pass.py: 26%

113 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""Forward pass benchmarks for TransformerBridge.""" 

2 

3from typing import Optional, Union 

4 

5import torch 

6 

7from transformer_lens import HookedTransformer 

8from transformer_lens.benchmarks.utils import ( 

9 BenchmarkResult, 

10 BenchmarkSeverity, 

11 bridge_self_target_loss, 

12 compare_scalars, 

13 compare_tensors, 

14) 

15from transformer_lens.model_bridge import TransformerBridge 

16 

17 

18def _compute_self_target_loss(bridge: TransformerBridge, test_text: str) -> torch.Tensor: 

19 """Compute loss with the tokenized input supplied as explicit labels.""" 

20 return bridge_self_target_loss(bridge, test_text) 

21 

22 

23def _is_encoder_decoder(model: torch.nn.Module) -> bool: 

24 """Check if a model is an encoder-decoder architecture.""" 

25 config = getattr(model, "config", None) 

26 if config is None: 26 ↛ 27line 26 didn't jump to line 27 because the condition on line 26 was never true

27 return False 

28 return getattr(config, "is_encoder_decoder", False) 

29 

30 

31def _get_decoder_input_ids(model: torch.nn.Module, batch_size: int = 1) -> torch.Tensor: 

32 """Get decoder_input_ids for encoder-decoder models. 

33 

34 Args: 

35 model: The model to get decoder_start_token_id from 

36 batch_size: Batch size for the decoder_input_ids 

37 

38 Returns: 

39 Tensor of shape [batch_size, 1] with decoder_start_token_id 

40 """ 

41 config = getattr(model, "config", None) 

42 decoder_start_token_id = getattr(config, "decoder_start_token_id", None) if config else None 

43 if decoder_start_token_id is None: 

44 # HF fallback chain: bos, then eos (MBart-family checkpoints leave 

45 # decoder_start unset and start from EOS). 

46 decoder_start_token_id = getattr(config, "bos_token_id", None) if config else None 

47 if decoder_start_token_id is None: 

48 decoder_start_token_id = getattr(config, "eos_token_id", None) if config else None 

49 if isinstance(decoder_start_token_id, (list, tuple)): 

50 decoder_start_token_id = decoder_start_token_id[0] 

51 if decoder_start_token_id is None: 

52 decoder_start_token_id = 0 

53 return torch.tensor([[decoder_start_token_id]] * batch_size) 

54 

55 

56def benchmark_forward_pass( 

57 bridge: TransformerBridge, 

58 test_input: Union[str, torch.Tensor], 

59 reference_model: Optional[Union[HookedTransformer, torch.nn.Module]] = None, 

60 reference_logits: Optional[torch.Tensor] = None, 

61 atol: float = 1e-3, 

62 rtol: float = 3e-2, 

63) -> BenchmarkResult: 

64 """Benchmark forward pass between TransformerBridge and reference model. 

65 

66 Args: 

67 bridge: TransformerBridge model to test 

68 test_input: Input text string or audio waveform tensor for testing 

69 reference_model: Optional reference model (HookedTransformer or HF model) 

70 reference_logits: Optional pre-computed reference logits/hidden states tensor 

71 (e.g., saved from a prior HF forward pass to avoid needing both models in memory) 

72 atol: Absolute tolerance for comparison 

73 rtol: Relative tolerance for comparison 

74 

75 Returns: 

76 BenchmarkResult with comparison details 

77 """ 

78 try: 

79 _is_audio = getattr(bridge.cfg, "is_audio_model", False) 

80 

81 is_enc_dec = _is_encoder_decoder(bridge.original_model) 

82 

83 # Prepare extra kwargs for encoder-decoder models 

84 extra_kwargs = {} 

85 if is_enc_dec and isinstance(test_input, str): 85 ↛ 86line 85 didn't jump to line 86 because the condition on line 85 was never true

86 tokens = bridge.to_tokens(test_input) 

87 batch_size = tokens.shape[0] 

88 decoder_input_ids = _get_decoder_input_ids(bridge.original_model, batch_size) 

89 decoder_input_ids = decoder_input_ids.to(tokens.device) 

90 extra_kwargs["decoder_input_ids"] = decoder_input_ids 

91 

92 # Run bridge forward pass (use no_grad to match HF reference context — 

93 # MPS SDPA can produce different results with vs without gradient tracking) 

94 with torch.no_grad(): 

95 if _is_audio and isinstance(test_input, torch.Tensor): 95 ↛ 97line 95 didn't jump to line 97 because the condition on line 95 was never true

96 # Audio models: pass waveform, extract tensor from output 

97 bridge_output_raw = bridge(test_input, return_type="logits") 

98 if isinstance(bridge_output_raw, torch.Tensor): 

99 bridge_output = bridge_output_raw 

100 elif hasattr(bridge_output_raw, "logits") and bridge_output_raw.logits is not None: 

101 bridge_output = bridge_output_raw.logits 

102 elif hasattr(bridge_output_raw, "last_hidden_state"): 

103 bridge_output = bridge_output_raw.last_hidden_state 

104 else: 

105 bridge_output = bridge_output_raw 

106 else: 

107 bridge_output = bridge(test_input, return_type="logits", **extra_kwargs) 

108 

109 if reference_model is None and reference_logits is None: 109 ↛ 111line 109 didn't jump to line 111 because the condition on line 109 was never true

110 # No reference model or logits - just verify output shape and validity 

111 if not isinstance(bridge_output, torch.Tensor): 

112 return BenchmarkResult( 

113 name="forward_pass", 

114 severity=BenchmarkSeverity.DANGER, 

115 message="Bridge output is not a tensor", 

116 passed=False, 

117 ) 

118 

119 if bridge_output.numel() == 0: 

120 return BenchmarkResult( 

121 name="forward_pass", 

122 severity=BenchmarkSeverity.DANGER, 

123 message="Bridge output is empty", 

124 passed=False, 

125 ) 

126 

127 return BenchmarkResult( 

128 name="forward_pass", 

129 severity=BenchmarkSeverity.INFO, 

130 message=f"Bridge forward pass successful (shape: {bridge_output.shape})", 

131 details={"output_shape": str(bridge_output.shape)}, 

132 ) 

133 

134 # Get reference logits from pre-computed tensor or live model 

135 if reference_logits is not None: 135 ↛ 137line 135 didn't jump to line 137 because the condition on line 135 was always true

136 reference_output = reference_logits.to(bridge_output.device) 

137 elif isinstance(reference_model, HookedTransformer): 

138 reference_output = reference_model(test_input, return_type="logits") 

139 elif _is_audio and isinstance(test_input, torch.Tensor): 

140 # Audio HF reference model: pass waveform directly 

141 assert reference_model is not None 

142 with torch.no_grad(): 

143 hf_output = reference_model(input_values=test_input) 

144 if hasattr(hf_output, "logits") and hf_output.logits is not None: 

145 reference_output = hf_output.logits 

146 else: 

147 reference_output = hf_output.last_hidden_state 

148 else: 

149 # HuggingFace model (reference_model is guaranteed non-None here 

150 # because we returned early at line 80 when both are None) 

151 assert reference_model is not None 

152 assert isinstance(test_input, str), "Text model requires string input" 

153 tokens = bridge.to_tokens(test_input) 

154 with torch.no_grad(): 

155 if is_enc_dec: 

156 # Encoder-decoder models need decoder_input_ids 

157 batch_size = tokens.shape[0] 

158 decoder_input_ids = _get_decoder_input_ids(reference_model, batch_size) 

159 decoder_input_ids = decoder_input_ids.to(tokens.device) 

160 hf_output = reference_model(tokens, decoder_input_ids=decoder_input_ids) 

161 else: 

162 hf_output = reference_model(tokens) 

163 reference_output = hf_output.logits 

164 

165 return compare_tensors( 

166 bridge_output, 

167 reference_output, 

168 atol=atol, 

169 rtol=rtol, 

170 name="forward_pass_logits", 

171 ) 

172 

173 except Exception as e: 

174 return BenchmarkResult( 

175 name="forward_pass", 

176 severity=BenchmarkSeverity.ERROR, 

177 message=f"Forward pass failed: {str(e)}", 

178 passed=False, 

179 ) 

180 

181 

182def benchmark_loss_equivalence( 

183 bridge: TransformerBridge, 

184 test_text: str, 

185 reference_model: Optional[HookedTransformer] = None, 

186 reference_loss: Optional[float] = None, 

187 atol: float = 1e-3, 

188) -> BenchmarkResult: 

189 """Benchmark loss computation between TransformerBridge and HookedTransformer. 

190 

191 Args: 

192 bridge: TransformerBridge model to test 

193 test_text: Input text for testing 

194 reference_model: Optional HookedTransformer reference model 

195 reference_loss: Optional pre-computed reference loss value (e.g., from Phase 1) 

196 atol: Absolute tolerance for comparison 

197 

198 Returns: 

199 BenchmarkResult with comparison details 

200 """ 

201 try: 

202 bridge_loss = _compute_self_target_loss(bridge, test_text) 

203 

204 if reference_model is None and reference_loss is None: 204 ↛ 206line 204 didn't jump to line 206 because the condition on line 204 was never true

205 # No reference - just verify loss is valid 

206 if not isinstance(bridge_loss, torch.Tensor): 

207 return BenchmarkResult( 

208 name="loss_equivalence", 

209 severity=BenchmarkSeverity.DANGER, 

210 message="Bridge loss is not a tensor", 

211 passed=False, 

212 ) 

213 

214 loss_value = bridge_loss.item() 

215 if torch.isnan(bridge_loss) or torch.isinf(bridge_loss): 

216 return BenchmarkResult( 

217 name="loss_equivalence", 

218 severity=BenchmarkSeverity.DANGER, 

219 message=f"Bridge loss is invalid: {loss_value}", 

220 passed=False, 

221 ) 

222 

223 return BenchmarkResult( 

224 name="loss_equivalence", 

225 severity=BenchmarkSeverity.INFO, 

226 message=f"Bridge loss computed successfully: {loss_value:.6f}", 

227 details={"loss": loss_value}, 

228 ) 

229 

230 # Get reference loss from model or pre-computed value 

231 if reference_loss is not None: 231 ↛ 233line 231 didn't jump to line 233 because the condition on line 231 was always true

232 ref_loss_val = reference_loss 

233 elif reference_model is not None: 

234 ref_loss_tensor = reference_model(test_text, return_type="loss") 

235 ref_loss_val = ref_loss_tensor.item() 

236 else: 

237 raise ValueError("Either reference_logits or reference_model must be provided") 

238 

239 return compare_scalars( 

240 bridge_loss.item(), 

241 ref_loss_val, 

242 atol=atol, 

243 name="loss_equivalence", 

244 ) 

245 

246 except Exception as e: 

247 return BenchmarkResult( 

248 name="loss_equivalence", 

249 severity=BenchmarkSeverity.ERROR, 

250 message=f"Loss computation failed: {str(e)}", 

251 passed=False, 

252 ) 

253 

254 

255def benchmark_logits_equivalence( 

256 bridge: TransformerBridge, 

257 test_text: str, 

258 reference_model: Optional[HookedTransformer] = None, 

259 reference_logits: Optional[torch.Tensor] = None, 

260 atol: float = 3e-2, 

261 rtol: float = 3e-2, 

262) -> BenchmarkResult: 

263 """Benchmark logits output between TransformerBridge and HookedTransformer. 

264 

265 Note: Uses relaxed tolerance (3e-2) as forward pass implementations differ 

266 slightly, leading to accumulated numerical precision differences. 

267 

268 Args: 

269 bridge: TransformerBridge model to test 

270 test_text: Input text for testing 

271 reference_model: Optional HookedTransformer reference model 

272 reference_logits: Optional pre-computed reference logits tensor (e.g., from Phase 1) 

273 atol: Absolute tolerance for comparison 

274 rtol: Relative tolerance for comparison 

275 

276 Returns: 

277 BenchmarkResult with comparison details 

278 """ 

279 try: 

280 bridge_logits = bridge(test_text, return_type="logits") 

281 

282 if reference_model is None and reference_logits is None: 

283 # No reference - just verify logits shape and validity 

284 if not isinstance(bridge_logits, torch.Tensor): 

285 return BenchmarkResult( 

286 name="logits_equivalence", 

287 severity=BenchmarkSeverity.DANGER, 

288 message="Bridge logits is not a tensor", 

289 passed=False, 

290 ) 

291 

292 if bridge_logits.numel() == 0: 

293 return BenchmarkResult( 

294 name="logits_equivalence", 

295 severity=BenchmarkSeverity.DANGER, 

296 message="Bridge logits is empty", 

297 passed=False, 

298 ) 

299 

300 return BenchmarkResult( 

301 name="logits_equivalence", 

302 severity=BenchmarkSeverity.INFO, 

303 message=f"Bridge logits computed successfully (shape: {bridge_logits.shape})", 

304 details={"output_shape": str(bridge_logits.shape)}, 

305 ) 

306 

307 # Get reference logits from model or pre-computed tensor 

308 if reference_logits is not None: 

309 ref_logits = reference_logits.to(bridge_logits.device) 

310 elif reference_model is not None: 

311 ref_logits = reference_model(test_text, return_type="logits") 

312 else: 

313 raise ValueError("Either reference_logits or reference_model must be provided") 

314 

315 return compare_tensors( 

316 bridge_logits, 

317 ref_logits, 

318 atol=atol, 

319 rtol=rtol, 

320 name="logits_equivalence", 

321 ) 

322 

323 except Exception as e: 

324 return BenchmarkResult( 

325 name="logits_equivalence", 

326 severity=BenchmarkSeverity.ERROR, 

327 message=f"Logits computation failed: {str(e)}", 

328 passed=False, 

329 )