Skip to main content

taconite_embeddinggemma2/
lib.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 Brishen Hawkins
2// SPDX-License-Identifier: Apache-2.0
3
4//! EmbeddingGemma 2 (`google/embeddinggemma-2`) image and text embeddings
5//! on an AMD XDNA NPU.
6//!
7//! The bundle `iron/applications/embeddinggemma2/export_eg2.py` writes
8//! holds every compiled IRON kernel and the weights (the NPU ones
9//! pre-packed); this crate replays the forward the Python app runs
10//! (`eg2_common.py` / `eg2_npu.py`):
11//!
12//! | | NPU | host (here) |
13//! |---|---|---|
14//! | vision tower (16 layers, 768 wide) | every projection (`flm.GEMM`s: patch embedding, qkv, o, GeGLU gate+up, down, embed_vision) and the attention (the MHA operator) | resize + patchify (`preprocess.rs`), position embeddings, RMSNorms, q / k / v norms, 2D RoPE, residual adds, 3 x 3 pooling |
15//! | text encoder (24 layers, 512 wide) | every projection (per-layer-input projection, qkv, o, GeGLU, down, PLE gate + projection) | norms, RoPE, attention (~270 tokens), per-layer gating, mean pooling, the 512 -> 768 projection |
16//!
17//! [`EmbeddingGemma2::embed_rgb`] and [`EmbeddingGemma2::embed_text`]
18//! give the L2-normalized 768-d embeddings sentence-transformers computes
19//! for an image and for a text (with its task prompt, e.g. `SearchQuery`
20//! or `Document`); both live in the model's one space, compared by cosine.
21//! Text runs through the same text encoder as an image's soft tokens:
22//! Gemma's tokenizer (`tokenizer.rs`), the token embeddings, then every
23//! projection on the NPU and the attention (a 512-token sliding window on
24//! 20 of the 24 layers) on the host.
25
26use 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/// The model's constants (the manifest's params).
76#[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    /// patches the device buffers hold
88    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    /// per layer: a global (full-attention) layer
97    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    /// the longest text (tokens, with `<bos>` / `<eos>`) the bundle takes;
108    /// 0 for an image-only bundle
109    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    /// Patches the processor makes at most of an image.
158    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    /// None for an image-only bundle
170    pub tokenizer: Option<Tokenizer>,
171    /// where the last call spent its time (`npu:<kernel>`, host stages)
172    pub timing: Timing,
173}
174
175impl EmbeddingGemma2 {
176    /// Loads a bundle: opens the NPU, loads every kernel and uploads the
177    /// packed weights.
178    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    /// Hardware contexts the bundle's kernels use.
189    pub fn contexts(&self) -> usize {
190        self.npu.contexts
191    }
192
193    /// RGB8 `[h, w, 3]` -> the image's patches, as the HF processor makes
194    /// them.
195    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    /// RGB8 `[h, w, 3]` -> its unit 768-d embedding.
202    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    /// An image's patches -> its unit embedding.
208    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    /// The prompt names the model was trained with ("SearchQuery",
216    /// "Document", "Classification", ...) and their text.
217    pub fn prompts(&self) -> &[(String, String)] {
218        self.tokenizer.as_ref().map_or(&[], |t| t.prompts.as_slice())
219    }
220
221    /// `text`'s token ids, with prompt `prompt` (a name from
222    /// [`prompts`](Self::prompts)) prepended, as sentence-transformers
223    /// makes them.
224    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    /// A text -> its unit 768-d embedding. `prompt` names the task prefix:
240    /// "SearchQuery" for a query, "Document" for what it searches (see the
241    /// model card); None embeds the text as it is.
242    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    /// Token ids -> their unit embedding.
248    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    /// An image's patches -> its soft tokens `[s, 512]` (the text model's
269    /// input embeddings of the image), for checks.
270    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
275/// Cosine similarity of two embeddings (dot product of unit vectors).
276pub 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
281/// An embedding truncated to its first `dim` values (Matryoshka: 768, 512,
282/// 256 or 128) and re-normalized.
283pub 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}