Coverage for transformer_lens/model_bridge/supported_architectures/gemma3.py: 48%

32 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-21 19:27 +0000

1"""Gemma3 architecture adapter.""" 

2 

3 

4from typing import Any 

5 

6from transformer_lens.conversion_utils.conversion_steps import ( 

7 ArithmeticTensorConversion, 

8 TransposeTensorConversion, 

9) 

10from transformer_lens.conversion_utils.conversion_steps.arithmetic_tensor_conversion import ( 

11 OperationTypes, 

12) 

13from transformer_lens.conversion_utils.param_processing_conversion import ( 

14 ParamProcessingConversion, 

15) 

16from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

17from transformer_lens.model_bridge.generalized_components import ( 

18 BlockBridge, 

19 EmbeddingBridge, 

20 LinearBridge, 

21 RMSNormalizationBridge, 

22 RotaryEmbeddingBridge, 

23 UnembeddingBridge, 

24) 

25from transformer_lens.model_bridge.generalized_components.position_embeddings_attention import ( 

26 PositionEmbeddingsAttentionBridge, 

27) 

28 

29 

30def _enable_native_qk_norm_autograd(bridge_model: Any) -> None: 

31 """Component tests grad-check through the wrapped Q/K norms; delegate them 

32 to HF's native layernorm autograd (shared by text and multimodal Gemma-3).""" 

33 if bridge_model is None or not hasattr(bridge_model, "blocks"): 

34 return 

35 for block in bridge_model.blocks: 

36 hf_attn = getattr(getattr(block, "attn", None), "original_component", None) 

37 if hf_attn is None: 

38 continue 

39 if hasattr(hf_attn, "q_norm"): 

40 hf_attn.q_norm.use_native_layernorm_autograd = True 

41 if hasattr(hf_attn, "k_norm"): 

42 hf_attn.k_norm.use_native_layernorm_autograd = True 

43 

44 

45class Gemma3ArchitectureAdapter(ArchitectureAdapter): 

46 """Architecture adapter for Gemma3 models.""" 

47 

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

49 """Initialize the Gemma3 architecture adapter.""" 

50 super().__init__(cfg) 

51 

52 self.cfg.gated_mlp = True 

53 

54 self.cfg.uses_rms_norm = True 

55 self.cfg.normalization_type = "RMS" 

56 # Gemma models use (1.0 + weight) in RMSNorm instead of just weight 

57 # See: https://github.com/huggingface/transformers/pull/29402 

58 self.cfg.rmsnorm_uses_offset = True 

59 

60 # Gemma 3 uses rotary positional embeddings (dual RoPE) 

61 self.cfg.positional_embedding_type = "rotary" 

62 

63 # Use eager attention to support output_attentions for hook_attn_scores and hook_pattern 

64 # SDPA doesn't support output_attentions, which is required for HookedTransformer compatibility 

65 self.cfg.attn_implementation = "eager" 

66 

67 self.weight_processing_conversions = { 

68 # Note: Gemma3TextScaledWordEmbedding scales by sqrt(d_model) inside 

69 # its own forward(). Bridge.embed wraps that layer, so embed.hook_out 

70 # already captures the scaled value — no weight pre-scaling and no 

71 # hook_conversion needed (setup_hook_compatibility is a no-op). 

72 # 

73 # Q/K/V weight conversions 

74 **self._qkvo_weight_conversions(), 

75 # RMSNorm weight conversions - Gemma adds 1.0 to weights before applying 

76 # See: https://github.com/huggingface/transformers/pull/29402 

77 "blocks.{i}.ln1.weight": ParamProcessingConversion( 

78 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0), 

79 ), 

80 "blocks.{i}.ln1_post.weight": ParamProcessingConversion( 

81 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0), 

82 ), 

83 "blocks.{i}.ln2.weight": ParamProcessingConversion( 

84 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0), 

85 ), 

86 "blocks.{i}.ln2_post.weight": ParamProcessingConversion( 

87 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0), 

88 ), 

89 "ln_final.weight": ParamProcessingConversion( 

90 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0), 

91 ), 

92 # Gemma-3 also has q_norm and k_norm in attention 

93 "blocks.{i}.attn.q_norm.weight": ParamProcessingConversion( 

94 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0), 

95 ), 

96 "blocks.{i}.attn.k_norm.weight": ParamProcessingConversion( 

97 tensor_conversion=ArithmeticTensorConversion(OperationTypes.ADDITION, 1.0), 

98 ), 

99 # MLP weight conversions - transpose from [out, in] to [in, out] 

100 "blocks.{i}.mlp.gate.weight": ParamProcessingConversion( 

101 tensor_conversion=TransposeTensorConversion(), 

102 ), 

103 "blocks.{i}.mlp.in.weight": ParamProcessingConversion( 

104 tensor_conversion=TransposeTensorConversion(), 

105 ), 

106 "blocks.{i}.mlp.out.weight": ParamProcessingConversion( 

107 tensor_conversion=TransposeTensorConversion(), 

108 ), 

109 # Unembed weight conversion - transpose from [vocab, d_model] to [d_model, vocab] 

110 "unembed.weight": ParamProcessingConversion( 

111 tensor_conversion=TransposeTensorConversion(), 

112 ), 

113 # Note: Gemma-3 does NOT have biases on attention projections (q/k/v/o_proj.bias are all None) 

114 # No bias conversions needed 

115 } 

116 

117 # Set up component mapping with actual bridge instances 

118 self.component_mapping = { 

119 "embed": EmbeddingBridge(name="model.embed_tokens"), 

120 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"), 

121 "blocks": BlockBridge( 

122 name="model.layers", 

123 submodules={ 

124 # All Gemma-3 normalizations use simple RMSNorm pass-through 

125 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg), 

126 "ln1_post": RMSNormalizationBridge( 

127 name="post_attention_layernorm", config=self.cfg 

128 ), 

129 "ln2": RMSNormalizationBridge( 

130 name="pre_feedforward_layernorm", config=self.cfg 

131 ), 

132 "ln2_post": RMSNormalizationBridge( 

133 name="post_feedforward_layernorm", config=self.cfg 

134 ), 

135 "attn": PositionEmbeddingsAttentionBridge( 

136 name="self_attn", 

137 config=self.cfg, 

138 submodules={ 

139 "q": LinearBridge(name="q_proj"), 

140 "k": LinearBridge(name="k_proj"), 

141 "v": LinearBridge(name="v_proj"), 

142 "o": LinearBridge(name="o_proj"), 

143 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg), 

144 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg), 

145 }, 

146 ), 

147 "mlp": self._gated_mlp(), 

148 }, 

149 ), 

150 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg), 

151 "unembed": UnembeddingBridge(name="lm_head"), 

152 } 

153 

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

155 """Wire local RoPE + eager attention; q/k norms delegate to HF autograd. 

156 

157 Gemma-3 uses dual RoPE (global + local); component tests share the local 

158 instance across all layers (layers on global RoPE accept the tradeoff). 

159 """ 

160 super().setup_component_testing(hf_model, bridge_model) 

161 _enable_native_qk_norm_autograd(bridge_model)