transformer_lens.model_bridge.supported_architectures.pretrain module

Architecture adapter for a lightweight decoder-only pretraining model.

Maps a decoder-only transformer using RoPE, RMSNorm, gated SwiGLU MLPs, and optional sparse mixture-of-experts feed-forward layers into TransformerBridge, by wrapping the source module and delegating to its own forward rather than translating parameters into a second implementation.

Usage: build_pretrain_bridge(model, cfg) – the public entry point. PretrainModelContainer and direct build_bridge_from_module use are internal/advanced details (see PretrainModelContainer’s docstring).

Scope: maps a live module into TransformerBridge. Does not load checkpoints, merge tensor-parallel shards, or depend on a training framework.

Required module protocol – “lightweight decoder-only pretraining models” describes intent, not a generality guarantee. The wrapped model must expose:

model.embed (embedding lookup) model.blocks[i].norm1 (pre-attention norm) model.blocks[i].attn (called as attn(x, …)) model.blocks[i].norm2 (pre-MLP norm) model.blocks[i].mlp (gate/up/down, or router/experts) model.norm_f (final norm) model.lm_head (unembedding)

gate/up/down and router/experts name the supported protocol. DenseOrMoEFeedForwardBridge checks these structurally – attribute presence plus basic type (each is a module, experts is a registered module collection) – and raises clearly on a mismatch, but that is structural validation only: it does not and cannot validate that a module satisfying the shape actually implements matching forward semantics. Blocks must take more than the bare hidden state (this target passes cos/sin) – see PretrainModelContainer.

class transformer_lens.model_bridge.supported_architectures.pretrain.DenseOrMoEFeedForwardBridge(name: str, config: Any)

Bases: GeneralizedComponent

Wraps a dense SwiGLU MLP or sparse MoE layer behind one interface. Dispatch is by structural inspection (router/experts vs gate/up/down), not config, so dense/MoE/mixed architectures all share the same component mapping.

forward(*args: Any, **kwargs: Any) → Any

Generic forward pass for bridge components with input/output hooks.

set_original_component(component: Module) → None

Set the original component that this bridge wraps.

Parameters:

original_component – The original transformer component to wrap

class transformer_lens.model_bridge.supported_architectures.pretrain.NativeForwardAttentionBridge(name: str | None, config: Any, submodules: Dict[str, GeneralizedComponent] | None = None, conversion_rule: BaseTensorConversion | None = None, pattern_conversion_rule: BaseTensorConversion | None = None, maintain_native_attention: bool = False, requires_position_embeddings: bool = False, requires_attention_mask: bool = False, attention_mask_4d: bool = False, requires_relative_position_bias: bool = False, is_cross_attention: bool = False, is_causal: bool = True, optional: bool = False, fused_qkv: bool = False)

Bases: AttentionBridge

Opaque attention bridge that delegates to the source attention.

This adapter intentionally exposes only input/output attention hooks. It has no mapped Q/K/V/O projection components, so the standard per-head aliases and weight aliases do not apply.

hook_aliases: Dict[str, str | List[str]] = {}
property_aliases: Dict[str, str] = {}
supports_split_qkv_fork: bool = False
class transformer_lens.model_bridge.supported_architectures.pretrain.PretrainArchitectureAdapter(cfg: Any)

Bases: ArchitectureAdapter

Adapter for a decoder-only transformer using RoPE, RMSNorm, gated SwiGLU MLPs, and optional sparse MoE feed-forward layers.

Uses an opaque NativeForwardAttentionBridge (delegates to Attention.forward) so RoPE runs under the source’s adjacent-pair convention rather than HF’s rotate-half – at the cost of only block-level hooks, no per-head hooks.

class transformer_lens.model_bridge.supported_architectures.pretrain.PretrainModelContainer(model: Module)

Bases: Module

Wraps the source model one level deeper (container.inner) so its own embed/blocks attrs don’t collide with TransformerBridge component_mapping keys, normalizes the forward return to the .logits contract, and strips _BRIDGE_COMPAT_KWARGS. Applied automatically by build_pretrain_bridge.

forward(*args: Any, **kwargs: Any) → Any

Define the computation performed at every call.

Should be overridden by all subclasses.

Note

Although the recipe for forward pass needs to be defined within this function, one should call the Module instance afterwards instead of this since the former takes care of running the registered hooks while the latter silently ignores them.

transformer_lens.model_bridge.supported_architectures.pretrain.build_pretrain_bridge(model: Module, cfg: TransformerBridgeConfig, *, device: Any = None, dtype: dtype | None = None, model_name: str | None = None) → TransformerBridge

Public entry point: wraps model in PretrainModelContainer and builds a TransformerBridge around it. Prefer this over calling build_bridge_from_module directly – the container is easy to forget.

device/dtype/model_name forward to build_bridge_from_module only when explicitly given.

bridge.train()/.eval() propagate to model via TransformerBridge.train() itself, which sets mode on original_model in addition to the registered module tree (original_model is deliberately not a registered submodule, so nn.Module.train()’s own recursion never reaches it). This adapter needs nothing extra for mode propagation.

Setting mode on model directly still works too and stays in sync.