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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +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)
35class Gemma3MultimodalArchitectureAdapter(ArchitectureAdapter):
36 """Architecture adapter for Gemma3 multimodal models (Gemma3ForConditionalGeneration).
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
45 The language model component follows the same patterns as Gemma3ArchitectureAdapter.
46 """
48 _testing_lm_attr = "model.language_model"
50 def __init__(self, cfg: Any) -> None:
51 """Initialize the Gemma3 multimodal architecture adapter."""
52 super().__init__(cfg)
54 self.cfg.is_multimodal = True
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"
66 # Store vision-related config
67 self._extract_vision_dims(cfg)
69 # Store multimodal projection config
70 self.cfg.mm_tokens_per_image = getattr(cfg, "mm_tokens_per_image", 256)
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 }
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 }
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