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

33 statements  

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

33 

34 

35class Gemma3MultimodalArchitectureAdapter(ArchitectureAdapter): 

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

37 

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

39 The model structure is: 

40 - model.vision_tower: SigLIP vision encoder 

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

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

43 - lm_head: Output projection 

44 

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

46 """ 

47 

48 _testing_lm_attr = "model.language_model" 

49 

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

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

52 super().__init__(cfg) 

53 

54 self.cfg.is_multimodal = True 

55 

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

57 self.cfg.gated_mlp = True 

58 self.cfg.uses_rms_norm = True 

59 self.cfg.normalization_type = "RMS" 

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

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

62 self.cfg.rmsnorm_uses_offset = True 

63 self.cfg.positional_embedding_type = "rotary" 

64 self.cfg.attn_implementation = "eager" 

65 

66 # Store vision-related config 

67 self._extract_vision_dims(cfg) 

68 

69 # Store multimodal projection config 

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

71 

72 # Weight processing conversions for the language model 

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

74 self.weight_processing_conversions = { 

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

76 **self._qkvo_weight_conversions(), 

77 # RMSNorm weight conversions - Gemma adds 1.0 to weights 

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

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

80 ), 

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

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

83 ), 

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

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

86 ), 

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

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

89 ), 

90 "ln_final.weight": ParamProcessingConversion( 

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

92 ), 

93 # Gemma-3 q_norm and k_norm in attention 

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

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

96 ), 

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

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

99 ), 

100 # MLP weight conversions 

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

102 tensor_conversion=TransposeTensorConversion(), 

103 ), 

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

105 tensor_conversion=TransposeTensorConversion(), 

106 ), 

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

108 tensor_conversion=TransposeTensorConversion(), 

109 ), 

110 # Unembed weight conversion 

111 "unembed.weight": ParamProcessingConversion( 

112 tensor_conversion=TransposeTensorConversion(), 

113 ), 

114 } 

115 

116 # Component mapping for the full multimodal model 

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

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

119 self.component_mapping = { 

120 # Vision components 

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

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

123 # Language model components (under model.language_model) 

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

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

126 "blocks": BlockBridge( 

127 name="model.language_model.layers", 

128 submodules={ 

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

130 "ln1_post": RMSNormalizationBridge( 

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

132 ), 

133 "ln2": RMSNormalizationBridge( 

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

135 ), 

136 "ln2_post": RMSNormalizationBridge( 

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

138 ), 

139 "attn": PositionEmbeddingsAttentionBridge( 

140 name="self_attn", 

141 config=self.cfg, 

142 submodules={ 

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

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

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

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

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

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

149 }, 

150 ), 

151 "mlp": self._gated_mlp(), 

152 }, 

153 ), 

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

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

156 } 

157 

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

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

160 super().setup_component_testing(hf_model, bridge_model) 

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

162 for block in bridge_model.blocks: 

163 hf_attn = getattr(getattr(block, "attn", None), "original_component", None) 

164 if hf_attn is None: 

165 continue 

166 if hasattr(hf_attn, "q_norm"): 

167 hf_attn.q_norm.use_native_layernorm_autograd = True 

168 if hasattr(hf_attn, "k_norm"): 

169 hf_attn.k_norm.use_native_layernorm_autograd = True