Coverage for transformer_lens/model_bridge/supported_architectures/gemma1.py: 100%

14 statements  

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

1"""Gemma1 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 Gemma1ArchitectureAdapter(ArchitectureAdapter): 

28 """Architecture adapter for Gemma1 models.""" 

29 

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

31 """Initialize the Gemma1 architecture adapter.""" 

32 super().__init__(cfg) 

33 

34 self._set_rms_rotary_defaults() 

35 

36 # Gemma models use BOS tokens (tokenizer prepends BOS by default) 

37 # Matches HookedTransformer behavior (default_prepend_bos = True) 

38 self.cfg.default_prepend_bos = True 

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

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

41 self.cfg.rmsnorm_uses_offset = True 

42 

43 self.weight_processing_conversions = { 

44 # NOTE: Gemma1 scales embeddings by sqrt(d_model) at RUNTIME inside 

45 # GemmaTextScaledWordEmbedding.forward() (HF transformers >= 5.0). 

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

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

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

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

50 # 

51 # Attention weight conversions 

52 **self._qkvo_weight_conversions(), 

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

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

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

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

57 ), 

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

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

60 ), 

61 "ln_final.weight": ParamProcessingConversion( 

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

63 ), 

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

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

66 tensor_conversion=TransposeTensorConversion(), 

67 ), 

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

69 tensor_conversion=TransposeTensorConversion(), 

70 ), 

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

72 tensor_conversion=TransposeTensorConversion(), 

73 ), 

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

75 "unembed.weight": ParamProcessingConversion( 

76 tensor_conversion=TransposeTensorConversion(), 

77 ), 

78 } 

79 

80 self.component_mapping = { 

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

82 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg), 

83 "blocks": BlockBridge( 

84 name="model.layers", 

85 submodules={ 

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

87 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg), 

88 "attn": PositionEmbeddingsAttentionBridge( 

89 name="self_attn", 

90 config=self.cfg, 

91 submodules={ 

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

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

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

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

96 }, 

97 requires_attention_mask=True, 

98 requires_position_embeddings=True, 

99 ), 

100 "mlp": self._gated_mlp(), 

101 }, 

102 ), 

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

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

105 }