transformer_lens.model_bridge.sources package

Subpackages

Submodules

Module contents

Sources module.

This module provides functionality to load and convert models from HuggingFace to TransformerLens format.

transformer_lens.model_bridge.sources.boot(model_name: str, hf_config_overrides: dict | None = None, device: str | device | None = None, dtype: dtype = torch.float32, tokenizer: PreTrainedTokenizerBase | None = None, load_weights: bool = True, trust_remote_code: bool = False, model_class: Any | None = None, hf_model: Any | None = None, n_ctx: int | None = None, revision: str | None = None, checkpoint_index: int | None = None, checkpoint_value: int | None = None, device_map: str | dict[str, str | int] | None = None, n_devices: int | None = None, max_memory: dict[str | int, str | int] | None = None, offload_folder: str | None = None) → TransformerBridge

Boot a model from HuggingFace (exposed as TransformerBridge.boot_transformers).

Returns raw HF weights by default — logits/activations match HF, not legacy HookedTransformer (which folds LayerNorm + centers weights). Call enable_compatibility_mode() on the result for HookedTransformer- equivalent numerics. Generation, argmax, and CE loss are unaffected.

Attention implementation is forced to "eager" so hooks can capture scores and patterns. For an apples-to-apples HF comparison, load the HF model with attn_implementation="eager" too; comparing against the default "sdpa" shows ~1e-3 fp32 drift from kernel-level op reordering, not a bridge bug.

Parameters:
  • model_name – The name of the model to load.

  • hf_config_overrides – Optional overrides applied to the HuggingFace config before model load.

  • device – The device to use. If None, will be determined automatically. Mutually exclusive with device_map.

  • dtype – The dtype to use for the model.

  • tokenizer – Optional pre-initialized tokenizer to use; if not provided one will be created.

  • load_weights – If False, load model without weights (on meta device) for config inspection only.

  • model_class – Optional HuggingFace model class to use instead of the default auto-detected class. When the class name matches a key in SUPPORTED_ARCHITECTURES, the corresponding adapter is selected automatically (e.g., BertForNextSentencePrediction).

  • hf_model – Optional pre-loaded HuggingFace model to use instead of loading one. Useful for models loaded with custom configurations (e.g., quantization via BitsAndBytesConfig). When provided, load_weights is ignored.

  • device_map – HuggingFace-style device map ("auto", "balanced", dict, etc.) for dispatched inference. Explicit maps may include CPU and disk targets; meta targets are still rejected when load_weights=True (meta has no real data to offload from, unlike disk/cpu). Mixed CPU/disk + GPU maps are rejected too, not because they’re known to be broken but because CPU/disk offload has only been verified on CPU-only hardware — no GPU to mix in. bridge.enable_compatibility_mode() (with weight processing, i.e. not no_processing=True) is unsupported on a CPU/disk-offloaded bridge and raises immediately rather than mid-fold; the default (non-compat-mode) forward pass, run_with_cache, and hooks all work normally under offload. Mutually exclusive with device.

  • n_devices – Convenience: split the model across this many CUDA devices (translated to a max_memory dict internally). Requires CUDA with at least this many visible devices.

  • max_memory – Optional per-device memory budget for HF’s dispatcher.

  • offload_folder – Directory for disk-offloaded weight shards when device_map includes a "disk" target. Defaults to a temporary directory (HF’s own default) if omitted.

  • n_ctx – Optional context length override. The bridge normally uses the model’s documented max context from the HF config. Setting this writes to whichever HF field the model uses (n_positions / max_position_embeddings / etc.), so callers don’t need to know the field name. If larger than the model’s default, a warning is emitted — quality may degrade past the trained length for rotary models.

  • revision – Optional HF revision string (branch, tag, or commit). Forwarded to config, model, and tokenizer loading. Mutually exclusive with checkpoint_index and checkpoint_value.

  • checkpoint_index – Index into the available training checkpoints for the model family. Convenience over revision for checkpointed models like EleutherAI/pythia* and stanford-crfm/*. Resolved to a revision string via the known per-family naming conventions (step{value} for Pythia, checkpoint-{value} for stanford-crfm).

  • checkpoint_value – Training step or token count of the desired checkpoint. Alternative to checkpoint_index; must be one of the labels returned by get_checkpoint_labels.

Returns:

The bridge to the loaded model.

transformer_lens.model_bridge.sources.boot_native(config: Any, tokenizer: Any | None = None, device: str | device | None = None, dtype: dtype | None = None, model_name: str = 'native') → TransformerBridge

Build a bridge around a small, randomly-initialized TL-native model.

No HuggingFace Hub call, no transformers import. config.init_mode and config.seed control reproducibility.

transformer_lens.model_bridge.sources.boot_tl_legacy(model_name: str, checkpoint_index: int | None = None, checkpoint_value: int | None = None, device: str | device | None = None, dtype: dtype = torch.float32, tokenizer: Any | None = None) → TransformerBridge

Build a bridge for a legacy TransformerLens-format HF repo.

checkpoint_index / checkpoint_value select a training checkpoint (checkpoints/*_<label>.pth); by default the final weights load. The resolved values are stamped on cfg.checkpoint_index / cfg.checkpoint_value, mirroring the legacy loader.

transformer_lens.model_bridge.sources.boot_vllm(model_name: str, tokenizer: Any | None = None, dtype: dtype | None = None, gpu_memory_utilization: float = 0.5, max_model_len: int | None = None, max_num_batched_tokens: int = 2048, enable_batching: bool = False, enable_position_interventions: bool = False, tensor_parallel_size: int = 1, pipeline_parallel_size: int = 1, **vllm_kwargs: Any) → RemoteBridge

Boot a model via vLLM and wrap it in a RemoteBridge via VLLMDriver.

vLLM drives the forward pass (PagedAttention + torch.compile + CUDA graphs). Capture buffers are populated by hooks the plugin installs pre-compile inside the worker; they come back via collective_rpc and replay through the bridge’s HookPoint tree.

Scope vs vllm-lens: vllm-lens is observation-only. This source extends to observation + spec-vocabulary mutation — each capture hook also applies an affine transform output = output * scale + bias (default identity), so interventions (suppress / scale / add / set) propagate to downstream layers. The hook’s return value replaces the module output per PyTorch register_forward_hook semantics. The mutation path under torch.compile + CUDA graphs is exercised end-to-end by demos/vLLM_Bridge_Integration_Test.ipynb (a manual GPU run, not CI); unit tests cover the dispatch protocol only.

Some captures use vLLM-native conventions that differ from HF/HT; see transformer_lens.model_bridge.sources.vllm.overlays.decoder_only for which hooks diverge and the conversion to apply for HT-equivalent values.

Returned logits are reconstructed full-sequence logits. vLLM’s sampler bypasses lm_head, so the driver rebuilds real logits host-side as ln_final @ lm_head.weight.T (+ bias, + Gemma soft-cap) from the captured final-norm activation — valid at every position, so return_type in {"loss", "both"} works. If the unembedding weight is unreachable the driver falls back to the sampler’s final-position log-probs (earlier positions -inf), declares provides_sequence_logits=False, and the bridge then rejects loss.

GPU memory cost: each capture buffer is max_num_batched_tokens × width at the model’s dtype. For Llama-3.2-1B at fp16 with max_num_batched_tokens=2048, the unembed buffer alone is ~525 MB (2048 × 128256 × 2 bytes); residual-stream buffers add ~8 MB per hook. The affine intervention hook also allocates a transient output-shape tensor per forward (even in identity mode), so peak forward memory is ~1.5× the capture buffers’ resident size.

KV-cache footprint: vLLM reserves KV cache sized for max_model_len × layers × heads × head_dim. If max_model_len is left as None, vLLM uses the model’s native context (e.g. 131072 for Llama-3.2-1B) — easily 4+ GiB even on a 1B model. Pass an explicit max_model_len (e.g. 2048 for typical mech-interp prompts) to keep the budget on smaller GPUs.

enable_batching switches to the eager batched path (enforce_eager, batch_size > 1) — the throughput path for SAE/probe data collection. Default False keeps the compile-validated single-prompt path. Batched caches are right-padded with zeros to the longest sequence.

enable_position_interventions widens each hook’s affine scale/bias buffers from (width,) to (max_num_batched_tokens, width) so an intervention spec can carry a pos field (int or list[int]) that scopes the edit to specific sequence positions — position-scoped activation patching / tensor injection. Costs ~2× extra resident GPU memory across all hooks (the scale and bias buffers join the already-(max_n, width) capture buffer), so it is opt-in and defaults False. Compiled-path only — incompatible with enable_batching.

tensor_parallel_size > 1 enables single-node tensor parallelism. Every capture point the overlay hooks is post-all-reduce and therefore replicated across a stage’s TP ranks; pipeline_parallel_size > 1 shards layers across stages, each owning the hooks that actually fire on it. Capture reads merge across ranks with a first-forward layout check that fails loud on installation drift or non-replicated hook points; sharded/stage-local weights (lm_head, final norm) are gathered for logit reconstruction and the ln_final un-fold. Single-node only (the spec channel is an env var, which never reaches Ray remote workers); both are incompatible with enable_batching (per-rank chunk boundaries are unvalidated).

transformer_lens.model_bridge.sources.build_bridge_config_from_hf(hf_config: Any, architecture: str, model_name: str, dtype: dtype) → TransformerBridgeConfig

Translate an HF config into a TransformerBridgeConfig.

transformer_lens.model_bridge.sources.build_bridge_from_module(model: Module, architecture: str, *, hf_config: Any | None = None, tl_config: TransformerBridgeConfig | None = None, tokenizer: Any | None = None, dtype: dtype | None = None, device: Any | None = None, model_name: str = 'external', post_adapter_hook: Callable[[ArchitectureAdapter], None] | None = None) → TransformerBridge

Build a TransformerBridge around a pre-loaded model.

The bridge never moves, casts, or mutates the supplied model.

Parameters:
  • model – Any nn.Module whose submodule tree matches the adapter’s expected dot-paths for architecture.

  • architecture – Architecture identifier registered in the ArchitectureAdapterFactory (e.g. "LlamaForCausalLM", "TransformerLensNative").

  • hf_config – Optional HF-style config; translated via build_bridge_config_from_hf(). Mutually exclusive with tl_config.

  • tl_config – Optional pre-built TransformerBridgeConfig; bypasses HF translation. Mutually exclusive with hf_config.

  • tokenizer – Optional tokenizer. If supplied, passes through setup_tokenizer and detects BOS/EOS behavior.

  • dtype – Recorded on cfg.dtype. Default None reads from the model’s first parameter; explicit values override.

  • device – Recorded on cfg.device. Default None reads from the model’s first parameter.

  • model_name – Recorded on cfg.model_name.

  • post_adapter_hook – Optional callback invoked after adapter selection and before adapter.prepare_model(). Source-specific overlays mutate component_mapping here.

Returns:

A TransformerBridge wrapping the supplied model.

transformer_lens.model_bridge.sources.check_model_support(model_id: str) → dict

Detailed support info for a model: is_supported, architecture_id, verified, suggestion.

transformer_lens.model_bridge.sources.detect_tokenizer_bos_eos(tokenizer: Any) → tuple[bool, bool]

Detect whether the tokenizer prepends BOS and/or appends EOS.

transformer_lens.model_bridge.sources.list_supported_models(architecture: str | None = None, verified_only: bool = False) → list[str]

List all models supported by TransformerLens.

Parameters:
  • architecture – Filter by architecture ID (e.g., “GPT2LMHeadModel”).

  • verified_only – If True, only return verified-to-work models.

Returns:

List of model IDs.