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