1use 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
13const PARALLEL_THRESHOLD: usize = 64;
17
18const UNK_CANDIDATES: [&str; 4] = ["[UNK]", "<unk>", "<|unk|>", "[unk]"];
20
21#[derive(Debug, Clone, Default)]
23pub struct LoadOptions {
24 pub max_tokens: Option<usize>,
28 pub unk_token: Option<String>,
30}
31
32#[derive(Debug, Clone, PartialEq)]
34pub struct Embedding {
35 pub vector: Vec<f32>,
36 pub tokens_seen: usize,
38 pub tokens_known: usize,
40}
41
42impl Embedding {
43 pub fn is_empty(&self) -> bool {
45 self.tokens_known == 0
46 }
47}
48
49#[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
64pub 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 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 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 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 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 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 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
230fn 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 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}