Coverage for transformer_lens/model_bridge/supported_architectures/gemma2.py: 83%

19 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""Gemma2 architecture adapter.""" 

2 

3from typing import Any 

4 

5from transformer_lens.conversion_utils.conversion_steps import ( 

6 ArithmeticTensorConversion, 

7 TransposeTensorConversion, 

8) 

9from transformer_lens.conversion_utils.conversion_steps.arithmetic_tensor_conversion import ( 

10 OperationTypes, 

11) 

12from transformer_lens.conversion_utils.param_processing_conversion import ( 

13 ParamProcessingConversion, 

14) 

15from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

16from transformer_lens.model_bridge.generalized_components import ( 

17 BlockBridge, 

18 EmbeddingBridge, 

19 LinearBridge, 

20 PositionEmbeddingsAttentionBridge, 

21 RMSNormalizationBridge, 

22 RotaryEmbeddingBridge, 

23 UnembeddingBridge, 

24) 

25 

26 

27class Gemma2ArchitectureAdapter(ArchitectureAdapter): 

28 """Architecture adapter for Gemma2 models.""" 

29 

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

31 """Initialize the Gemma2 architecture adapter.""" 

32 super().__init__(cfg) 

33 

34 self._set_rms_rotary_defaults() 

35 

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

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

38 self.cfg.rmsnorm_uses_offset = True 

39 

40 # Gemma2 uses logit softcapping 

41 if hasattr(self.cfg, "final_logit_softcapping"): 41 ↛ 42line 41 didn't jump to line 42 because the condition on line 41 was never true

42 self.cfg.output_logits_soft_cap = self.cfg.final_logit_softcapping 

43 if hasattr(self.cfg, "attn_logit_softcapping"): 43 ↛ 44line 43 didn't jump to line 44 because the condition on line 43 was never true

44 self.cfg.attn_scores_soft_cap = self.cfg.attn_logit_softcapping 

45 

46 # Note: n_key_value_heads is now automatically mapped from num_key_value_heads 

47 # by map_default_transformer_lens_config() in sources/transformers.py 

48 

49 self.weight_processing_conversions = { 

50 # NOTE: Gemma2 scales embeddings by sqrt(d_model) at RUNTIME inside 

51 # Gemma2TextScaledWordEmbedding.forward() (HF transformers >= 5.0). 

52 # That layer is what bridge.embed wraps, so embed.hook_out already 

53 # captures the scaled value — matching HookedTransformer's hook_embed 

54 # (which uses pre-scaled W_E). We must NOT pre-scale weights here and 

55 # we must NOT install a runtime hook_conversion that re-scales. 

56 **self._qkvo_weight_conversions(), 

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

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

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

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

61 ), 

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

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

64 ), 

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

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

67 ), 

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

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

70 ), 

71 "ln_final.weight": ParamProcessingConversion( 

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

73 ), 

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

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

76 tensor_conversion=TransposeTensorConversion(), 

77 ), 

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

79 tensor_conversion=TransposeTensorConversion(), 

80 ), 

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

82 tensor_conversion=TransposeTensorConversion(), 

83 ), 

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

85 "unembed.weight": ParamProcessingConversion( 

86 tensor_conversion=TransposeTensorConversion(), 

87 ), 

88 } 

89 

90 self.component_mapping = { 

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

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

93 "blocks": BlockBridge( 

94 name="model.layers", 

95 config=self.cfg, 

96 submodules={ 

97 # Gemma 2 uses RMSNorm for all normalization layers 

98 **self._block_norms(), 

99 # Gemma 2 uses PositionEmbeddingsAttentionBridge like Gemma 3 

100 "attn": PositionEmbeddingsAttentionBridge( 

101 name="self_attn", 

102 config=self.cfg, 

103 submodules={ 

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

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

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

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

108 }, 

109 requires_attention_mask=True, 

110 requires_position_embeddings=True, 

111 ), 

112 "mlp": self._gated_mlp(), 

113 }, 

114 ), 

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

116 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg), 

117 } 

118 

119 def _block_norms(self): 

120 """Norm-layout seam; VaultGemma drops the two post-norms.""" 

121 return { 

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

123 "ln1_post": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg), 

124 "ln2": RMSNormalizationBridge(name="pre_feedforward_layernorm", config=self.cfg), 

125 "ln2_post": RMSNormalizationBridge(name="post_feedforward_layernorm", config=self.cfg), 

126 }