transformer_lens.benchmarks.main_benchmark module¶
Main benchmark runner for TransformerBridge.
This module provides the main benchmark suite that compares TransformerBridge against reference implementations in an optimized multi-phase approach: Phase 1: HF + Bridge (unprocessed) - Compare against raw HuggingFace model Phase 2: Bridge (unprocessed) - Runtime self-checks + HF logits/loss equivalence Phase 3: Bridge (processed) - Compatibility mode + HF logits/loss equivalence Phase 4: Text Quality - profile prompts scored by a pinned judge’s perplexity ratio Phase 5: Granular Weight Processing Tests (optional, individual flags) Phase 6: Granular Weight Processing Tests (optional, combined flags) Phase 7: Multimodal Tests (only for multimodal models with pixel_values support) Phase 8: Audio Tests (only for audio encoder models / audio-conditioned decoders) Phase 9: Vision Tests (only for vision-only encoder models, e.g. ViT/DeiT)
- transformer_lens.benchmarks.main_benchmark.get_auto_model_class(model_name: str, trust_remote_code: bool = False)¶
Delegates to the bridge’s architecture detection for consistency.
- transformer_lens.benchmarks.main_benchmark.main()¶
Run benchmarks from command line.
- transformer_lens.benchmarks.main_benchmark.run_benchmark_suite(model_name: str, device: str = 'cpu', dtype: dtype = torch.float32, test_text: str | None = None, use_hf_reference: bool = True, enable_compatibility_mode: bool = True, verbose: bool = True, track_memory: bool = False, phases: list[int] | None = None, trust_remote_code: bool = False, judge_model: PreTrainedModel | None = None, judge_tokenizer: PreTrainedTokenizerBase | None = None, prompt_profile: str | None = None) List[BenchmarkResult]¶
Run comprehensive benchmark suite for TransformerBridge.
This function implements an optimized multi-phase approach to minimize model reloading: Phase 1: HF + Bridge (unprocessed) - Compare against raw HuggingFace model Phase 2: Bridge (unprocessed) - Runtime self-checks + HF logits/loss equivalence Phase 3: Bridge (processed) - Compatibility mode + HF logits/loss equivalence Phase 4: Text Quality - profile prompts scored by a pinned judge’s perplexity ratio
- Parameters:
model_name – Name of the model to benchmark (e.g., “gpt2”)
device – Device to run on (“cpu” or “cuda”)
dtype – Precision for model loading (default: torch.float32). Use torch.bfloat16 to halve memory for larger models. Phase 2/3 comparisons automatically upcast to float32 for precision.
test_text – Optional test text (default: standard test prompt)
use_hf_reference – Whether to compare against HuggingFace model
enable_compatibility_mode – Whether to enable compatibility mode on bridge
verbose – Whether to print results to console
track_memory – Whether to track and report memory usage (requires psutil)
phases – Optional list of phase numbers to run (e.g., [1, 2, 3]). If None, runs all phases.
trust_remote_code – Whether to trust remote code for custom architectures.
judge_model – Optional pre-loaded Phase-4 judge. When provided with judge_tokenizer, avoids reloading for each model in batch.
judge_tokenizer – Optional pre-loaded tokenizer for the Phase-4 judge.
prompt_profile – Optional Phase-4 prompt profile (e.g. “chat”, “task:translation@en-de”). Resolved from curation + the registry when None.
- Returns:
List of BenchmarkResult objects
- transformer_lens.benchmarks.main_benchmark.run_comparison_benchmarks(bridge_model: TransformerBridge, test_text: str, phase_name: str, is_processed: bool, verbose: bool = True, phase1_reference: PhaseReferenceData | None = None, restore_dtype_after_equivalence: dtype | None = None) List[BenchmarkResult]¶
Run standardized runtime benchmarks on the bridge.
This function runs the same comprehensive test suite for both unprocessed (Phase 2) and processed (Phase 3) modes: HF-anchored logits/loss equivalence (via the saved Phase 1 reference) plus reference-free structural self-checks for hooks, cache, and gradients.
- Parameters:
bridge_model – TransformerBridge model to test
test_text – Input text for testing
phase_name – Name of the phase (“Phase 2” or “Phase 3”) for logging
is_processed – Whether models have processed weights (for weight-specific tests)
verbose – Whether to print detailed results
phase1_reference – Optional saved Phase 1 HF reference data for equivalence testing
restore_dtype_after_equivalence – If set, downcast bridge_model to this dtype after the equivalence comparison but before hook/cache/gradient tests. Used when the bridge was upcast to float32 for precise equivalence testing.
- Returns:
List of BenchmarkResult objects
- transformer_lens.benchmarks.main_benchmark.update_model_registry(model_name: str, results: List[BenchmarkResult], use_hf_reference: bool = False) bool¶
Update the model registry with benchmark results.
- Parameters:
model_name – The model that was benchmarked
results – List of benchmark results
use_hf_reference – Whether the run numerically compared against an HF reference. Defaults to False so an unstated reference state records a passing run as PROVISIONAL, never VERIFIED.
- Returns:
True if registry was updated successfully