use std::fmt;
use std::path::Path;
use taconite_bundle::{Manifest, Store};
pub use taconite::Timing;
pub mod model;
pub mod npu;
pub mod preprocess;
pub mod tokenizer;
use model::Model;
use npu::Npu;
use tokenizer::Tokenizer;
pub const VERSION: u32 = 1;
#[derive(Debug)]
pub enum Error {
Bundle(String),
Npu(String),
Input(String),
}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Error::Bundle(m) => write!(f, "bundle: {m}"),
Error::Npu(m) => write!(f, "NPU: {m}"),
Error::Input(m) => write!(f, "input: {m}"),
}
}
}
impl std::error::Error for Error {}
impl From<taconite_bundle::Error> for Error {
fn from(e: taconite_bundle::Error) -> Self {
Error::Bundle(e.to_string())
}
}
impl From<taconite::Error> for Error {
fn from(e: taconite::Error) -> Self {
Error::Npu(e.to_string())
}
}
#[derive(Debug, Clone)]
pub struct Config {
pub v_d: usize,
pub v_i: usize,
pub v_layers: usize,
pub v_heads: usize,
pub v_hd: usize,
pub patch: usize,
pub pool: usize,
pub v_theta: f32,
pub v_eps: f32,
pub v_rows: usize,
pub d: usize,
pub i: usize,
pub layers: usize,
pub heads: usize,
pub eps: f32,
pub hd: Vec<usize>,
pub kv_heads: Vec<usize>,
pub global: Vec<bool>,
pub theta_s: f32,
pub theta_g: f32,
pub ple: usize,
pub out: usize,
pub window: usize,
pub t_rows: usize,
pub o_width: usize,
pub max_soft_tokens: usize,
pub no_window: u32,
pub max_text_tokens: usize,
}
impl Config {
fn load(m: &Manifest) -> Result<Self, Error> {
let p = |k: &str| m.param_as::<usize>(k);
let f = |k: &str| m.param_as::<f32>(k);
let global = m
.param("layer_types")?
.split(',')
.map(|t| match t {
"g" => Ok(true),
"s" => Ok(false),
_ => Err(Error::Bundle(format!("layer type {t}"))),
})
.collect::<Result<Vec<_>, _>>()?;
Ok(Config {
v_d: p("v.D")?,
v_i: p("v.I")?,
v_layers: p("v.layers")?,
v_heads: p("v.heads")?,
v_hd: p("v.hd")?,
patch: p("v.patch")?,
pool: p("v.pool")?,
v_theta: f("v.theta")?,
v_eps: f("v.eps")?,
v_rows: p("v.rows")?,
d: p("D")?,
i: p("I")?,
layers: p("layers")?,
heads: p("heads")?,
eps: f("eps")?,
hd: m.list("hd")?,
kv_heads: m.list("kv_heads")?,
global,
theta_s: f("theta_s")?,
theta_g: f("theta_g")?,
ple: p("ple")?,
out: p("out")?,
window: p("window")?,
t_rows: p("t.rows")?,
o_width: p("o_width")?,
max_soft_tokens: p("max_soft_tokens")?,
no_window: m.param_as("no_window")?,
max_text_tokens: if m.has_param("max_text_tokens") { p("max_text_tokens")? } else { 0 },
})
}
pub fn max_patches(&self) -> usize {
self.max_soft_tokens * self.pool * self.pool
}
}
pub struct EmbeddingGemma2 {
pub cfg: Config,
pub npu: Npu,
pub model: Model,
pub manifest: Manifest,
pub store: Store,
pub tokenizer: Option<Tokenizer>,
pub timing: Timing,
}
impl EmbeddingGemma2 {
pub fn load(dir: &Path) -> Result<Self, Error> {
let manifest = Manifest::load(dir, VERSION)?;
let store = Store::load(dir)?;
let cfg = Config::load(&manifest)?;
let npu = Npu::open(&manifest)?;
let model = Model::load(&cfg, &store, &npu)?;
let tokenizer = Tokenizer::load(&manifest, &store)?;
Ok(EmbeddingGemma2 { cfg, npu, model, manifest, store, tokenizer, timing: Timing::default() })
}
pub fn contexts(&self) -> usize {
self.npu.contexts
}
pub fn preprocess(&self, rgb: &[u8], w: usize, h: usize) -> Result<preprocess::Patches, Error> {
let c = &self.cfg;
preprocess::patches(rgb, w, h, c.patch, c.pool, c.max_soft_tokens)
.ok_or_else(|| Error::Input(format!("a {w} x {h} image is too thin")))
}
pub fn embed_rgb(&mut self, rgb: &[u8], w: usize, h: usize) -> Result<Vec<f32>, Error> {
let p = self.preprocess(rgb, w, h)?;
self.embed_patches(&p)
}
pub fn embed_patches(&mut self, p: &preprocess::Patches) -> Result<Vec<f32>, Error> {
self.timing.clear();
let soft = self.model.vision(&self.cfg, &self.npu, p, &mut self.timing)?;
let x = self.model.image_sequence(&self.cfg, &soft);
self.model.text(&self.cfg, &self.npu, x, &mut self.timing)
}
pub fn prompts(&self) -> &[(String, String)] {
self.tokenizer.as_ref().map_or(&[], |t| t.prompts.as_slice())
}
pub fn tokenize(&self, text: &str, prompt: Option<&str>) -> Result<Vec<u32>, Error> {
let tok = self.tokenizer.as_ref().ok_or_else(|| Error::Bundle("this bundle has no text path".into()))?;
let full = match prompt {
Some(p) => {
let pre = tok.prompt(p).ok_or_else(|| {
let names: Vec<&str> = tok.prompts.iter().map(|(n, _)| n.as_str()).collect();
Error::Input(format!("no prompt {p} (the model's: {})", names.join(", ")))
})?;
format!("{pre}{text}")
}
None => text.to_string(),
};
Ok(tok.encode(&full))
}
pub fn embed_text(&mut self, text: &str, prompt: Option<&str>) -> Result<Vec<f32>, Error> {
let ids = self.tokenize(text, prompt)?;
self.embed_ids(&ids)
}
pub fn embed_ids(&mut self, ids: &[u32]) -> Result<Vec<f32>, Error> {
let (d, max) = (self.cfg.d, self.cfg.max_text_tokens);
if ids.len() > max {
return Err(Error::Input(format!("{} tokens: this bundle takes at most {max}", ids.len())));
}
self.timing.clear();
let table = self.store.bf16("t.embed")?;
let vocab = table.len() / d;
let scale = (d as f32).sqrt();
let mut x = Vec::with_capacity(ids.len() * d);
for &id in ids {
let id = id as usize;
if id >= vocab {
return Err(Error::Input(format!("token id {id} (vocabulary {vocab})")));
}
x.extend(table[id * d..(id + 1) * d].iter().map(|&b| taconite::bf16_to_f32(b) * scale));
}
self.model.text(&self.cfg, &self.npu, x, &mut self.timing)
}
pub fn soft_tokens(&mut self, p: &preprocess::Patches) -> Result<Vec<f32>, Error> {
self.model.vision(&self.cfg, &self.npu, p, &mut self.timing)
}
}
pub fn cosine(a: &[f32], b: &[f32]) -> f32 {
let d = |x: &[f32], y: &[f32]| x.iter().zip(y).map(|(p, q)| (*p as f64) * (*q as f64)).sum::<f64>();
(d(a, b) / (d(a, a) * d(b, b)).sqrt()) as f32
}
pub fn truncate(e: &[f32], dim: usize) -> Vec<f32> {
let v = &e[..dim.min(e.len())];
let n = v.iter().map(|x| x * x).sum::<f32>().sqrt().max(1e-12);
v.iter().map(|x| x / n).collect()
}