Coverage for transformer_lens/model_bridge/supported_architectures/t5gemma.py: 46%
36 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"""T5Gemma architecture adapter.
3T5GemmaForConditionalGeneration is an encoder-decoder model combining:
4- Gemma-style RoPE, GQA, gated MLP, and RMSNorm with offset (+1.0)
5- Encoder-decoder cross-attention in the decoder stack
6- Nested config: encoder/decoder dims live in cfg.encoder / cfg.decoder
8Key differences from plain T5:
9- Uses model.encoder.layers / model.decoder.layers (not .block)
10- No relative position bias; uses RoPE instead
11- All norms are Gemma-style (weight + 1.0)
12- lm_head is T5GemmaLMHead wrapping out_proj (no .weight at the top level)
13"""
15from typing import Any
17from transformer_lens.conversion_utils.conversion_steps import (
18 ArithmeticTensorConversion,
19 RearrangeTensorConversion,
20 TransposeTensorConversion,
21)
22from transformer_lens.conversion_utils.conversion_steps.arithmetic_tensor_conversion import (
23 OperationTypes,
24)
25from transformer_lens.conversion_utils.param_processing_conversion import (
26 ParamProcessingConversion,
27)
28from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
29from transformer_lens.model_bridge.generalized_components import (
30 AttentionBridge,
31 BlockBridge,
32 EmbeddingBridge,
33 LinearBridge,
34 PositionEmbeddingsAttentionBridge,
35 RMSNormalizationBridge,
36 RotaryEmbeddingBridge,
37 UnembeddingBridge,
38)
39from transformer_lens.model_bridge.generalized_components.t5gemma_decoder_block import (
40 T5GemmaDecoderBlockBridge,
41)
44class T5GemmaArchitectureAdapter(ArchitectureAdapter):
45 """Architecture adapter for T5GemmaForConditionalGeneration.
47 Encoder: BlockBridge over model.encoder.layers (Gemma-style, no cross-attn)
48 Decoder: T5GemmaDecoderBlockBridge over model.decoder.layers (adds cross-attn hooks)
49 """
51 def __init__(self, cfg: Any) -> None:
52 super().__init__(cfg)
54 self.supports_fold_ln = False
56 # Config flags used by bridge weight processing
57 self._set_rms_rotary_defaults()
58 # Gemma-family GELU; the nested enc/dec config defeats the auto-mapper,
59 # which would otherwise leave act_fn at the "relu" default.
60 self.cfg.act_fn = "gelu_pytorch_tanh"
61 # T5Gemma uses Gemma-style (1.0 + weight) RMSNorm offset
62 self.cfg.rmsnorm_uses_offset = True
64 # n_heads/n_kv are decoder-effective; unbalanced pairs (t5gemma-9b-2b)
65 # set different encoder counts, surfaced by the builder.
66 n_heads = self.cfg.n_heads
67 n_kv = getattr(self.cfg, "n_key_value_heads", None) or n_heads
68 enc_heads = getattr(self.cfg, "encoder_attention_heads", None) or n_heads
69 enc_kv = getattr(self.cfg, "encoder_key_value_heads", None) or n_kv
71 self.weight_processing_conversions = {
72 # Encoder self-attention
73 "encoder_blocks.{i}.self_attn.q_proj.weight": ParamProcessingConversion(
74 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_heads),
75 ),
76 "encoder_blocks.{i}.self_attn.k_proj.weight": ParamProcessingConversion(
77 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_kv),
78 ),
79 "encoder_blocks.{i}.self_attn.v_proj.weight": ParamProcessingConversion(
80 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=enc_kv),
81 ),
82 "encoder_blocks.{i}.self_attn.o_proj.weight": ParamProcessingConversion(
83 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=enc_heads),
84 ),
85 # Encoder RMSNorm offset - HF stores raw weight; Gemma applies weight+1
86 "encoder_blocks.{i}.pre_self_attn_layernorm.weight": ParamProcessingConversion(
87 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
88 ),
89 "encoder_blocks.{i}.post_self_attn_layernorm.weight": ParamProcessingConversion(
90 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
91 ),
92 "encoder_blocks.{i}.pre_feedforward_layernorm.weight": ParamProcessingConversion(
93 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
94 ),
95 "encoder_blocks.{i}.post_feedforward_layernorm.weight": ParamProcessingConversion(
96 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
97 ),
98 # Encoder MLP (gated)
99 "encoder_blocks.{i}.mlp.gate_proj.weight": ParamProcessingConversion(
100 tensor_conversion=TransposeTensorConversion(),
101 ),
102 "encoder_blocks.{i}.mlp.up_proj.weight": ParamProcessingConversion(
103 tensor_conversion=TransposeTensorConversion(),
104 ),
105 "encoder_blocks.{i}.mlp.down_proj.weight": ParamProcessingConversion(
106 tensor_conversion=TransposeTensorConversion(),
107 ),
108 # Decoder self-attention
109 "decoder_blocks.{i}.self_attn.q_proj.weight": ParamProcessingConversion(
110 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_heads),
111 ),
112 "decoder_blocks.{i}.self_attn.k_proj.weight": ParamProcessingConversion(
113 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_kv),
114 ),
115 "decoder_blocks.{i}.self_attn.v_proj.weight": ParamProcessingConversion(
116 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_kv),
117 ),
118 "decoder_blocks.{i}.self_attn.o_proj.weight": ParamProcessingConversion(
119 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=n_heads),
120 ),
121 # Decoder cross-attention
122 "decoder_blocks.{i}.cross_attn.q_proj.weight": ParamProcessingConversion(
123 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_heads),
124 ),
125 "decoder_blocks.{i}.cross_attn.k_proj.weight": ParamProcessingConversion(
126 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_kv),
127 ),
128 "decoder_blocks.{i}.cross_attn.v_proj.weight": ParamProcessingConversion(
129 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=n_kv),
130 ),
131 "decoder_blocks.{i}.cross_attn.o_proj.weight": ParamProcessingConversion(
132 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=n_heads),
133 ),
134 # Decoder RMSNorm offset
135 "decoder_blocks.{i}.pre_self_attn_layernorm.weight": ParamProcessingConversion(
136 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
137 ),
138 "decoder_blocks.{i}.post_self_attn_layernorm.weight": ParamProcessingConversion(
139 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
140 ),
141 "decoder_blocks.{i}.pre_cross_attn_layernorm.weight": ParamProcessingConversion(
142 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
143 ),
144 "decoder_blocks.{i}.post_cross_attn_layernorm.weight": ParamProcessingConversion(
145 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
146 ),
147 "decoder_blocks.{i}.pre_feedforward_layernorm.weight": ParamProcessingConversion(
148 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
149 ),
150 "decoder_blocks.{i}.post_feedforward_layernorm.weight": ParamProcessingConversion(
151 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
152 ),
153 # Decoder MLP (gated)
154 "decoder_blocks.{i}.mlp.gate_proj.weight": ParamProcessingConversion(
155 tensor_conversion=TransposeTensorConversion(),
156 ),
157 "decoder_blocks.{i}.mlp.up_proj.weight": ParamProcessingConversion(
158 tensor_conversion=TransposeTensorConversion(),
159 ),
160 "decoder_blocks.{i}.mlp.down_proj.weight": ParamProcessingConversion(
161 tensor_conversion=TransposeTensorConversion(),
162 ),
163 # Final layer norms
164 "encoder_ln_final.weight": ParamProcessingConversion(
165 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
166 ),
167 "decoder_ln_final.weight": ParamProcessingConversion(
168 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0),
169 ),
170 # Unembed
171 "unembed.weight": ParamProcessingConversion(
172 tensor_conversion=TransposeTensorConversion(),
173 ),
174 }
176 self.component_mapping = {
177 # Encoder embedding and positional
178 "encoder_embed": EmbeddingBridge(name="model.encoder.embed_tokens"),
179 "encoder_rotary_emb": RotaryEmbeddingBridge(name="model.encoder.rotary_emb"),
180 # Encoder layers - Gemma-style BlockBridge (pre/post norms, RoPE attention, gated MLP)
181 "encoder_blocks": BlockBridge(
182 name="model.encoder.layers",
183 config=self.cfg,
184 submodules={
185 "ln1": RMSNormalizationBridge(name="pre_self_attn_layernorm", config=self.cfg),
186 "ln1_post": RMSNormalizationBridge(
187 name="post_self_attn_layernorm", config=self.cfg
188 ),
189 "attn": PositionEmbeddingsAttentionBridge(
190 name="self_attn",
191 config=self.cfg,
192 submodules={
193 "q": LinearBridge(name="q_proj"),
194 "k": LinearBridge(name="k_proj"),
195 "v": LinearBridge(name="v_proj"),
196 "o": LinearBridge(name="o_proj"),
197 },
198 requires_attention_mask=True,
199 requires_position_embeddings=True,
200 is_causal=False, # T5Gemma encoder is bidirectional
201 ),
202 "ln2": RMSNormalizationBridge(
203 name="pre_feedforward_layernorm", config=self.cfg
204 ),
205 "ln2_post": RMSNormalizationBridge(
206 name="post_feedforward_layernorm", config=self.cfg
207 ),
208 "mlp": self._gated_mlp(),
209 },
210 ),
211 # Encoder final norm
212 "encoder_ln_final": RMSNormalizationBridge(name="model.encoder.norm", config=self.cfg),
213 # Decoder embedding and positional
214 "decoder_embed": EmbeddingBridge(name="model.decoder.embed_tokens"),
215 "decoder_rotary_emb": RotaryEmbeddingBridge(name="model.decoder.rotary_emb"),
216 # Decoder layers — T5GemmaDecoderBlockBridge (adds cross-attn + two mid hooks)
217 "decoder_blocks": T5GemmaDecoderBlockBridge(
218 name="model.decoder.layers",
219 config=self.cfg,
220 submodules={
221 # Self-attention norms
222 "ln1": RMSNormalizationBridge(name="pre_self_attn_layernorm", config=self.cfg),
223 "ln1_post": RMSNormalizationBridge(
224 name="post_self_attn_layernorm", config=self.cfg
225 ),
226 "self_attn": PositionEmbeddingsAttentionBridge(
227 name="self_attn",
228 config=self.cfg,
229 submodules={
230 "q": LinearBridge(name="q_proj"),
231 "k": LinearBridge(name="k_proj"),
232 "v": LinearBridge(name="v_proj"),
233 "o": LinearBridge(name="o_proj"),
234 },
235 requires_attention_mask=True,
236 requires_position_embeddings=True,
237 ),
238 # Cross-attention norms
239 "ln2": RMSNormalizationBridge(name="pre_cross_attn_layernorm", config=self.cfg),
240 "ln2_post": RMSNormalizationBridge(
241 name="post_cross_attn_layernorm", config=self.cfg
242 ),
243 "cross_attn": AttentionBridge(
244 name="cross_attn",
245 config=self.cfg,
246 submodules={
247 "q": LinearBridge(name="q_proj"),
248 "k": LinearBridge(name="k_proj"),
249 "v": LinearBridge(name="v_proj"),
250 "o": LinearBridge(name="o_proj"),
251 },
252 is_cross_attention=True,
253 ),
254 # MLP norms
255 "ln3": RMSNormalizationBridge(
256 name="pre_feedforward_layernorm", config=self.cfg
257 ),
258 "ln3_post": RMSNormalizationBridge(
259 name="post_feedforward_layernorm", config=self.cfg
260 ),
261 "mlp": self._gated_mlp(),
262 },
263 ),
264 # Decoder final norm
265 "decoder_ln_final": RMSNormalizationBridge(name="model.decoder.norm", config=self.cfg),
266 # lm_head is T5GemmaLMHead; the weight lives on its inner out_proj Linear
267 "unembed": UnembeddingBridge(name="lm_head.out_proj", config=self.cfg),
268 }
270 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None:
271 """Set up rotary embedding references for T5Gemma component testing.
273 Both the encoder and decoder carry their own rotary_emb. We set the
274 reference on all PositionEmbeddingsAttentionBridge instances so that
275 component-level forward calls can compute RoPE correctly.
276 """
277 encoder_rotary = hf_model.model.encoder.rotary_emb
278 decoder_rotary = hf_model.model.decoder.rotary_emb
280 if hasattr(hf_model, "config") and hasattr(hf_model.config, "_attn_implementation"):
281 hf_model.config._attn_implementation = "eager"
283 if bridge_model is not None:
284 for block in getattr(bridge_model, "encoder_blocks", []):
285 if hasattr(block, "attn"):
286 block.attn.set_rotary_emb(encoder_rotary)
287 for block in getattr(bridge_model, "decoder_blocks", []):
288 if hasattr(block, "self_attn"):
289 block.self_attn.set_rotary_emb(decoder_rotary)
291 enc_attn = self.get_generalized_component("encoder_blocks.0.attn")
292 enc_attn.set_rotary_emb(encoder_rotary)
293 dec_self_attn = self.get_generalized_component("decoder_blocks.0.self_attn")
294 dec_self_attn.set_rotary_emb(decoder_rotary)