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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
1"""Gemma3 Multimodal architecture adapter.
3This adapter supports Gemma3ForConditionalGeneration, the vision-language
4variant of Gemma 3 used by models like MedGemma.
5"""
7from typing import Any
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)
38class Gemma3MultimodalArchitectureAdapter(ArchitectureAdapter):
39 """Architecture adapter for Gemma3 multimodal models (Gemma3ForConditionalGeneration).
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
48 The language model component follows the same patterns as Gemma3ArchitectureAdapter.
49 """
51 _testing_lm_attr = "model.language_model"
53 def __init__(self, cfg: Any) -> None:
54 """Initialize the Gemma3 multimodal architecture adapter."""
55 super().__init__(cfg)
57 self.cfg.is_multimodal = True
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"
69 # Store vision-related config
70 self._extract_vision_dims(cfg)
72 # Store multimodal projection config
73 self.cfg.mm_tokens_per_image = getattr(cfg, "mm_tokens_per_image", 256)
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 }
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 }
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)