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

1"""T5Gemma architecture adapter. 

2 

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 

7 

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

14 

15from typing import Any 

16 

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) 

42 

43 

44class T5GemmaArchitectureAdapter(ArchitectureAdapter): 

45 """Architecture adapter for T5GemmaForConditionalGeneration. 

46 

47 Encoder: BlockBridge over model.encoder.layers (Gemma-style, no cross-attn) 

48 Decoder: T5GemmaDecoderBlockBridge over model.decoder.layers (adds cross-attn hooks) 

49 """ 

50 

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

52 super().__init__(cfg) 

53 

54 self.supports_fold_ln = False 

55 

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 

63 

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 

70 

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 } 

175 

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 } 

269 

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

271 """Set up rotary embedding references for T5Gemma component testing. 

272 

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 

279 

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

281 hf_model.config._attn_implementation = "eager" 

282 

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) 

290 

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)