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

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) 

49from transformer_lens.utilities.attn_implementation import force_eager_attention 

50 

51 

52class T5Gemma2ArchitectureAdapter(ArchitectureAdapter): 

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

54 

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

58 

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

60 super().__init__(cfg) 

61 

62 self.supports_fold_ln = False 

63 

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 

76 

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 

83 

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 } 

183 

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 } 

287 

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

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

290 

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 

298 

299 force_eager_attention(hf_model) 

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)