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, read_verified};
5use crate::pool;
6use crate::weights::{Matrix, load_matrix};
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_bytes = read_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(&w_bytes)?;
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        // Everything that changes the meaning of a stored vector goes in here: both file
133        // digests (weights and tokenizer), the dimension and the pooling name.
134        let profile_id = {
135            let mut h = blake3::Hasher::new();
136            h.update(manifest.weights_blake3.as_bytes());
137            h.update(manifest.tokenizer_blake3.as_bytes());
138            let short = &h.finalize().to_hex()[..16];
139            format!("m2v-{POOLING}-d{}-{short}", matrix.dim)
140        };
141
142        Ok(Self {
143            tokenizer,
144            matrix,
145            unk_id,
146            max_tokens: opts.max_tokens,
147            manifest: manifest.clone(),
148            profile_id,
149        })
150    }
151
152    /// Embeds one input. Errors only on a tokenizer failure, which no ordinary text causes.
153    pub fn embed_one(&self, text: &str) -> Result<Embedding> {
154        let enc = self
155            .tokenizer
156            .encode_fast(text, false)
157            .map_err(|e| Error::Embed(format!("tokenizer failed on input: {e}")))?;
158        let mut ids = enc.get_ids();
159        if let Some(max) = self.max_tokens
160            && ids.len() > max
161        {
162            ids = &ids[..max];
163        }
164        Ok(self.pool_ids(ids))
165    }
166
167    /// Embeds a batch, in input order, splitting large batches across rayon's pool.
168    pub fn embed_all(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
169        self.embed_all_impl(texts, texts.len() >= PARALLEL_THRESHOLD)
170    }
171
172    pub(crate) fn embed_all_impl(&self, texts: &[&str], parallel: bool) -> Result<Vec<Embedding>> {
173        if parallel {
174            texts.par_iter().map(|t| self.embed_one(t)).collect()
175        } else {
176            texts.iter().map(|t| self.embed_one(t)).collect()
177        }
178    }
179
180    fn pool_ids(&self, ids: &[u32]) -> Embedding {
181        let dim = self.matrix.dim;
182        let rows = self.matrix.rows;
183        let data: &[f32] = &self.matrix.data;
184        let mut vector = vec![0.0f32; dim];
185        let known_rows = ids.iter().filter_map(|&id| {
186            let i = id as usize;
187            // An id at or past the last row cannot happen after the load-time vocab check,
188            // but an out-of-bounds slice would panic; treat it as unknown instead.
189            if Some(id) == self.unk_id || i >= rows {
190                None
191            } else {
192                Some(&data[i * dim..(i + 1) * dim])
193            }
194        });
195        let tokens_known = pool::mean_pool_normalise(&mut vector, known_rows);
196        Embedding {
197            vector,
198            tokens_seen: ids.len(),
199            tokens_known,
200        }
201    }
202
203    pub fn vocab_rows(&self) -> usize {
204        self.matrix.rows
205    }
206
207    pub fn unk_id(&self) -> Option<u32> {
208        self.unk_id
209    }
210
211    pub fn manifest(&self) -> &ArtefactManifest {
212        &self.manifest
213    }
214
215    pub fn info(&self) -> ModelInfo {
216        ModelInfo {
217            profile_id: self.profile_id.clone(),
218            dim: self.matrix.dim,
219            vocab_rows: self.matrix.rows,
220            weights_dtype: self.matrix.dtype,
221            pooling: POOLING,
222            weights_blake3: self.manifest.weights_blake3.clone(),
223            tokenizer_blake3: self.manifest.tokenizer_blake3.clone(),
224            unk_id: self.unk_id,
225            max_tokens: self.max_tokens,
226        }
227    }
228}
229
230/// Finds the id of the unknown token so it can be skipped in pooling (model2vec drops it
231/// rather than averaging a meaningless row into every vector that has a typo).
232fn resolve_unk(tokenizer: &Tokenizer, override_name: Option<&str>) -> Option<u32> {
233    if let Some(name) = override_name {
234        return tokenizer.token_to_id(name);
235    }
236    let from_model: Option<String> = match tokenizer.get_model() {
237        ModelWrapper::WordPiece(m) => Some(m.unk_token.clone()),
238        ModelWrapper::WordLevel(m) => Some(m.unk_token.clone()),
239        ModelWrapper::BPE(m) => m.unk_token.clone(),
240        // Unigram keeps its unk id private; fall through to the name probe.
241        ModelWrapper::Unigram(_) => None,
242    };
243    if let Some(id) = from_model.and_then(|n| tokenizer.token_to_id(&n)) {
244        return Some(id);
245    }
246    UNK_CANDIDATES.iter().find_map(|n| tokenizer.token_to_id(n))
247}
248
249impl Embedder for StaticEmbedder {
250    fn dim(&self) -> usize {
251        self.matrix.dim
252    }
253
254    fn profile_id(&self) -> &str {
255        &self.profile_id
256    }
257
258    fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
259        Ok(self
260            .embed_all(texts)?
261            .into_iter()
262            .map(|e| e.vector)
263            .collect())
264    }
265}