Coverage for transformer_lens/model_bridge/supported_architectures/olmo.py: 100%
15 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
1"""OLMo architecture adapter."""
3from typing import Any
5from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
6from transformer_lens.model_bridge.generalized_components import (
7 BlockBridge,
8 EmbeddingBridge,
9 LinearBridge,
10 NormalizationBridge,
11 PositionEmbeddingsAttentionBridge,
12 RotaryEmbeddingBridge,
13 UnembeddingBridge,
14)
17class OlmoArchitectureAdapter(ArchitectureAdapter):
18 """Architecture adapter for OLMo (v1) models.
20 OLMo v1 uses a pre-norm architecture with a custom non-learnable LayerNorm
21 (fixed weight=1, bias=0), rotary position embeddings (RoPE), and gated MLP
22 (SwiGLU). Key differences from later OLMo variants:
24 - Pre-norm: LayerNorm is applied BEFORE attention and BEFORE MLP.
25 - Non-learnable LayerNorm: Weight and bias are not trainable parameters.
26 Delegating to HF's native forward via NormalizationBridge handles this correctly.
27 - No Q/K normalization in attention.
28 - Optional QKV clipping (applied out-of-place by the reconstructed
29 attention forward when config.clip_qkv is set).
31 Optional Parameters (may not exist in state_dict):
32 -------------------------------------------------
33 - blocks.{i}.attn.b_Q - No bias on query projection
34 - blocks.{i}.attn.b_K - No bias on key projection
35 - blocks.{i}.attn.b_V - No bias on value projection
36 - blocks.{i}.attn.b_O - No bias on output projection
37 - blocks.{i}.mlp.b_in - No bias on MLP up_proj
38 - blocks.{i}.mlp.b_gate - No bias on MLP gate_proj
39 - blocks.{i}.mlp.b_out - No bias on MLP down_proj
40 """
42 def __init__(self, cfg: Any) -> None:
43 """Initialize the OLMo architecture adapter."""
44 super().__init__(cfg)
46 # Set config variables for weight processing
47 self.cfg.normalization_type = "LN"
48 self.cfg.positional_embedding_type = "rotary"
49 self.cfg.final_rms = False
50 self.cfg.gated_mlp = True
51 self.cfg.attn_only = False
52 self.cfg.uses_rms_norm = False
53 # Force eager attention for numerical consistency with benchmark reference
54 self.cfg.attn_implementation = "eager"
56 self.weight_processing_conversions = {
57 **self._qkvo_weight_conversions(),
58 }
60 # Component mapping — PRE-NORM architecture:
61 # ln1 = input_layernorm (applied BEFORE attention)
62 # ln2 = post_attention_layernorm (applied BEFORE MLP)
63 self.component_mapping = {
64 "embed": EmbeddingBridge(name="model.embed_tokens"),
65 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg),
66 "blocks": BlockBridge(
67 name="model.layers",
68 submodules={
69 "ln1": NormalizationBridge(
70 name="input_layernorm",
71 config=self.cfg,
72 use_native_layernorm_autograd=True,
73 ),
74 "ln2": NormalizationBridge(
75 name="post_attention_layernorm",
76 config=self.cfg,
77 use_native_layernorm_autograd=True,
78 ),
79 "attn": PositionEmbeddingsAttentionBridge(
80 name="self_attn",
81 config=self.cfg,
82 submodules={
83 "q": LinearBridge(name="q_proj"),
84 "k": LinearBridge(name="k_proj"),
85 "v": LinearBridge(name="v_proj"),
86 "o": LinearBridge(name="o_proj"),
87 },
88 requires_attention_mask=True,
89 requires_position_embeddings=True,
90 ),
91 "mlp": self._gated_mlp(),
92 },
93 ),
94 "ln_final": NormalizationBridge(
95 name="model.norm",
96 config=self.cfg,
97 use_native_layernorm_autograd=True,
98 ),
99 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
100 }