1use 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
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_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 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 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 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 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 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 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
246fn 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 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#[derive(Debug, Clone, PartialEq, Eq)]
285pub struct Description {
286 pub profile_id: String,
287 pub dim: usize,
288}
289
290fn 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}