Coverage for transformer_lens/model_bridge/supported_architectures/gemma3_multimodal.py: 92%

26 statements  

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

1"""Gemma3 Multimodal architecture adapter. 

2 

3This adapter supports Gemma3ForConditionalGeneration, the vision-language 

4variant of Gemma 3 used by models like MedGemma. 

5""" 

6 

7from typing import Any 

8 

9from transformer_lens.conversion_utils.conversion_steps import ( 

10 ArithmeticTensorConversion, 

11 TransposeTensorConversion, 

12) 

13from transformer_lens.conversion_utils.conversion_steps.arithmetic_tensor_conversion import ( 

14 OperationTypes, 

15) 

16from transformer_lens.conversion_utils.param_processing_conversion import ( 

17 ParamProcessingConversion, 

18) 

19from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

20from transformer_lens.model_bridge.generalized_components import ( 

21 BlockBridge, 

22 EmbeddingBridge, 

23 LinearBridge, 

24 RMSNormalizationBridge, 

25 RotaryEmbeddingBridge, 

26 SiglipVisionEncoderBridge, 

27 UnembeddingBridge, 

28 VisionProjectionBridge, 

29) 

30from transformer_lens.model_bridge.generalized_components.position_embeddings_attention import ( 

31 PositionEmbeddingsAttentionBridge, 

32) 

33from transformer_lens.model_bridge.supported_architectures.gemma3 import ( 

34 _enable_native_qk_norm_autograd, 

35) 

36 

37 

38class Gemma3MultimodalArchitectureAdapter(ArchitectureAdapter): 

39 """Architecture adapter for Gemma3 multimodal models (Gemma3ForConditionalGeneration). 

40 

41 This adapter handles vision-language models like Gemma 3 4B/12B/27B and MedGemma. 

42 The model structure is: 

43 - model.vision_tower: SigLIP vision encoder 

44 - model.multi_modal_projector: Projects vision embeddings to language space 

45 - model.language_model: Gemma3TextModel (same as text-only Gemma 3) 

46 - lm_head: Output projection 

47 

48 The language model component follows the same patterns as Gemma3ArchitectureAdapter. 

49 """ 

50 

51 _testing_lm_attr = "model.language_model" 

52 

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

54 """Initialize the Gemma3 multimodal architecture adapter.""" 

55 super().__init__(cfg) 

56 

57 self.cfg.is_multimodal = True 

58 

59 # Language model configuration (same as text-only Gemma 3) 

60 self.cfg.gated_mlp = True 

61 self.cfg.uses_rms_norm = True 

62 self.cfg.normalization_type = "RMS" 

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

64 # Without this, fold_ln sets identity to 1.0 instead of 0.0, causing 2x scaling. 

65 self.cfg.rmsnorm_uses_offset = True 

66 self.cfg.positional_embedding_type = "rotary" 

67 self.cfg.attn_implementation = "eager" 

68 

69 # Store vision-related config 

70 self._extract_vision_dims(cfg) 

71 

72 # Store multimodal projection config 

73 self.cfg.mm_tokens_per_image = getattr(cfg, "mm_tokens_per_image", 256) 

74 

75 # Weight processing conversions for the language model 

76 # Note: The language model weights are under "model.language_model.*" 

77 self.weight_processing_conversions = { 

78 # Q/K/V weight conversions for language model 

79 **self._qkvo_weight_conversions(), 

80 # RMSNorm weight conversions - Gemma adds 1.0 to weights 

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

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

83 ), 

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

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

86 ), 

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

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

89 ), 

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

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

92 ), 

93 "ln_final.weight": ParamProcessingConversion( 

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

95 ), 

96 # Gemma-3 q_norm and k_norm in attention 

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

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

99 ), 

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

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

102 ), 

103 # MLP weight conversions 

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

105 tensor_conversion=TransposeTensorConversion(), 

106 ), 

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

108 tensor_conversion=TransposeTensorConversion(), 

109 ), 

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

111 tensor_conversion=TransposeTensorConversion(), 

112 ), 

113 # Unembed weight conversion 

114 "unembed.weight": ParamProcessingConversion( 

115 tensor_conversion=TransposeTensorConversion(), 

116 ), 

117 } 

118 

119 # Component mapping for the full multimodal model 

120 # Note: We use distinct TL names (vision_encoder, vision_projector) to avoid 

121 # conflicting with HF model attribute names (vision_tower, multi_modal_projector) 

122 self.component_mapping = { 

123 # Vision components 

124 "vision_encoder": SiglipVisionEncoderBridge(name="model.vision_tower", config=self.cfg), 

125 "vision_projector": VisionProjectionBridge(name="model.multi_modal_projector"), 

126 # Language model components (under model.language_model) 

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

128 "rotary_emb": RotaryEmbeddingBridge(name="model.language_model.rotary_emb"), 

129 "blocks": BlockBridge( 

130 name="model.language_model.layers", 

131 submodules={ 

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

133 "ln1_post": RMSNormalizationBridge( 

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

135 ), 

136 "ln2": RMSNormalizationBridge( 

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

138 ), 

139 "ln2_post": RMSNormalizationBridge( 

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

141 ), 

142 "attn": PositionEmbeddingsAttentionBridge( 

143 name="self_attn", 

144 config=self.cfg, 

145 submodules={ 

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

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

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

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

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

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

152 }, 

153 ), 

154 "mlp": self._gated_mlp(), 

155 }, 

156 ), 

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

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

159 } 

160 

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

162 """Wire rotary + eager, then enable native autograd on the Q/K norms.""" 

163 super().setup_component_testing(hf_model, bridge_model) 

164 _enable_native_qk_norm_autograd(bridge_model)