Skip to main content

candle_transformers/models/gemma4/
mod.rs

1//! Gemma 4 multimodal model (text + vision + audio).
2//!
3//! See:
4//! - [Google Blog](https://blog.google/technology/developers/gemma-4/)
5
6pub mod audio;
7pub mod config;
8pub mod multimodal_embedding;
9pub mod text;
10pub mod vision;
11
12use candle::{DType, Result, Tensor, D};
13
14use config::Gemma4Config;
15use multimodal_embedding::MultimodalEmbedder;
16use text::TextModel;
17use vision::VisionTower;
18
19pub use audio::AudioModel;
20pub use config::{Gemma4AudioConfig, Gemma4TextConfig, Gemma4VisionConfig};
21
22/// Full Gemma4 multimodal model.
23pub struct Model {
24    pub language_model: TextModel,
25    pub vision_tower: VisionTower,
26    pub embed_vision: MultimodalEmbedder,
27    pub audio_tower: Option<AudioModel>,
28    pub embed_audio: Option<MultimodalEmbedder>,
29    pub cfg: Gemma4Config,
30}
31
32impl Model {
33    pub fn new(cfg: &Gemma4Config, vb: candle_nn::VarBuilder) -> Result<Self> {
34        let vb = vb.pp("model");
35
36        let vision_tower = VisionTower::new(&cfg.vision_config, vb.pp("vision_tower"))?;
37
38        let vis_hidden = cfg.vision_config.hidden_size;
39        let text_hidden = cfg.text_config.hidden_size;
40        let embed_vision = MultimodalEmbedder::new(
41            vis_hidden,
42            text_hidden,
43            cfg.vision_config.rms_norm_eps,
44            vb.pp("embed_vision"),
45        )?;
46
47        let (audio_tower, embed_audio) = if let Some(ref audio_cfg) = cfg.audio_config {
48            let tower = AudioModel::new(audio_cfg, vb.pp("audio_tower"))?;
49            let audio_hidden = audio_cfg.output_proj_dims.unwrap_or(audio_cfg.hidden_size);
50            let embed = MultimodalEmbedder::new(
51                audio_hidden,
52                text_hidden,
53                audio_cfg.rms_norm_eps,
54                vb.pp("embed_audio"),
55            )?;
56            (Some(tower), Some(embed))
57        } else {
58            (None, None)
59        };
60
61        let language_model = TextModel::new(&cfg.text_config, vb.pp("language_model"))?;
62
63        Ok(Self {
64            language_model,
65            vision_tower,
66            embed_vision,
67            audio_tower,
68            embed_audio,
69            cfg: cfg.clone(),
70        })
71    }
72
73    /// Text-only forward pass.
74    pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
75        self.language_model.forward(input_ids, seqlen_offset)
76    }
77
78    /// Forward with multimodal inputs.
79    ///
80    /// `pixel_values`: optional batch of images, each `(1, C, H, W)`.
81    /// `audio_mel`: optional `(batch, time, mel_bins)` mel spectrogram.
82    /// `audio_mel_mask`: optional `(batch, time)` mask (1.0 = padding).
83    #[allow(clippy::too_many_arguments)]
84    pub fn forward_multimodal(
85        &mut self,
86        input_ids: &Tensor,
87        pixel_values: Option<&[Tensor]>,
88        audio_mel: Option<&Tensor>,
89        audio_mel_mask: Option<&Tensor>,
90        seqlen_offset: usize,
91    ) -> Result<Tensor> {
92        let (b_size, seq_len) = input_ids.dims2()?;
93        let mut input_embeds = self.language_model.embed_tokens(input_ids)?;
94
95        // ── Vision embedding injection ──────────────────────────────────
96        if let Some(pixel_values) = pixel_values {
97            let image_mask = input_ids
98                .to_dtype(DType::F32)?
99                .eq(self.cfg.image_token_id as f64)?;
100
101            let vision_features = self.vision_tower.forward(pixel_values)?;
102            let image_embeds = self
103                .embed_vision
104                .forward(&vision_features)?
105                .to_dtype(input_embeds.dtype())?;
106
107            // Replace image token positions with vision embeddings
108            let image_embeds_flat = image_embeds.squeeze(0)?;
109            let mask_expanded = image_mask
110                .unsqueeze(D::Minus1)?
111                .broadcast_as(input_embeds.shape())?
112                .to_dtype(input_embeds.dtype())?;
113            let image_embeds_broadcast = broadcast_embed_to_mask(&image_embeds_flat, &image_mask)?;
114            input_embeds = ((mask_expanded.clone() * image_embeds_broadcast)?
115                + ((1.0 - mask_expanded)? * input_embeds)?)?;
116        }
117
118        // ── Audio embedding injection ───────────────────────────────────
119        if let (
120            Some(audio_mel),
121            Some(audio_mel_mask),
122            Some(ref audio_tower),
123            Some(ref embed_audio),
124        ) = (
125            audio_mel,
126            audio_mel_mask,
127            &self.audio_tower,
128            &self.embed_audio,
129        ) {
130            let audio_mask = input_ids
131                .to_dtype(DType::F32)?
132                .eq(self.cfg.audio_token_id as f64)?;
133
134            let (audio_features, enc_mask) = audio_tower.forward(audio_mel, audio_mel_mask)?;
135            // Filter valid frames: where enc_mask == 0
136            let valid = enc_mask.eq(0.0)?;
137            let batch = audio_features.dim(0)?;
138            let mut all_feats = Vec::new();
139            for b in 0..batch {
140                let valid_b = valid.get(b)?;
141                // Count valid frames
142                let valid_sum = valid_b
143                    .to_dtype(DType::F32)?
144                    .sum_all()?
145                    .to_scalar::<f32>()? as usize;
146                if valid_sum > 0 {
147                    // Take the first valid_sum frames (they are contiguous after masking)
148                    all_feats.push(audio_features.get(b)?.narrow(0, 0, valid_sum)?);
149                }
150            }
151            if !all_feats.is_empty() {
152                let audio_feats = Tensor::cat(&all_feats, 0)?.unsqueeze(0)?;
153                let audio_embeds = embed_audio
154                    .forward(&audio_feats)?
155                    .to_dtype(input_embeds.dtype())?;
156
157                let audio_embeds_flat = audio_embeds.squeeze(0)?;
158                let mask_expanded = audio_mask
159                    .unsqueeze(D::Minus1)?
160                    .broadcast_as(input_embeds.shape())?
161                    .to_dtype(input_embeds.dtype())?;
162                let audio_embeds_broadcast =
163                    broadcast_embed_to_mask(&audio_embeds_flat, &audio_mask)?;
164                input_embeds = ((mask_expanded.clone() * audio_embeds_broadcast)?
165                    + ((1.0 - mask_expanded)? * input_embeds)?)?;
166            }
167        }
168
169        self.language_model
170            .forward_embeds(&input_embeds, seqlen_offset, b_size, seq_len)
171    }
172
173    pub fn clear_kv_cache(&mut self) {
174        self.language_model.clear_kv_cache()
175    }
176}
177
178/// Broadcast encoder embeddings (num_tokens, hidden) into positions marked by
179/// a boolean mask (batch, seq_len), producing (batch, seq_len, hidden).
180/// Token embeddings are placed sequentially where the mask is true.
181fn broadcast_embed_to_mask(embeds: &Tensor, mask: &Tensor) -> Result<Tensor> {
182    let (b_sz, seq_len) = mask.dims2()?;
183    let hidden = embeds.dim(D::Minus1)?;
184
185    // Count masked positions per batch, fill them in sequence from embeds
186    let mask_f32 = mask.to_dtype(DType::F32)?;
187    // cumsum along seq dimension to assign embed indices
188    // Since candle doesn't have cumsum, we use a broadcast approach:
189    // Create output tensor of zeros, then use where_cond
190    let zeros = Tensor::zeros((b_sz, seq_len, hidden), embeds.dtype(), embeds.device())?;
191
192    // For single-batch simple case, just expand embeds to the output shape
193    // and let the caller do the masking.
194    if b_sz == 1 {
195        let num_tokens = mask_f32.sum_all()?.to_scalar::<f32>()? as usize;
196        if num_tokens == 0 {
197            return Ok(zeros);
198        }
199        // Pad or truncate embeds to seq_len
200        let embed_len = embeds.dim(0)?;
201        if embed_len >= seq_len {
202            return embeds.narrow(0, 0, seq_len)?.unsqueeze(0);
203        }
204        let padding = Tensor::zeros(
205            (seq_len - embed_len, hidden),
206            embeds.dtype(),
207            embeds.device(),
208        )?;
209        let padded = Tensor::cat(&[embeds, &padding], 0)?;
210        return padded.unsqueeze(0);
211    }
212
213    Ok(zeros)
214}