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

20 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-21 19:27 +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 self.cfg.uses_rms_norm = True 

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

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

39 self.cfg.rmsnorm_uses_offset = True 

40 

41 # Gemma2 uses logit softcapping 

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

43 self.cfg.output_logits_soft_cap = self.cfg.final_logit_softcapping 

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

45 self.cfg.attn_scores_soft_cap = self.cfg.attn_logit_softcapping 

46 

47 # Note: n_key_value_heads is now automatically mapped from num_key_value_heads 

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

49 

50 self.weight_processing_conversions = { 

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

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

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

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

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

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

57 **self._qkvo_weight_conversions(), 

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

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

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

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

62 ), 

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

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

65 ), 

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

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

68 ), 

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

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

71 ), 

72 "ln_final.weight": ParamProcessingConversion( 

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

74 ), 

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

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

77 tensor_conversion=TransposeTensorConversion(), 

78 ), 

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

80 tensor_conversion=TransposeTensorConversion(), 

81 ), 

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

83 tensor_conversion=TransposeTensorConversion(), 

84 ), 

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

86 "unembed.weight": ParamProcessingConversion( 

87 tensor_conversion=TransposeTensorConversion(), 

88 ), 

89 } 

90 

91 self.component_mapping = { 

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

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

94 "blocks": BlockBridge( 

95 name="model.layers", 

96 config=self.cfg, 

97 submodules={ 

98 # Gemma 2 uses RMSNorm for all normalization layers 

99 **self._block_norms(), 

100 # Gemma 2 uses PositionEmbeddingsAttentionBridge like Gemma 3 

101 "attn": PositionEmbeddingsAttentionBridge( 

102 name="self_attn", 

103 config=self.cfg, 

104 submodules={ 

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

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

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

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

109 }, 

110 requires_attention_mask=True, 

111 requires_position_embeddings=True, 

112 ), 

113 "mlp": self._gated_mlp(), 

114 }, 

115 ), 

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

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

118 } 

119 

120 def _block_norms(self): 

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

122 return { 

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

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

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

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

127 }