taconite_embeddinggemma2/
lib.rs1use std::fmt;
27use std::path::Path;
28
29use taconite_bundle::{Manifest, Store};
30
31pub use taconite::Timing;
32
33pub mod model;
34pub mod npu;
35pub mod preprocess;
36pub mod tokenizer;
37
38use model::Model;
39use npu::Npu;
40use tokenizer::Tokenizer;
41
42pub const VERSION: u32 = 1;
43
44#[derive(Debug)]
45pub enum Error {
46 Bundle(String),
47 Npu(String),
48 Input(String),
49}
50
51impl fmt::Display for Error {
52 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
53 match self {
54 Error::Bundle(m) => write!(f, "bundle: {m}"),
55 Error::Npu(m) => write!(f, "NPU: {m}"),
56 Error::Input(m) => write!(f, "input: {m}"),
57 }
58 }
59}
60
61impl std::error::Error for Error {}
62
63impl From<taconite_bundle::Error> for Error {
64 fn from(e: taconite_bundle::Error) -> Self {
65 Error::Bundle(e.to_string())
66 }
67}
68
69impl From<taconite::Error> for Error {
70 fn from(e: taconite::Error) -> Self {
71 Error::Npu(e.to_string())
72 }
73}
74
75#[derive(Debug, Clone)]
77pub struct Config {
78 pub v_d: usize,
79 pub v_i: usize,
80 pub v_layers: usize,
81 pub v_heads: usize,
82 pub v_hd: usize,
83 pub patch: usize,
84 pub pool: usize,
85 pub v_theta: f32,
86 pub v_eps: f32,
87 pub v_rows: usize,
89 pub d: usize,
90 pub i: usize,
91 pub layers: usize,
92 pub heads: usize,
93 pub eps: f32,
94 pub hd: Vec<usize>,
95 pub kv_heads: Vec<usize>,
96 pub global: Vec<bool>,
98 pub theta_s: f32,
99 pub theta_g: f32,
100 pub ple: usize,
101 pub out: usize,
102 pub window: usize,
103 pub t_rows: usize,
104 pub o_width: usize,
105 pub max_soft_tokens: usize,
106 pub no_window: u32,
107 pub max_text_tokens: usize,
110}
111
112impl Config {
113 fn load(m: &Manifest) -> Result<Self, Error> {
114 let p = |k: &str| m.param_as::<usize>(k);
115 let f = |k: &str| m.param_as::<f32>(k);
116 let global = m
117 .param("layer_types")?
118 .split(',')
119 .map(|t| match t {
120 "g" => Ok(true),
121 "s" => Ok(false),
122 _ => Err(Error::Bundle(format!("layer type {t}"))),
123 })
124 .collect::<Result<Vec<_>, _>>()?;
125 Ok(Config {
126 v_d: p("v.D")?,
127 v_i: p("v.I")?,
128 v_layers: p("v.layers")?,
129 v_heads: p("v.heads")?,
130 v_hd: p("v.hd")?,
131 patch: p("v.patch")?,
132 pool: p("v.pool")?,
133 v_theta: f("v.theta")?,
134 v_eps: f("v.eps")?,
135 v_rows: p("v.rows")?,
136 d: p("D")?,
137 i: p("I")?,
138 layers: p("layers")?,
139 heads: p("heads")?,
140 eps: f("eps")?,
141 hd: m.list("hd")?,
142 kv_heads: m.list("kv_heads")?,
143 global,
144 theta_s: f("theta_s")?,
145 theta_g: f("theta_g")?,
146 ple: p("ple")?,
147 out: p("out")?,
148 window: p("window")?,
149 t_rows: p("t.rows")?,
150 o_width: p("o_width")?,
151 max_soft_tokens: p("max_soft_tokens")?,
152 no_window: m.param_as("no_window")?,
153 max_text_tokens: if m.has_param("max_text_tokens") { p("max_text_tokens")? } else { 0 },
154 })
155 }
156
157 pub fn max_patches(&self) -> usize {
159 self.max_soft_tokens * self.pool * self.pool
160 }
161}
162
163pub struct EmbeddingGemma2 {
164 pub cfg: Config,
165 pub npu: Npu,
166 pub model: Model,
167 pub manifest: Manifest,
168 pub store: Store,
169 pub tokenizer: Option<Tokenizer>,
171 pub timing: Timing,
173}
174
175impl EmbeddingGemma2 {
176 pub fn load(dir: &Path) -> Result<Self, Error> {
179 let manifest = Manifest::load(dir, VERSION)?;
180 let store = Store::load(dir)?;
181 let cfg = Config::load(&manifest)?;
182 let npu = Npu::open(&manifest)?;
183 let model = Model::load(&cfg, &store, &npu)?;
184 let tokenizer = Tokenizer::load(&manifest, &store)?;
185 Ok(EmbeddingGemma2 { cfg, npu, model, manifest, store, tokenizer, timing: Timing::default() })
186 }
187
188 pub fn contexts(&self) -> usize {
190 self.npu.contexts
191 }
192
193 pub fn preprocess(&self, rgb: &[u8], w: usize, h: usize) -> Result<preprocess::Patches, Error> {
196 let c = &self.cfg;
197 preprocess::patches(rgb, w, h, c.patch, c.pool, c.max_soft_tokens)
198 .ok_or_else(|| Error::Input(format!("a {w} x {h} image is too thin")))
199 }
200
201 pub fn embed_rgb(&mut self, rgb: &[u8], w: usize, h: usize) -> Result<Vec<f32>, Error> {
203 let p = self.preprocess(rgb, w, h)?;
204 self.embed_patches(&p)
205 }
206
207 pub fn embed_patches(&mut self, p: &preprocess::Patches) -> Result<Vec<f32>, Error> {
209 self.timing.clear();
210 let soft = self.model.vision(&self.cfg, &self.npu, p, &mut self.timing)?;
211 let x = self.model.image_sequence(&self.cfg, &soft);
212 self.model.text(&self.cfg, &self.npu, x, &mut self.timing)
213 }
214
215 pub fn prompts(&self) -> &[(String, String)] {
218 self.tokenizer.as_ref().map_or(&[], |t| t.prompts.as_slice())
219 }
220
221 pub fn tokenize(&self, text: &str, prompt: Option<&str>) -> Result<Vec<u32>, Error> {
225 let tok = self.tokenizer.as_ref().ok_or_else(|| Error::Bundle("this bundle has no text path".into()))?;
226 let full = match prompt {
227 Some(p) => {
228 let pre = tok.prompt(p).ok_or_else(|| {
229 let names: Vec<&str> = tok.prompts.iter().map(|(n, _)| n.as_str()).collect();
230 Error::Input(format!("no prompt {p} (the model's: {})", names.join(", ")))
231 })?;
232 format!("{pre}{text}")
233 }
234 None => text.to_string(),
235 };
236 Ok(tok.encode(&full))
237 }
238
239 pub fn embed_text(&mut self, text: &str, prompt: Option<&str>) -> Result<Vec<f32>, Error> {
243 let ids = self.tokenize(text, prompt)?;
244 self.embed_ids(&ids)
245 }
246
247 pub fn embed_ids(&mut self, ids: &[u32]) -> Result<Vec<f32>, Error> {
249 let (d, max) = (self.cfg.d, self.cfg.max_text_tokens);
250 if ids.len() > max {
251 return Err(Error::Input(format!("{} tokens: this bundle takes at most {max}", ids.len())));
252 }
253 self.timing.clear();
254 let table = self.store.bf16("t.embed")?;
255 let vocab = table.len() / d;
256 let scale = (d as f32).sqrt();
257 let mut x = Vec::with_capacity(ids.len() * d);
258 for &id in ids {
259 let id = id as usize;
260 if id >= vocab {
261 return Err(Error::Input(format!("token id {id} (vocabulary {vocab})")));
262 }
263 x.extend(table[id * d..(id + 1) * d].iter().map(|&b| taconite::bf16_to_f32(b) * scale));
264 }
265 self.model.text(&self.cfg, &self.npu, x, &mut self.timing)
266 }
267
268 pub fn soft_tokens(&mut self, p: &preprocess::Patches) -> Result<Vec<f32>, Error> {
271 self.model.vision(&self.cfg, &self.npu, p, &mut self.timing)
272 }
273}
274
275pub fn cosine(a: &[f32], b: &[f32]) -> f32 {
277 let d = |x: &[f32], y: &[f32]| x.iter().zip(y).map(|(p, q)| (*p as f64) * (*q as f64)).sum::<f64>();
278 (d(a, b) / (d(a, a) * d(b, b)).sqrt()) as f32
279}
280
281pub fn truncate(e: &[f32], dim: usize) -> Vec<f32> {
284 let v = &e[..dim.min(e.len())];
285 let n = v.iter().map(|x| x * x).sum::<f32>().sqrt().max(1e-12);
286 v.iter().map(|x| x / n).collect()
287}