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

100 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-21 19:27 +0000

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

2 

3from typing import Optional, Union 

4 

5import torch 

6 

7from transformer_lens.benchmarks.utils import ( 

8 BenchmarkResult, 

9 BenchmarkSeverity, 

10 bridge_self_target_loss, 

11 compare_scalars, 

12 compare_tensors, 

13) 

14from transformer_lens.model_bridge import TransformerBridge 

15 

16 

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

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

19 return bridge_self_target_loss(bridge, test_text) 

20 

21 

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

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

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

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

26 return False 

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

28 

29 

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

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

32 

33 Args: 

34 model: The model to get decoder_start_token_id from 

35 batch_size: Batch size for the decoder_input_ids 

36 

37 Returns: 

38 Tensor of shape [batch_size, 1] with decoder_start_token_id 

39 """ 

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

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

42 if decoder_start_token_id is None: 

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

44 # decoder_start unset and start from EOS). 

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

46 if decoder_start_token_id is None: 

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

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

49 decoder_start_token_id = decoder_start_token_id[0] 

50 if decoder_start_token_id is None: 

51 decoder_start_token_id = 0 

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

53 

54 

55def benchmark_forward_pass( 

56 bridge: TransformerBridge, 

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

58 reference_model: Optional[torch.nn.Module] = None, 

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

60 atol: float = 1e-3, 

61 rtol: float = 3e-2, 

62) -> BenchmarkResult: 

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

64 

65 Args: 

66 bridge: TransformerBridge model to test 

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

68 reference_model: Optional live HF reference model (audio / encoder-decoder paths) 

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

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

71 atol: Absolute tolerance for comparison 

72 rtol: Relative tolerance for comparison 

73 

74 Returns: 

75 BenchmarkResult with comparison details 

76 """ 

77 try: 

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

79 

80 is_enc_dec = _is_encoder_decoder(bridge.original_model) 

81 

82 # Prepare extra kwargs for encoder-decoder models 

83 extra_kwargs = {} 

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

85 tokens = bridge.to_tokens(test_input) 

86 batch_size = tokens.shape[0] 

87 decoder_input_ids = _get_decoder_input_ids(bridge.original_model, batch_size) 

88 decoder_input_ids = decoder_input_ids.to(tokens.device) 

89 extra_kwargs["decoder_input_ids"] = decoder_input_ids 

90 

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

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

93 with torch.no_grad(): 

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

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

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

97 if isinstance(bridge_output_raw, torch.Tensor): 

98 bridge_output = bridge_output_raw 

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

100 bridge_output = bridge_output_raw.logits 

101 elif hasattr(bridge_output_raw, "last_hidden_state"): 

102 bridge_output = bridge_output_raw.last_hidden_state 

103 else: 

104 bridge_output = bridge_output_raw 

105 else: 

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

107 

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

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

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

111 return BenchmarkResult( 

112 name="forward_pass", 

113 severity=BenchmarkSeverity.DANGER, 

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

115 passed=False, 

116 ) 

117 

118 if bridge_output.numel() == 0: 

119 return BenchmarkResult( 

120 name="forward_pass", 

121 severity=BenchmarkSeverity.DANGER, 

122 message="Bridge output is empty", 

123 passed=False, 

124 ) 

125 

126 return BenchmarkResult( 

127 name="forward_pass", 

128 severity=BenchmarkSeverity.INFO, 

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

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

131 ) 

132 

133 # Get reference logits from a pre-computed tensor or a live HF model 

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

135 reference_output = reference_logits.to(bridge_output.device) 

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

137 # Audio HF reference model: pass the prepared audio input positionally 

138 # (input_values for wav2vec2-style, input_features for AST-style) 

139 assert reference_model is not None 

140 with torch.no_grad(): 

141 hf_output = reference_model(test_input) 

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

143 reference_output = hf_output.logits 

144 else: 

145 reference_output = hf_output.last_hidden_state 

146 else: 

147 # reference_model is non-None here: the both-None case returns early above 

148 assert reference_model is not None 

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

150 tokens = bridge.to_tokens(test_input) 

151 with torch.no_grad(): 

152 if is_enc_dec: 

153 # Encoder-decoder models need decoder_input_ids 

154 batch_size = tokens.shape[0] 

155 decoder_input_ids = _get_decoder_input_ids(reference_model, batch_size) 

156 decoder_input_ids = decoder_input_ids.to(tokens.device) 

157 hf_output = reference_model(tokens, decoder_input_ids=decoder_input_ids) 

158 else: 

159 hf_output = reference_model(tokens) 

160 reference_output = hf_output.logits 

161 

162 return compare_tensors( 

163 bridge_output, 

164 reference_output, 

165 atol=atol, 

166 rtol=rtol, 

167 name="forward_pass_logits", 

168 ) 

169 

170 except Exception as e: 

171 return BenchmarkResult( 

172 name="forward_pass", 

173 severity=BenchmarkSeverity.ERROR, 

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

175 passed=False, 

176 ) 

177 

178 

179def benchmark_loss_equivalence( 

180 bridge: TransformerBridge, 

181 test_text: str, 

182 reference_loss: Optional[float] = None, 

183 atol: float = 1e-3, 

184) -> BenchmarkResult: 

185 """Benchmark loss computation against a pre-computed reference value. 

186 

187 Args: 

188 bridge: TransformerBridge model to test 

189 test_text: Input text for testing 

190 reference_loss: Optional pre-computed reference loss value (e.g., the HF 

191 loss captured in Phase 1, or a golden fixture). Self-check only if None. 

192 atol: Absolute tolerance for comparison 

193 

194 Returns: 

195 BenchmarkResult with comparison details 

196 """ 

197 try: 

198 bridge_loss = _compute_self_target_loss(bridge, test_text) 

199 

200 if reference_loss is None: 200 ↛ 202line 200 didn't jump to line 202 because the condition on line 200 was never true

201 # No reference - just verify loss is valid 

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

203 return BenchmarkResult( 

204 name="loss_equivalence", 

205 severity=BenchmarkSeverity.DANGER, 

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

207 passed=False, 

208 ) 

209 

210 loss_value = bridge_loss.item() 

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

212 return BenchmarkResult( 

213 name="loss_equivalence", 

214 severity=BenchmarkSeverity.DANGER, 

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

216 passed=False, 

217 ) 

218 

219 return BenchmarkResult( 

220 name="loss_equivalence", 

221 severity=BenchmarkSeverity.INFO, 

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

223 details={"loss": loss_value}, 

224 ) 

225 

226 return compare_scalars( 

227 bridge_loss.item(), 

228 reference_loss, 

229 atol=atol, 

230 name="loss_equivalence", 

231 ) 

232 

233 except Exception as e: 

234 return BenchmarkResult( 

235 name="loss_equivalence", 

236 severity=BenchmarkSeverity.ERROR, 

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

238 passed=False, 

239 ) 

240 

241 

242def benchmark_logits_equivalence( 

243 bridge: TransformerBridge, 

244 test_text: str, 

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

246 atol: float = 3e-2, 

247 rtol: float = 3e-2, 

248) -> BenchmarkResult: 

249 """Benchmark logits output against a pre-computed reference tensor. 

250 

251 Args: 

252 bridge: TransformerBridge model to test 

253 test_text: Input text for testing 

254 reference_logits: Optional pre-computed reference logits tensor (e.g., the 

255 HF logits captured in Phase 1, or a golden fixture). Self-check only if None. 

256 atol: Absolute tolerance for comparison 

257 rtol: Relative tolerance for comparison 

258 

259 Returns: 

260 BenchmarkResult with comparison details 

261 """ 

262 try: 

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

264 

265 if reference_logits is None: 

266 # No reference - just verify logits shape and validity 

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

268 return BenchmarkResult( 

269 name="logits_equivalence", 

270 severity=BenchmarkSeverity.DANGER, 

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

272 passed=False, 

273 ) 

274 

275 if bridge_logits.numel() == 0: 

276 return BenchmarkResult( 

277 name="logits_equivalence", 

278 severity=BenchmarkSeverity.DANGER, 

279 message="Bridge logits is empty", 

280 passed=False, 

281 ) 

282 

283 return BenchmarkResult( 

284 name="logits_equivalence", 

285 severity=BenchmarkSeverity.INFO, 

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

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

288 ) 

289 

290 ref_logits = reference_logits.to(bridge_logits.device) 

291 

292 return compare_tensors( 

293 bridge_logits, 

294 ref_logits, 

295 atol=atol, 

296 rtol=rtol, 

297 name="logits_equivalence", 

298 ) 

299 

300 except Exception as e: 

301 return BenchmarkResult( 

302 name="logits_equivalence", 

303 severity=BenchmarkSeverity.ERROR, 

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

305 passed=False, 

306 )