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