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

1"""T5Gemma2 architecture adapter (text-only). 

2 

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. 

8 

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""" 

18 

19from typing import Any 

20 

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) 

49 

50 

51class T5Gemma2ArchitectureAdapter(ArchitectureAdapter): 

52 """Architecture adapter for T5Gemma2ForConditionalGeneration (text-only). 

53 

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 """ 

57 

58 def __init__(self, cfg: Any) -> None: 

59 super().__init__(cfg) 

60 

61 self.supports_fold_ln = False 

62 

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 

75 

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 

82 

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 } 

182 

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 } 

286 

287 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None: 

288 """Set up rotary embedding references for T5Gemma2 component testing. 

289 

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 

297 

298 if hasattr(hf_model, "config") and hasattr(hf_model.config, "_attn_implementation"): 

299 hf_model.config._attn_implementation = "eager" 

300 

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 

311 

312 _enable_qk_native_autograd(hf_model.model.encoder.text_model.layers) 

313 _enable_qk_native_autograd(hf_model.model.decoder.layers) 

314 

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. 

321 

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)