use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder;
use std::time::Instant;
const PROMPT: usize = 1142;
const STEPS: usize = 24;
fn main() -> Result<(), Box<dyn std::error::Error>> {
let home = std::env::var("USERPROFILE").or_else(|_| std::env::var("HOME"))?;
let snap = std::fs::read_dir(
std::path::Path::new(&home)
.join(".cache/huggingface/hub/models--HuggingFaceTB--SmolVLM-256M-Instruct/snapshots"),
)?
.flatten()
.next()
.ok_or("no snapshot")?
.path();
let weights = snap.join("model.safetensors");
let cj = std::fs::read_to_string(snap.join("config.json"))?;
let d = Device::Cpu;
let v: serde_json::Value = serde_json::from_str(&cj)?;
let t = v.get("text_config").unwrap_or(&v);
let g = |k: &str, dv: u64| t.get(k).and_then(serde_json::Value::as_u64).unwrap_or(dv);
let gf = |k: &str, dv: f64| t.get(k).and_then(serde_json::Value::as_f64).unwrap_or(dv);
let heads = g("num_attention_heads", 9) as usize;
let hidden = g("hidden_size", 576) as usize;
let cfg = ffai_argus::text::Cfg {
layers: g("num_hidden_layers", 30) as usize,
hidden,
heads,
kv_heads: g("num_key_value_heads", 3) as usize,
head_dim: hidden / heads,
inter: g("intermediate_size", 1536) as usize,
eps: gf("rms_norm_eps", 1e-5),
rope_theta: gf("rope_theta", 100_000.0) as f32,
max_pos: g("max_position_embeddings", 8192) as usize,
};
let layers = cfg.layers;
let inter = cfg.inter;
let kv_heads = cfg.kv_heads;
let head_dim = cfg.head_dim;
let vb = unsafe {
VarBuilder::from_mmaped_safetensors(std::slice::from_ref(&weights), DType::F32, &d)?
};
let mut tower = ffai_argus::text::TextTower::load(&vb, cfg, &d)?;
let prompt = Tensor::rand(-1.0f32, 1.0, (1, PROMPT, hidden), &d)?;
let step = Tensor::rand(-1.0f32, 1.0, (1, 1, hidden), &d)?;
tower.reset();
let _ = tower.forward(&prompt, 0)?;
let _ = tower.forward(&step, PROMPT)?;
let _ = ffai_argus::text::prof::take();
tower.reset();
let _ = tower.forward(&prompt, 0)?;
let _ = ffai_argus::text::prof::take();
let t0 = Instant::now();
for i in 0..STEPS {
let _ = tower.forward(&step, PROMPT + i)?;
}
let whole = t0.elapsed().as_secs_f64() * 1e3;
let rows = ffai_argus::text::prof::take();
let per = whole / STEPS as f64;
if rows.is_empty() {
println!("
DECODE with profiling OFF: {per:.2} ms/token over {STEPS} steps");
println!(" (compare against the profiled run to price the instrument itself)");
return Ok(());
}
let sum: f64 = rows.iter().map(|r| r.1).sum();
println!("\nOUR text DECODE, {STEPS} steps at seq 1, cache primed to {PROMPT}, {layers} layers\n");
println!(" {:<24} {:>10} {:>10} {:>8}", "op (all layers)", "ms total", "ms/token", "share");
println!(" {:-<24} {:->10} {:->10} {:->8}", "", "", "", "");
for (n, ms) in &rows {
println!(
" {n:<24} {ms:>10.1} {:>10.2} {:>7.1}%",
ms / STEPS as f64,
100.0 * ms / whole
);
}
println!(" {:-<24} {:->10} {:->10} {:->8}", "", "", "", "");
println!(" {:<24} {sum:>10.1} {:>10.2} {:>7.1}%", "accounted", sum / STEPS as f64, 100.0 * sum / whole);
println!(
" {:<24} {:>10.1} {:>10.2} {:>7.1}%",
"UNACCOUNTED",
whole - sum,
(whole - sum) / STEPS as f64,
100.0 * (whole - sum) / whole
);
println!(" {:<24} {whole:>10.1} {per:>10.2}", "whole decode");
let per_layer = hidden * (heads * head_dim) + 2 * hidden * (kv_heads * head_dim) + hidden * (heads * head_dim) + 3 * hidden * inter + 2 * hidden; let params = layers * per_layer;
let bytes = params as f64 * 4.0;
println!("\n Weight bytes read per token: {:.0} MB ({:.1} M params, f32)", bytes / 1e6, params as f64 / 1e6);
for bw in [15.0f64, 20.0, 25.0] {
println!(
" at {bw:>4.0} GB/s -> {:>6.2} ms/token floor (we are at {per:.2}, {:.2}x the floor)",
bytes / (bw * 1e9) * 1e3,
per / (bytes / (bw * 1e9) * 1e3)
);
}
println!(
"\n If we sit near 1.0x the floor, decode is BANDWIDTH-bound and the only\n \
lever is fewer bytes (dtype), not better code. Well above it means overhead."
);
Ok(())
}