1use 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
28struct 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
48struct SiglipEncoder<B: Backend> {
50 cfg: VisionConfig,
51 patch_w: Tensor<B, 2>,
53 patch_b: Tensor<B, 1>,
54 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
62fn 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 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 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 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 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 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 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 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
203struct Connector<B: Backend> {
206 scale: usize,
207 proj: Tensor<B, 2>, }
209
210impl<B: Backend> Connector<B> {
211 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) .reshape([b, side / s, side / s, c * s * s])
223 .swap_dims(1, 2) .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
232pub 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 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 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
392pub 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#[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
416pub 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 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 let expected: Vec<f32> = vec![
450 0.0, 1.0, 4.0, 5.0, 2.0, 3.0, 6.0, 7.0, 8.0, 9.0, 12.0, 13.0, 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 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}