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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
1"""Forward pass benchmarks for TransformerBridge."""
3from typing import Optional, Union
5import torch
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
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)
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)
30def _get_decoder_input_ids(model: torch.nn.Module, batch_size: int = 1) -> torch.Tensor:
31 """Get decoder_input_ids for encoder-decoder models.
33 Args:
34 model: The model to get decoder_start_token_id from
35 batch_size: Batch size for the decoder_input_ids
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)
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.
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
74 Returns:
75 BenchmarkResult with comparison details
76 """
77 try:
78 _is_audio = getattr(bridge.cfg, "is_audio_model", False)
80 is_enc_dec = _is_encoder_decoder(bridge.original_model)
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
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)
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 )
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 )
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 )
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
162 return compare_tensors(
163 bridge_output,
164 reference_output,
165 atol=atol,
166 rtol=rtol,
167 name="forward_pass_logits",
168 )
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 )
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.
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
194 Returns:
195 BenchmarkResult with comparison details
196 """
197 try:
198 bridge_loss = _compute_self_target_loss(bridge, test_text)
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 )
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 )
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 )
226 return compare_scalars(
227 bridge_loss.item(),
228 reference_loss,
229 atol=atol,
230 name="loss_equivalence",
231 )
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 )
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.
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
259 Returns:
260 BenchmarkResult with comparison details
261 """
262 try:
263 bridge_logits = bridge(test_text, return_type="logits")
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 )
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 )
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 )
290 ref_logits = reference_logits.to(bridge_logits.device)
292 return compare_tensors(
293 bridge_logits,
294 ref_logits,
295 atol=atol,
296 rtol=rtol,
297 name="logits_equivalence",
298 )
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 )