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