Coverage for transformer_lens/model_bridge/supported_architectures/gemma1.py: 100%
14 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Gemma1 architecture adapter."""
3from typing import Any
5from transformer_lens.conversion_utils.conversion_steps import (
6 ArithmeticTensorConversion,
7 TransposeTensorConversion,
8)
9from transformer_lens.conversion_utils.conversion_steps.arithmetic_tensor_conversion import (
10 OperationTypes,
11)
12from transformer_lens.conversion_utils.param_processing_conversion import (
13 ParamProcessingConversion,
14)
15from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
16from transformer_lens.model_bridge.generalized_components import (
17 BlockBridge,
18 EmbeddingBridge,
19 LinearBridge,
20 PositionEmbeddingsAttentionBridge,
21 RMSNormalizationBridge,
22 RotaryEmbeddingBridge,
23 UnembeddingBridge,
24)
27class Gemma1ArchitectureAdapter(ArchitectureAdapter):
28 """Architecture adapter for Gemma1 models."""
30 def __init__(self, cfg: Any) -> None:
31 """Initialize the Gemma1 architecture adapter."""
32 super().__init__(cfg)
34 self._set_rms_rotary_defaults()
36 # Gemma models use BOS tokens (tokenizer prepends BOS by default)
37 # Matches HookedTransformer behavior (default_prepend_bos = True)
38 self.cfg.default_prepend_bos = True
39 # Gemma models use (1.0 + weight) in RMSNorm instead of just weight
40 # See: https://github.com/huggingface/transformers/pull/29402
41 self.cfg.rmsnorm_uses_offset = True
43 self.weight_processing_conversions = {
44 # NOTE: Gemma1 scales embeddings by sqrt(d_model) at RUNTIME inside
45 # GemmaTextScaledWordEmbedding.forward() (HF transformers >= 5.0).
46 # That layer is what bridge.embed wraps, so embed.hook_out already
47 # captures the scaled value — matching HookedTransformer's hook_embed
48 # (which uses pre-scaled W_E). We must NOT pre-scale weights here and
49 # we must NOT install a runtime hook_conversion that re-scales.
50 #
51 # Attention weight conversions
52 **self._qkvo_weight_conversions(),
53 # RMSNorm weight conversions - Gemma adds 1.0 to weights before applying
54 # See: https://github.com/huggingface/transformers/pull/29402
55 "blocks.{i}.ln1.weight": ParamProcessingConversion(
56 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
57 ),
58 "blocks.{i}.ln2.weight": ParamProcessingConversion(
59 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
60 ),
61 "ln_final.weight": ParamProcessingConversion(
62 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
63 ),
64 # MLP weight conversions - transpose from [out, in] to [in, out]
65 "blocks.{i}.mlp.gate.weight": ParamProcessingConversion(
66 tensor_conversion=TransposeTensorConversion(),
67 ),
68 "blocks.{i}.mlp.in.weight": ParamProcessingConversion(
69 tensor_conversion=TransposeTensorConversion(),
70 ),
71 "blocks.{i}.mlp.out.weight": ParamProcessingConversion(
72 tensor_conversion=TransposeTensorConversion(),
73 ),
74 # Unembed weight conversion - transpose from [vocab, d_model] to [d_model, vocab]
75 "unembed.weight": ParamProcessingConversion(
76 tensor_conversion=TransposeTensorConversion(),
77 ),
78 }
80 self.component_mapping = {
81 "embed": EmbeddingBridge(name="model.embed_tokens"),
82 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg),
83 "blocks": BlockBridge(
84 name="model.layers",
85 submodules={
86 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
87 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
88 "attn": PositionEmbeddingsAttentionBridge(
89 name="self_attn",
90 config=self.cfg,
91 submodules={
92 "q": LinearBridge(name="q_proj"),
93 "k": LinearBridge(name="k_proj"),
94 "v": LinearBridge(name="v_proj"),
95 "o": LinearBridge(name="o_proj"),
96 },
97 requires_attention_mask=True,
98 requires_position_embeddings=True,
99 ),
100 "mlp": self._gated_mlp(),
101 },
102 ),
103 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
104 "unembed": UnembeddingBridge(name="lm_head"),
105 }