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

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) 

42from transformer_lens.utilities.attn_implementation import force_eager_attention 

43 

44 

45class T5GemmaArchitectureAdapter(ArchitectureAdapter): 

46 """Architecture adapter for T5GemmaForConditionalGeneration. 

47 

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

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

50 """ 

51 

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

53 super().__init__(cfg) 

54 

55 self.supports_fold_ln = False 

56 

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 

64 

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 

71 

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 } 

176 

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 } 

270 

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

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

273 

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 

280 

281 force_eager_attention(hf_model) 

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)