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