Skip to main content

combs_models/
smolvlm.rs

1//! Idefics3 / SmolVLM architecture: SigLIP vision encoder + pixel-shuffle
2//! connector + Llama-family text decoder (SmolLM2), on the same
3//! [`GenerativeModel`] contract. The text stack is the shared Llama trunk
4//! (weights under `model.text_model.*`); the vision tower is stateless and
5//! runs once per image inside `embed_multimodal`, whose output replaces the
6//! `<image>` token spans in the embedded prompt.
7//!
8//! Weight names (HF safetensors):
9//! `model.vision_model.embeddings.{patch_embedding.weight,bias}`,
10//! `model.vision_model.embeddings.position_embedding.weight`,
11//! `model.vision_model.encoder.layers.{i}.{layer_norm1,self_attn,layer_norm2,mlp}.*`,
12//! `model.vision_model.post_layernorm.{weight,bias}`,
13//! `model.connector.modality_projection.proj.weight`,
14//! `model.text_model.*` (Llama layout), `lm_head.weight`.
15
16use std::ops::Range;
17
18use burn::tensor::{Device, Int, Tensor, TensorData, activation::softmax, backend::Backend};
19use combs_formats::{ModelMetadata, ModelSource, VisionConfig};
20
21use crate::kv::{CacheConfig, KVCache};
22use crate::llama::{LlamaModel, linear, load_tensor};
23use crate::matmul::safe_matmul;
24use crate::norm::layer_norm;
25use crate::traits::GenerativeModel;
26use crate::{ModelError, Result};
27
28/// One SigLIP encoder layer's weights (all projections carry biases).
29struct SiglipLayer<B: Backend> {
30    ln1_w: Tensor<B, 1>,
31    ln1_b: Tensor<B, 1>,
32    q_w: Tensor<B, 2>,
33    q_b: Tensor<B, 1>,
34    k_w: Tensor<B, 2>,
35    k_b: Tensor<B, 1>,
36    v_w: Tensor<B, 2>,
37    v_b: Tensor<B, 1>,
38    o_w: Tensor<B, 2>,
39    o_b: Tensor<B, 1>,
40    ln2_w: Tensor<B, 1>,
41    ln2_b: Tensor<B, 1>,
42    fc1_w: Tensor<B, 2>,
43    fc1_b: Tensor<B, 1>,
44    fc2_w: Tensor<B, 2>,
45    fc2_b: Tensor<B, 1>,
46}
47
48/// SigLIP vision transformer (fixed square input, full self-attention).
49struct SiglipEncoder<B: Backend> {
50    cfg: VisionConfig,
51    /// Patch-embed conv weight flattened to `[hidden, channels*patch²]`.
52    patch_w: Tensor<B, 2>,
53    patch_b: Tensor<B, 1>,
54    /// Learned absolute position embeddings `[num_patches, hidden]`.
55    pos_embed: Tensor<B, 2>,
56    layers: Vec<SiglipLayer<B>>,
57    post_ln_w: Tensor<B, 1>,
58    post_ln_b: Tensor<B, 1>,
59    scale: f64,
60}
61
62/// GELU (tanh approximation), SigLIP's `gelu_pytorch_tanh`.
63fn gelu_tanh<B: Backend, const D: usize>(x: Tensor<B, D>) -> Tensor<B, D> {
64    let inner = (x.clone() + x.clone().powf_scalar(3.0).mul_scalar(0.044715))
65        .mul_scalar((2.0f64 / std::f64::consts::PI).sqrt());
66    // 0.5 * x * (1 + tanh(inner))
67    x * inner.tanh().add_scalar(1.0).mul_scalar(0.5)
68}
69
70impl<B: Backend> SiglipEncoder<B> {
71    fn load(source: &dyn ModelSource, device: &Device<B>, cfg: &VisionConfig) -> Result<Self> {
72        let p = "model.vision_model";
73        // Conv weight [hidden, channels, patch, patch] -> [hidden, channels*patch²].
74        let conv: Tensor<B, 4> = load_tensor(
75            source,
76            device,
77            &format!("{p}.embeddings.patch_embedding.weight"),
78        )?;
79        let patch_w = conv.reshape([
80            cfg.hidden_size,
81            3 * cfg.patch_size * cfg.patch_size,
82        ]);
83        let patch_b = load_tensor(source, device, &format!("{p}.embeddings.patch_embedding.bias"))?;
84        let pos_embed: Tensor<B, 2> = load_tensor(
85            source,
86            device,
87            &format!("{p}.embeddings.position_embedding.weight"),
88        )?;
89
90        let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
91        for i in 0..cfg.num_hidden_layers {
92            let lp = format!("{p}.encoder.layers.{i}");
93            layers.push(SiglipLayer {
94                ln1_w: load_tensor(source, device, &format!("{lp}.layer_norm1.weight"))?,
95                ln1_b: load_tensor(source, device, &format!("{lp}.layer_norm1.bias"))?,
96                q_w: load_tensor(source, device, &format!("{lp}.self_attn.q_proj.weight"))?,
97                q_b: load_tensor(source, device, &format!("{lp}.self_attn.q_proj.bias"))?,
98                k_w: load_tensor(source, device, &format!("{lp}.self_attn.k_proj.weight"))?,
99                k_b: load_tensor(source, device, &format!("{lp}.self_attn.k_proj.bias"))?,
100                v_w: load_tensor(source, device, &format!("{lp}.self_attn.v_proj.weight"))?,
101                v_b: load_tensor(source, device, &format!("{lp}.self_attn.v_proj.bias"))?,
102                o_w: load_tensor(source, device, &format!("{lp}.self_attn.out_proj.weight"))?,
103                o_b: load_tensor(source, device, &format!("{lp}.self_attn.out_proj.bias"))?,
104                ln2_w: load_tensor(source, device, &format!("{lp}.layer_norm2.weight"))?,
105                ln2_b: load_tensor(source, device, &format!("{lp}.layer_norm2.bias"))?,
106                fc1_w: load_tensor(source, device, &format!("{lp}.mlp.fc1.weight"))?,
107                fc1_b: load_tensor(source, device, &format!("{lp}.mlp.fc1.bias"))?,
108                fc2_w: load_tensor(source, device, &format!("{lp}.mlp.fc2.weight"))?,
109                fc2_b: load_tensor(source, device, &format!("{lp}.mlp.fc2.bias"))?,
110            });
111        }
112
113        Ok(SiglipEncoder {
114            scale: 1.0 / (cfg.head_dim() as f64).sqrt(),
115            cfg: cfg.clone(),
116            patch_w,
117            patch_b,
118            pos_embed,
119            layers,
120            post_ln_w: load_tensor(source, device, &format!("{p}.post_layernorm.weight"))?,
121            post_ln_b: load_tensor(source, device, &format!("{p}.post_layernorm.bias"))?,
122        })
123    }
124
125    /// Full (non-causal) self-attention over the patch sequence.
126    fn attention(&self, layer: &SiglipLayer<B>, x: Tensor<B, 3>) -> Tensor<B, 3> {
127        let cfg = &self.cfg;
128        let [batch, seq, _] = x.dims();
129        let heads = cfg.num_attention_heads;
130        let head_dim = cfg.head_dim();
131
132        let q = linear(x.clone(), &layer.q_w, Some(&layer.q_b))
133            .reshape([batch, seq, heads, head_dim])
134            .swap_dims(1, 2);
135        let k = linear(x.clone(), &layer.k_w, Some(&layer.k_b))
136            .reshape([batch, seq, heads, head_dim])
137            .swap_dims(1, 2);
138        let v = linear(x, &layer.v_w, Some(&layer.v_b))
139            .reshape([batch, seq, heads, head_dim])
140            .swap_dims(1, 2);
141
142        // K dims hit the broken wgpu/Metal matmul region (>=512) — safe_matmul.
143        let scores = safe_matmul(q, k.transpose()).mul_scalar(self.scale);
144        let ctx = safe_matmul(softmax(scores, 3), v);
145        let ctx = ctx.swap_dims(1, 2).reshape([batch, seq, heads * head_dim]);
146        linear(ctx, &layer.o_w, Some(&layer.o_b))
147    }
148
149    /// `[1, channels, image, image] -> [1, num_patches, hidden]`.
150    fn forward(&self, pixels: Tensor<B, 4>) -> Tensor<B, 3> {
151        let cfg = &self.cfg;
152        let p = cfg.patch_size;
153        let [_, c, h, w] = pixels.dims();
154        let (gh, gw) = (h / p, w / p);
155        debug_assert_eq!(c, 3);
156        debug_assert_eq!(h % p + w % p, 0);
157
158        // Unfold into patches (equivalent to the stride-p conv): reshape to
159        // [c, gh, p, gw, p] -> [gh, gw, c, p, p] -> [gh*gw, c*p*p].
160        let patches = pixels
161            .reshape([c, gh, p, gw, p])
162            .swap_dims(0, 1)
163            .swap_dims(1, 3)
164            .swap_dims(2, 3)
165            .reshape([gh * gw, c * p * p])
166            .unsqueeze_dim::<3>(0);
167        let mut x = linear(patches, &self.patch_w, Some(&self.patch_b));
168
169        // Fixed square input: positional ids are exactly 0..num_patches.
170        let np = gh * gw;
171        x = x + self
172            .pos_embed
173            .clone()
174            .narrow(0, 0, np)
175            .reshape([1, np, cfg.hidden_size]);
176
177        for layer in &self.layers {
178            let h = layer_norm(
179                x.clone(),
180                layer.ln1_w.clone(),
181                layer.ln1_b.clone(),
182                cfg.layer_norm_eps,
183            );
184            x = x + self.attention(layer, h);
185            let h = layer_norm(
186                x.clone(),
187                layer.ln2_w.clone(),
188                layer.ln2_b.clone(),
189                cfg.layer_norm_eps,
190            );
191            let mlp = linear(
192                gelu_tanh(linear(h, &layer.fc1_w, Some(&layer.fc1_b))),
193                &layer.fc2_w,
194                Some(&layer.fc2_b),
195            );
196            x = x + mlp;
197        }
198
199        layer_norm(x, self.post_ln_w.clone(), self.post_ln_b.clone(), cfg.layer_norm_eps)
200    }
201}
202
203/// Idefics3 connector: pixel-shuffle (space-to-depth, HF ordering) followed
204/// by a bias-free linear projection into the text hidden size.
205struct Connector<B: Backend> {
206    scale: usize,
207    proj: Tensor<B, 2>, // [text_hidden, vision_hidden * scale²]
208}
209
210impl<B: Backend> Connector<B> {
211    /// `[1, patches, vision_hidden] -> [1, patches/s², vision_hidden*s²]`
212    /// (matches HF `Idefics3Connector.pixel_shuffle` channel ordering:
213    /// channel = ((row_in_group * s) + col_in_group) * C + c).
214    fn pixel_shuffle(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
215        let s = self.scale;
216        let [b, seq, c] = x.dims();
217        let side = (seq as f64).sqrt() as usize;
218        assert_eq!(side * side, seq, "patch grid must be square");
219        x.reshape([b, side, side, c])
220            .reshape([b, side, side / s, c * s])
221            .swap_dims(1, 2) // [b, side/s, side, c*s]
222            .reshape([b, side / s, side / s, c * s * s])
223            .swap_dims(1, 2) // [b, side/s, side/s, c*s²]
224            .reshape([b, seq / (s * s), c * s * s])
225    }
226
227    fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
228        linear(self.pixel_shuffle(x), &self.proj, None)
229    }
230}
231
232/// SmolVLM (Idefics3): SigLIP + connector + Llama-family text decoder.
233pub struct SmolVlmModel<B: Backend> {
234    metadata: ModelMetadata,
235    vision_cfg: VisionConfig,
236    vision: SiglipEncoder<B>,
237    connector: Connector<B>,
238    text: LlamaModel<B>,
239}
240
241impl<B: Backend> SmolVlmModel<B> {
242    /// Runs one image through the vision tower + connector:
243    /// `[1, 3, H, W] -> [1, image_seq_len, text_hidden]`.
244    fn image_features(&self, pixels: Tensor<B, 4>) -> Tensor<B, 3> {
245        self.connector.forward(self.vision.forward(pixels))
246    }
247}
248
249impl<B: Backend> GenerativeModel<B> for SmolVlmModel<B> {
250    fn metadata(&self) -> &ModelMetadata {
251        &self.metadata
252    }
253
254    fn load(source: &dyn ModelSource, device: &Device<B>) -> Result<Self> {
255        let metadata = source.metadata().clone();
256        let vision_cfg = metadata
257            .vision
258            .clone()
259            .ok_or_else(|| ModelError::MissingTensor("vision_config".to_string()))?;
260        let vision = SiglipEncoder::load(source, device, &vision_cfg)?;
261        let proj: Tensor<B, 2> = load_tensor(
262            source,
263            device,
264            "model.connector.modality_projection.proj.weight",
265        )?;
266        LlamaModel::<B>::expect_shape(
267            "model.connector.modality_projection.proj.weight",
268            &proj.dims(),
269            &[
270                metadata.hidden_size,
271                vision_cfg.hidden_size * vision_cfg.scale_factor * vision_cfg.scale_factor,
272            ],
273        )?;
274        let text = LlamaModel::<B>::load_with_prefix(source, device, "model.text_model")?;
275        Ok(SmolVlmModel {
276            connector: Connector {
277                scale: vision_cfg.scale_factor,
278                proj,
279            },
280            metadata,
281            vision_cfg,
282            vision,
283            text,
284        })
285    }
286
287    fn create_kv_cache(&self, config: &CacheConfig) -> Box<dyn KVCache<B>> {
288        self.text.create_kv_cache(config)
289    }
290
291    fn embed(&self, tokens: Tensor<B, 2, Int>) -> Tensor<B, 3> {
292        self.text.embed(tokens)
293    }
294
295    fn embed_multimodal(
296        &self,
297        tokens: Tensor<B, 2, Int>,
298        images: &[Tensor<B, 4>],
299    ) -> Result<Tensor<B, 3>> {
300        if images.is_empty() {
301            return Ok(self.text.embed(tokens));
302        }
303        let [_, seq] = tokens.dims();
304        let ids: Vec<i64> = tokens
305            .clone()
306            .into_data()
307            .convert::<i64>()
308            .to_vec()
309            .map_err(|e| ModelError::BadShape {
310                tensor: "tokens".to_string(),
311                expected: vec![1, seq],
312                got: vec![],
313            })
314            .unwrap_or_default();
315        if ids.len() != seq {
316            return Err(ModelError::BadShape {
317                tensor: "tokens".to_string(),
318                expected: vec![1, seq],
319                got: vec![ids.len()],
320            });
321        }
322
323        // Consecutive spans of the image token; one span per image, in order.
324        let image_id = self.vision_cfg.image_token_id as i64;
325        let span_len = self.vision_cfg.image_seq_len();
326        let mut spans: Vec<(usize, usize)> = Vec::new();
327        let mut i = 0;
328        while i < seq {
329            if ids[i] == image_id {
330                let start = i;
331                while i < seq && ids[i] == image_id {
332                    i += 1;
333                }
334                spans.push((start, i));
335            } else {
336                i += 1;
337            }
338        }
339        if spans.len() != images.len() {
340            return Err(ModelError::UnsupportedMedia(format!(
341                "found {} image-token span(s) of len {span_len} but {} image(s) were provided",
342                spans.len(),
343                images.len()
344            )));
345        }
346
347        let base = self.text.embed(tokens);
348        let mut pieces: Vec<Tensor<B, 3>> = Vec::new();
349        let mut cursor = 0;
350        for (idx, (start, end)) in spans.iter().enumerate() {
351            if end - start != span_len {
352                return Err(ModelError::UnsupportedMedia(format!(
353                    "image-token span has length {}, expected {span_len}",
354                    end - start
355                )));
356            }
357            if *start > cursor {
358                pieces.push(base.clone().narrow(1, cursor, start - cursor));
359            }
360            pieces.push(self.image_features(images[idx].clone()));
361            cursor = *end;
362        }
363        if cursor < seq {
364            pieces.push(base.narrow(1, cursor, seq - cursor));
365        }
366        Ok(Tensor::cat(pieces, 1))
367    }
368
369    fn prefill(
370        &mut self,
371        input: Tensor<B, 3>,
372        cache: &mut dyn KVCache<B>,
373        pos: Range<u32>,
374    ) -> Tensor<B, 2> {
375        let [_, seq, _] = input.dims();
376        assert_eq!(
377            seq,
378            (pos.end - pos.start) as usize,
379            "prefill pos range must match the input sequence length"
380        );
381        let hidden = self.text.forward_hidden(input, cache, pos.start as usize);
382        self.text.last_logits(hidden)
383    }
384
385    fn decode(&mut self, input: Tensor<B, 3>, cache: &mut dyn KVCache<B>) -> Tensor<B, 2> {
386        let pos = cache.seq_len();
387        let hidden = self.text.forward_hidden(input, cache, pos);
388        self.text.last_logits(hidden)
389    }
390}
391
392/// Builds the Idefics3 prompt expansion for one image:
393/// `<fake_token_around_image><global-img><image>×image_seq_len<fake_token_around_image>`.
394/// (`image_seq_len` = 64 for SmolVLM-256M.) The returned string is meant to
395/// replace each `<image>` placeholder in the chat text.
396pub fn image_prompt_expansion(image_seq_len: usize) -> String {
397    let mut s = String::with_capacity(image_seq_len * 8 + 64);
398    s.push_str("<fake_token_around_image><global-img>");
399    for _ in 0..image_seq_len {
400        s.push_str("<image>");
401    }
402    s.push_str("<fake_token_around_image>");
403    s
404}
405
406/// Reads token data back to host ids (helper for tests).
407#[allow(dead_code)]
408fn token_ids<B: Backend>(tokens: Tensor<B, 2, Int>) -> Vec<i64> {
409    tokens
410        .into_data()
411        .convert::<i64>()
412        .to_vec()
413        .unwrap_or_default()
414}
415
416/// Builds a `[1, 3, H, W]` pixel tensor from planar CHW f32 data (used by the
417/// runtime to hand media to `embed_multimodal`).
418pub fn pixels_to_tensor<B: Backend>(
419    data: Vec<f32>,
420    shape: [usize; 4],
421    device: &Device<B>,
422) -> Tensor<B, 4> {
423    Tensor::from_data(TensorData::new(data, shape), device)
424}
425
426#[cfg(test)]
427mod tests {
428    use super::*;
429    type TestBackend = burn::backend::NdArray<f32>;
430
431    #[test]
432    fn pixel_shuffle_matches_hf_ordering() {
433        // seq=16 (4x4 grid), s=2, C=1: channel groups must follow
434        // ((row_in_group * s) + col_in_group) * C + c ordering.
435        let device = Default::default();
436        let data: Vec<f32> = (0..16).map(|v| v as f32).collect();
437        let x = Tensor::<TestBackend, 3>::from_data(TensorData::new(data, [1, 16, 1]), &device);
438        let conn = Connector::<TestBackend> {
439            scale: 2,
440            proj: Tensor::eye(4, &device),
441        };
442        let out = conn.pixel_shuffle(x);
443        let got: Vec<f32> = out.into_data().to_vec().unwrap();
444        // Grid (row-major ids): 0 1 2 3 / 4 5 6 7 / 8 9 10 11 / 12 13 14 15
445        // Token 0 (top-left 2x2 group): rows 0-1, cols 0-1 → [0,1,4,5]
446        // Token 1: rows 0-1, cols 2-3 → [2,3,6,7]
447        // Token 2: rows 2-3, cols 0-1 → [8,9,12,13]
448        // Token 3: rows 2-3, cols 2-3 → [10,11,14,15]
449        let expected: Vec<f32> = vec![
450            0.0, 1.0, 4.0, 5.0, //
451            2.0, 3.0, 6.0, 7.0, //
452            8.0, 9.0, 12.0, 13.0, //
453            10.0, 11.0, 14.0, 15.0,
454        ];
455        assert_eq!(got, expected);
456    }
457
458    #[test]
459    fn image_prompt_expansion_shape() {
460        let s = image_prompt_expansion(64);
461        assert!(s.starts_with("<fake_token_around_image><global-img>"));
462        assert!(s.ends_with("<fake_token_around_image>"));
463        assert_eq!(s.matches("<image>").count(), 64);
464    }
465
466    #[test]
467    fn gelu_tanh_reference() {
468        let device = Default::default();
469        let x = Tensor::<TestBackend, 1>::from_data(TensorData::new(vec![0.0f32, 1.0, -1.0], [3]), &device);
470        let y: Vec<f32> = gelu_tanh(x).into_data().to_vec().unwrap();
471        // gelu(0)=0, gelu(1)≈0.8412, gelu(-1)≈-0.1588 (tanh approx).
472        assert!(y[0].abs() < 1e-5);
473        assert!((y[1] - 0.8412).abs() < 1e-3, "gelu(1) = {}", y[1]);
474        assert!((y[2] + 0.1588).abs() < 1e-3, "gelu(-1) = {}", y[2]);
475    }
476}