Skip to main content

cyberbrain_embed/
model.rs

1//! The static embedder: tokenizer + embedding matrix, mean pooling, L2 normalisation.
2
3use crate::POOLING;
4use crate::artefact::{ArtefactManifest, ModelPaths, map_verified, read_verified, record_verified};
5use crate::pool;
6use crate::weights::{Matrix, load_matrix_mapped, tensor_shape};
7use cyberbrain_core::{Embedder, Error, Result, Slash};
8use rayon::prelude::*;
9use serde::Serialize;
10use tokenizers::Tokenizer;
11use tokenizers::models::ModelWrapper;
12
13/// Batches at least this large are split across rayon's thread pool. Below it the
14/// per-item work (a few microseconds) is smaller than the scheduling overhead, and the
15/// hook path embeds exactly one query at a time.
16const PARALLEL_THRESHOLD: usize = 64;
17
18/// Unknown-token names probed when the tokenizer model does not expose its own.
19const UNK_CANDIDATES: [&str; 4] = ["[UNK]", "<unk>", "<|unk|>", "[unk]"];
20
21/// Knobs for loading. `Default` is what production uses.
22#[derive(Debug, Clone, Default)]
23pub struct LoadOptions {
24    /// Keep at most this many token ids per input before pooling. `None` pools every
25    /// token. Truncation is applied to the id list after tokenisation, never by the
26    /// tokenizer's own truncation setting, which is disabled on load.
27    pub max_tokens: Option<usize>,
28    /// Override the unknown-token string. Normally derived from the tokenizer model.
29    pub unk_token: Option<String>,
30}
31
32/// One embedded input. The vector is L2-normalised, or all zero when nothing was pooled.
33#[derive(Debug, Clone, PartialEq)]
34pub struct Embedding {
35    pub vector: Vec<f32>,
36    /// Token ids the tokenizer produced for the input (after `max_tokens`).
37    pub tokens_seen: usize,
38    /// Of those, how many had a row in the matrix. `0` means the vector is all zero.
39    pub tokens_known: usize,
40}
41
42impl Embedding {
43    /// Nothing was pooled: empty input, or every token unknown. The vector is all zero.
44    pub fn is_empty(&self) -> bool {
45        self.tokens_known == 0
46    }
47}
48
49/// Technical facts about the loaded model, for `status` and the model card (SPEC §12.7).
50/// Source and licence are not known here; the policy crate's registry carries those.
51#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
52pub struct ModelInfo {
53    pub profile_id: String,
54    pub dim: usize,
55    pub vocab_rows: usize,
56    pub weights_dtype: &'static str,
57    pub pooling: &'static str,
58    pub weights_blake3: String,
59    pub tokenizer_blake3: String,
60    pub unk_id: Option<u32>,
61    pub max_tokens: Option<usize>,
62}
63
64/// A loaded model2vec-format model. Cheap to share behind an `Arc`; `embed` takes `&self`.
65pub struct StaticEmbedder {
66    tokenizer: Tokenizer,
67    matrix: Matrix,
68    unk_id: Option<u32>,
69    max_tokens: Option<usize>,
70    manifest: ArtefactManifest,
71    profile_id: String,
72}
73
74impl std::fmt::Debug for StaticEmbedder {
75    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
76        f.debug_struct("StaticEmbedder")
77            .field("profile_id", &self.profile_id)
78            .field("dim", &self.matrix.dim)
79            .field("vocab_rows", &self.matrix.rows)
80            .field("dtype", &self.matrix.dtype)
81            .field("unk_id", &self.unk_id)
82            .finish_non_exhaustive()
83    }
84}
85
86impl StaticEmbedder {
87    /// Loads and verifies a model from local files. Fails, never warns, when either file's
88    /// blake3 digest differs from the manifest, when the files do not parse, when the
89    /// matrix is not `[vocab, dim]`, contains a non-finite value, or has fewer rows than
90    /// the tokenizer has ids.
91    pub fn load(paths: &ModelPaths, manifest: &ArtefactManifest) -> Result<Self> {
92        Self::load_with(paths, manifest, LoadOptions::default())
93    }
94
95    pub fn load_with(
96        paths: &ModelPaths,
97        manifest: &ArtefactManifest,
98        opts: LoadOptions,
99    ) -> Result<Self> {
100        manifest.validate()?;
101
102        let tok_bytes = read_verified(&paths.tokenizer, &manifest.tokenizer_blake3, "tokenizer")?;
103        let (w_map, w_known) = map_verified(&paths.weights, &manifest.weights_blake3, "weights")?;
104
105        let mut tokenizer = Tokenizer::from_bytes(&tok_bytes).map_err(|e| {
106            Error::Embed(format!(
107                "tokenizer {} is not a valid tokenizers file: {e}",
108                Slash(&paths.tokenizer)
109            ))
110        })?;
111        // The artefact may carry the source model's padding/truncation. Neither belongs in
112        // a mean-pooled lookup: padding would pool the pad row, truncation would drop text
113        // without a trace. `max_tokens` is applied explicitly instead.
114        tokenizer.with_padding(None);
115        tokenizer
116            .with_truncation(None)
117            .map_err(|e| Error::Embed(format!("cannot disable tokenizer truncation: {e}")))?;
118
119        let matrix = load_matrix_mapped(w_map, !w_known)?;
120
121        let vocab = tokenizer.get_vocab_size(true);
122        if vocab > matrix.rows {
123            return Err(Error::Embed(format!(
124                "tokenizer knows {vocab} ids but the embedding matrix has only {} rows; the \
125                 two files do not belong to the same model",
126                matrix.rows
127            )));
128        }
129
130        let unk_id = resolve_unk(&tokenizer, opts.unk_token.as_deref());
131
132        let profile_id = profile_id(manifest, matrix.dim);
133        if !w_known {
134            // Only a load that got this far — digest, shape, vocabulary, every value finite —
135            // may let the next one skip those checks.
136            record_verified(&paths.weights, &manifest.weights_blake3);
137        }
138
139        Ok(Self {
140            tokenizer,
141            matrix,
142            unk_id,
143            max_tokens: opts.max_tokens,
144            manifest: manifest.clone(),
145            profile_id,
146        })
147    }
148
149    /// What a loaded model would report as its profile and dimension, without loading it.
150    ///
151    /// For `status`, `doctor` and a `scan` with nothing to embed (2026-09-25: each of them
152    /// parsed the tokenizer and read the matrix for these two values, 2.6 s). Both files are
153    /// verified against the manifest exactly as `load` does, the weights through the same
154    /// staleness record; only the tokenizer is not parsed and the matrix not read beyond its
155    /// header. A tokenizer that does not parse is therefore found by the first `recall`, not
156    /// here.
157    pub fn describe(paths: &ModelPaths, manifest: &ArtefactManifest) -> Result<Description> {
158        manifest.validate()?;
159        read_verified(&paths.tokenizer, &manifest.tokenizer_blake3, "tokenizer")?;
160        let (w_map, _) = map_verified(&paths.weights, &manifest.weights_blake3, "weights")?;
161        let (_, dim) = tensor_shape(&w_map)?;
162        Ok(Description {
163            profile_id: profile_id(manifest, dim),
164            dim,
165        })
166    }
167
168    /// Embeds one input. Errors only on a tokenizer failure, which no ordinary text causes.
169    pub fn embed_one(&self, text: &str) -> Result<Embedding> {
170        let enc = self
171            .tokenizer
172            .encode_fast(text, false)
173            .map_err(|e| Error::Embed(format!("tokenizer failed on input: {e}")))?;
174        let mut ids = enc.get_ids();
175        if let Some(max) = self.max_tokens
176            && ids.len() > max
177        {
178            ids = &ids[..max];
179        }
180        Ok(self.pool_ids(ids))
181    }
182
183    /// Embeds a batch, in input order, splitting large batches across rayon's pool.
184    pub fn embed_all(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
185        self.embed_all_impl(texts, texts.len() >= PARALLEL_THRESHOLD)
186    }
187
188    pub(crate) fn embed_all_impl(&self, texts: &[&str], parallel: bool) -> Result<Vec<Embedding>> {
189        if parallel {
190            texts.par_iter().map(|t| self.embed_one(t)).collect()
191        } else {
192            texts.iter().map(|t| self.embed_one(t)).collect()
193        }
194    }
195
196    fn pool_ids(&self, ids: &[u32]) -> Embedding {
197        let dim = self.matrix.dim;
198        let rows = self.matrix.rows;
199        let data: &[f32] = &self.matrix.data;
200        let mut vector = vec![0.0f32; dim];
201        let known_rows = ids.iter().filter_map(|&id| {
202            let i = id as usize;
203            // An id at or past the last row cannot happen after the load-time vocab check,
204            // but an out-of-bounds slice would panic; treat it as unknown instead.
205            if Some(id) == self.unk_id || i >= rows {
206                None
207            } else {
208                Some(&data[i * dim..(i + 1) * dim])
209            }
210        });
211        let tokens_known = pool::mean_pool_normalise(&mut vector, known_rows);
212        Embedding {
213            vector,
214            tokens_seen: ids.len(),
215            tokens_known,
216        }
217    }
218
219    pub fn vocab_rows(&self) -> usize {
220        self.matrix.rows
221    }
222
223    pub fn unk_id(&self) -> Option<u32> {
224        self.unk_id
225    }
226
227    pub fn manifest(&self) -> &ArtefactManifest {
228        &self.manifest
229    }
230
231    pub fn info(&self) -> ModelInfo {
232        ModelInfo {
233            profile_id: self.profile_id.clone(),
234            dim: self.matrix.dim,
235            vocab_rows: self.matrix.rows,
236            weights_dtype: self.matrix.dtype,
237            pooling: POOLING,
238            weights_blake3: self.manifest.weights_blake3.clone(),
239            tokenizer_blake3: self.manifest.tokenizer_blake3.clone(),
240            unk_id: self.unk_id,
241            max_tokens: self.max_tokens,
242        }
243    }
244}
245
246/// Finds the id of the unknown token so it can be skipped in pooling (model2vec drops it
247/// rather than averaging a meaningless row into every vector that has a typo).
248fn resolve_unk(tokenizer: &Tokenizer, override_name: Option<&str>) -> Option<u32> {
249    if let Some(name) = override_name {
250        return tokenizer.token_to_id(name);
251    }
252    let from_model: Option<String> = match tokenizer.get_model() {
253        ModelWrapper::WordPiece(m) => Some(m.unk_token.clone()),
254        ModelWrapper::WordLevel(m) => Some(m.unk_token.clone()),
255        ModelWrapper::BPE(m) => m.unk_token.clone(),
256        // Unigram keeps its unk id private; fall through to the name probe.
257        ModelWrapper::Unigram(_) => None,
258    };
259    if let Some(id) = from_model.and_then(|n| tokenizer.token_to_id(&n)) {
260        return Some(id);
261    }
262    UNK_CANDIDATES.iter().find_map(|n| tokenizer.token_to_id(n))
263}
264
265impl Embedder for StaticEmbedder {
266    fn dim(&self) -> usize {
267        self.matrix.dim
268    }
269
270    fn profile_id(&self) -> &str {
271        &self.profile_id
272    }
273
274    fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
275        Ok(self
276            .embed_all(texts)?
277            .into_iter()
278            .map(|e| e.vector)
279            .collect())
280    }
281}
282
283/// A model's identity without the model: see [`StaticEmbedder::describe`].
284#[derive(Debug, Clone, PartialEq, Eq)]
285pub struct Description {
286    pub profile_id: String,
287    pub dim: usize,
288}
289
290/// Everything that changes the meaning of a stored vector goes in here: both file digests
291/// (weights and tokenizer), the dimension and the pooling name.
292fn profile_id(manifest: &ArtefactManifest, dim: usize) -> String {
293    let mut h = blake3::Hasher::new();
294    h.update(manifest.weights_blake3.as_bytes());
295    h.update(manifest.tokenizer_blake3.as_bytes());
296    let short = &h.finalize().to_hex()[..16];
297    format!("m2v-{POOLING}-d{dim}-{short}")
298}