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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Forward pass benchmarks for TransformerBridge."""
3from typing import Optional, Union
5import torch
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
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)
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)
31def _get_decoder_input_ids(model: torch.nn.Module, batch_size: int = 1) -> torch.Tensor:
32 """Get decoder_input_ids for encoder-decoder models.
34 Args:
35 model: The model to get decoder_start_token_id from
36 batch_size: Batch size for the decoder_input_ids
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)
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.
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
75 Returns:
76 BenchmarkResult with comparison details
77 """
78 try:
79 _is_audio = getattr(bridge.cfg, "is_audio_model", False)
81 is_enc_dec = _is_encoder_decoder(bridge.original_model)
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
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)
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 )
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 )
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 )
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
165 return compare_tensors(
166 bridge_output,
167 reference_output,
168 atol=atol,
169 rtol=rtol,
170 name="forward_pass_logits",
171 )
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 )
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.
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
198 Returns:
199 BenchmarkResult with comparison details
200 """
201 try:
202 bridge_loss = _compute_self_target_loss(bridge, test_text)
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 )
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 )
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 )
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")
239 return compare_scalars(
240 bridge_loss.item(),
241 ref_loss_val,
242 atol=atol,
243 name="loss_equivalence",
244 )
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 )
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.
265 Note: Uses relaxed tolerance (3e-2) as forward pass implementations differ
266 slightly, leading to accumulated numerical precision differences.
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
276 Returns:
277 BenchmarkResult with comparison details
278 """
279 try:
280 bridge_logits = bridge(test_text, return_type="logits")
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 )
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 )
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 )
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")
315 return compare_tensors(
316 bridge_logits,
317 ref_logits,
318 atol=atol,
319 rtol=rtol,
320 name="logits_equivalence",
321 )
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 )