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 sp_prepend: bool,
41 metaspace: bool,
44 nfc: bool,
47 byte_to_char: [char; 256],
49 char_to_byte: HashMap<char, u8>,
51 pub bos_token_id: Option<u32>,
53 pub eos_token_id: Option<u32>,
54 pub pad_token_id: Option<u32>,
55 pub im_start_id: Option<u32>,
57 pub im_end_id: Option<u32>,
58 pub chat_template: Option<String>,
61 pub extra_eos: HashSet<u32>,
63 pub add_bos: bool,
65}
66
67impl std::fmt::Debug for Tokenizer {
68 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
69 f.debug_struct("Tokenizer")
70 .field("vocab", &self.vocab.len())
71 .field("merges", &self.ranks.len())
72 .field("added", &self.added.len())
73 .finish()
74 }
75}
76
77fn bytes_to_unicode() -> ([char; 256], HashMap<char, u8>) {
80 let mut b2c = ['\0'; 256];
81 let mut c2b = HashMap::with_capacity(256);
82 let mut n = 0u32;
83 for b in 0..=255u16 {
84 let printable =
85 (0x21..=0x7E).contains(&b) || (0xA1..=0xAC).contains(&b) || (0xAE..=0xFF).contains(&b);
86 let c = if printable {
87 char::from_u32(b as u32).unwrap()
88 } else {
89 let c = char::from_u32(256 + n).unwrap();
90 n += 1;
91 c
92 };
93 b2c[b as usize] = c;
94 c2b.insert(c, b as u8);
95 }
96 (b2c, c2b)
97}
98
99#[derive(Deserialize)]
101struct HfTokenizerJson {
102 model: HfModel,
103 #[serde(default)]
104 added_tokens: Vec<HfAddedToken>,
105 #[serde(default)]
106 pre_tokenizer: Option<serde_json::Value>,
107 #[serde(default)]
108 normalizer: Option<serde_json::Value>,
109 #[serde(default)]
110 post_processor: Option<serde_json::Value>,
111}
112
113#[derive(Deserialize)]
114struct HfModel {
115 vocab: HashMap<String, u32>,
116 #[serde(default)]
117 merges: Vec<HfMerge>,
118 #[serde(default)]
119 byte_fallback: bool,
120}
121
122#[derive(Deserialize)]
125#[serde(untagged)]
126enum HfMerge {
127 Pair([String; 2]),
128 Text(String),
129}
130
131#[derive(Deserialize)]
132struct HfAddedToken {
133 id: u32,
134 content: String,
135 special: bool,
136}
137
138fn find_split_pattern(pt: &serde_json::Value) -> Option<String> {
141 if pt.get("type").and_then(|t| t.as_str()) == Some("Split") {
142 return pt
143 .get("pattern")
144 .and_then(|p| p.get("Regex"))
145 .and_then(|r| r.as_str())
146 .map(String::from);
147 }
148 if let Some(list) = pt.get("pretokenizers").and_then(|l| l.as_array()) {
149 return list.iter().find_map(find_split_pattern);
150 }
151 None
152}
153
154impl Tokenizer {
155 pub fn from_file(path: impl AsRef<Path>) -> Result<Self, TokenizerError> {
157 let data = std::fs::read_to_string(path.as_ref())
158 .map_err(|e| TokenizerError::Io(e.to_string()))?;
159 Self::from_json(&data)
160 }
161
162 pub fn from_bytes(bytes: &[u8]) -> Result<Self, TokenizerError> {
164 let s = std::str::from_utf8(bytes)
165 .map_err(|e| TokenizerError::Parse(format!("vocab is not UTF-8: {e}")))?;
166 Self::from_json(s)
167 }
168
169 pub fn from_json(json: &str) -> Result<Self, TokenizerError> {
171 let hf: HfTokenizerJson =
172 serde_json::from_str(json).map_err(|e| TokenizerError::Parse(e.to_string()))?;
173
174 let mut vocab = hf.model.vocab;
175 let mut ranks = HashMap::new();
176 for (rank, m) in hf.model.merges.into_iter().enumerate() {
177 let (a, b) = match m {
178 HfMerge::Pair([a, b]) => (a, b),
179 HfMerge::Text(s) => {
180 let mut it = s.splitn(2, ' ');
181 match (it.next(), it.next()) {
182 (Some(a), Some(b)) => (a.to_string(), b.to_string()),
183 _ => continue,
184 }
185 }
186 };
187 ranks.insert((a, b), rank as u32);
188 }
189
190 let mut saw_gemma_bos = false;
195 let add_bos_detected = hf
196 .post_processor
197 .as_ref()
198 .map(|p| {
199 let pp = p.to_string();
200 pp.contains("\"<s>\"") || pp.contains("\"<bos>\"")
201 })
202 .unwrap_or(false);
203 let nfc = hf
204 .normalizer
205 .as_ref()
206 .map(|n| n.to_string().contains("NFC"))
207 .unwrap_or(false);
208 let metaspace = hf.model.byte_fallback
209 || hf
210 .normalizer
211 .as_ref()
212 .map(|n| n.to_string().contains("\u{2581}") || n.to_string().contains("▁"))
213 .unwrap_or(false);
214 let sp_prepend = hf
215 .normalizer
216 .as_ref()
217 .map(|n| n.to_string().contains("Prepend"))
218 .unwrap_or(false);
219 let split_re = if metaspace {
220 None
221 } else {
222 let pattern = hf
223 .pre_tokenizer
224 .as_ref()
225 .and_then(find_split_pattern)
226 .unwrap_or_else(|| DEFAULT_SPLIT.to_string());
227 Some(
228 fancy_regex::Regex::new(&pattern)
229 .map_err(|e| TokenizerError::Parse(format!("pre-tokenizer regex: {e}")))?,
230 )
231 };
232
233 let mut bos_token_id = None;
235 let mut eos_token_id = None;
236 let mut pad_token_id = None;
237 let mut im_start_id = None;
238 let mut im_end_id = None;
239 let mut special_ids = HashSet::new();
240 let mut added_ids = HashSet::new();
241 let mut added = Vec::new();
242
243 for at in &hf.added_tokens {
244 vocab.insert(at.content.clone(), at.id);
245 added.push((at.content.clone(), at.id));
246 added_ids.insert(at.id);
247 if at.special {
248 special_ids.insert(at.id);
249 }
250 match at.content.as_str() {
251 "<|endoftext|>" | "</s>" | "[EOS]" => eos_token_id = Some(at.id),
252 "<|im_start|>" => im_start_id = Some(at.id),
253 "<|im_end|>" => im_end_id = Some(at.id),
254 "<s>" | "[BOS]" => bos_token_id = Some(at.id),
255 "<bos>" => {
259 bos_token_id = Some(at.id);
260 saw_gemma_bos = true;
261 }
262 "<pad>" => pad_token_id = Some(at.id),
263 _ => {}
264 }
265 }
266 added.sort_by_key(|(c, _)| std::cmp::Reverse(c.len()));
267
268 let gemma_family = saw_gemma_bos
275 || vocab.contains_key("<start_of_turn>")
276 || added.iter().any(|(c, _)| c == "<start_of_turn>");
277
278 if let Some(pp) = hf.post_processor.as_ref() {
282 let pp = pp.to_string();
283 for name in ["<bos>", "<s>"] {
284 if pp.contains(&format!("\"{name}\"")) {
285 if let Some(&id) = vocab.get(name) {
286 bos_token_id = Some(id);
287 }
288 break;
289 }
290 }
291 }
292
293 let max_id = vocab.values().copied().max().unwrap_or(0) as usize;
295 let mut id_to_token = vec![String::new(); max_id + 1];
296 for (token, &id) in &vocab {
297 if (id as usize) < id_to_token.len() {
298 id_to_token[id as usize] = token.clone();
299 }
300 }
301
302 let (byte_to_char, char_to_byte) = bytes_to_unicode();
303
304 tracing::info!(
305 "Tokenizer loaded: {} vocab, {} merges, {} added, eos={:?}",
306 vocab.len(),
307 ranks.len(),
308 added.len(),
309 eos_token_id
310 );
311
312 Ok(Self {
313 vocab,
314 id_to_token,
315 ranks,
316 added,
317 added_ids,
318 special_ids,
319 split_re,
320 metaspace,
321 sp_prepend,
322 nfc,
323 byte_to_char,
324 char_to_byte,
325 bos_token_id,
326 eos_token_id,
327 pad_token_id,
328 im_start_id,
329 im_end_id,
330 chat_template: None,
331 extra_eos: HashSet::new(),
332 add_bos: add_bos_detected || gemma_family,
333 })
334 }
335
336 pub fn byte_level() -> Self {
338 let mut vocab = HashMap::new();
339 let mut id_to_token = Vec::with_capacity(256);
340 for i in 0..256u32 {
341 let tok = format!("<0x{:02X}>", i);
342 vocab.insert(tok.clone(), i);
343 id_to_token.push(tok);
344 }
345 let (byte_to_char, char_to_byte) = bytes_to_unicode();
346 Self {
347 vocab,
348 id_to_token,
349 ranks: HashMap::new(),
350 added: Vec::new(),
351 added_ids: HashSet::new(),
352 special_ids: HashSet::new(),
353 split_re: None,
354 metaspace: false,
355 sp_prepend: false,
356 nfc: false,
357 byte_to_char,
358 char_to_byte,
359 bos_token_id: None,
360 eos_token_id: None,
361 pad_token_id: None,
362 im_start_id: None,
363 im_end_id: None,
364 chat_template: None,
365 extra_eos: HashSet::new(),
366 add_bos: false,
367 }
368 }
369
370 pub fn encode(&self, text: &str) -> Vec<u32> {
372 let mut ids = Vec::new();
373 let mut rest = text;
375 'outer: while !rest.is_empty() {
376 let mut best: Option<(usize, usize, u32)> = None; for (content, id) in &self.added {
378 if let Some(pos) = rest.find(content.as_str()) {
379 let better = match best {
380 None => true,
381 Some((bp, bl, _)) => pos < bp || (pos == bp && content.len() > bl),
382 };
383 if better {
384 best = Some((pos, content.len(), *id));
385 }
386 if pos == 0 {
387 break; }
389 }
390 }
391 match best {
392 Some((pos, len, id)) => {
393 self.encode_segment(&rest[..pos], &mut ids);
394 ids.push(id);
395 rest = &rest[pos + len..];
396 }
397 None => {
398 self.encode_segment(rest, &mut ids);
399 break 'outer;
400 }
401 }
402 }
403 ids
404 }
405
406 fn encode_segment(&self, segment: &str, out: &mut Vec<u32>) {
408 if segment.is_empty() {
409 return;
410 }
411 let norm: String = if self.nfc {
412 segment.nfc().collect()
413 } else {
414 segment.to_string()
415 };
416 if self.metaspace {
417 let sp = if self.sp_prepend {
421 format!("\u{2581}{}", norm).replace(' ', "\u{2581}")
422 } else {
423 norm.replace(' ', "\u{2581}")
424 };
425 self.bpe_piece_sp(&sp, out);
426 return;
427 }
428 match &self.split_re {
429 Some(re) => {
430 let mut last = 0;
431 for m in re.find_iter(&norm) {
432 let m = match m {
433 Ok(m) => m,
434 Err(e) => {
435 tracing::error!("pre-tokenizer regex failed: {e}");
436 break;
437 }
438 };
439 if m.start() > last {
440 self.bpe_piece(&norm[last..m.start()], out);
442 }
443 self.bpe_piece(m.as_str(), out);
444 last = m.end();
445 }
446 if last < norm.len() {
447 self.bpe_piece(&norm[last..], out);
448 }
449 }
450 None => {
451 for b in norm.bytes() {
453 let tok = format!("<0x{:02X}>", b);
454 if let Some(&id) = self.vocab.get(&tok) {
455 out.push(id);
456 }
457 }
458 }
459 }
460 }
461
462 fn bpe_piece_sp(&self, piece: &str, out: &mut Vec<u32>) {
465 if piece.is_empty() {
466 return;
467 }
468 let mut sym: Vec<String> = piece.chars().map(|c| c.to_string()).collect();
469 loop {
470 let mut best: Option<(u32, usize)> = None;
471 for i in 0..sym.len().saturating_sub(1) {
472 if let Some(&r) = self.ranks.get(&(sym[i].clone(), sym[i + 1].clone())) {
473 if best.map(|(br, _)| r < br).unwrap_or(true) {
474 best = Some((r, i));
475 }
476 }
477 }
478 let Some((_, i)) = best else { break };
479 let merged = format!("{}{}", sym[i], sym[i + 1]);
480 let (left, right) = (sym[i].clone(), sym[i + 1].clone());
481 let mut j = 0;
482 while j + 1 < sym.len() {
483 if sym[j] == left && sym[j + 1] == right {
484 sym[j] = merged.clone();
485 sym.remove(j + 1);
486 }
487 j += 1;
488 }
489 }
490 for t in &sym {
491 if let Some(&id) = self.vocab.get(t) {
492 out.push(id);
493 } else {
494 let mut ok = true;
495 for byte in t.bytes() {
496 let tok = format!("<0x{:02X}>", byte);
497 match self.vocab.get(&tok) {
498 Some(&id) => out.push(id),
499 None => {
500 ok = false;
501 break;
502 }
503 }
504 }
505 if !ok {
506 tracing::error!("tokenizer: no id for SP symbol {t:?} — dropped");
507 }
508 }
509 }
510 }
511
512 fn bpe_piece(&self, piece: &str, out: &mut Vec<u32>) {
514 if piece.is_empty() {
515 return;
516 }
517 let mapped: Vec<String> = piece
518 .bytes()
519 .map(|b| self.byte_to_char[b as usize].to_string())
520 .collect();
521 let mut sym = mapped;
522
523 loop {
525 let mut best: Option<(u32, usize)> = None;
526 for i in 0..sym.len().saturating_sub(1) {
527 if let Some(&r) = self.ranks.get(&(sym[i].clone(), sym[i + 1].clone())) {
528 if best.map(|(br, _)| r < br).unwrap_or(true) {
529 best = Some((r, i));
530 }
531 }
532 }
533 let Some((_, i)) = best else { break };
534 let merged = format!("{}{}", sym[i], sym[i + 1]);
535 let (left, right) = (sym[i].clone(), sym[i + 1].clone());
537 let mut j = 0;
538 while j + 1 < sym.len() {
539 if sym[j] == left && sym[j + 1] == right {
540 sym[j] = merged.clone();
541 sym.remove(j + 1);
542 }
543 j += 1;
544 }
545 }
546
547 for s in &sym {
548 if let Some(&id) = self.vocab.get(s) {
549 out.push(id);
550 } else {
551 let mut ok = true;
553 for ch in s.chars() {
554 let Some(&b) = self.char_to_byte.get(&ch) else {
555 ok = false;
556 break;
557 };
558 let tok = format!("<0x{:02X}>", b);
559 if let Some(&id) = self.vocab.get(&tok) {
560 out.push(id);
561 } else {
562 ok = false;
563 break;
564 }
565 }
566 if !ok {
567 tracing::error!("tokenizer: no id for symbol {s:?} — dropped");
568 }
569 }
570 }
571 }
572
573 pub fn decode(&self, ids: &[u32]) -> String {
576 let mut bytes: Vec<u8> = Vec::new();
577 for &id in ids {
578 if self.special_ids.contains(&id) {
579 continue;
580 }
581 let idx = id as usize;
582 if idx >= self.id_to_token.len() {
583 continue;
584 }
585 let tok = &self.id_to_token[idx];
586 if self.added_ids.contains(&id) {
587 if self.metaspace && tok.contains('\u{2581}') {
590 bytes.extend_from_slice(tok.replace('\u{2581}', " ").as_bytes());
591 } else {
592 bytes.extend_from_slice(tok.as_bytes());
593 }
594 continue;
595 }
596 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 bytes.push(b);
600 continue;
601 }
602 }
603 if self.metaspace {
604 for ch in tok.chars() {
606 if ch == '\u{2581}' {
607 bytes.push(b' ');
608 } else {
609 let mut buf = [0u8; 4];
610 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
611 }
612 }
613 continue;
614 }
615 for ch in tok.chars() {
616 match self.char_to_byte.get(&ch) {
617 Some(&b) => bytes.push(b),
618 None => {
621 let mut buf = [0u8; 4];
622 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
623 }
624 }
625 }
626 }
627 let text = String::from_utf8_lossy(&bytes).into_owned();
628 if self.metaspace && self.sp_prepend {
629 if let Some(stripped) = text.strip_prefix(' ') {
631 return stripped.to_string();
632 }
633 }
634 text
635 }
636
637 pub fn decode_token(&self, id: u32) -> String {
640 if self.special_ids.contains(&id) {
641 return String::new();
642 }
643 let idx = id as usize;
644 if idx >= self.id_to_token.len() {
645 return String::new();
646 }
647 let tok = &self.id_to_token[idx];
648 if self.added_ids.contains(&id) {
649 if self.metaspace && tok.contains('\u{2581}') {
650 return tok.replace('\u{2581}', " ");
651 }
652 return tok.clone();
653 }
654 if tok.starts_with("<0x") && tok.ends_with('>') && tok.len() == 6 {
655 if let Ok(b) = u8::from_str_radix(&tok[3..5], 16) {
656 return String::from_utf8_lossy(&[b]).into_owned();
657 }
658 }
659 if self.metaspace {
660 return tok.replace('\u{2581}', " ");
661 }
662 let mut bytes = Vec::new();
663 for ch in tok.chars() {
664 match self.char_to_byte.get(&ch) {
665 Some(&b) => bytes.push(b),
666 None => {
667 let mut buf = [0u8; 4];
668 bytes.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
669 }
670 }
671 }
672 String::from_utf8_lossy(&bytes).into_owned()
673 }
674
675 pub fn apply_chat_template(&self, messages: &[(String, String)]) -> Vec<u32> {
679 self.apply_chat_template_opts(messages, None)
680 }
681
682 pub fn apply_chat_template_opts(
687 &self,
688 messages: &[(String, String)],
689 enable_thinking: Option<bool>,
690 ) -> Vec<u32> {
691 if let Some(tpl) = &self.chat_template {
692 match self.render_template(tpl, messages, enable_thinking) {
693 Ok(text) => return self.with_bos(self.encode(&text)),
694 Err(e) => {
695 tracing::error!("chat template render failed ({e}); ChatML fallback");
696 }
697 }
698 }
699 self.with_bos(self.chatml_fallback_opts(messages, enable_thinking))
700 }
701
702 pub fn with_bos(&self, mut ids: Vec<u32>) -> Vec<u32> {
704 if self.add_bos {
705 if let Some(b) = self.bos_token_id {
706 if ids.first() != Some(&b) {
707 ids.insert(0, b);
708 }
709 }
710 }
711 ids
712 }
713
714 pub fn render_chat(&self, messages: &[(String, String)]) -> Option<String> {
716 self.render_chat_opts(messages, None)
717 }
718
719 pub fn render_chat_opts(
721 &self,
722 messages: &[(String, String)],
723 enable_thinking: Option<bool>,
724 ) -> Option<String> {
725 let tpl = self.chat_template.as_ref()?;
726 match self.render_template(tpl, messages, enable_thinking) {
727 Ok(t) => Some(t),
728 Err(e) => {
729 tracing::error!("chat template render: {e:#}");
730 None
731 }
732 }
733 }
734
735 fn render_template(
736 &self,
737 tpl: &str,
738 messages: &[(String, String)],
739 enable_thinking: Option<bool>,
740 ) -> Result<String, minijinja::Error> {
741 let mut env = minijinja::Environment::new();
742 env.set_trim_blocks(true);
743 env.set_lstrip_blocks(true);
744 env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
746 env.add_template("chat", tpl)?;
747 let msgs: Vec<minijinja::Value> = messages
748 .iter()
749 .map(|(role, content)| {
750 minijinja::context! { role => role, content => content }
751 })
752 .collect();
753 let rendered = match enable_thinking {
756 Some(v) => env.get_template("chat")?.render(minijinja::context! {
757 messages => msgs,
758 add_generation_prompt => true,
759 enable_thinking => v,
760 })?,
761 None => env.get_template("chat")?.render(minijinja::context! {
762 messages => msgs,
763 add_generation_prompt => true,
764 })?,
765 };
766 if enable_thinking == Some(false) && !rendered.contains("</think>") {
770 if let Some(pos) = rendered.rfind("assistant") {
771 let mut insert_at = pos + "assistant".len();
772 if let Some(idx) = rendered[insert_at..].find('\n') {
773 insert_at += idx + 1;
774 }
775 let mut out = String::with_capacity(rendered.len() + 24);
776 out.push_str(&rendered[..insert_at]);
777 if !out.ends_with('\n') {
778 out.push('\n');
779 }
780 out.push_str("<think>\n\n</think>\n\n");
781 out.push_str(&rendered[insert_at..]);
782 return Ok(out);
783 }
784 }
785 Ok(rendered)
786 }
787
788 fn chatml_fallback(&self, messages: &[(String, String)]) -> Vec<u32> {
790 self.chatml_fallback_opts(messages, None)
791 }
792
793 fn chatml_fallback_opts(
795 &self,
796 messages: &[(String, String)],
797 enable_thinking: Option<bool>,
798 ) -> Vec<u32> {
799 let mut tokens = Vec::new();
800
801 for (role, content) in messages {
802 if let Some(start_id) = self.im_start_id {
804 tokens.push(start_id);
805 }
806 tokens.extend(self.encode(&format!("{}\n{}", role, content)));
807 if let Some(end_id) = self.im_end_id {
808 tokens.push(end_id);
809 }
810 tokens.extend(self.encode("\n"));
811 }
812
813 if let Some(start_id) = self.im_start_id {
815 tokens.push(start_id);
816 }
817 tokens.extend(self.encode("assistant\n"));
818 if enable_thinking == Some(false) {
819 tokens.extend(self.encode("<think>\n\n</think>\n\n"));
820 }
821
822 tokens
823 }
824
825 pub fn vocab_size(&self) -> usize {
827 self.id_to_token.len()
828 }
829
830 pub fn is_eos(&self, id: u32) -> bool {
832 self.eos_token_id == Some(id) || self.im_end_id == Some(id) || self.extra_eos.contains(&id)
833 }
834}
835
836#[derive(Debug, thiserror::Error)]
837pub enum TokenizerError {
838 #[error("IO error: {0}")]
839 Io(String),
840 #[error("Parse error: {0}")]
841 Parse(String),
842}
843
844#[cfg(test)]
845mod tests {
846 use super::*;
847
848 #[test]
849 fn byte_unicode_bijection() {
850 let (b2c, c2b) = bytes_to_unicode();
851 for b in 0..=255u8 {
852 assert_eq!(c2b[&b2c[b as usize]], b);
853 }
854 assert_eq!(b2c[b' ' as usize], 'Ġ');
856 assert_eq!(b2c[b'\n' as usize], 'Ċ');
857 }
858
859 #[test]
860 fn byte_level_roundtrip_utf8() {
861 let tok = Tokenizer::byte_level();
862 let text = "hello 🌍 hi\n";
863 let ids = tok.encode(text);
864 assert_eq!(ids.len(), text.len()); assert_eq!(tok.decode(&ids), text);
866 }
867
868 fn mini_json() -> String {
871 let vocab: Vec<(&str, u32)> = vec![
873 ("h", 0),
874 ("e", 1),
875 ("l", 2),
876 ("o", 3),
877 ("Ġ", 4),
878 ("w", 5),
879 ("r", 6),
880 ("d", 7),
881 ("he", 8),
882 ("Ġw", 9),
883 ];
884 let vocab_json: String = vocab
885 .iter()
886 .map(|(t, i)| format!("\"{t}\": {i}"))
887 .collect::<Vec<_>>()
888 .join(", ");
889 format!(
890 r#"{{
891 "model": {{
892 "type": "BPE",
893 "vocab": {{ {vocab_json} }},
894 "merges": [["h", "e"], ["Ġ", "w"]]
895 }},
896 "added_tokens": [
897 {{"id": 10, "content": "<|eot|>", "special": true}}
898 ]
899 }}"#
900 )
901 }
902
903 #[test]
904 fn full_pipeline_merges_and_added_tokens() {
905 let tok = Tokenizer::from_json(&mini_json()).unwrap();
906 let ids = tok.encode("hello world");
908 assert_eq!(ids, vec![8, 2, 2, 3, 9, 3, 6, 2, 7]);
909 assert_eq!(tok.decode(&ids), "hello world");
910 let ids2 = tok.encode("he<|eot|>he");
912 assert_eq!(ids2, vec![8, 10, 8]);
913 assert_eq!(tok.decode(&ids2), "hehe");
914 }
915
916 #[test]
917 fn non_ascii_is_never_silently_dropped() {
918 let tok = Tokenizer::from_json(&mini_json()).unwrap();
919 let ids = tok.encode("hello");
922 assert!(!ids.is_empty());
923 }
924}