Coverage for transformer_lens/model_bridge/supported_architectures/gemma2.py: 83%
19 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"""Gemma2 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 Gemma2ArchitectureAdapter(ArchitectureAdapter):
28 """Architecture adapter for Gemma2 models."""
30 def __init__(self, cfg: Any) -> None:
31 """Initialize the Gemma2 architecture adapter."""
32 super().__init__(cfg)
34 self._set_rms_rotary_defaults()
36 # Gemma models use (1.0 + weight) in RMSNorm instead of just weight
37 # See: https://github.com/huggingface/transformers/pull/29402
38 self.cfg.rmsnorm_uses_offset = True
40 # Gemma2 uses logit softcapping
41 if hasattr(self.cfg, "final_logit_softcapping"): 41 ↛ 42line 41 didn't jump to line 42 because the condition on line 41 was never true
42 self.cfg.output_logits_soft_cap = self.cfg.final_logit_softcapping
43 if hasattr(self.cfg, "attn_logit_softcapping"): 43 ↛ 44line 43 didn't jump to line 44 because the condition on line 43 was never true
44 self.cfg.attn_scores_soft_cap = self.cfg.attn_logit_softcapping
46 # Note: n_key_value_heads is now automatically mapped from num_key_value_heads
47 # by map_default_transformer_lens_config() in sources/transformers.py
49 self.weight_processing_conversions = {
50 # NOTE: Gemma2 scales embeddings by sqrt(d_model) at RUNTIME inside
51 # Gemma2TextScaledWordEmbedding.forward() (HF transformers >= 5.0).
52 # That layer is what bridge.embed wraps, so embed.hook_out already
53 # captures the scaled value — matching HookedTransformer's hook_embed
54 # (which uses pre-scaled W_E). We must NOT pre-scale weights here and
55 # we must NOT install a runtime hook_conversion that re-scales.
56 **self._qkvo_weight_conversions(),
57 # RMSNorm weight conversions - Gemma adds 1.0 to weights before applying
58 # See: https://github.com/huggingface/transformers/pull/29402
59 "blocks.{i}.ln1.weight": ParamProcessingConversion(
60 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
61 ),
62 "blocks.{i}.ln1_post.weight": ParamProcessingConversion(
63 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
64 ),
65 "blocks.{i}.ln2.weight": ParamProcessingConversion(
66 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
67 ),
68 "blocks.{i}.ln2_post.weight": ParamProcessingConversion(
69 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
70 ),
71 "ln_final.weight": ParamProcessingConversion(
72 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
73 ),
74 # MLP weight conversions - transpose from [out, in] to [in, out]
75 "blocks.{i}.mlp.gate.weight": ParamProcessingConversion(
76 tensor_conversion=TransposeTensorConversion(),
77 ),
78 "blocks.{i}.mlp.in.weight": ParamProcessingConversion(
79 tensor_conversion=TransposeTensorConversion(),
80 ),
81 "blocks.{i}.mlp.out.weight": ParamProcessingConversion(
82 tensor_conversion=TransposeTensorConversion(),
83 ),
84 # # Unembed weight conversion - transpose from [vocab, d_model] to [d_model, vocab]
85 "unembed.weight": ParamProcessingConversion(
86 tensor_conversion=TransposeTensorConversion(),
87 ),
88 }
90 self.component_mapping = {
91 "embed": EmbeddingBridge(name="model.embed_tokens"),
92 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
93 "blocks": BlockBridge(
94 name="model.layers",
95 config=self.cfg,
96 submodules={
97 # Gemma 2 uses RMSNorm for all normalization layers
98 **self._block_norms(),
99 # Gemma 2 uses PositionEmbeddingsAttentionBridge like Gemma 3
100 "attn": PositionEmbeddingsAttentionBridge(
101 name="self_attn",
102 config=self.cfg,
103 submodules={
104 "q": LinearBridge(name="q_proj"),
105 "k": LinearBridge(name="k_proj"),
106 "v": LinearBridge(name="v_proj"),
107 "o": LinearBridge(name="o_proj"),
108 },
109 requires_attention_mask=True,
110 requires_position_embeddings=True,
111 ),
112 "mlp": self._gated_mlp(),
113 },
114 ),
115 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
116 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
117 }
119 def _block_norms(self):
120 """Norm-layout seam; VaultGemma drops the two post-norms."""
121 return {
122 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
123 "ln1_post": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
124 "ln2": RMSNormalizationBridge(name="pre_feedforward_layernorm", config=self.cfg),
125 "ln2_post": RMSNormalizationBridge(name="post_feedforward_layernorm", config=self.cfg),
126 }