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 const TOOL_MARKUP_TOKENS: &[&str] = &[
33 "<tool_call>",
34 "</tool_call>",
35 "<function",
36 "</function>",
37 "<param",
38 "</param>",
39];
40
41pub struct Tokenizer {
43 vocab: HashMap<String, u32>,
45 id_to_token: Vec<String>,
47 ranks: HashMap<(String, String), u32>,
49 added: Vec<(String, u32)>,
51 added_ids: HashSet<u32>,
53 special_ids: HashSet<u32>,
55 split_res: Vec<fancy_regex::Regex>,
59 sp_prepend: bool,
62 sp_prepend_first: bool,
68 metaspace: bool,
71 nfc: bool,
74 byte_to_char: [char; 256],
76 char_to_byte: HashMap<char, u8>,
78 pub bos_token_id: Option<u32>,
80 pub eos_token_id: Option<u32>,
81 pub pad_token_id: Option<u32>,
82 pub im_start_id: Option<u32>,
84 pub im_end_id: Option<u32>,
85 pub chat_template: Option<String>,
88 pub extra_eos: HashSet<u32>,
90 pub add_bos: bool,
92}
93
94impl std::fmt::Debug for Tokenizer {
95 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
96 f.debug_struct("Tokenizer")
97 .field("vocab", &self.vocab.len())
98 .field("merges", &self.ranks.len())
99 .field("added", &self.added.len())
100 .finish()
101 }
102}
103
104fn bytes_to_unicode() -> ([char; 256], HashMap<char, u8>) {
107 let mut b2c = ['\0'; 256];
108 let mut c2b = HashMap::with_capacity(256);
109 let mut n = 0u32;
110 for b in 0..=255u16 {
111 let printable =
112 (0x21..=0x7E).contains(&b) || (0xA1..=0xAC).contains(&b) || (0xAE..=0xFF).contains(&b);
113 let c = if printable {
114 char::from_u32(b as u32).unwrap()
115 } else {
116 let c = char::from_u32(256 + n).unwrap();
117 n += 1;
118 c
119 };
120 b2c[b as usize] = c;
121 c2b.insert(c, b as u8);
122 }
123 (b2c, c2b)
124}
125
126#[derive(Deserialize)]
128struct HfTokenizerJson {
129 model: HfModel,
130 #[serde(default)]
131 added_tokens: Vec<HfAddedToken>,
132 #[serde(default)]
133 pre_tokenizer: Option<serde_json::Value>,
134 #[serde(default)]
135 normalizer: Option<serde_json::Value>,
136 #[serde(default)]
137 post_processor: Option<serde_json::Value>,
138}
139
140#[derive(Deserialize)]
141struct HfModel {
142 vocab: HashMap<String, u32>,
143 #[serde(default)]
144 merges: Vec<HfMerge>,
145 #[serde(default)]
146 byte_fallback: bool,
147}
148
149#[derive(Deserialize)]
152#[serde(untagged)]
153enum HfMerge {
154 Pair([String; 2]),
155 Text(String),
156}
157
158#[derive(Deserialize)]
159struct HfAddedToken {
160 id: u32,
161 content: String,
162 special: bool,
163}
164
165fn collect_split_patterns(pt: &serde_json::Value, out: &mut Vec<String>) {
174 if pt.get("type").and_then(|t| t.as_str()) == Some("Split") {
175 if let Some(r) = pt
176 .get("pattern")
177 .and_then(|p| p.get("Regex"))
178 .and_then(|r| r.as_str())
179 {
180 out.push(r.to_string());
181 }
182 return;
183 }
184 if let Some(list) = pt.get("pretokenizers").and_then(|l| l.as_array()) {
185 for p in list {
186 collect_split_patterns(p, out);
187 }
188 }
189}
190
191fn find_prepend_scheme(pt: &serde_json::Value) -> Option<String> {
194 if pt.get("type").and_then(|t| t.as_str()) == Some("Metaspace") {
195 return pt
196 .get("prepend_scheme")
197 .and_then(|p| p.as_str())
198 .map(String::from);
199 }
200 if let Some(list) = pt.get("pretokenizers").and_then(|l| l.as_array()) {
201 return list.iter().find_map(find_prepend_scheme);
202 }
203 None
204}
205
206pub(crate) fn strip_generation_tags(tpl: &str) -> std::borrow::Cow<'_, str> {
222 if !tpl.contains("generation") {
223 return std::borrow::Cow::Borrowed(tpl);
224 }
225 let mut out = String::with_capacity(tpl.len());
226 let mut rest = tpl;
227 let mut touched = false;
228 while let Some(open) = rest.find("{%") {
229 let Some(close_rel) = rest[open..].find("%}") else {
230 break;
231 };
232 let close = open + close_rel + 2;
233 let tag = &rest[open..close];
234 let inner = tag[2..tag.len() - 2].trim();
235 let lead = inner.starts_with('-');
236 let trail = inner.ends_with('-');
237 let name = inner.trim_matches('-').trim();
238 out.push_str(&rest[..open]);
239 if name == "generation" || name == "endgeneration" {
240 out.push_str(if lead { "{%-" } else { "{%" });
241 out.push_str(" set _generation_span = true ");
242 out.push_str(if trail { "-%}" } else { "%}" });
243 touched = true;
244 } else {
245 out.push_str(tag);
246 }
247 rest = &rest[close..];
248 }
249 if !touched {
250 return std::borrow::Cow::Borrowed(tpl);
251 }
252 out.push_str(rest);
253 std::borrow::Cow::Owned(out)
254}
255
256fn close_direct_think(rendered: String, enable_thinking: Option<bool>) -> String {
265 if enable_thinking != Some(false) || rendered.contains("</think>") {
266 return rendered;
267 }
268 if let Some(pos) = rendered.rfind("<|assistant|>") {
269 let marker_end = pos + "<|assistant|>".len();
270 let after = &rendered[marker_end..];
271 let mut out = String::with_capacity(rendered.len() + 24);
272 out.push_str(&rendered[..marker_end]);
273 if let Some(rest) = after.strip_prefix("<think>") {
274 out.push_str("<think></think>\n\n");
275 out.push_str(rest);
276 } else {
277 out.push_str(after);
278 if !out.ends_with('\n') {
279 out.push('\n');
280 }
281 out.push_str("<think>\n\n</think>\n\n");
282 }
283 return out;
284 }
285 if let Some(pos) = rendered.rfind("\nassistant") {
289 let mut insert_at = pos + 1 + "assistant".len();
290 if let Some(idx) = rendered[insert_at..].find('\n') {
291 insert_at += idx + 1;
292 }
293 let mut out = String::with_capacity(rendered.len() + 24);
294 out.push_str(&rendered[..insert_at]);
295 if !out.ends_with('\n') {
296 out.push('\n');
297 }
298 out.push_str("<think>\n\n</think>\n\n");
299 out.push_str(&rendered[insert_at..]);
300 return out;
301 }
302 rendered
303}
304
305impl Tokenizer {
306 pub fn from_file(path: impl AsRef<Path>) -> Result<Self, TokenizerError> {
308 let data = std::fs::read_to_string(path.as_ref())
309 .map_err(|e| TokenizerError::Io(e.to_string()))?;
310 Self::from_json(&data)
311 }
312
313 pub fn from_bytes(bytes: &[u8]) -> Result<Self, TokenizerError> {
315 let s = std::str::from_utf8(bytes)
316 .map_err(|e| TokenizerError::Parse(format!("vocab is not UTF-8: {e}")))?;
317 Self::from_json(s)
318 }
319
320 pub fn from_json(json: &str) -> Result<Self, TokenizerError> {
322 let hf: HfTokenizerJson =
323 serde_json::from_str(json).map_err(|e| TokenizerError::Parse(e.to_string()))?;
324
325 let mut vocab = hf.model.vocab;
326 let mut ranks = HashMap::new();
327 for (rank, m) in hf.model.merges.into_iter().enumerate() {
328 let (a, b) = match m {
329 HfMerge::Pair([a, b]) => (a, b),
330 HfMerge::Text(s) => {
331 let mut it = s.splitn(2, ' ');
332 match (it.next(), it.next()) {
333 (Some(a), Some(b)) => (a.to_string(), b.to_string()),
334 _ => continue,
335 }
336 }
337 };
338 ranks.insert((a, b), rank as u32);
339 }
340
341 let mut saw_gemma_bos = false;
346 let add_bos_detected = hf
347 .post_processor
348 .as_ref()
349 .map(|p| {
350 let pp = p.to_string();
351 pp.contains("\"<s>\"") || pp.contains("\"<bos>\"")
352 })
353 .unwrap_or(false);
354 let nfc = hf
355 .normalizer
356 .as_ref()
357 .map(|n| n.to_string().contains("NFC"))
358 .unwrap_or(false);
359 let metaspace = hf.model.byte_fallback
360 || hf
361 .normalizer
362 .as_ref()
363 .map(|n| n.to_string().contains("\u{2581}") || n.to_string().contains("▁"))
364 .unwrap_or(false);
365 let sp_prepend = hf
366 .normalizer
367 .as_ref()
368 .map(|n| n.to_string().contains("Prepend"))
369 .unwrap_or(false);
370 let (sp_prepend, sp_prepend_first) = if sp_prepend {
374 (true, false)
375 } else {
376 match hf.pre_tokenizer.as_ref().and_then(find_prepend_scheme) {
377 Some(s) if s == "always" => (true, false),
378 Some(s) if s == "first" => (false, true),
379 _ => (false, false),
380 }
381 };
382 let split_res = if metaspace {
383 Vec::new()
384 } else {
385 let mut pats = Vec::new();
386 if let Some(pt) = hf.pre_tokenizer.as_ref() {
387 collect_split_patterns(pt, &mut pats);
388 }
389 if pats.is_empty() {
390 pats.push(DEFAULT_SPLIT.to_string());
391 }
392 pats.iter()
393 .map(|p| {
394 fancy_regex::Regex::new(p)
395 .map_err(|e| TokenizerError::Parse(format!("pre-tokenizer regex: {e}")))
396 })
397 .collect::<Result<Vec<_>, _>>()?
398 };
399
400 let mut bos_token_id = None;
402 let mut eos_token_id = None;
403 let mut pad_token_id = None;
404 let mut im_start_id = None;
405 let mut im_end_id = None;
406 let mut special_ids = HashSet::new();
407 let mut added_ids = HashSet::new();
408 let mut added = Vec::new();
409
410 for at in &hf.added_tokens {
411 vocab.insert(at.content.clone(), at.id);
412 added.push((at.content.clone(), at.id));
413 added_ids.insert(at.id);
414 if at.special && !TOOL_MARKUP_TOKENS.contains(&at.content.as_str()) {
415 special_ids.insert(at.id);
416 }
417 match at.content.as_str() {
418 "<|endoftext|>" | "</s>" | "[EOS]" => eos_token_id = Some(at.id),
419 "<|im_start|>" => im_start_id = Some(at.id),
420 "<|im_end|>" => im_end_id = Some(at.id),
421 "<s>" | "[BOS]" => bos_token_id = Some(at.id),
422 "<bos>" => {
426 bos_token_id = Some(at.id);
427 saw_gemma_bos = true;
428 }
429 "<pad>" => pad_token_id = Some(at.id),
430 _ => {}
431 }
432 }
433
434 const DSV41_SPECIALS: &[&str] = &[
439 "<|begin▁of▁sentence|>",
440 "<|end▁of▁sentence|>",
441 "<|User|>",
442 "<|Assistant|>",
443 "<|System|>",
444 "<|latest_reminder|>",
445 "<|deepseek_image|>",
446 "<|action|>",
447 "<|query|>",
448 "<|authority|>",
449 "<|domain|>",
450 "<|title|>",
451 "<|read_url|>",
452 "<think>",
453 "</think>",
454 "|DSML|",
455 ];
456 for token in DSV41_SPECIALS {
457 if let Some(&id) = vocab.get(*token) {
458 if !added.iter().any(|(content, _)| content.as_str() == *token) {
459 added.push(((*token).to_string(), id));
460 }
461 added_ids.insert(id);
462 match *token {
463 "<|begin▁of▁sentence|>" => bos_token_id = Some(id),
464 "<|end▁of▁sentence|>" => eos_token_id = Some(id),
465 _ => {}
466 }
467 }
468 }
469 added.sort_by_key(|(c, _)| std::cmp::Reverse(c.len()));
470
471 let gemma_family = saw_gemma_bos
478 || vocab.contains_key("<start_of_turn>")
479 || added.iter().any(|(c, _)| c == "<start_of_turn>");
480
481 if let Some(pp) = hf.post_processor.as_ref() {
485 let pp = pp.to_string();
486 for name in ["<bos>", "<s>"] {
487 if pp.contains(&format!("\"{name}\"")) {
488 if let Some(&id) = vocab.get(name) {
489 bos_token_id = Some(id);
490 }
491 break;
492 }
493 }
494 }
495
496 let max_id = vocab.values().copied().max().unwrap_or(0) as usize;
498 let mut id_to_token = vec![String::new(); max_id + 1];
499 for (token, &id) in &vocab {
500 if (id as usize) < id_to_token.len() {
501 id_to_token[id as usize] = token.clone();
502 }
503 }
504
505 let (byte_to_char, char_to_byte) = bytes_to_unicode();
506
507 tracing::info!(
508 "Tokenizer loaded: {} vocab, {} merges, {} added, eos={:?}",
509 vocab.len(),
510 ranks.len(),
511 added.len(),
512 eos_token_id
513 );
514
515 Ok(Self {
516 vocab,
517 id_to_token,
518 ranks,
519 added,
520 added_ids,
521 special_ids,
522 split_res,
523 metaspace,
524 sp_prepend,
525 sp_prepend_first,
526 nfc,
527 byte_to_char,
528 char_to_byte,
529 bos_token_id,
530 eos_token_id,
531 pad_token_id,
532 im_start_id,
533 im_end_id,
534 chat_template: None,
535 extra_eos: HashSet::new(),
536 add_bos: add_bos_detected || gemma_family,
537 })
538 }
539
540 pub fn byte_level() -> Self {
542 let mut vocab = HashMap::new();
543 let mut id_to_token = Vec::with_capacity(256);
544 for i in 0..256u32 {
545 let tok = format!("<0x{:02X}>", i);
546 vocab.insert(tok.clone(), i);
547 id_to_token.push(tok);
548 }
549 let (byte_to_char, char_to_byte) = bytes_to_unicode();
550 Self {
551 vocab,
552 id_to_token,
553 ranks: HashMap::new(),
554 added: Vec::new(),
555 added_ids: HashSet::new(),
556 special_ids: HashSet::new(),
557 split_res: Vec::new(),
558 metaspace: false,
559 sp_prepend: false,
560 sp_prepend_first: false,
561 nfc: false,
562 byte_to_char,
563 char_to_byte,
564 bos_token_id: None,
565 eos_token_id: None,
566 pad_token_id: None,
567 im_start_id: None,
568 im_end_id: None,
569 chat_template: None,
570 extra_eos: HashSet::new(),
571 add_bos: false,
572 }
573 }
574
575 pub fn encode(&self, text: &str) -> Vec<u32> {
577 let mut ids = Vec::new();
578 let mut rest = text;
580 let mut head = true;
585 'outer: while !rest.is_empty() {
586 let mut best: Option<(usize, usize, u32)> = None; for (content, id) in &self.added {
588 if let Some(pos) = rest.find(content.as_str()) {
589 let better = match best {
590 None => true,
591 Some((bp, bl, _)) => pos < bp || (pos == bp && content.len() > bl),
592 };
593 if better {
594 best = Some((pos, content.len(), *id));
595 }
596 if pos == 0 {
597 break; }
599 }
600 }
601 match best {
602 Some((pos, len, id)) => {
603 self.encode_segment_at(&rest[..pos], head, &mut ids);
604 ids.push(id);
605 rest = &rest[pos + len..];
606 head = false;
607 }
608 None => {
609 self.encode_segment_at(rest, head, &mut ids);
610 break 'outer;
611 }
612 }
613 }
614 ids
615 }
616
617 pub fn encode_plain(&self, text: &str) -> Vec<u32> {
623 let mut ids = Vec::new();
624 self.encode_segment_at(text, true, &mut ids);
625 ids
626 }
627
628 fn encode_segment_at(&self, segment: &str, head: bool, out: &mut Vec<u32>) {
632 if segment.is_empty() {
633 return;
634 }
635 let norm: String = if self.nfc {
636 segment.nfc().collect()
637 } else {
638 segment.to_string()
639 };
640 if self.metaspace {
641 let sp = if self.sp_prepend {
645 format!("\u{2581}{}", norm).replace(' ', "\u{2581}")
648 } else {
649 let replaced = norm.replace(' ', "\u{2581}");
653 if self.sp_prepend_first && head && !replaced.starts_with('\u{2581}') {
654 format!("\u{2581}{replaced}")
655 } else {
656 replaced
657 }
658 };
659 self.bpe_piece_sp(&sp, out);
660 return;
661 }
662 if !self.split_res.is_empty() {
663 let mut pieces: Vec<(usize, usize)> = vec![(0, norm.len())];
666 for re in &self.split_res {
667 let mut next: Vec<(usize, usize)> = Vec::with_capacity(pieces.len() * 2);
668 for (ps, pe) in pieces {
669 let seg = &norm[ps..pe];
670 let mut last = 0usize;
671 for m in re.find_iter(seg) {
672 let m = match m {
673 Ok(m) => m,
674 Err(e) => {
675 tracing::error!("pre-tokenizer regex failed: {e}");
676 break;
677 }
678 };
679 if m.start() > last {
680 next.push((ps + last, ps + m.start()));
681 }
682 if m.end() > m.start() {
683 next.push((ps + m.start(), ps + m.end()));
684 }
685 last = m.end();
686 }
687 if last < seg.len() {
688 next.push((ps + last, pe));
689 }
690 }
691 pieces = next;
692 }
693 for (ps, pe) in pieces {
694 self.bpe_piece(&norm[ps..pe], out);
695 }
696 } else {
697 {
698 for b in norm.bytes() {
700 let tok = format!("<0x{:02X}>", b);
701 if let Some(&id) = self.vocab.get(&tok) {
702 out.push(id);
703 }
704 }
705 }
706 }
707 }
708
709 fn bpe_piece_sp(&self, piece: &str, out: &mut Vec<u32>) {
712 if piece.is_empty() {
713 return;
714 }
715 let mut sym: Vec<String> = piece.chars().map(|c| c.to_string()).collect();
716 loop {
717 let mut best: Option<(u32, usize)> = None;
718 for i in 0..sym.len().saturating_sub(1) {
719 if let Some(&r) = self.ranks.get(&(sym[i].clone(), sym[i + 1].clone())) {
720 if best.map(|(br, _)| r < br).unwrap_or(true) {
721 best = Some((r, i));
722 }
723 }
724 }
725 let Some((_, i)) = best else { break };
726 let merged = format!("{}{}", sym[i], sym[i + 1]);
727 let (left, right) = (sym[i].clone(), sym[i + 1].clone());
728 let mut j = 0;
729 while j + 1 < sym.len() {
730 if sym[j] == left && sym[j + 1] == right {
731 sym[j] = merged.clone();
732 sym.remove(j + 1);
733 }
734 j += 1;
735 }
736 }
737 for t in &sym {
738 if let Some(&id) = self.vocab.get(t) {
739 out.push(id);
740 } else {
741 let mut ok = true;
742 for byte in t.bytes() {
743 let tok = format!("<0x{:02X}>", byte);
744 match self.vocab.get(&tok) {
745 Some(&id) => out.push(id),
746 None => {
747 ok = false;
748 break;
749 }
750 }
751 }
752 if !ok {
753 tracing::error!("tokenizer: no id for SP symbol {t:?} — dropped");
754 }
755 }
756 }
757 }
758
759 fn bpe_piece(&self, piece: &str, out: &mut Vec<u32>) {
761 if piece.is_empty() {
762 return;
763 }
764 let mapped: Vec<String> = piece
765 .bytes()
766 .map(|b| self.byte_to_char[b as usize].to_string())
767 .collect();
768 let mut sym = mapped;
769
770 loop {
772 let mut best: Option<(u32, usize)> = None;
773 for i in 0..sym.len().saturating_sub(1) {
774 if let Some(&r) = self.ranks.get(&(sym[i].clone(), sym[i + 1].clone())) {
775 if best.map(|(br, _)| r < br).unwrap_or(true) {
776 best = Some((r, i));
777 }
778 }
779 }
780 let Some((_, i)) = best else { break };
781 let merged = format!("{}{}", sym[i], sym[i + 1]);
782 let (left, right) = (sym[i].clone(), sym[i + 1].clone());
784 let mut j = 0;
785 while j + 1 < sym.len() {
786 if sym[j] == left && sym[j + 1] == right {
787 sym[j] = merged.clone();
788 sym.remove(j + 1);
789 }
790 j += 1;
791 }
792 }
793
794 for s in &sym {
795 if let Some(&id) = self.vocab.get(s) {
796 out.push(id);
797 } else {
798 let mut ok = true;
800 for ch in s.chars() {
801 let Some(&b) = self.char_to_byte.get(&ch) else {
802 ok = false;
803 break;
804 };
805 let tok = format!("<0x{:02X}>", b);
806 if let Some(&id) = self.vocab.get(&tok) {
807 out.push(id);
808 } else {
809 ok = false;
810 break;
811 }
812 }
813 if !ok {
814 tracing::error!("tokenizer: no id for symbol {s:?} — dropped");
815 }
816 }
817 }
818 }
819
820 pub fn decode(&self, ids: &[u32]) -> String {
823 let mut bytes: Vec<u8> = Vec::new();
824 for &id in ids {
825 if self.special_ids.contains(&id) {
826 continue;
827 }
828 let idx = id as usize;
829 if idx >= self.id_to_token.len() {
830 continue;
831 }
832 let tok = &self.id_to_token[idx];
833 if self.added_ids.contains(&id) {
834 if self.metaspace && tok.contains('\u{2581}') {
837 bytes.extend_from_slice(tok.replace('\u{2581}', " ").as_bytes());
838 } else {
839 bytes.extend_from_slice(tok.as_bytes());
840 }
841 continue;
842 }
843 if tok.starts_with("<0x") && tok.ends_with('>') && tok.len() == 6 {
845 if let Ok(b) = u8::from_str_radix(&tok[3..5], 16) {
846 bytes.push(b);
847 continue;
848 }
849 }
850 if self.metaspace {
851 for ch in tok.chars() {
853 if ch == '\u{2581}' {
854 bytes.push(b' ');
855 } else {
856 let mut buf = [0u8; 4];
857 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
858 }
859 }
860 continue;
861 }
862 for ch in tok.chars() {
863 match self.char_to_byte.get(&ch) {
864 Some(&b) => bytes.push(b),
865 None => {
868 let mut buf = [0u8; 4];
869 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
870 }
871 }
872 }
873 }
874 let text = String::from_utf8_lossy(&bytes).into_owned();
875 if self.metaspace && (self.sp_prepend || self.sp_prepend_first) {
876 if let Some(stripped) = text.strip_prefix(' ') {
878 return stripped.to_string();
879 }
880 }
881 text
882 }
883
884 pub fn decode_token(&self, id: u32) -> String {
887 if self.special_ids.contains(&id) {
888 return String::new();
889 }
890 let idx = id as usize;
891 if idx >= self.id_to_token.len() {
892 return String::new();
893 }
894 let tok = &self.id_to_token[idx];
895 if self.added_ids.contains(&id) {
896 if self.metaspace && tok.contains('\u{2581}') {
897 return tok.replace('\u{2581}', " ");
898 }
899 return tok.clone();
900 }
901 if tok.starts_with("<0x") && tok.ends_with('>') && tok.len() == 6 {
902 if let Ok(b) = u8::from_str_radix(&tok[3..5], 16) {
903 return String::from_utf8_lossy(&[b]).into_owned();
904 }
905 }
906 if self.metaspace {
907 return tok.replace('\u{2581}', " ");
908 }
909 let mut bytes = Vec::new();
910 for ch in tok.chars() {
911 match self.char_to_byte.get(&ch) {
912 Some(&b) => bytes.push(b),
913 None => {
914 let mut buf = [0u8; 4];
915 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
916 }
917 }
918 }
919 String::from_utf8_lossy(&bytes).into_owned()
920 }
921
922 pub fn decode_token_for_hash(&self, id: u32) -> String {
925 let idx = id as usize;
926 if idx >= self.id_to_token.len() {
927 return String::new();
928 }
929 if self.special_ids.contains(&id) {
930 return self.id_to_token[idx].clone();
931 }
932 self.decode_token(id)
933 }
934
935 pub fn decode_for_protocol(&self, ids: &[u32]) -> String {
938 let mut out = String::new();
939 for &id in ids {
940 let idx = id as usize;
941 if self.special_ids.contains(&id) {
942 if let Some(token) = self.id_to_token.get(idx) {
943 out.push_str(token);
944 }
945 } else {
946 out.push_str(&self.decode_token(id));
947 }
948 }
949 out
950 }
951
952 pub fn raw_token_for_hash(&self, id: u32) -> String {
954 self.id_to_token
955 .get(id as usize)
956 .cloned()
957 .unwrap_or_default()
958 }
959
960 pub fn apply_chat_template(&self, messages: &[(String, String)]) -> Vec<u32> {
964 self.apply_chat_template_opts(messages, None)
965 }
966
967 pub fn apply_chat_template_json(
984 &self,
985 messages: &[serde_json::Value],
986 tools: Option<&[serde_json::Value]>,
987 enable_thinking: Option<bool>,
988 ) -> Vec<u32> {
989 match self.try_apply_chat_template_json(messages, tools, enable_thinking) {
990 Ok(ids) => ids,
991 Err(e) => {
992 tracing::error!("chat template render failed ({e}); ChatML fallback");
993 self.chatml_json_fallback(messages, enable_thinking)
994 }
995 }
996 }
997
998 pub fn try_apply_chat_template_json(
1009 &self,
1010 messages: &[serde_json::Value],
1011 tools: Option<&[serde_json::Value]>,
1012 enable_thinking: Option<bool>,
1013 ) -> Result<Vec<u32>, String> {
1014 if let Some(tpl) = &self.chat_template {
1015 return self
1016 .render_template_json(tpl, messages, tools, enable_thinking)
1017 .map(|text| self.with_bos(self.encode(&text)))
1018 .map_err(|e| format!("{e:#}"));
1019 }
1020 Ok(self.chatml_json_fallback(messages, enable_thinking))
1021 }
1022
1023 fn chatml_json_fallback(
1024 &self,
1025 messages: &[serde_json::Value],
1026 enable_thinking: Option<bool>,
1027 ) -> Vec<u32> {
1028 let pairs: Vec<(String, String)> = messages
1029 .iter()
1030 .map(|m| {
1031 (
1032 m.get("role")
1033 .and_then(|v| v.as_str())
1034 .unwrap_or("user")
1035 .to_string(),
1036 m.get("content")
1037 .and_then(|v| v.as_str())
1038 .unwrap_or("")
1039 .to_string(),
1040 )
1041 })
1042 .collect();
1043 self.with_bos(self.chatml_fallback_opts(&pairs, enable_thinking))
1044 }
1045
1046 pub fn render_chat_json(
1048 &self,
1049 messages: &[serde_json::Value],
1050 tools: Option<&[serde_json::Value]>,
1051 enable_thinking: Option<bool>,
1052 ) -> Option<String> {
1053 let tpl = self.chat_template.as_ref()?;
1054 match self.render_template_json(tpl, messages, tools, enable_thinking) {
1055 Ok(t) => Some(t),
1056 Err(e) => {
1057 tracing::error!("chat template render (json): {e:#}");
1058 eprintln!("chat template render (json): {e:#}");
1059 None
1060 }
1061 }
1062 }
1063
1064 fn render_template_json(
1065 &self,
1066 tpl: &str,
1067 messages: &[serde_json::Value],
1068 tools: Option<&[serde_json::Value]>,
1069 enable_thinking: Option<bool>,
1070 ) -> Result<String, minijinja::Error> {
1071 let mut env = crate::chat_template::environment();
1072 let tpl_src = strip_generation_tags(tpl);
1073 env.add_template("chat", &tpl_src)?;
1074 let msgs: Vec<minijinja::Value> = messages
1075 .iter()
1076 .map(minijinja::Value::from_serialize)
1077 .collect();
1078 let tools_v: Option<Vec<minijinja::Value>> =
1079 tools.map(|ts| ts.iter().map(minijinja::Value::from_serialize).collect());
1080 let tpl = env.get_template("chat")?;
1081 let rendered = match (tools_v, enable_thinking) {
1088 (Some(ts), Some(v)) => tpl.render(minijinja::context! {
1089 messages => msgs, tools => ts, add_generation_prompt => true, enable_thinking => v,
1090 reasoning_effort => if !v { Some("low") } else { None::<&str> },
1095 tool_call_format => "json",
1096 })?,
1097 (Some(ts), None) => tpl.render(minijinja::context! {
1098 messages => msgs, tools => ts, add_generation_prompt => true,
1099 tool_call_format => "json",
1100 })?,
1101 (None, Some(v)) => tpl.render(minijinja::context! {
1102 messages => msgs, add_generation_prompt => true, enable_thinking => v,
1103 reasoning_effort => if !v { Some("low") } else { None::<&str> },
1105 tool_call_format => "json",
1106 })?,
1107 (None, None) => tpl.render(minijinja::context! {
1108 messages => msgs, add_generation_prompt => true,
1109 tool_call_format => "json",
1110 })?,
1111 };
1112 Ok(close_direct_think(rendered, enable_thinking))
1113 }
1114
1115 pub fn apply_chat_template_opts(
1116 &self,
1117 messages: &[(String, String)],
1118 enable_thinking: Option<bool>,
1119 ) -> Vec<u32> {
1120 if let Some(tpl) = &self.chat_template {
1121 match self.render_template(tpl, messages, enable_thinking) {
1122 Ok(text) => return self.with_bos(self.encode(&text)),
1123 Err(e) => {
1124 tracing::error!("chat template render failed ({e}); ChatML fallback");
1125 }
1126 }
1127 }
1128 self.with_bos(self.chatml_fallback_opts(messages, enable_thinking))
1129 }
1130
1131 pub fn with_bos(&self, mut ids: Vec<u32>) -> Vec<u32> {
1133 if self.add_bos {
1134 if let Some(b) = self.bos_token_id {
1135 if ids.first() != Some(&b) {
1136 ids.insert(0, b);
1137 }
1138 }
1139 }
1140 ids
1141 }
1142
1143 pub fn render_chat(&self, messages: &[(String, String)]) -> Option<String> {
1145 self.render_chat_opts(messages, None)
1146 }
1147
1148 pub fn render_chat_opts(
1150 &self,
1151 messages: &[(String, String)],
1152 enable_thinking: Option<bool>,
1153 ) -> Option<String> {
1154 let tpl = self.chat_template.as_ref()?;
1155 match self.render_template(tpl, messages, enable_thinking) {
1156 Ok(t) => Some(t),
1157 Err(e) => {
1158 tracing::error!("chat template render: {e:#}");
1159 None
1160 }
1161 }
1162 }
1163
1164 fn render_template(
1165 &self,
1166 tpl: &str,
1167 messages: &[(String, String)],
1168 enable_thinking: Option<bool>,
1169 ) -> Result<String, minijinja::Error> {
1170 let mut env = crate::chat_template::environment();
1171 let tpl_src = strip_generation_tags(tpl);
1172 env.add_template("chat", &tpl_src)?;
1173 let msgs: Vec<minijinja::Value> = messages
1174 .iter()
1175 .map(|(role, content)| {
1176 minijinja::context! { role => role, content => content }
1177 })
1178 .collect();
1179 let rendered = match enable_thinking {
1182 Some(v) => env.get_template("chat")?.render(minijinja::context! {
1183 messages => msgs,
1184 add_generation_prompt => true,
1185 enable_thinking => v,
1186 reasoning_effort => if !v { Some("low") } else { None::<&str> },
1189 })?,
1190 None => env.get_template("chat")?.render(minijinja::context! {
1191 messages => msgs,
1192 add_generation_prompt => true,
1193 })?,
1194 };
1195 Ok(close_direct_think(rendered, enable_thinking))
1196 }
1197
1198 fn chatml_fallback(&self, messages: &[(String, String)]) -> Vec<u32> {
1200 self.chatml_fallback_opts(messages, None)
1201 }
1202
1203 fn chatml_fallback_opts(
1205 &self,
1206 messages: &[(String, String)],
1207 enable_thinking: Option<bool>,
1208 ) -> Vec<u32> {
1209 let mut tokens = Vec::new();
1210
1211 for (role, content) in messages {
1212 if let Some(start_id) = self.im_start_id {
1214 tokens.push(start_id);
1215 }
1216 tokens.extend(self.encode(&format!("{}\n{}", role, content)));
1217 if let Some(end_id) = self.im_end_id {
1218 tokens.push(end_id);
1219 }
1220 tokens.extend(self.encode("\n"));
1221 }
1222
1223 if let Some(start_id) = self.im_start_id {
1225 tokens.push(start_id);
1226 }
1227 tokens.extend(self.encode("assistant\n"));
1228 if enable_thinking == Some(false) {
1229 tokens.extend(self.encode("<think>\n\n</think>\n\n"));
1230 }
1231
1232 tokens
1233 }
1234
1235 pub fn vocab_size(&self) -> usize {
1237 self.id_to_token.len()
1238 }
1239
1240 pub fn token_to_id(&self, token: &str) -> Option<u32> {
1244 self.vocab.get(token).copied()
1245 }
1246
1247 pub fn convert_tokens_to_ids(&self, token: &str) -> Option<u32> {
1250 self.token_to_id(token)
1251 }
1252
1253 pub fn is_eos(&self, id: u32) -> bool {
1255 self.eos_token_id == Some(id) || self.im_end_id == Some(id) || self.extra_eos.contains(&id)
1256 }
1257}
1258
1259#[derive(Debug, thiserror::Error)]
1260pub enum TokenizerError {
1261 #[error("IO error: {0}")]
1262 Io(String),
1263 #[error("Parse error: {0}")]
1264 Parse(String),
1265}
1266
1267#[cfg(test)]
1268mod tests {
1269 use super::*;
1270
1271 #[test]
1272 fn byte_unicode_bijection() {
1273 let (b2c, c2b) = bytes_to_unicode();
1274 for b in 0..=255u8 {
1275 assert_eq!(c2b[&b2c[b as usize]], b);
1276 }
1277 assert_eq!(b2c[b' ' as usize], 'Ġ');
1279 assert_eq!(b2c[b'\n' as usize], 'Ċ');
1280 }
1281
1282 #[test]
1283 fn byte_level_roundtrip_utf8() {
1284 let tok = Tokenizer::byte_level();
1285 let text = "hello 🌍 hi\n";
1286 let ids = tok.encode(text);
1287 assert_eq!(ids.len(), text.len()); assert_eq!(tok.decode(&ids), text);
1289 }
1290
1291 fn mini_json() -> String {
1294 let vocab: Vec<(&str, u32)> = vec![
1296 ("h", 0),
1297 ("e", 1),
1298 ("l", 2),
1299 ("o", 3),
1300 ("Ġ", 4),
1301 ("w", 5),
1302 ("r", 6),
1303 ("d", 7),
1304 ("he", 8),
1305 ("Ġw", 9),
1306 ];
1307 let vocab_json: String = vocab
1308 .iter()
1309 .map(|(t, i)| format!("\"{t}\": {i}"))
1310 .collect::<Vec<_>>()
1311 .join(", ");
1312 format!(
1313 r#"{{
1314 "model": {{
1315 "type": "BPE",
1316 "vocab": {{ {vocab_json} }},
1317 "merges": [["h", "e"], ["Ġ", "w"]]
1318 }},
1319 "added_tokens": [
1320 {{"id": 10, "content": "<|eot|>", "special": true}}
1321 ]
1322 }}"#
1323 )
1324 }
1325
1326 #[test]
1330 fn tool_markup_decodes_even_when_special() {
1331 let json = r#"{
1332 "model": {"type": "BPE", "vocab": {"h": 0, "e": 1, "l": 2, "o": 3}, "merges": []},
1333 "added_tokens": [
1334 {"id": 10, "content": "<|im_end|>", "special": true},
1335 {"id": 11, "content": "<function", "special": true},
1336 {"id": 12, "content": "</function>", "special": true},
1337 {"id": 13, "content": "<param", "special": true},
1338 {"id": 14, "content": "</param>", "special": true},
1339 {"id": 15, "content": "<tool_call>", "special": true}
1340 ]
1341 }"#;
1342 let t = Tokenizer::from_json(json).unwrap();
1343 let ids = [11, 0, 1, 13, 2, 14, 12, 15, 10];
1344 assert_eq!(
1345 t.decode(&ids),
1346 "<functionhe<paraml</param></function><tool_call>"
1347 );
1348 let streamed: String = ids.iter().map(|&i| t.decode_token(i)).collect();
1349 assert_eq!(streamed, t.decode(&ids), "streaming must agree with decode");
1350 assert!(
1351 !t.decode(&[10]).contains("im_end"),
1352 "control tokens stay hidden"
1353 );
1354 }
1355
1356 #[test]
1359 fn real_tokenizer_parity_when_available() {
1360 let Ok(path) = std::env::var("CMF_TOK_PARITY") else {
1361 return;
1362 };
1363 let t = Tokenizer::from_file(&path).expect("load");
1364 for (text, want) in [
1365 (
1366 "The capital of France is",
1367 vec![671u32, 6102, 294, 8760, 344],
1368 ),
1369 ("2 + 2 =", vec![20, 940, 223, 20, 438]),
1370 ] {
1371 let got = t.encode(text);
1372 assert_eq!(got, want, "«{text}»");
1373 }
1374 }
1375
1376 #[test]
1381 fn granite_42_chat_template_when_available() {
1382 let Ok(path) = std::env::var("CMF_GRANITE_CHAT_TEMPLATE") else {
1383 return;
1384 };
1385 let mut tok = Tokenizer::byte_level();
1386 tok.chat_template = Some(std::fs::read_to_string(path).expect("read Granite template"));
1387 let messages = vec![("user".to_string(), "Hello".to_string())];
1388
1389 let thinking = tok
1390 .render_chat_opts(&messages, Some(true))
1391 .expect("render Granite thinking prompt");
1392 assert_eq!(
1393 thinking,
1394 "<|im_start|>system\n<|im_end|>\n<|im_start|>user\nHello<|im_end|>\n<|im_start|>assistant\n<think>\n"
1395 );
1396
1397 let direct = tok
1398 .render_chat_opts(&messages, Some(false))
1399 .expect("render Granite direct prompt");
1400 assert_eq!(
1401 direct,
1402 "<|im_start|>system\n<|im_end|>\n<|im_start|>user\nHello<|im_end|>\n<|im_start|>assistant\n<think></think>"
1403 );
1404 }
1405
1406 #[test]
1410 fn glm_direct_template_keeps_assistant_special_token_intact() {
1411 let mut tok = Tokenizer::byte_level();
1412 tok.chat_template = Some(
1413 "[gMASK]<sop>{%- set effort = reasoning_effort if reasoning_effort is defined and reasoning_effort in ['low', 'high'] else 'max' -%}<|system|>Reasoning Effort: {{ effort | capitalize }}{%- for m in messages -%}<|user|>{{ m.content }}{%- endfor -%}<|assistant|><think>"
1414 .to_string(),
1415 );
1416 let messages = vec![("user".to_string(), "2+2?".to_string())];
1417 let rendered = tok
1418 .render_chat_opts(&messages, Some(false))
1419 .expect("render GLM direct prompt");
1420 assert!(rendered.contains("Reasoning Effort: Low"));
1421 assert!(rendered.contains("<|assistant|><think></think>"));
1422 assert!(!rendered.contains("<|assistant\n"));
1423
1424 let thinking = tok
1427 .render_chat_opts(&messages, Some(true))
1428 .expect("render GLM thinking prompt");
1429 assert!(thinking.contains("Reasoning Effort: Max"));
1430
1431 let json_messages = vec![serde_json::json!({
1432 "role": "user",
1433 "content": "2+2?"
1434 })];
1435 let rendered_json = tok
1436 .render_chat_json(&json_messages, None, Some(false))
1437 .expect("render GLM JSON direct prompt");
1438 assert!(rendered_json.contains("Reasoning Effort: Low"));
1439 assert!(rendered_json.contains("<|assistant|><think></think>"));
1440 assert!(!rendered_json.contains("<|assistant\n"));
1441 }
1442
1443 #[test]
1449 fn every_split_in_a_sequence_is_applied() {
1450 let pt = serde_json::json!({
1451 "type": "Sequence",
1452 "pretokenizers": [
1453 {"type": "Split", "behavior": "Isolated",
1454 "pattern": {"Regex": r"\p{N}{1,3}"}},
1455 {"type": "Split", "behavior": "Isolated",
1456 "pattern": {"Regex": r" ?[\p{L}]+"}},
1457 {"type": "ByteLevel", "add_prefix_space": false, "use_regex": false}
1458 ]
1459 });
1460 let mut pats = Vec::new();
1461 collect_split_patterns(&pt, &mut pats);
1462 assert_eq!(
1463 pats.len(),
1464 2,
1465 "both Split stages must be collected: {pats:?}"
1466 );
1467 assert!(pats[0].contains("p{N}"), "digit rule first");
1468 assert!(pats[1].contains("p{L}"), "word rule second");
1469
1470 let re: Vec<fancy_regex::Regex> = pats
1474 .iter()
1475 .map(|p| fancy_regex::Regex::new(p).unwrap())
1476 .collect();
1477 let norm = "ab cd12";
1478 let mut pieces: Vec<(usize, usize)> = vec![(0, norm.len())];
1479 for r in &re {
1480 let mut next = Vec::new();
1481 for (ps, pe) in pieces {
1482 let seg = &norm[ps..pe];
1483 let mut last = 0;
1484 for m in r.find_iter(seg).flatten() {
1485 if m.start() > last {
1486 next.push((ps + last, ps + m.start()));
1487 }
1488 if m.end() > m.start() {
1489 next.push((ps + m.start(), ps + m.end()));
1490 }
1491 last = m.end();
1492 }
1493 if last < seg.len() {
1494 next.push((ps + last, pe));
1495 }
1496 }
1497 pieces = next;
1498 }
1499 let got: Vec<&str> = pieces.iter().map(|(a, b)| &norm[*a..*b]).collect();
1500 assert_eq!(
1501 got,
1502 vec!["ab", " cd", "12"],
1503 "staged split produced {got:?}"
1504 );
1505 }
1506
1507 #[test]
1508 fn full_pipeline_merges_and_added_tokens() {
1509 let tok = Tokenizer::from_json(&mini_json()).unwrap();
1510 let ids = tok.encode("hello world");
1512 assert_eq!(ids, vec![8, 2, 2, 3, 9, 3, 6, 2, 7]);
1513 assert_eq!(tok.decode(&ids), "hello world");
1514 let ids2 = tok.encode("he<|eot|>he");
1516 assert_eq!(ids2, vec![8, 10, 8]);
1517 assert_eq!(tok.decode(&ids2), "hehe");
1518 }
1519
1520 #[test]
1521 fn non_ascii_is_never_silently_dropped() {
1522 let tok = Tokenizer::from_json(&mini_json()).unwrap();
1523 let ids = tok.encode("hello");
1526 assert!(!ids.is_empty());
1527 }
1528}
1529
1530#[cfg(test)]
1531mod generation_tag_tests {
1532 use super::strip_generation_tags;
1533
1534 #[test]
1539 fn a_generation_block_becomes_a_no_op_keeping_its_whitespace_control() {
1540 let tpl = "a{%- generation -%}b{%- endgeneration -%}c";
1541 let out = strip_generation_tags(tpl);
1542 assert!(!out.contains("{%- generation"));
1543 assert!(!out.contains("endgeneration"));
1544 assert_eq!(out.matches("{%-").count(), 2);
1546 assert_eq!(out.matches("-%}").count(), 2);
1547 assert!(out.starts_with('a') && out.ends_with('c'));
1548 }
1549
1550 #[test]
1553 fn each_side_keeps_its_own_dash() {
1554 let out = strip_generation_tags("{% generation %}x{%- endgeneration %}");
1555 assert!(out.starts_with("{% set"), "no dash added on the left");
1556 assert!(out.contains("{%- set"), "the right tag keeps its dash");
1557 assert!(!out.contains("-%}"), "no trailing dash invented");
1558 }
1559
1560 #[test]
1563 fn everything_else_is_left_alone() {
1564 let plain = "{%- if x -%}{{ y }}{%- endif -%}";
1565 assert_eq!(strip_generation_tags(plain), plain);
1566 let prose = "{{ 'the generation of tokens' }}";
1568 assert_eq!(strip_generation_tags(prose), prose);
1569 }
1570}