Coverage for transformer_lens/model_bridge/supported_architectures/gemma3.py: 51%
29 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"""Gemma3 architecture adapter."""
4from typing import Any
6from transformer_lens.conversion_utils.conversion_steps import (
7 ArithmeticTensorConversion,
8 TransposeTensorConversion,
9)
10from transformer_lens.conversion_utils.conversion_steps.arithmetic_tensor_conversion import (
11 OperationTypes,
12)
13from transformer_lens.conversion_utils.param_processing_conversion import (
14 ParamProcessingConversion,
15)
16from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
17from transformer_lens.model_bridge.generalized_components import (
18 BlockBridge,
19 EmbeddingBridge,
20 LinearBridge,
21 RMSNormalizationBridge,
22 RotaryEmbeddingBridge,
23 UnembeddingBridge,
24)
25from transformer_lens.model_bridge.generalized_components.position_embeddings_attention import (
26 PositionEmbeddingsAttentionBridge,
27)
30class Gemma3ArchitectureAdapter(ArchitectureAdapter):
31 """Architecture adapter for Gemma3 models."""
33 def __init__(self, cfg: Any) -> None:
34 """Initialize the Gemma3 architecture adapter."""
35 super().__init__(cfg)
37 self.cfg.gated_mlp = True
39 self.cfg.uses_rms_norm = True
40 self.cfg.normalization_type = "RMS"
41 # Gemma models use (1.0 + weight) in RMSNorm instead of just weight
42 # See: https://github.com/huggingface/transformers/pull/29402
43 self.cfg.rmsnorm_uses_offset = True
45 # Gemma 3 uses rotary positional embeddings (dual RoPE)
46 self.cfg.positional_embedding_type = "rotary"
48 # Use eager attention to support output_attentions for hook_attn_scores and hook_pattern
49 # SDPA doesn't support output_attentions, which is required for HookedTransformer compatibility
50 self.cfg.attn_implementation = "eager"
52 self.weight_processing_conversions = {
53 # Note: Gemma3TextScaledWordEmbedding scales by sqrt(d_model) inside
54 # its own forward(). Bridge.embed wraps that layer, so embed.hook_out
55 # already captures the scaled value — no weight pre-scaling and no
56 # hook_conversion needed (setup_hook_compatibility is a no-op).
57 #
58 # Q/K/V weight conversions
59 **self._qkvo_weight_conversions(),
60 # RMSNorm weight conversions - Gemma adds 1.0 to weights before applying
61 # See: https://github.com/huggingface/transformers/pull/29402
62 "blocks.{i}.ln1.weight": ParamProcessingConversion(
63 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
64 ),
65 "blocks.{i}.ln1_post.weight": ParamProcessingConversion(
66 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
67 ),
68 "blocks.{i}.ln2.weight": ParamProcessingConversion(
69 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
70 ),
71 "blocks.{i}.ln2_post.weight": ParamProcessingConversion(
72 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
73 ),
74 "ln_final.weight": ParamProcessingConversion(
75 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
76 ),
77 # Gemma-3 also has q_norm and k_norm in attention
78 "blocks.{i}.attn.q_norm.weight": ParamProcessingConversion(
79 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
80 ),
81 "blocks.{i}.attn.k_norm.weight": ParamProcessingConversion(
82 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
83 ),
84 # MLP weight conversions - transpose from [out, in] to [in, out]
85 "blocks.{i}.mlp.gate.weight": ParamProcessingConversion(
86 tensor_conversion=TransposeTensorConversion(),
87 ),
88 "blocks.{i}.mlp.in.weight": ParamProcessingConversion(
89 tensor_conversion=TransposeTensorConversion(),
90 ),
91 "blocks.{i}.mlp.out.weight": ParamProcessingConversion(
92 tensor_conversion=TransposeTensorConversion(),
93 ),
94 # Unembed weight conversion - transpose from [vocab, d_model] to [d_model, vocab]
95 "unembed.weight": ParamProcessingConversion(
96 tensor_conversion=TransposeTensorConversion(),
97 ),
98 # Note: Gemma-3 does NOT have biases on attention projections (q/k/v/o_proj.bias are all None)
99 # No bias conversions needed
100 }
102 # Set up component mapping with actual bridge instances
103 self.component_mapping = {
104 "embed": EmbeddingBridge(name="model.embed_tokens"),
105 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
106 "blocks": BlockBridge(
107 name="model.layers",
108 submodules={
109 # All Gemma-3 normalizations use simple RMSNorm pass-through
110 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
111 "ln1_post": RMSNormalizationBridge(
112 name="post_attention_layernorm", config=self.cfg
113 ),
114 "ln2": RMSNormalizationBridge(
115 name="pre_feedforward_layernorm", config=self.cfg
116 ),
117 "ln2_post": RMSNormalizationBridge(
118 name="post_feedforward_layernorm", config=self.cfg
119 ),
120 "attn": PositionEmbeddingsAttentionBridge(
121 name="self_attn",
122 config=self.cfg,
123 submodules={
124 "q": LinearBridge(name="q_proj"),
125 "k": LinearBridge(name="k_proj"),
126 "v": LinearBridge(name="v_proj"),
127 "o": LinearBridge(name="o_proj"),
128 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
129 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
130 },
131 ),
132 "mlp": self._gated_mlp(),
133 },
134 ),
135 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
136 "unembed": UnembeddingBridge(name="lm_head"),
137 }
139 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None:
140 """Wire local RoPE + eager attention; q/k norms delegate to HF autograd.
142 Gemma-3 uses dual RoPE (global + local); component tests share the local
143 instance across all layers (layers on global RoPE accept the tradeoff).
144 """
145 super().setup_component_testing(hf_model, bridge_model)
146 if bridge_model is not None and hasattr(bridge_model, "blocks"):
147 for block in bridge_model.blocks:
148 hf_attn = getattr(block, "attn", None) and getattr(
149 block.attn, "original_component", None
150 )
151 if hf_attn is None:
152 continue
153 if hasattr(hf_attn, "q_norm"):
154 hf_attn.q_norm.use_native_layernorm_autograd = True
155 if hasattr(hf_attn, "k_norm"):
156 hf_attn.k_norm.use_native_layernorm_autograd = True