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

29 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +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 

30class Gemma3ArchitectureAdapter(ArchitectureAdapter): 

31 """Architecture adapter for Gemma3 models.""" 

32 

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

34 """Initialize the Gemma3 architecture adapter.""" 

35 super().__init__(cfg) 

36 

37 self.cfg.gated_mlp = True 

38 

39 self.cfg.uses_rms_norm = True 

40 self.cfg.normalization_type = "RMS" 

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

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

43 self.cfg.rmsnorm_uses_offset = True 

44 

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

46 self.cfg.positional_embedding_type = "rotary" 

47 

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

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

50 self.cfg.attn_implementation = "eager" 

51 

52 self.weight_processing_conversions = { 

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

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

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

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

57 # 

58 # Q/K/V weight conversions 

59 **self._qkvo_weight_conversions(), 

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

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

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

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

64 ), 

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

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

67 ), 

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

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

70 ), 

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

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

73 ), 

74 "ln_final.weight": ParamProcessingConversion( 

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

76 ), 

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

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

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

80 ), 

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

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

83 ), 

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

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

86 tensor_conversion=TransposeTensorConversion(), 

87 ), 

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

89 tensor_conversion=TransposeTensorConversion(), 

90 ), 

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

92 tensor_conversion=TransposeTensorConversion(), 

93 ), 

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

95 "unembed.weight": ParamProcessingConversion( 

96 tensor_conversion=TransposeTensorConversion(), 

97 ), 

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

99 # No bias conversions needed 

100 } 

101 

102 # Set up component mapping with actual bridge instances 

103 self.component_mapping = { 

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

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

106 "blocks": BlockBridge( 

107 name="model.layers", 

108 submodules={ 

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

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

111 "ln1_post": RMSNormalizationBridge( 

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

113 ), 

114 "ln2": RMSNormalizationBridge( 

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

116 ), 

117 "ln2_post": RMSNormalizationBridge( 

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

119 ), 

120 "attn": PositionEmbeddingsAttentionBridge( 

121 name="self_attn", 

122 config=self.cfg, 

123 submodules={ 

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

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

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

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

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

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

130 }, 

131 ), 

132 "mlp": self._gated_mlp(), 

133 }, 

134 ), 

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

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

137 } 

138 

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

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

141 

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

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

144 """ 

145 super().setup_component_testing(hf_model, bridge_model) 

146 if bridge_model is not None and hasattr(bridge_model, "blocks"): 

147 for block in bridge_model.blocks: 

148 hf_attn = getattr(block, "attn", None) and getattr( 

149 block.attn, "original_component", None 

150 ) 

151 if hf_attn is None: 

152 continue 

153 if hasattr(hf_attn, "q_norm"): 

154 hf_attn.q_norm.use_native_layernorm_autograd = True 

155 if hasattr(hf_attn, "k_norm"): 

156 hf_attn.k_norm.use_native_layernorm_autograd = True