lean_ctx/core/embeddings/
tokenizer.rs1use std::collections::HashMap;
13use std::path::Path;
14
15pub struct WordPieceTokenizer {
16 vocab: HashMap<String, i32>,
17 cls_id: i32,
18 sep_id: i32,
19 pad_id: i32,
20 unk_id: i32,
21 max_word_chars: usize,
22}
23
24#[derive(Debug, Clone)]
25pub struct TokenizedInput {
26 pub input_ids: Vec<i32>,
27 pub attention_mask: Vec<i32>,
28 pub token_type_ids: Vec<i32>,
29}
30
31impl TokenizedInput {
32 pub fn pad_to(&mut self, target_len: usize, pad_id: i32) {
34 while self.input_ids.len() < target_len {
35 self.input_ids.push(pad_id);
36 self.attention_mask.push(0);
37 self.token_type_ids.push(0);
38 }
39 }
40}
41
42impl WordPieceTokenizer {
43 pub fn from_file(path: &Path) -> anyhow::Result<Self> {
45 let content = std::fs::read_to_string(path)
46 .map_err(|e| anyhow::anyhow!("Failed to read vocab file {}: {}", path.display(), e))?;
47 Self::from_vocab_str(&content)
48 }
49
50 pub fn from_vocab_str(vocab_str: &str) -> anyhow::Result<Self> {
52 let vocab: HashMap<String, i32> = vocab_str
53 .lines()
54 .enumerate()
55 .map(|(i, line)| (line.to_string(), i as i32))
56 .collect();
57
58 let cls_id = *vocab
59 .get("[CLS]")
60 .ok_or_else(|| anyhow::anyhow!("Vocabulary missing [CLS] token"))?;
61 let sep_id = *vocab
62 .get("[SEP]")
63 .ok_or_else(|| anyhow::anyhow!("Vocabulary missing [SEP] token"))?;
64 let pad_id = *vocab
65 .get("[PAD]")
66 .ok_or_else(|| anyhow::anyhow!("Vocabulary missing [PAD] token"))?;
67 let unk_id = *vocab
68 .get("[UNK]")
69 .ok_or_else(|| anyhow::anyhow!("Vocabulary missing [UNK] token"))?;
70
71 Ok(Self {
72 vocab,
73 cls_id,
74 sep_id,
75 pad_id,
76 unk_id,
77 max_word_chars: 200,
78 })
79 }
80
81 pub fn encode(&self, text: &str, max_len: usize) -> TokenizedInput {
83 let words = self.pre_tokenize(text);
84 let mut ids = vec![self.cls_id];
85
86 for word in &words {
87 if ids.len() >= max_len - 1 {
88 break;
89 }
90 let subword_ids = self.wordpiece_encode(word);
91 for id in subword_ids {
92 if ids.len() >= max_len - 1 {
93 break;
94 }
95 ids.push(id);
96 }
97 }
98
99 ids.push(self.sep_id);
100
101 let len = ids.len();
102 TokenizedInput {
103 input_ids: ids,
104 attention_mask: vec![1; len],
105 token_type_ids: vec![0; len],
106 }
107 }
108
109 pub fn pad_id(&self) -> i32 {
110 self.pad_id
111 }
112
113 pub fn vocab_size(&self) -> usize {
114 self.vocab.len()
115 }
116
117 fn pre_tokenize(&self, text: &str) -> Vec<String> {
121 let mut words = Vec::new();
122 let mut current = String::new();
123
124 for ch in text.chars() {
125 if ch.is_whitespace() {
126 if !current.is_empty() {
127 words.extend(self.split_identifier(¤t));
128 current.clear();
129 }
130 } else if is_bert_punctuation(ch) {
131 if !current.is_empty() {
132 words.extend(self.split_identifier(¤t));
133 current.clear();
134 }
135 words.push(ch.to_string());
136 } else {
137 current.push(ch);
138 }
139 }
140 if !current.is_empty() {
141 words.extend(self.split_identifier(¤t));
142 }
143
144 words.iter().map(|w| w.to_lowercase()).collect()
145 }
146
147 fn split_identifier(&self, word: &str) -> Vec<String> {
150 let lower = word.to_lowercase();
151 if self.vocab.contains_key(&lower) {
152 return vec![word.to_string()];
153 }
154
155 let mut parts = Vec::new();
156 let mut current = String::new();
157 let chars: Vec<char> = word.chars().collect();
158
159 for (i, &ch) in chars.iter().enumerate() {
160 if ch == '_' || ch == '-' {
161 if !current.is_empty() {
162 parts.push(current.clone());
163 current.clear();
164 }
165 } else if i > 0 && ch.is_ascii_uppercase() && chars[i - 1].is_ascii_lowercase() {
166 if !current.is_empty() {
167 parts.push(current.clone());
168 current.clear();
169 }
170 current.push(ch);
171 } else {
172 current.push(ch);
173 }
174 }
175 if !current.is_empty() {
176 parts.push(current);
177 }
178
179 if parts.is_empty() {
180 vec![word.to_string()]
181 } else {
182 parts
183 }
184 }
185
186 fn wordpiece_encode(&self, word: &str) -> Vec<i32> {
188 if word.chars().count() > self.max_word_chars {
189 return vec![self.unk_id];
190 }
191
192 let chars: Vec<char> = word.chars().collect();
193 let mut tokens = Vec::new();
194 let mut start = 0;
195
196 while start < chars.len() {
197 let mut end = chars.len();
198 let mut matched = false;
199
200 while start < end {
201 let substr: String = chars[start..end].iter().collect();
202 let candidate = if start > 0 {
203 format!("##{substr}")
204 } else {
205 substr
206 };
207
208 if let Some(&id) = self.vocab.get(&candidate) {
209 tokens.push(id);
210 matched = true;
211 start = end;
212 break;
213 }
214 end -= 1;
215 }
216
217 if !matched {
218 tokens.push(self.unk_id);
219 start += 1;
220 }
221 }
222
223 tokens
224 }
225}
226
227fn is_bert_punctuation(ch: char) -> bool {
229 if ch.is_ascii() {
230 matches!(
231 ch,
232 '!' | '"'
233 | '#'
234 | '$'
235 | '%'
236 | '&'
237 | '\''
238 | '('
239 | ')'
240 | '*'
241 | '+'
242 | ','
243 | '-'
244 | '.'
245 | '/'
246 | ':'
247 | ';'
248 | '<'
249 | '='
250 | '>'
251 | '?'
252 | '@'
253 | '['
254 | '\\'
255 | ']'
256 | '^'
257 | '_'
258 | '`'
259 | '{'
260 | '|'
261 | '}'
262 | '~'
263 )
264 } else {
265 ch.is_ascii_punctuation()
266 }
267}
268
269pub struct HfTokenizerWrapper {
275 inner: WordPieceTokenizer,
276}
277
278impl HfTokenizerWrapper {
279 pub fn from_file(path: &Path) -> anyhow::Result<Self> {
281 let content = std::fs::read_to_string(path).map_err(|e| {
282 anyhow::anyhow!("Failed to read tokenizer.json {}: {}", path.display(), e)
283 })?;
284 Self::from_json(&content)
285 }
286
287 fn from_json(json_str: &str) -> anyhow::Result<Self> {
288 let parsed: serde_json::Value = serde_json::from_str(json_str)
289 .map_err(|e| anyhow::anyhow!("Invalid tokenizer.json: {e}"))?;
290
291 let vocab_obj = parsed
292 .get("model")
293 .and_then(|m| m.get("vocab"))
294 .and_then(|v| v.as_object())
295 .ok_or_else(|| anyhow::anyhow!("tokenizer.json missing model.vocab object"))?;
296
297 let mut vocab_lines: Vec<(String, i32)> = vocab_obj
298 .iter()
299 .filter_map(|(token, id)| id.as_i64().map(|id| (token.clone(), id as i32)))
300 .collect();
301 vocab_lines.sort_by_key(|(_, id)| *id);
302
303 for (token, _) in &mut vocab_lines {
309 let mapped: &str = match token.as_str() {
310 "<s>" => "[CLS]",
311 "</s>" => "[SEP]",
312 "<pad>" => "[PAD]",
313 "<unk>" => "[UNK]",
314 "<mask>" => "[MASK]",
316 _ => continue,
317 };
318 *token = mapped.to_string();
319 }
320
321 let vocab_str: String = vocab_lines
322 .into_iter()
323 .map(|(token, _)| token)
324 .collect::<Vec<_>>()
325 .join("\n");
326
327 let inner = WordPieceTokenizer::from_vocab_str(&vocab_str)?;
328 Ok(Self { inner })
329 }
330
331 pub fn encode(&self, text: &str, max_len: usize) -> TokenizedInput {
332 self.inner.encode(text, max_len)
333 }
334}
335
336#[cfg(test)]
337mod tests {
338 use super::*;
339
340 fn test_vocab() -> WordPieceTokenizer {
341 let vocab = "[PAD]\n[UNK]\n[CLS]\n[SEP]\nhello\nworld\nfn\nvalidate\ntoken\n##s\n##ing\nauth\n##enticate\nuser\nhandle\nrequest\n##er\nprocess\ndata\n.\n,\n(\n)\n{";
342 WordPieceTokenizer::from_vocab_str(vocab).unwrap()
343 }
344
345 #[test]
346 fn encode_basic() {
347 let tok = test_vocab();
348 let input = tok.encode("hello world", 512);
349 assert_eq!(input.input_ids[0], tok.cls_id);
350 assert_eq!(*input.input_ids.last().unwrap(), tok.sep_id);
351 assert!(input.input_ids.len() >= 4); }
353
354 #[test]
355 fn encode_attention_mask() {
356 let tok = test_vocab();
357 let input = tok.encode("hello", 512);
358 assert!(input.attention_mask.iter().all(|&m| m == 1));
359 assert_eq!(input.attention_mask.len(), input.input_ids.len());
360 }
361
362 #[test]
363 fn encode_token_type_ids_are_zero() {
364 let tok = test_vocab();
365 let input = tok.encode("hello", 512);
366 assert!(input.token_type_ids.iter().all(|&t| t == 0));
367 }
368
369 #[test]
370 fn encode_respects_max_len() {
371 let tok = test_vocab();
372 let input = tok.encode("hello world hello world hello world", 6);
373 assert!(input.input_ids.len() <= 6);
374 assert_eq!(input.input_ids[0], tok.cls_id);
375 assert_eq!(*input.input_ids.last().unwrap(), tok.sep_id);
376 }
377
378 #[test]
379 fn wordpiece_subwords() {
380 let tok = test_vocab();
381 let ids = tok.wordpiece_encode("tokens");
383 assert_eq!(ids.len(), 2);
384 assert_eq!(ids[0], *tok.vocab.get("token").unwrap());
385 assert_eq!(ids[1], *tok.vocab.get("##s").unwrap());
386 }
387
388 #[test]
389 fn wordpiece_unknown() {
390 let tok = test_vocab();
391 let ids = tok.wordpiece_encode("xyzzyplugh");
392 assert!(ids.contains(&tok.unk_id));
393 }
394
395 #[test]
396 fn pre_tokenize_camel_case() {
397 let tok = test_vocab();
398 let words = tok.pre_tokenize("handleRequest");
399 assert!(words.contains(&"handle".to_string()));
400 assert!(words.contains(&"request".to_string()));
401 }
402
403 #[test]
404 fn pre_tokenize_snake_case() {
405 let tok = test_vocab();
406 let words = tok.pre_tokenize("validate_token");
407 assert!(words.contains(&"validate".to_string()));
408 assert!(words.contains(&"token".to_string()));
409 }
410
411 #[test]
412 fn pre_tokenize_punctuation() {
413 let tok = test_vocab();
414 let words = tok.pre_tokenize("fn(x)");
415 assert!(words.contains(&"fn".to_string()));
416 assert!(words.contains(&"(".to_string()));
417 assert!(words.contains(&")".to_string()));
418 }
419
420 #[test]
421 fn pad_to_extends() {
422 let tok = test_vocab();
423 let mut input = tok.encode("hello", 512);
424 let original_len = input.input_ids.len();
425 input.pad_to(10, tok.pad_id);
426 assert_eq!(input.input_ids.len(), 10);
427 assert_eq!(input.attention_mask[original_len], 0);
428 }
429
430 #[test]
431 fn vocab_size() {
432 let tok = test_vocab();
433 assert_eq!(tok.vocab_size(), 24);
434 }
435
436 #[test]
437 fn empty_input() {
438 let tok = test_vocab();
439 let input = tok.encode("", 512);
440 assert_eq!(input.input_ids.len(), 2); }
442
443 #[test]
444 fn bert_punctuation_detection() {
445 assert!(is_bert_punctuation('.'));
446 assert!(is_bert_punctuation('('));
447 assert!(is_bert_punctuation('{'));
448 assert!(!is_bert_punctuation('a'));
449 assert!(!is_bert_punctuation('0'));
450 }
451
452 #[test]
453 fn hf_tokenizer_remaps_bpe_special_tokens() {
454 let json = r#"{
456 "version": "1.0",
457 "model": {
458 "type": "BPE",
459 "vocab": {
460 "<s>": 0, "<pad>": 1, "</s>": 2, "<unk>": 3,
461 "hello": 4, "world": 5, "fn": 6
462 }
463 }
464 }"#;
465 let tok = HfTokenizerWrapper::from_json(json).unwrap();
466
467 let input = tok.encode("hello world", 512);
470 assert_eq!(
471 input.input_ids[0], 0,
472 "first token should be [CLS] (remapped from <s>)"
473 );
474 assert_eq!(
475 *input.input_ids.last().unwrap(),
476 2,
477 "last token should be [SEP] (remapped from </s>)"
478 );
479 assert_eq!(input.input_ids.len(), 4); }
481
482 #[test]
483 fn hf_tokenizer_from_json() {
484 let json = r#"{
485 "version": "1.0",
486 "model": {
487 "type": "WordPiece",
488 "vocab": {
489 "[PAD]": 0, "[UNK]": 1, "[CLS]": 2, "[SEP]": 3,
490 "hello": 4, "world": 5, "fn": 6
491 }
492 }
493 }"#;
494 let tok = HfTokenizerWrapper::from_json(json).unwrap();
495 let input = tok.encode("hello world", 512);
496 assert_eq!(input.input_ids[0], 2); assert_eq!(*input.input_ids.last().unwrap(), 3); assert!(input.input_ids.len() >= 4);
499 }
500
501 #[test]
502 fn hf_tokenizer_invalid_json() {
503 assert!(HfTokenizerWrapper::from_json("not json").is_err());
504 }
505
506 #[test]
507 fn hf_tokenizer_missing_vocab() {
508 let json = r#"{"model": {"type": "WordPiece"}}"#;
509 assert!(HfTokenizerWrapper::from_json(json).is_err());
510 }
511}