Coverage for transformer_lens/model_bridge/supported_architectures/t5gemma2.py: 42%
48 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +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)
51class T5Gemma2ArchitectureAdapter(ArchitectureAdapter):
52 """Architecture adapter for T5Gemma2ForConditionalGeneration (text-only).
54 Encoder: BlockBridge over model.encoder.text_model.layers (Gemma-style, QK-norm, no cross-attn)
55 Decoder: T5Gemma2DecoderBlockBridge over model.decoder.layers (merged self+cross attention)
56 """
58 def __init__(self, cfg: Any) -> None:
59 super().__init__(cfg)
61 self.supports_fold_ln = False
63 # Config flags used by bridge weight processing
64 self.cfg.normalization_type = "RMS"
65 self.cfg.positional_embedding_type = "rotary"
66 self.cfg.final_rms = True
67 self.cfg.gated_mlp = True
68 self.cfg.attn_only = False
69 # Gemma-family GELU; the nested enc/dec config defeats the auto-mapper,
70 # which would otherwise leave act_fn at the "relu" default.
71 self.cfg.act_fn = "gelu_pytorch_tanh"
72 self.cfg.uses_rms_norm = True
73 # T5Gemma2 uses Gemma-style (1.0 + weight) RMSNorm offset
74 self.cfg.rmsnorm_uses_offset = True
76 # n_heads/n_kv are decoder-effective; the builder surfaces the encoder
77 # text stack's own counts for unbalanced pairs.
78 n_heads = self.cfg.n_heads
79 n_kv = getattr(self.cfg, "n_key_value_heads", None) or n_heads
80 enc_heads = getattr(self.cfg, "encoder_attention_heads", None) or n_heads
81 enc_kv = getattr(self.cfg, "encoder_key_value_heads", None) or n_kv
83 self.weight_processing_conversions = {
84 # Encoder self-attention
85 "encoder_blocks.{i}.self_attn.q_proj.weight": ParamProcessingConversion(
86 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_heads),
87 ),
88 "encoder_blocks.{i}.self_attn.k_proj.weight": ParamProcessingConversion(
89 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_kv),
90 ),
91 "encoder_blocks.{i}.self_attn.v_proj.weight": ParamProcessingConversion(
92 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_kv),
93 ),
94 "encoder_blocks.{i}.self_attn.o_proj.weight": ParamProcessingConversion(
95 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=enc_heads),
96 ),
97 # Encoder QK-norm (Gemma-style +1 offset)
98 "encoder_blocks.{i}.self_attn.q_norm.weight": ParamProcessingConversion(
99 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
100 ),
101 "encoder_blocks.{i}.self_attn.k_norm.weight": ParamProcessingConversion(
102 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
103 ),
104 # Encoder RMSNorm offset - HF stores raw weight; Gemma applies weight+1
105 "encoder_blocks.{i}.pre_self_attn_layernorm.weight": ParamProcessingConversion(
106 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
107 ),
108 "encoder_blocks.{i}.post_self_attn_layernorm.weight": ParamProcessingConversion(
109 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
110 ),
111 "encoder_blocks.{i}.pre_feedforward_layernorm.weight": ParamProcessingConversion(
112 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
113 ),
114 "encoder_blocks.{i}.post_feedforward_layernorm.weight": ParamProcessingConversion(
115 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
116 ),
117 # Encoder MLP (gated)
118 "encoder_blocks.{i}.mlp.gate_proj.weight": ParamProcessingConversion(
119 tensor_conversion=TransposeTensorConversion(),
120 ),
121 "encoder_blocks.{i}.mlp.up_proj.weight": ParamProcessingConversion(
122 tensor_conversion=TransposeTensorConversion(),
123 ),
124 "encoder_blocks.{i}.mlp.down_proj.weight": ParamProcessingConversion(
125 tensor_conversion=TransposeTensorConversion(),
126 ),
127 # Decoder merged attention (self + cross share these projections)
128 "decoder_blocks.{i}.self_attn.q_proj.weight": ParamProcessingConversion(
129 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_heads),
130 ),
131 "decoder_blocks.{i}.self_attn.k_proj.weight": ParamProcessingConversion(
132 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_kv),
133 ),
134 "decoder_blocks.{i}.self_attn.v_proj.weight": ParamProcessingConversion(
135 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_kv),
136 ),
137 "decoder_blocks.{i}.self_attn.o_proj.weight": ParamProcessingConversion(
138 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=n_heads),
139 ),
140 # Decoder QK-norm (Gemma-style +1 offset)
141 "decoder_blocks.{i}.self_attn.q_norm.weight": ParamProcessingConversion(
142 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
143 ),
144 "decoder_blocks.{i}.self_attn.k_norm.weight": ParamProcessingConversion(
145 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
146 ),
147 # Decoder RMSNorm offset
148 "decoder_blocks.{i}.pre_self_attn_layernorm.weight": ParamProcessingConversion(
149 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
150 ),
151 "decoder_blocks.{i}.post_self_attn_layernorm.weight": ParamProcessingConversion(
152 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
153 ),
154 "decoder_blocks.{i}.pre_feedforward_layernorm.weight": ParamProcessingConversion(
155 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
156 ),
157 "decoder_blocks.{i}.post_feedforward_layernorm.weight": ParamProcessingConversion(
158 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
159 ),
160 # Decoder MLP (gated)
161 "decoder_blocks.{i}.mlp.gate_proj.weight": ParamProcessingConversion(
162 tensor_conversion=TransposeTensorConversion(),
163 ),
164 "decoder_blocks.{i}.mlp.up_proj.weight": ParamProcessingConversion(
165 tensor_conversion=TransposeTensorConversion(),
166 ),
167 "decoder_blocks.{i}.mlp.down_proj.weight": ParamProcessingConversion(
168 tensor_conversion=TransposeTensorConversion(),
169 ),
170 # Final layer norms
171 "encoder_ln_final.weight": ParamProcessingConversion(
172 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
173 ),
174 "decoder_ln_final.weight": ParamProcessingConversion(
175 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
176 ),
177 # Unembed
178 "unembed.weight": ParamProcessingConversion(
179 tensor_conversion=TransposeTensorConversion(),
180 ),
181 }
183 self.component_mapping = {
184 # Encoder embedding and positional (text stack lives under text_model)
185 "encoder_embed": EmbeddingBridge(name="model.encoder.text_model.embed_tokens"),
186 "encoder_rotary_emb": RotaryEmbeddingBridge(name="model.encoder.text_model.rotary_emb"),
187 # Encoder layers - Gemma-style BlockBridge (pre/post norms, QK-norm attention, gated MLP)
188 "encoder_blocks": BlockBridge(
189 name="model.encoder.text_model.layers",
190 config=self.cfg,
191 submodules={
192 "ln1": RMSNormalizationBridge(name="pre_self_attn_layernorm", config=self.cfg),
193 "ln1_post": RMSNormalizationBridge(
194 name="post_self_attn_layernorm", config=self.cfg
195 ),
196 # Native delegation: the encoder uses per-layer sliding/full
197 # bidirectional windows carried by HF's per-layer mask, which the
198 # manual attention path does not apply (it drifts materially past
199 # the sliding_window length). Delegating keeps sliding correct.
200 "attn": AttentionBridge(
201 name="self_attn",
202 config=self.cfg,
203 # HF's T5Gemma2SelfAttention unpacks position_embeddings
204 # unconditionally, so component testing must supply it.
205 requires_position_embeddings=True,
206 submodules={
207 "q": LinearBridge(name="q_proj"),
208 "k": LinearBridge(name="k_proj"),
209 "v": LinearBridge(name="v_proj"),
210 "o": LinearBridge(name="o_proj"),
211 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
212 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
213 },
214 ),
215 "ln2": RMSNormalizationBridge(
216 name="pre_feedforward_layernorm", config=self.cfg
217 ),
218 "ln2_post": RMSNormalizationBridge(
219 name="post_feedforward_layernorm", config=self.cfg
220 ),
221 "mlp": GatedMLPBridge(
222 name="mlp",
223 config=self.cfg,
224 submodules={
225 "gate": LinearBridge(name="gate_proj"),
226 "in": LinearBridge(name="up_proj"),
227 "out": LinearBridge(name="down_proj"),
228 },
229 ),
230 },
231 ),
232 # Encoder final norm
233 "encoder_ln_final": RMSNormalizationBridge(
234 name="model.encoder.text_model.norm", config=self.cfg
235 ),
236 # Decoder embedding and positional
237 "decoder_embed": EmbeddingBridge(name="model.decoder.embed_tokens"),
238 "decoder_rotary_emb": RotaryEmbeddingBridge(name="model.decoder.rotary_emb"),
239 # Decoder layers — T5Gemma2DecoderBlockBridge (merged self+cross attention)
240 "decoder_blocks": T5Gemma2DecoderBlockBridge(
241 name="model.decoder.layers",
242 config=self.cfg,
243 submodules={
244 "ln1": RMSNormalizationBridge(name="pre_self_attn_layernorm", config=self.cfg),
245 "ln1_post": RMSNormalizationBridge(
246 name="post_self_attn_layernorm", config=self.cfg
247 ),
248 # Delegates to the native T5Gemma2MergedAttention (self+cross with
249 # shared q/k/v/o); the merged/cross logic, QK-norm, RoPE, and scaling
250 # cannot be reimplemented by the manual attention path. Exposes the
251 # self pattern (hook_pattern) and cross pattern (hook_cross_pattern).
252 "self_attn": T5Gemma2MergedAttentionBridge(
253 name="self_attn",
254 config=self.cfg,
255 submodules={
256 "q": LinearBridge(name="q_proj"),
257 "k": LinearBridge(name="k_proj"),
258 "v": LinearBridge(name="v_proj"),
259 "o": LinearBridge(name="o_proj"),
260 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
261 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
262 },
263 ),
264 "ln2": RMSNormalizationBridge(
265 name="pre_feedforward_layernorm", config=self.cfg
266 ),
267 "ln2_post": RMSNormalizationBridge(
268 name="post_feedforward_layernorm", config=self.cfg
269 ),
270 "mlp": GatedMLPBridge(
271 name="mlp",
272 config=self.cfg,
273 submodules={
274 "gate": LinearBridge(name="gate_proj"),
275 "in": LinearBridge(name="up_proj"),
276 "out": LinearBridge(name="down_proj"),
277 },
278 ),
279 },
280 ),
281 # Decoder final norm
282 "decoder_ln_final": RMSNormalizationBridge(name="model.decoder.norm", config=self.cfg),
283 # lm_head is T5Gemma2LMHead; the weight lives on its inner out_proj Linear
284 "unembed": UnembeddingBridge(name="lm_head.out_proj", config=self.cfg),
285 }
287 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None:
288 """Set up rotary embedding references for T5Gemma2 component testing.
290 Both the encoder text stack and the decoder carry their own rotary_emb. We
291 set the reference on all PositionEmbeddingsAttentionBridge instances so that
292 component-level forward calls can compute RoPE correctly, force eager
293 attention (so patterns are hookable), and enable native layernorm autograd
294 on QK-norm so the manual encoder path matches HF exactly.
295 """
296 encoder_rotary = hf_model.model.encoder.text_model.rotary_emb
298 if hasattr(hf_model, "config") and hasattr(hf_model.config, "_attn_implementation"):
299 hf_model.config._attn_implementation = "eager"
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)