Coverage for transformer_lens/model_bridge/supported_architectures/t5gemma2.py: 44%
48 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"""T5Gemma2 architecture adapter (text-only).
3T5Gemma2ForConditionalGeneration is a multimodal encoder-decoder model. This
4adapter bridges the text path only:
5- Encoder text stack under model.encoder.text_model (the SigLIP vision_tower and
6 multi_modal_projector are intentionally left unmapped).
7- Decoder stack under model.decoder.
9Key differences from T5Gemma:
10- Encoder text lives at model.encoder.text_model.* (not model.encoder.*).
11- The decoder uses a single T5Gemma2MergedAttention that fuses self- and
12 cross-attention with shared q/k/v/o projections; there is no separate
13 cross-attention module and no cross-attention layernorms.
14- Both encoder and decoder attention add Gemma-style QK-norm (q_norm/k_norm).
15- Per-layer sliding/full attention with dual RoPE and per-head QK-norm are all
16 handled natively by HF — the bridge only routes inputs and fires hooks.
17"""
19from typing import Any
21from transformer_lens.conversion_utils.conversion_steps import (
22 ArithmeticTensorConversion,
23 RearrangeTensorConversion,
24 TransposeTensorConversion,
25)
26from transformer_lens.conversion_utils.conversion_steps.arithmetic_tensor_conversion import (
27 OperationTypes,
28)
29from transformer_lens.conversion_utils.param_processing_conversion import (
30 ParamProcessingConversion,
31)
32from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
33from transformer_lens.model_bridge.generalized_components import (
34 AttentionBridge,
35 BlockBridge,
36 EmbeddingBridge,
37 GatedMLPBridge,
38 LinearBridge,
39 RMSNormalizationBridge,
40 RotaryEmbeddingBridge,
41 UnembeddingBridge,
42)
43from transformer_lens.model_bridge.generalized_components.t5gemma2_decoder_block import (
44 T5Gemma2DecoderBlockBridge,
45)
46from transformer_lens.model_bridge.generalized_components.t5gemma2_merged_attention import (
47 T5Gemma2MergedAttentionBridge,
48)
49from transformer_lens.utilities.attn_implementation import force_eager_attention
52class T5Gemma2ArchitectureAdapter(ArchitectureAdapter):
53 """Architecture adapter for T5Gemma2ForConditionalGeneration (text-only).
55 Encoder: BlockBridge over model.encoder.text_model.layers (Gemma-style, QK-norm, no cross-attn)
56 Decoder: T5Gemma2DecoderBlockBridge over model.decoder.layers (merged self+cross attention)
57 """
59 def __init__(self, cfg: Any) -> None:
60 super().__init__(cfg)
62 self.supports_fold_ln = False
64 # Config flags used by bridge weight processing
65 self.cfg.normalization_type = "RMS"
66 self.cfg.positional_embedding_type = "rotary"
67 self.cfg.final_rms = True
68 self.cfg.gated_mlp = True
69 self.cfg.attn_only = False
70 # Gemma-family GELU; the nested enc/dec config defeats the auto-mapper,
71 # which would otherwise leave act_fn at the "relu" default.
72 self.cfg.act_fn = "gelu_pytorch_tanh"
73 self.cfg.uses_rms_norm = True
74 # T5Gemma2 uses Gemma-style (1.0 + weight) RMSNorm offset
75 self.cfg.rmsnorm_uses_offset = True
77 # n_heads/n_kv are decoder-effective; the builder surfaces the encoder
78 # text stack's own counts for unbalanced pairs.
79 n_heads = self.cfg.n_heads
80 n_kv = getattr(self.cfg, "n_key_value_heads", None) or n_heads
81 enc_heads = getattr(self.cfg, "encoder_attention_heads", None) or n_heads
82 enc_kv = getattr(self.cfg, "encoder_key_value_heads", None) or n_kv
84 self.weight_processing_conversions = {
85 # Encoder self-attention
86 "encoder_blocks.{i}.self_attn.q_proj.weight": ParamProcessingConversion(
87 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_heads),
88 ),
89 "encoder_blocks.{i}.self_attn.k_proj.weight": ParamProcessingConversion(
90 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_kv),
91 ),
92 "encoder_blocks.{i}.self_attn.v_proj.weight": ParamProcessingConversion(
93 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_kv),
94 ),
95 "encoder_blocks.{i}.self_attn.o_proj.weight": ParamProcessingConversion(
96 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=enc_heads),
97 ),
98 # Encoder QK-norm (Gemma-style +1 offset)
99 "encoder_blocks.{i}.self_attn.q_norm.weight": ParamProcessingConversion(
100 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
101 ),
102 "encoder_blocks.{i}.self_attn.k_norm.weight": ParamProcessingConversion(
103 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
104 ),
105 # Encoder RMSNorm offset - HF stores raw weight; Gemma applies weight+1
106 "encoder_blocks.{i}.pre_self_attn_layernorm.weight": ParamProcessingConversion(
107 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
108 ),
109 "encoder_blocks.{i}.post_self_attn_layernorm.weight": ParamProcessingConversion(
110 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
111 ),
112 "encoder_blocks.{i}.pre_feedforward_layernorm.weight": ParamProcessingConversion(
113 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
114 ),
115 "encoder_blocks.{i}.post_feedforward_layernorm.weight": ParamProcessingConversion(
116 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
117 ),
118 # Encoder MLP (gated)
119 "encoder_blocks.{i}.mlp.gate_proj.weight": ParamProcessingConversion(
120 tensor_conversion=TransposeTensorConversion(),
121 ),
122 "encoder_blocks.{i}.mlp.up_proj.weight": ParamProcessingConversion(
123 tensor_conversion=TransposeTensorConversion(),
124 ),
125 "encoder_blocks.{i}.mlp.down_proj.weight": ParamProcessingConversion(
126 tensor_conversion=TransposeTensorConversion(),
127 ),
128 # Decoder merged attention (self + cross share these projections)
129 "decoder_blocks.{i}.self_attn.q_proj.weight": ParamProcessingConversion(
130 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_heads),
131 ),
132 "decoder_blocks.{i}.self_attn.k_proj.weight": ParamProcessingConversion(
133 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_kv),
134 ),
135 "decoder_blocks.{i}.self_attn.v_proj.weight": ParamProcessingConversion(
136 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_kv),
137 ),
138 "decoder_blocks.{i}.self_attn.o_proj.weight": ParamProcessingConversion(
139 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=n_heads),
140 ),
141 # Decoder QK-norm (Gemma-style +1 offset)
142 "decoder_blocks.{i}.self_attn.q_norm.weight": ParamProcessingConversion(
143 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
144 ),
145 "decoder_blocks.{i}.self_attn.k_norm.weight": ParamProcessingConversion(
146 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
147 ),
148 # Decoder RMSNorm offset
149 "decoder_blocks.{i}.pre_self_attn_layernorm.weight": ParamProcessingConversion(
150 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
151 ),
152 "decoder_blocks.{i}.post_self_attn_layernorm.weight": ParamProcessingConversion(
153 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
154 ),
155 "decoder_blocks.{i}.pre_feedforward_layernorm.weight": ParamProcessingConversion(
156 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
157 ),
158 "decoder_blocks.{i}.post_feedforward_layernorm.weight": ParamProcessingConversion(
159 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
160 ),
161 # Decoder MLP (gated)
162 "decoder_blocks.{i}.mlp.gate_proj.weight": ParamProcessingConversion(
163 tensor_conversion=TransposeTensorConversion(),
164 ),
165 "decoder_blocks.{i}.mlp.up_proj.weight": ParamProcessingConversion(
166 tensor_conversion=TransposeTensorConversion(),
167 ),
168 "decoder_blocks.{i}.mlp.down_proj.weight": ParamProcessingConversion(
169 tensor_conversion=TransposeTensorConversion(),
170 ),
171 # Final layer norms
172 "encoder_ln_final.weight": ParamProcessingConversion(
173 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
174 ),
175 "decoder_ln_final.weight": ParamProcessingConversion(
176 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
177 ),
178 # Unembed
179 "unembed.weight": ParamProcessingConversion(
180 tensor_conversion=TransposeTensorConversion(),
181 ),
182 }
184 self.component_mapping = {
185 # Encoder embedding and positional (text stack lives under text_model)
186 "encoder_embed": EmbeddingBridge(name="model.encoder.text_model.embed_tokens"),
187 "encoder_rotary_emb": RotaryEmbeddingBridge(name="model.encoder.text_model.rotary_emb"),
188 # Encoder layers - Gemma-style BlockBridge (pre/post norms, QK-norm attention, gated MLP)
189 "encoder_blocks": BlockBridge(
190 name="model.encoder.text_model.layers",
191 config=self.cfg,
192 submodules={
193 "ln1": RMSNormalizationBridge(name="pre_self_attn_layernorm", config=self.cfg),
194 "ln1_post": RMSNormalizationBridge(
195 name="post_self_attn_layernorm", config=self.cfg
196 ),
197 # Native delegation: the encoder uses per-layer sliding/full
198 # bidirectional windows carried by HF's per-layer mask, which the
199 # manual attention path does not apply (it drifts materially past
200 # the sliding_window length). Delegating keeps sliding correct.
201 "attn": AttentionBridge(
202 name="self_attn",
203 config=self.cfg,
204 # HF's T5Gemma2SelfAttention unpacks position_embeddings
205 # unconditionally, so component testing must supply it.
206 requires_position_embeddings=True,
207 submodules={
208 "q": LinearBridge(name="q_proj"),
209 "k": LinearBridge(name="k_proj"),
210 "v": LinearBridge(name="v_proj"),
211 "o": LinearBridge(name="o_proj"),
212 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
213 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
214 },
215 ),
216 "ln2": RMSNormalizationBridge(
217 name="pre_feedforward_layernorm", config=self.cfg
218 ),
219 "ln2_post": RMSNormalizationBridge(
220 name="post_feedforward_layernorm", config=self.cfg
221 ),
222 "mlp": GatedMLPBridge(
223 name="mlp",
224 config=self.cfg,
225 submodules={
226 "gate": LinearBridge(name="gate_proj"),
227 "in": LinearBridge(name="up_proj"),
228 "out": LinearBridge(name="down_proj"),
229 },
230 ),
231 },
232 ),
233 # Encoder final norm
234 "encoder_ln_final": RMSNormalizationBridge(
235 name="model.encoder.text_model.norm", config=self.cfg
236 ),
237 # Decoder embedding and positional
238 "decoder_embed": EmbeddingBridge(name="model.decoder.embed_tokens"),
239 "decoder_rotary_emb": RotaryEmbeddingBridge(name="model.decoder.rotary_emb"),
240 # Decoder layers — T5Gemma2DecoderBlockBridge (merged self+cross attention)
241 "decoder_blocks": T5Gemma2DecoderBlockBridge(
242 name="model.decoder.layers",
243 config=self.cfg,
244 submodules={
245 "ln1": RMSNormalizationBridge(name="pre_self_attn_layernorm", config=self.cfg),
246 "ln1_post": RMSNormalizationBridge(
247 name="post_self_attn_layernorm", config=self.cfg
248 ),
249 # Delegates to the native T5Gemma2MergedAttention (self+cross with
250 # shared q/k/v/o); the merged/cross logic, QK-norm, RoPE, and scaling
251 # cannot be reimplemented by the manual attention path. Exposes the
252 # self pattern (hook_pattern) and cross pattern (hook_cross_pattern).
253 "self_attn": T5Gemma2MergedAttentionBridge(
254 name="self_attn",
255 config=self.cfg,
256 submodules={
257 "q": LinearBridge(name="q_proj"),
258 "k": LinearBridge(name="k_proj"),
259 "v": LinearBridge(name="v_proj"),
260 "o": LinearBridge(name="o_proj"),
261 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
262 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
263 },
264 ),
265 "ln2": RMSNormalizationBridge(
266 name="pre_feedforward_layernorm", config=self.cfg
267 ),
268 "ln2_post": RMSNormalizationBridge(
269 name="post_feedforward_layernorm", config=self.cfg
270 ),
271 "mlp": GatedMLPBridge(
272 name="mlp",
273 config=self.cfg,
274 submodules={
275 "gate": LinearBridge(name="gate_proj"),
276 "in": LinearBridge(name="up_proj"),
277 "out": LinearBridge(name="down_proj"),
278 },
279 ),
280 },
281 ),
282 # Decoder final norm
283 "decoder_ln_final": RMSNormalizationBridge(name="model.decoder.norm", config=self.cfg),
284 # lm_head is T5Gemma2LMHead; the weight lives on its inner out_proj Linear
285 "unembed": UnembeddingBridge(name="lm_head.out_proj", config=self.cfg),
286 }
288 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None:
289 """Set up rotary embedding references for T5Gemma2 component testing.
291 Both the encoder text stack and the decoder carry their own rotary_emb. We
292 set the reference on all PositionEmbeddingsAttentionBridge instances so that
293 component-level forward calls can compute RoPE correctly, force eager
294 attention (so patterns are hookable), and enable native layernorm autograd
295 on QK-norm so the manual encoder path matches HF exactly.
296 """
297 encoder_rotary = hf_model.model.encoder.text_model.rotary_emb
299 force_eager_attention(hf_model)
301 # QK-norm must delegate to HF's exact RMSNorm autograd to avoid manual drift.
302 def _enable_qk_native_autograd(layers: Any) -> None:
303 for layer in layers:
304 attn = getattr(layer, "self_attn", None)
305 if attn is None:
306 continue
307 if hasattr(attn, "q_norm"):
308 attn.q_norm.use_native_layernorm_autograd = True
309 if hasattr(attn, "k_norm"):
310 attn.k_norm.use_native_layernorm_autograd = True
312 _enable_qk_native_autograd(hf_model.model.encoder.text_model.layers)
313 _enable_qk_native_autograd(hf_model.model.decoder.layers)
315 if bridge_model is not None:
316 for block in getattr(bridge_model, "encoder_blocks", []):
317 if hasattr(block, "attn") and hasattr(block.attn, "set_rotary_emb"):
318 block.attn.set_rotary_emb(encoder_rotary)
319 # Decoder self_attn delegates to native (which owns its RoPE), so it
320 # has no set_rotary_emb; nothing to wire.
322 enc_attn = self.get_generalized_component("encoder_blocks.0.attn")
323 if hasattr(enc_attn, "set_rotary_emb"):
324 enc_attn.set_rotary_emb(encoder_rotary)