candle_transformers/models/gemma4/
mod.rs1pub 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
22pub 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 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 #[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 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 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 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 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 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 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
178fn 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 let mask_f32 = mask.to_dtype(DType::F32)?;
187 let zeros = Tensor::zeros((b_sz, seq_len, hidden), embeds.dtype(), embeds.device())?;
191
192 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 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}