#[derive(Debug, Clone)]
pub struct PromptLayout {
pub image_token: String,
pub fake_token: String,
pub global_token: String,
pub row_col_fmt: fn(usize, usize) -> String,
pub tokens_per_tile: usize,
}
impl Default for PromptLayout {
fn default() -> Self {
Self {
image_token: "<image>".into(),
fake_token: "<fake_token_around_image>".into(),
global_token: "<global-img>".into(),
row_col_fmt: |r, c| format!("<row_{r}_col_{c}>"),
tokens_per_tile: 64,
}
}
}
impl PromptLayout {
#[must_use]
pub const fn with_geometry(
mut self,
image_size: usize,
patch_size: usize,
scale_factor: usize,
) -> Self {
let side = image_size / patch_size;
self.tokens_per_tile = (side * side) / (scale_factor * scale_factor);
self
}
#[must_use]
pub fn image_block(&self, rows: usize, cols: usize) -> String {
let imgs = self.image_token.repeat(self.tokens_per_tile);
let mut s = String::new();
for r in 1..=rows {
for c in 1..=cols {
s.push_str(&self.fake_token);
s.push_str(&(self.row_col_fmt)(r, c));
s.push_str(&imgs);
}
s.push('\n');
}
if rows > 0 {
s.push('\n');
}
s.push_str(&self.fake_token);
s.push_str(&self.global_token);
s.push_str(&imgs);
s.push_str(&self.fake_token);
s
}
#[must_use]
pub fn user_turn(&self, question: &str, rows: usize, cols: usize) -> String {
format!(
"<|im_start|>User:{}{question}<end_of_utterance>\nAssistant:",
self.image_block(rows, cols)
)
}
}
#[must_use]
pub const fn expected_image_tokens(layout: &PromptLayout, rows: usize, cols: usize) -> usize {
(rows * cols + 1) * layout.tokens_per_tile
}
#[must_use]
pub const fn expected_fake_tokens(rows: usize, cols: usize) -> usize {
rows * cols + 2
}
pub fn merge_image_embeddings(
text_embeds: &candle_core::Tensor,
image_hidden: &candle_core::Tensor,
input_ids: &[i64],
image_token_id: i64,
) -> candle_core::Result<candle_core::Tensor> {
use candle_core::{IndexOp, Tensor};
let (batch, seq, dim) = text_embeds.dims3()?;
if batch != 1 {
candle_core::bail!("merge expects batch 1, got {batch}");
}
if seq != input_ids.len() {
candle_core::bail!(
"input_ids has {} tokens but text_embeds has {seq} positions",
input_ids.len()
);
}
let img = image_hidden.flatten_to(image_hidden.rank() - 2)?;
let (n_img, img_dim) = img.dims2()?;
if img_dim != dim {
candle_core::bail!("image vectors are {img_dim}-dim but text embeds are {dim}-dim");
}
let positions: Vec<usize> = input_ids
.iter()
.enumerate()
.filter(|&(_, &t)| t == image_token_id)
.map(|(i, _)| i)
.collect();
if positions.len() != n_img {
candle_core::bail!(
"{} image positions in the prompt but {n_img} image vectors supplied — \
a mismatch here would misalign every block that follows, so it is an \
error rather than a truncation",
positions.len()
);
}
let text = text_embeds.i(0)?;
let mut rows: Vec<Tensor> = Vec::with_capacity(seq);
let mut next = 0usize;
for (i, &tok) in input_ids.iter().enumerate() {
if tok == image_token_id {
rows.push(img.i(next)?);
next += 1;
} else {
rows.push(text.i(i)?);
}
}
debug_assert_eq!(next, n_img, "every image vector must be consumed");
Tensor::stack(&rows, 0)?.unsqueeze(0)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn geometry_gives_smolvlm_64_tokens_per_tile() {
let l = PromptLayout::default().with_geometry(512, 16, 4);
assert_eq!(l.tokens_per_tile, 64);
}
#[test]
fn the_block_has_the_counts_the_reference_reports() {
let l = PromptLayout::default();
let s = l.image_block(4, 4);
assert_eq!(
s.matches("<image>").count(),
expected_image_tokens(&l, 4, 4),
"17 blocks of 64"
);
assert_eq!(s.matches("<image>").count(), 1088);
assert_eq!(
s.matches("<fake_token_around_image>").count(),
expected_fake_tokens(4, 4),
"16 tiles + global + closing"
);
assert_eq!(s.matches("<fake_token_around_image>").count(), 18);
}
#[test]
fn the_global_thumbnail_comes_after_every_tile() {
let s = PromptLayout::default().image_block(4, 4);
let global = s.find("<global-img>").expect("global marker");
let last_tile = s.find("<row_4_col_4>").expect("last tile marker");
assert!(
global > last_tile,
"the thumbnail must follow the grid — reversing it produces fluent, \
degraded output rather than an error"
);
assert!(s[last_tile..global].contains("\n\n"));
}
#[test]
fn a_thumbnail_only_image_has_no_grid_and_no_separator() {
let s = PromptLayout::default().image_block(0, 0);
assert!(!s.contains("<row_"));
assert!(!s.contains("\n\n"));
assert_eq!(s.matches("<fake_token_around_image>").count(), 2);
}
#[test]
fn the_user_turn_carries_the_chat_markers() {
let t = PromptLayout::default().user_turn("What is written in this image?", 4, 4);
assert!(t.starts_with("<|im_start|>User:"));
assert!(t.ends_with("<end_of_utterance>\nAssistant:"));
}
}