1use serde::Deserialize;
13use std::collections::{HashMap, HashSet};
14use std::path::Path;
15use unicode_normalization::UnicodeNormalization;
16
17const DEFAULT_SPLIT: &str = r"(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+";
20
21pub struct Tokenizer {
23 vocab: HashMap<String, u32>,
25 id_to_token: Vec<String>,
27 ranks: HashMap<(String, String), u32>,
29 added: Vec<(String, u32)>,
31 added_ids: HashSet<u32>,
33 special_ids: HashSet<u32>,
35 split_re: Option<fancy_regex::Regex>,
38 metaspace: bool,
41 nfc: bool,
44 byte_to_char: [char; 256],
46 char_to_byte: HashMap<char, u8>,
48 pub bos_token_id: Option<u32>,
50 pub eos_token_id: Option<u32>,
51 pub pad_token_id: Option<u32>,
52 pub im_start_id: Option<u32>,
54 pub im_end_id: Option<u32>,
55 pub chat_template: Option<String>,
58 pub extra_eos: HashSet<u32>,
60 pub add_bos: bool,
62}
63
64impl std::fmt::Debug for Tokenizer {
65 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
66 f.debug_struct("Tokenizer")
67 .field("vocab", &self.vocab.len())
68 .field("merges", &self.ranks.len())
69 .field("added", &self.added.len())
70 .finish()
71 }
72}
73
74fn bytes_to_unicode() -> ([char; 256], HashMap<char, u8>) {
77 let mut b2c = ['\0'; 256];
78 let mut c2b = HashMap::with_capacity(256);
79 let mut n = 0u32;
80 for b in 0..=255u16 {
81 let printable = (0x21..=0x7E).contains(&b)
82 || (0xA1..=0xAC).contains(&b)
83 || (0xAE..=0xFF).contains(&b);
84 let c = if printable {
85 char::from_u32(b as u32).unwrap()
86 } else {
87 let c = char::from_u32(256 + n).unwrap();
88 n += 1;
89 c
90 };
91 b2c[b as usize] = c;
92 c2b.insert(c, b as u8);
93 }
94 (b2c, c2b)
95}
96
97#[derive(Deserialize)]
99struct HfTokenizerJson {
100 model: HfModel,
101 #[serde(default)]
102 added_tokens: Vec<HfAddedToken>,
103 #[serde(default)]
104 pre_tokenizer: Option<serde_json::Value>,
105 #[serde(default)]
106 normalizer: Option<serde_json::Value>,
107 #[serde(default)]
108 post_processor: Option<serde_json::Value>,
109}
110
111#[derive(Deserialize)]
112struct HfModel {
113 vocab: HashMap<String, u32>,
114 #[serde(default)]
115 merges: Vec<HfMerge>,
116 #[serde(default)]
117 byte_fallback: bool,
118}
119
120#[derive(Deserialize)]
123#[serde(untagged)]
124enum HfMerge {
125 Pair([String; 2]),
126 Text(String),
127}
128
129#[derive(Deserialize)]
130struct HfAddedToken {
131 id: u32,
132 content: String,
133 special: bool,
134}
135
136fn find_split_pattern(pt: &serde_json::Value) -> Option<String> {
139 if pt.get("type").and_then(|t| t.as_str()) == Some("Split") {
140 return pt
141 .get("pattern")
142 .and_then(|p| p.get("Regex"))
143 .and_then(|r| r.as_str())
144 .map(String::from);
145 }
146 if let Some(list) = pt.get("pretokenizers").and_then(|l| l.as_array()) {
147 return list.iter().find_map(find_split_pattern);
148 }
149 None
150}
151
152impl Tokenizer {
153 pub fn from_file(path: impl AsRef<Path>) -> Result<Self, TokenizerError> {
155 let data = std::fs::read_to_string(path.as_ref())
156 .map_err(|e| TokenizerError::Io(e.to_string()))?;
157 Self::from_json(&data)
158 }
159
160 pub fn from_bytes(bytes: &[u8]) -> Result<Self, TokenizerError> {
162 let s = std::str::from_utf8(bytes)
163 .map_err(|e| TokenizerError::Parse(format!("vocab is not UTF-8: {e}")))?;
164 Self::from_json(s)
165 }
166
167 pub fn from_json(json: &str) -> Result<Self, TokenizerError> {
169 let hf: HfTokenizerJson =
170 serde_json::from_str(json).map_err(|e| TokenizerError::Parse(e.to_string()))?;
171
172 let mut vocab = hf.model.vocab;
173 let mut ranks = HashMap::new();
174 for (rank, m) in hf.model.merges.into_iter().enumerate() {
175 let (a, b) = match m {
176 HfMerge::Pair([a, b]) => (a, b),
177 HfMerge::Text(s) => {
178 let mut it = s.splitn(2, ' ');
179 match (it.next(), it.next()) {
180 (Some(a), Some(b)) => (a.to_string(), b.to_string()),
181 _ => continue,
182 }
183 }
184 };
185 ranks.insert((a, b), rank as u32);
186 }
187
188 let add_bos = hf
193 .post_processor
194 .as_ref()
195 .map(|p| p.to_string().contains("\"<s>\""))
196 .unwrap_or(false);
197 let nfc = hf
198 .normalizer
199 .as_ref()
200 .map(|n| n.to_string().contains("NFC"))
201 .unwrap_or(false);
202 let metaspace = hf.model.byte_fallback
203 || hf
204 .normalizer
205 .as_ref()
206 .map(|n| n.to_string().contains("\u{2581}") || n.to_string().contains("▁"))
207 .unwrap_or(false);
208 let split_re = if metaspace {
209 None
210 } else {
211 let pattern = hf
212 .pre_tokenizer
213 .as_ref()
214 .and_then(find_split_pattern)
215 .unwrap_or_else(|| DEFAULT_SPLIT.to_string());
216 Some(fancy_regex::Regex::new(&pattern).map_err(|e| {
217 TokenizerError::Parse(format!("pre-tokenizer regex: {e}"))
218 })?)
219 };
220
221 let mut bos_token_id = None;
223 let mut eos_token_id = None;
224 let mut pad_token_id = None;
225 let mut im_start_id = None;
226 let mut im_end_id = None;
227 let mut special_ids = HashSet::new();
228 let mut added_ids = HashSet::new();
229 let mut added = Vec::new();
230
231 for at in &hf.added_tokens {
232 vocab.insert(at.content.clone(), at.id);
233 added.push((at.content.clone(), at.id));
234 added_ids.insert(at.id);
235 if at.special {
236 special_ids.insert(at.id);
237 }
238 match at.content.as_str() {
239 "<|endoftext|>" | "</s>" => eos_token_id = Some(at.id),
240 "<|im_start|>" => im_start_id = Some(at.id),
241 "<|im_end|>" => im_end_id = Some(at.id),
242 "<s>" => bos_token_id = Some(at.id),
243 "<pad>" => pad_token_id = Some(at.id),
244 _ => {}
245 }
246 }
247 added.sort_by_key(|(c, _)| std::cmp::Reverse(c.len()));
248
249 let max_id = vocab.values().copied().max().unwrap_or(0) as usize;
251 let mut id_to_token = vec![String::new(); max_id + 1];
252 for (token, &id) in &vocab {
253 if (id as usize) < id_to_token.len() {
254 id_to_token[id as usize] = token.clone();
255 }
256 }
257
258 let (byte_to_char, char_to_byte) = bytes_to_unicode();
259
260 tracing::info!(
261 "Tokenizer loaded: {} vocab, {} merges, {} added, eos={:?}",
262 vocab.len(),
263 ranks.len(),
264 added.len(),
265 eos_token_id
266 );
267
268 Ok(Self {
269 vocab,
270 id_to_token,
271 ranks,
272 added,
273 added_ids,
274 special_ids,
275 split_re,
276 metaspace,
277 nfc,
278 byte_to_char,
279 char_to_byte,
280 bos_token_id,
281 eos_token_id,
282 pad_token_id,
283 im_start_id,
284 im_end_id,
285 chat_template: None,
286 extra_eos: HashSet::new(),
287 add_bos,
288 })
289 }
290
291 pub fn byte_level() -> Self {
293 let mut vocab = HashMap::new();
294 let mut id_to_token = Vec::with_capacity(256);
295 for i in 0..256u32 {
296 let tok = format!("<0x{:02X}>", i);
297 vocab.insert(tok.clone(), i);
298 id_to_token.push(tok);
299 }
300 let (byte_to_char, char_to_byte) = bytes_to_unicode();
301 Self {
302 vocab,
303 id_to_token,
304 ranks: HashMap::new(),
305 added: Vec::new(),
306 added_ids: HashSet::new(),
307 special_ids: HashSet::new(),
308 split_re: None,
309 metaspace: false,
310 nfc: false,
311 byte_to_char,
312 char_to_byte,
313 bos_token_id: None,
314 eos_token_id: None,
315 pad_token_id: None,
316 im_start_id: None,
317 im_end_id: None,
318 chat_template: None,
319 extra_eos: HashSet::new(),
320 add_bos: false,
321 }
322 }
323
324 pub fn encode(&self, text: &str) -> Vec<u32> {
326 let mut ids = Vec::new();
327 let mut rest = text;
329 'outer: while !rest.is_empty() {
330 let mut best: Option<(usize, usize, u32)> = None; for (content, id) in &self.added {
332 if let Some(pos) = rest.find(content.as_str()) {
333 let better = match best {
334 None => true,
335 Some((bp, bl, _)) => pos < bp || (pos == bp && content.len() > bl),
336 };
337 if better {
338 best = Some((pos, content.len(), *id));
339 }
340 if pos == 0 {
341 break; }
343 }
344 }
345 match best {
346 Some((pos, len, id)) => {
347 self.encode_segment(&rest[..pos], &mut ids);
348 ids.push(id);
349 rest = &rest[pos + len..];
350 }
351 None => {
352 self.encode_segment(rest, &mut ids);
353 break 'outer;
354 }
355 }
356 }
357 ids
358 }
359
360 fn encode_segment(&self, segment: &str, out: &mut Vec<u32>) {
362 if segment.is_empty() {
363 return;
364 }
365 let norm: String = if self.nfc {
366 segment.nfc().collect()
367 } else {
368 segment.to_string()
369 };
370 if self.metaspace {
371 let sp = format!("\u{2581}{}", norm).replace(' ', "\u{2581}");
374 self.bpe_piece_sp(&sp, out);
375 return;
376 }
377 match &self.split_re {
378 Some(re) => {
379 let mut last = 0;
380 for m in re.find_iter(&norm) {
381 let m = match m {
382 Ok(m) => m,
383 Err(e) => {
384 tracing::error!("pre-tokenizer regex failed: {e}");
385 break;
386 }
387 };
388 if m.start() > last {
389 self.bpe_piece(&norm[last..m.start()], out);
391 }
392 self.bpe_piece(m.as_str(), out);
393 last = m.end();
394 }
395 if last < norm.len() {
396 self.bpe_piece(&norm[last..], out);
397 }
398 }
399 None => {
400 for b in norm.bytes() {
402 let tok = format!("<0x{:02X}>", b);
403 if let Some(&id) = self.vocab.get(&tok) {
404 out.push(id);
405 }
406 }
407 }
408 }
409 }
410
411 fn bpe_piece_sp(&self, piece: &str, out: &mut Vec<u32>) {
414 if piece.is_empty() {
415 return;
416 }
417 let mut sym: Vec<String> = piece.chars().map(|c| c.to_string()).collect();
418 loop {
419 let mut best: Option<(u32, usize)> = None;
420 for i in 0..sym.len().saturating_sub(1) {
421 if let Some(&r) = self.ranks.get(&(sym[i].clone(), sym[i + 1].clone())) {
422 if best.map(|(br, _)| r < br).unwrap_or(true) {
423 best = Some((r, i));
424 }
425 }
426 }
427 let Some((_, i)) = best else { break };
428 let merged = format!("{}{}", sym[i], sym[i + 1]);
429 let (left, right) = (sym[i].clone(), sym[i + 1].clone());
430 let mut j = 0;
431 while j + 1 < sym.len() {
432 if sym[j] == left && sym[j + 1] == right {
433 sym[j] = merged.clone();
434 sym.remove(j + 1);
435 }
436 j += 1;
437 }
438 }
439 for t in &sym {
440 if let Some(&id) = self.vocab.get(t) {
441 out.push(id);
442 } else {
443 let mut ok = true;
444 for byte in t.bytes() {
445 let tok = format!("<0x{:02X}>", byte);
446 match self.vocab.get(&tok) {
447 Some(&id) => out.push(id),
448 None => {
449 ok = false;
450 break;
451 }
452 }
453 }
454 if !ok {
455 tracing::error!("tokenizer: no id for SP symbol {t:?} — dropped");
456 }
457 }
458 }
459 }
460
461 fn bpe_piece(&self, piece: &str, out: &mut Vec<u32>) {
463 if piece.is_empty() {
464 return;
465 }
466 let mapped: Vec<String> = piece
467 .bytes()
468 .map(|b| self.byte_to_char[b as usize].to_string())
469 .collect();
470 let mut sym = mapped;
471
472 loop {
474 let mut best: Option<(u32, usize)> = None;
475 for i in 0..sym.len().saturating_sub(1) {
476 if let Some(&r) = self
477 .ranks
478 .get(&(sym[i].clone(), sym[i + 1].clone()))
479 {
480 if best.map(|(br, _)| r < br).unwrap_or(true) {
481 best = Some((r, i));
482 }
483 }
484 }
485 let Some((_, i)) = best else { break };
486 let merged = format!("{}{}", sym[i], sym[i + 1]);
487 let (left, right) = (sym[i].clone(), sym[i + 1].clone());
489 let mut j = 0;
490 while j + 1 < sym.len() {
491 if sym[j] == left && sym[j + 1] == right {
492 sym[j] = merged.clone();
493 sym.remove(j + 1);
494 }
495 j += 1;
496 }
497 }
498
499 for s in &sym {
500 if let Some(&id) = self.vocab.get(s) {
501 out.push(id);
502 } else {
503 let mut ok = true;
505 for ch in s.chars() {
506 let Some(&b) = self.char_to_byte.get(&ch) else {
507 ok = false;
508 break;
509 };
510 let tok = format!("<0x{:02X}>", b);
511 if let Some(&id) = self.vocab.get(&tok) {
512 out.push(id);
513 } else {
514 ok = false;
515 break;
516 }
517 }
518 if !ok {
519 tracing::error!("tokenizer: no id for symbol {s:?} — dropped");
520 }
521 }
522 }
523 }
524
525 pub fn decode(&self, ids: &[u32]) -> String {
528 let mut bytes: Vec<u8> = Vec::new();
529 for &id in ids {
530 if self.special_ids.contains(&id) {
531 continue;
532 }
533 let idx = id as usize;
534 if idx >= self.id_to_token.len() {
535 continue;
536 }
537 let tok = &self.id_to_token[idx];
538 if self.added_ids.contains(&id) {
539 bytes.extend_from_slice(tok.as_bytes());
540 continue;
541 }
542 if tok.starts_with("<0x") && tok.ends_with('>') && tok.len() == 6 {
544 if let Ok(b) = u8::from_str_radix(&tok[3..5], 16) {
545 bytes.push(b);
546 continue;
547 }
548 }
549 if self.metaspace {
550 for ch in tok.chars() {
552 if ch == '\u{2581}' {
553 bytes.push(b' ');
554 } else {
555 let mut buf = [0u8; 4];
556 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
557 }
558 }
559 continue;
560 }
561 for ch in tok.chars() {
562 match self.char_to_byte.get(&ch) {
563 Some(&b) => bytes.push(b),
564 None => {
567 let mut buf = [0u8; 4];
568 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
569 }
570 }
571 }
572 }
573 let text = String::from_utf8_lossy(&bytes).into_owned();
574 if self.metaspace {
575 if let Some(stripped) = text.strip_prefix(' ') {
577 return stripped.to_string();
578 }
579 }
580 text
581 }
582
583 pub fn decode_token(&self, id: u32) -> String {
586 if self.special_ids.contains(&id) {
587 return String::new();
588 }
589 let idx = id as usize;
590 if idx >= self.id_to_token.len() {
591 return String::new();
592 }
593 let tok = &self.id_to_token[idx];
594 if self.added_ids.contains(&id) {
595 return tok.clone();
596 }
597 if tok.starts_with("<0x") && tok.ends_with('>') && tok.len() == 6 {
598 if let Ok(b) = u8::from_str_radix(&tok[3..5], 16) {
599 return String::from_utf8_lossy(&[b]).into_owned();
600 }
601 }
602 if self.metaspace {
603 return tok.replace('\u{2581}', " ");
604 }
605 let mut bytes = Vec::new();
606 for ch in tok.chars() {
607 match self.char_to_byte.get(&ch) {
608 Some(&b) => bytes.push(b),
609 None => {
610 let mut buf = [0u8; 4];
611 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
612 }
613 }
614 }
615 String::from_utf8_lossy(&bytes).into_owned()
616 }
617
618 pub fn apply_chat_template(&self, messages: &[(String, String)]) -> Vec<u32> {
622 if let Some(tpl) = &self.chat_template {
623 match self.render_template(tpl, messages) {
624 Ok(text) => return self.with_bos(self.encode(&text)),
625 Err(e) => {
626 tracing::error!("chat template render failed ({e}); ChatML fallback");
627 }
628 }
629 }
630 self.with_bos(self.chatml_fallback(messages))
631 }
632
633 pub fn with_bos(&self, mut ids: Vec<u32>) -> Vec<u32> {
635 if self.add_bos {
636 if let Some(b) = self.bos_token_id {
637 if ids.first() != Some(&b) {
638 ids.insert(0, b);
639 }
640 }
641 }
642 ids
643 }
644
645 pub fn render_chat(&self, messages: &[(String, String)]) -> Option<String> {
647 let tpl = self.chat_template.as_ref()?;
648 match self.render_template(tpl, messages) {
649 Ok(t) => Some(t),
650 Err(e) => {
651 tracing::error!("chat template render: {e:#}");
652 None
653 }
654 }
655 }
656
657 fn render_template(
658 &self,
659 tpl: &str,
660 messages: &[(String, String)],
661 ) -> Result<String, minijinja::Error> {
662 let mut env = minijinja::Environment::new();
663 env.set_trim_blocks(true);
664 env.set_lstrip_blocks(true);
665 env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
667 env.add_template("chat", tpl)?;
668 let msgs: Vec<minijinja::Value> = messages
669 .iter()
670 .map(|(role, content)| {
671 minijinja::context! { role => role, content => content }
672 })
673 .collect();
674 env.get_template("chat")?.render(minijinja::context! {
675 messages => msgs,
676 add_generation_prompt => true,
677 })
678 }
679
680 fn chatml_fallback(&self, messages: &[(String, String)]) -> Vec<u32> {
682 let mut tokens = Vec::new();
683
684 for (role, content) in messages {
685 if let Some(start_id) = self.im_start_id {
687 tokens.push(start_id);
688 }
689 tokens.extend(self.encode(&format!("{}\n{}", role, content)));
690 if let Some(end_id) = self.im_end_id {
691 tokens.push(end_id);
692 }
693 tokens.extend(self.encode("\n"));
694 }
695
696 if let Some(start_id) = self.im_start_id {
698 tokens.push(start_id);
699 }
700 tokens.extend(self.encode("assistant\n"));
701
702 tokens
703 }
704
705 pub fn vocab_size(&self) -> usize {
707 self.id_to_token.len()
708 }
709
710 pub fn is_eos(&self, id: u32) -> bool {
712 self.eos_token_id == Some(id)
713 || self.im_end_id == Some(id)
714 || self.extra_eos.contains(&id)
715 }
716}
717
718#[derive(Debug, thiserror::Error)]
719pub enum TokenizerError {
720 #[error("IO error: {0}")]
721 Io(String),
722 #[error("Parse error: {0}")]
723 Parse(String),
724}
725
726#[cfg(test)]
727mod tests {
728 use super::*;
729
730 #[test]
731 fn byte_unicode_bijection() {
732 let (b2c, c2b) = bytes_to_unicode();
733 for b in 0..=255u8 {
734 assert_eq!(c2b[&b2c[b as usize]], b);
735 }
736 assert_eq!(b2c[b' ' as usize], 'Ġ');
738 assert_eq!(b2c[b'\n' as usize], 'Ċ');
739 }
740
741 #[test]
742 fn byte_level_roundtrip_utf8() {
743 let tok = Tokenizer::byte_level();
744 let text = "hello 🌍 hi\n";
745 let ids = tok.encode(text);
746 assert_eq!(ids.len(), text.len()); assert_eq!(tok.decode(&ids), text);
748 }
749
750 fn mini_json() -> String {
753 let vocab: Vec<(&str, u32)> = vec![
755 ("h", 0),
756 ("e", 1),
757 ("l", 2),
758 ("o", 3),
759 ("Ġ", 4),
760 ("w", 5),
761 ("r", 6),
762 ("d", 7),
763 ("he", 8),
764 ("Ġw", 9),
765 ];
766 let vocab_json: String = vocab
767 .iter()
768 .map(|(t, i)| format!("\"{t}\": {i}"))
769 .collect::<Vec<_>>()
770 .join(", ");
771 format!(
772 r#"{{
773 "model": {{
774 "type": "BPE",
775 "vocab": {{ {vocab_json} }},
776 "merges": [["h", "e"], ["Ġ", "w"]]
777 }},
778 "added_tokens": [
779 {{"id": 10, "content": "<|eot|>", "special": true}}
780 ]
781 }}"#
782 )
783 }
784
785 #[test]
786 fn full_pipeline_merges_and_added_tokens() {
787 let tok = Tokenizer::from_json(&mini_json()).unwrap();
788 let ids = tok.encode("hello world");
790 assert_eq!(ids, vec![8, 2, 2, 3, 9, 3, 6, 2, 7]);
791 assert_eq!(tok.decode(&ids), "hello world");
792 let ids2 = tok.encode("he<|eot|>he");
794 assert_eq!(ids2, vec![8, 10, 8]);
795 assert_eq!(tok.decode(&ids2), "hehe");
796 }
797
798 #[test]
799 fn non_ascii_is_never_silently_dropped() {
800 let tok = Tokenizer::from_json(&mini_json()).unwrap();
801 let ids = tok.encode("hello");
804 assert!(!ids.is_empty());
805 }
806}