use std::path::{Path, PathBuf};
use std::sync::{Mutex, OnceLock};
use candle_core::{Device, Tensor};
use ffai_core::engine::{
EngineInfo, EngineStatus, Task, VlmEngine, VlmOptions, VlmPart, VlmPrompt,
};
use ffai_core::error::{Error, Result};
use ffai_core::types::{ImageBuffer, PixelFormat, TimedSegment, VideoFrame};
use crate::decode::TextDecoder;
use crate::preprocess::preprocess_rgb8_opts;
use crate::prompt::{merge_image_embeddings, PromptLayout};
use crate::vision::SmolVlmVision;
pub const MODEL: &str = "smolvlm-256m-instruct";
const DEFAULT_MAX_NEW_TOKENS: usize = 256;
const DEFAULT_FRAMES_PER_WINDOW: usize = 8;
#[derive(Debug, Clone, Default, PartialEq)]
pub struct CaptionTrace {
pub preprocess_ms: f64,
pub tower_ms: f64,
pub tower_per_tile_ms: Vec<f64>,
pub assemble_ms: f64,
pub prefill_ms: f64,
pub step_ms: Vec<f64>,
pub detokenize_ms: f64,
pub tiles: usize,
pub rows: usize,
pub cols: usize,
pub tile: usize,
pub image_tokens: usize,
pub text_tokens: usize,
pub prompt_tokens: usize,
pub max_positions: usize,
pub split: bool,
pub resized_to: Vec<(usize, usize)>,
}
impl CaptionTrace {
#[must_use]
pub fn decode_ms(&self) -> f64 {
self.step_ms.iter().sum()
}
#[must_use]
pub fn total_ms(&self) -> f64 {
self.preprocess_ms
+ self.tower_ms
+ self.assemble_ms
+ self.prefill_ms
+ self.decode_ms()
+ self.detokenize_ms
}
#[must_use]
pub fn tokens_per_sec(&self) -> f64 {
let ms = self.decode_ms();
if ms <= 0.0 {
return 0.0;
}
self.step_ms.len() as f64 / (ms / 1e3)
}
}
pub struct SmolVlm {
manifest_dir: PathBuf,
model: OnceLock<std::result::Result<Model, String>>,
}
struct Model {
vision: SmolVlmVision,
decoder: Mutex<TextDecoder>,
tokenizer: tokenizers::Tokenizer,
layout: PromptLayout,
image_token_id: i64,
stop_ids: Vec<u32>,
max_positions: usize,
device: Device,
}
impl SmolVlm {
#[must_use]
pub fn new() -> Self {
Self::with_manifest_dir(PathBuf::from("models"))
}
#[must_use]
pub const fn with_manifest_dir(dir: PathBuf) -> Self {
Self {
manifest_dir: dir,
model: OnceLock::new(),
}
}
pub fn from_bytes(w: ArgusBytes) -> Result<Self> {
let tokenizer = tokenizers::Tokenizer::from_bytes(&w.tokenizer)
.map_err(|e| Error::Model(format!("tokenizer: {e}")))?;
let device = Device::Cpu;
let vb = candle_nn::VarBuilder::from_buffered_safetensors(
w.weights,
candle_core::DType::F32,
&device,
)
.map_err(|e| Error::Model(format!("reading safetensors: {e}")))?;
let model = build(vb, &w.config, tokenizer, device)?;
let cell = OnceLock::new();
let _ = cell.set(Ok(model));
Ok(Self {
manifest_dir: PathBuf::new(),
model: cell,
})
}
#[must_use]
pub fn is_loaded(&self) -> bool {
self.model.get().is_some()
}
pub fn warm(&self) -> Result<()> {
self.model().map(|_| ())
}
fn model(&self) -> Result<&Model> {
match self
.model
.get_or_init(|| load(&self.manifest_dir).map_err(|e| e.to_string()))
{
Ok(m) => Ok(m),
Err(e) => Err(Error::Model(e.clone())),
}
}
}
impl SmolVlm {
pub fn describe_image_unsplit(
&self,
image: &ImageBuffer,
opts: &VlmOptions,
) -> Result<String> {
let mut pieces = vec![Piece::Image(image)];
if let Some(t) = opts.prompt.as_deref() {
pieces.push(Piece::Text(t));
}
self.caption(&pieces, false, opts)
}
pub fn describe_image_traced(
&self,
image: &ImageBuffer,
opts: &VlmOptions,
) -> Result<(String, CaptionTrace)> {
let mut trace = CaptionTrace::default();
let mut pieces = vec![Piece::Image(image)];
if let Some(t) = opts.prompt.as_deref() {
pieces.push(Piece::Text(t));
}
let text = self.caption_traced(&pieces, true, opts, Some(&mut trace))?;
Ok((text, trace))
}
}
impl Default for SmolVlm {
fn default() -> Self {
Self::new()
}
}
impl VlmEngine for SmolVlm {
fn info(&self) -> EngineInfo {
EngineInfo {
name: "smolvlm".into(),
task: Task::Vlm,
status: EngineStatus::Stable,
description:
"SmolVLM-256M-Instruct on candle — SigLIP tower, pixel-shuffle connector, Llama decoder"
.into(),
}
}
fn describe(&self, prompt: &VlmPrompt<'_>, opts: &VlmOptions) -> Result<String> {
if prompt.image_count() == 0 {
return Err(Error::Other(
"argus: describe() needs at least one image — this is a VLM, not a chat model"
.into(),
));
}
let pieces: Vec<Piece<'_>> = prompt
.parts
.iter()
.map(|p| match p {
VlmPart::Text(t) => Piece::Text(t),
VlmPart::Image(i) => Piece::Image(i),
})
.collect();
self.caption(&pieces, true, opts)
}
fn describe_video(
&self,
frames: &[VideoFrame],
opts: &VlmOptions,
) -> Result<Vec<TimedSegment<String>>> {
if frames.is_empty() {
return Ok(Vec::new());
}
let window = opts
.frames_per_window
.unwrap_or(DEFAULT_FRAMES_PER_WINDOW)
.max(1);
let step = median_step(frames);
let mut out = Vec::with_capacity(frames.len().div_ceil(window));
for (w, chunk) in frames.chunks(window).enumerate() {
let mut pieces: Vec<Piece<'_>> =
chunk.iter().map(|f| Piece::Image(&f.image)).collect();
if let Some(t) = opts.prompt.as_deref() {
pieces.push(Piece::Text(t));
}
let value = self.caption(&pieces, false, opts)?;
let start = chunk[0].timestamp;
let end = frames
.get((w + 1) * window)
.map_or_else(|| chunk[chunk.len() - 1].timestamp + step, |f| f.timestamp);
out.push(TimedSegment {
start,
end,
value,
confidence: None,
});
}
Ok(out)
}
}
fn tile_workers(tiles: usize) -> usize {
if let Some(n) = std::env::var("FFAI_ARGUS_TILE_WORKERS")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|&n| n > 0)
{
return n.min(tiles).max(1);
}
let cores = std::thread::available_parallelism().map_or(1, std::num::NonZero::get);
(cores / 4).clamp(1, 6).min(tiles.max(1))
}
fn run_tower(
vision: &crate::vision::SmolVlmVision,
pre: &crate::preprocess::Preprocessed,
device: &Device,
) -> Result<(Vec<Tensor>, Vec<f64>)> {
let per = 3 * pre.tile * pre.tile;
let workers = tile_workers(pre.tiles);
if workers <= 1 || pre.tiles <= 1 {
let mut out = Vec::with_capacity(pre.tiles);
let mut ms = Vec::with_capacity(pre.tiles);
for t in 0..pre.tiles {
let t0 = crate::clock::Instant::now();
let px = pre.pixel_values[t * per..(t + 1) * per].to_vec();
let tensor = Tensor::from_vec(px, (1, 3, pre.tile, pre.tile), device)?;
out.push(vision.forward(&tensor)?.squeeze(0)?);
ms.push(t0.elapsed().as_secs_f64() * 1e3);
}
return Ok((out, ms));
}
let cores = std::thread::available_parallelism().map_or(1, std::num::NonZero::get);
let kernels_on = match std::env::var("FFAI_ARGUS_KERNELS_PARALLEL").ok().as_deref() {
Some("1") => true,
Some("0") => false,
_ => workers.saturating_mul(6) <= cores,
};
let prev = crate::siglip::set_kernels_parallel(kernels_on);
let next = std::sync::atomic::AtomicUsize::new(0);
let slots: Mutex<Vec<Option<(Tensor, f64)>>> = Mutex::new((0..pre.tiles).map(|_| None).collect());
let failed: Mutex<Option<String>> = Mutex::new(None);
std::thread::scope(|scope| {
for _ in 0..workers {
scope.spawn(|| loop {
let i = next.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if i >= pre.tiles {
break;
}
let t0 = crate::clock::Instant::now();
let px = pre.pixel_values[i * per..(i + 1) * per].to_vec();
let done = Tensor::from_vec(px, (1, 3, pre.tile, pre.tile), device)
.and_then(|t| vision.forward(&t))
.and_then(|o| o.squeeze(0));
match done {
Ok(block) => {
if let Ok(mut g) = slots.lock() {
g[i] = Some((block, t0.elapsed().as_secs_f64() * 1e3));
}
}
Err(e) => {
if let Ok(mut f) = failed.lock() {
f.get_or_insert_with(|| e.to_string());
}
}
}
});
}
});
crate::siglip::set_kernels_parallel(prev);
if let Some(e) = failed.into_inner().ok().flatten() {
return Err(Error::Model(format!("vision tower: {e}")));
}
let done = slots
.into_inner()
.map_err(|_| Error::Other("argus: tile results poisoned".into()))?;
let mut out = Vec::with_capacity(pre.tiles);
let mut ms = Vec::with_capacity(pre.tiles);
for (i, slot) in done.into_iter().enumerate() {
let (block, t) = slot.ok_or_else(|| Error::Model(format!("tile {i} produced nothing")))?;
out.push(block);
ms.push(t);
}
Ok((out, ms))
}
enum Piece<'a> {
Text(&'a str),
Image(&'a ImageBuffer),
}
impl SmolVlm {
fn caption(&self, pieces: &[Piece<'_>], split: bool, opts: &VlmOptions) -> Result<String> {
self.caption_traced(pieces, split, opts, None)
}
fn caption_traced(
&self,
pieces: &[Piece<'_>],
split: bool,
opts: &VlmOptions,
mut trace: Option<&mut CaptionTrace>,
) -> Result<String> {
let m = self.model()?;
let images: Vec<&ImageBuffer> = pieces
.iter()
.filter_map(|p| match p {
Piece::Image(i) => Some(*i),
Piece::Text(_) => None,
})
.collect();
let planned: usize = images
.iter()
.map(|i| {
crate::preprocess::tile_geometry(i.width as usize, i.height as usize, split).0
})
.sum();
let planned_tokens = planned * m.layout.tokens_per_tile;
if planned_tokens >= m.max_positions {
return Err(Error::Other(format!(
"argus: {} image(s) would contribute {planned_tokens} image tokens \
({planned} tiles x {}), and the text tower holds {}. {}",
images.len(),
m.layout.tokens_per_tile,
m.max_positions,
if split {
"Stills are split into 17 tiles each; pass fewer images."
} else {
"Reduce frames_per_window (--window)."
}
)));
}
let mut blocks: Vec<Tensor> = Vec::new();
let mut text = String::new();
for piece in pieces {
match piece {
Piece::Text(t) => text.push_str(t),
Piece::Image(img) => {
let t_pre = crate::clock::Instant::now();
let rgb = to_rgb8(img)?;
let pre = preprocess_rgb8_opts(
&rgb,
img.width as usize,
img.height as usize,
split,
);
if let Some(tr) = trace.as_deref_mut() {
tr.preprocess_ms += t_pre.elapsed().as_secs_f64() * 1e3;
tr.tiles += pre.tiles;
tr.rows = pre.rows;
tr.cols = pre.cols;
tr.tile = pre.tile;
tr.resized_to.push(crate::preprocess::resized_size(
img.width as usize,
img.height as usize,
));
}
let t_tower = crate::clock::Instant::now();
let (tiles, per_tile_ms) = run_tower(&m.vision, &pre, &m.device)?;
blocks.extend(tiles);
if let Some(tr) = trace.as_deref_mut() {
tr.tower_ms += t_tower.elapsed().as_secs_f64() * 1e3;
tr.tower_per_tile_ms.extend(per_tile_ms);
}
text.push_str(&m.layout.image_block(pre.rows, pre.cols));
}
}
}
let t_asm = crate::clock::Instant::now();
let templated = format!("<|im_start|>User:{text}<end_of_utterance>\nAssistant:");
let enc = m
.tokenizer
.encode(templated.as_str(), true)
.map_err(|e| Error::Model(format!("tokenize: {e}")))?;
let ids: Vec<i64> = enc.get_ids().iter().map(|&i| i64::from(i)).collect();
if ids.len() >= m.max_positions {
return Err(Error::Other(format!(
"argus: assembled prompt is {} tokens but the text tower holds {} — \
{planned_tokens} of those are image tokens, so the rest is text",
ids.len(),
m.max_positions
)));
}
let image_hidden = Tensor::stack(&blocks, 0)?;
let id_tensor = Tensor::from_vec(enc.get_ids().to_vec(), (1, ids.len()), &m.device)?;
let mut dec = m
.decoder
.lock()
.map_err(|_| Error::Other("argus: decoder mutex poisoned".into()))?;
let text_embeds = dec.embed(&id_tensor)?;
let merged = merge_image_embeddings(&text_embeds, &image_hidden, &ids, m.image_token_id)?;
let mut dtrace = crate::decode::DecodeTrace::default();
if let Some(tr) = trace.as_deref_mut() {
tr.assemble_ms += t_asm.elapsed().as_secs_f64() * 1e3;
tr.image_tokens = planned_tokens;
tr.prompt_tokens = ids.len();
tr.text_tokens = ids.len().saturating_sub(planned_tokens);
tr.max_positions = m.max_positions;
tr.split = split;
}
let out = dec.generate_traced(
&merged,
opts.max_new_tokens.unwrap_or(DEFAULT_MAX_NEW_TOKENS),
&m.stop_ids,
&opts.decoding,
opts.repetition_penalty,
Some(&mut dtrace),
)?;
drop(dec);
let t_detok = crate::clock::Instant::now();
let decoded = m
.tokenizer
.decode(&out, true)
.map_err(|e| Error::Model(format!("detokenize: {e}")))?;
let answer = truncate_at_stop(&decoded, &opts.stop).trim().to_string();
if let Some(tr) = &mut trace {
tr.prefill_ms = dtrace.prefill_ms;
tr.step_ms = dtrace.steps_ms;
tr.detokenize_ms = t_detok.elapsed().as_secs_f64() * 1e3;
}
Ok(answer)
}
}
fn median_step(frames: &[VideoFrame]) -> f64 {
if frames.len() < 2 {
return 1.0;
}
let mut steps: Vec<f64> = frames.windows(2).map(|w| w[1].timestamp - w[0].timestamp).collect();
steps.sort_by(f64::total_cmp);
steps[steps.len() / 2]
}
fn to_rgb8(img: &ImageBuffer) -> Result<Vec<u8>> {
if img.width == 0 || img.height == 0 {
return Err(Error::Media(format!(
"argus: image is {}x{} — nothing to describe",
img.width, img.height
)));
}
let n = img.width as usize * img.height as usize;
let want = n * img.format.bytes_per_pixel();
if img.data.len() < want {
return Err(Error::Media(format!(
"argus: {}x{} {:?} needs {want} bytes, buffer has {}",
img.width,
img.height,
img.format,
img.data.len()
)));
}
Ok(match img.format {
PixelFormat::Rgb8 => img.data[..want].to_vec(),
PixelFormat::Rgba8 => img.data[..want]
.chunks_exact(4)
.flat_map(|p| [p[0], p[1], p[2]])
.collect(),
PixelFormat::Gray8 => img.data[..want].iter().flat_map(|&g| [g, g, g]).collect(),
})
}
fn truncate_at_stop<'a>(text: &'a str, stops: &[String]) -> &'a str {
let cut = stops
.iter()
.filter(|s| !s.is_empty())
.filter_map(|s| text.find(s.as_str()))
.min();
cut.map_or(text, |i| &text[..i])
}
fn load(manifest_dir: &Path) -> Result<Model> {
let manifests = ffai_models::load_dir(manifest_dir)?;
let manifest = manifests
.iter()
.find(|m| m.name == MODEL)
.ok_or_else(|| {
Error::Model(format!(
"no model manifest named `{MODEL}` in {}",
manifest_dir.display()
))
})?;
let resolved = manifest.fetch()?;
let weights = resolved.file("model.safetensors")?.to_path_buf();
let config_path = resolved.file("config.json")?;
let tokenizer_path = resolved.file("tokenizer.json")?;
let config_json = std::fs::read_to_string(config_path)?;
let device = Device::Cpu;
#[allow(unsafe_code)]
let vb = unsafe {
candle_nn::VarBuilder::from_mmaped_safetensors(
std::slice::from_ref(&weights),
candle_core::DType::F32,
&device,
)
}
.map_err(|e| Error::Model(format!("load {}: {e}", weights.display())))?;
let tokenizer = tokenizers::Tokenizer::from_file(tokenizer_path)
.map_err(|e| Error::Model(format!("tokenizer: {e}")))?;
build(vb, &config_json, tokenizer, device)
}
pub struct ArgusBytes {
pub weights: Vec<u8>,
pub config: String,
pub tokenizer: Vec<u8>,
}
fn build(
vb: candle_nn::VarBuilder<'static>,
config_json: &str,
tokenizer: tokenizers::Tokenizer,
device: Device,
) -> Result<Model> {
let vision = crate::vision::load_vb(vb.clone(), config_json).map_err(Error::Model)?;
let decoder = TextDecoder::load_vb(vb, config_json, &device).map_err(Error::Model)?;
let cfg: serde_json::Value = serde_json::from_str(config_json)
.map_err(|e| Error::Model(format!("config.json: {e}")))?;
let vision_cfg = cfg.get("vision_config");
let get = |k: &str, d: usize| -> usize {
vision_cfg
.and_then(|v| v.get(k))
.and_then(serde_json::Value::as_u64)
.map_or(d, |x| x as usize)
};
let scale_factor = cfg
.get("scale_factor")
.and_then(serde_json::Value::as_u64)
.map_or(4, |x| x as usize);
let layout = PromptLayout::default().with_geometry(
get("image_size", 512),
get("patch_size", 16),
scale_factor,
);
let id_of = |t: &str| tokenizer.token_to_id(t).map(i64::from);
let image_token_id = id_of("<image>")
.ok_or_else(|| Error::Model("tokenizer has no `<image>` token".into()))?;
let mut stop_ids: Vec<u32> = Vec::new();
for t in ["<end_of_utterance>", "<|im_end|>", "<|endoftext|>"] {
if let Some(id) = tokenizer.token_to_id(t) {
stop_ids.push(id);
}
}
if let Some(id) = cfg
.get("text_config")
.and_then(|t| t.get("eos_token_id"))
.and_then(serde_json::Value::as_u64)
{
let id = id as u32;
if !stop_ids.contains(&id) {
stop_ids.push(id);
}
}
if stop_ids.is_empty() {
return Err(Error::Model(
"no end-of-turn token found in the tokenizer or config — every caption would run to the token budget".into(),
));
}
let max_positions = cfg
.get("text_config")
.and_then(|t| t.get("max_position_embeddings"))
.and_then(serde_json::Value::as_u64)
.map_or(8192, |x| x as usize);
Ok(Model {
vision,
decoder: Mutex::new(decoder),
tokenizer,
layout,
image_token_id,
stop_ids,
max_positions,
device,
})
}
#[cfg(test)]
mod tests {
use super::*;
use ffai_core::engine::Decoding;
fn img(format: PixelFormat, data: Vec<u8>) -> ImageBuffer {
ImageBuffer {
width: 2,
height: 1,
format,
data,
}
}
#[test]
fn grayscale_is_replicated_and_alpha_is_dropped() {
assert_eq!(
to_rgb8(&img(PixelFormat::Gray8, vec![10, 200])).unwrap(),
vec![10, 10, 10, 200, 200, 200]
);
assert_eq!(
to_rgb8(&img(PixelFormat::Rgba8, vec![1, 2, 3, 255, 4, 5, 6, 0])).unwrap(),
vec![1, 2, 3, 4, 5, 6],
"alpha is dropped, not composited — compositing needs a background \
colour and inventing one changes the picture"
);
assert_eq!(
to_rgb8(&img(PixelFormat::Rgb8, vec![1, 2, 3, 4, 5, 6])).unwrap(),
vec![1, 2, 3, 4, 5, 6]
);
}
#[test]
fn a_zero_dimension_image_is_refused() {
let e = to_rgb8(&ImageBuffer {
width: 0,
height: 0,
format: PixelFormat::Rgb8,
data: Vec::new(),
})
.unwrap_err();
assert!(format!("{e}").contains("nothing to describe"), "{e}");
}
#[test]
fn a_short_buffer_is_an_error_not_a_panic() {
let e = to_rgb8(&img(PixelFormat::Rgb8, vec![1, 2, 3])).unwrap_err();
assert!(format!("{e}").contains("needs 6 bytes"), "{e}");
}
#[test]
fn stops_cut_before_the_marker_and_take_the_earliest() {
let stops = vec!["\nUser:".to_string(), "###".to_string()];
assert_eq!(truncate_at_stop("a cat ### b", &stops), "a cat ");
assert_eq!(truncate_at_stop("x\nUser: y ### z", &stops), "x");
assert_eq!(truncate_at_stop("no marker", &stops), "no marker");
assert_eq!(truncate_at_stop("keep me", &[String::new()]), "keep me");
}
#[test]
fn the_last_video_segment_gets_the_median_spacing() {
let f = |t: f64| VideoFrame {
image: img(PixelFormat::Gray8, vec![0, 0]),
timestamp: t,
};
assert_eq!(median_step(&[f(0.0), f(0.5), f(1.0)]), 0.5);
assert_eq!(median_step(&[f(0.0)]), 1.0);
}
#[test]
fn the_default_decoding_is_greedy_so_a_caption_is_reproducible() {
assert_eq!(VlmOptions::default().decoding, Decoding::Greedy);
}
}