1use std::collections::HashSet;
5use std::path::Path;
6
7use base64::Engine as _;
8use rayon::prelude::*;
9use rustc_hash::FxHashMap;
10use tiktoken_rs::CoreBPE;
11
12use super::{
13 Encoding, Error, Result, TokenIdType,
14 traits::{DecodeResult, Decoder, Encoder, Tokenizer},
15};
16
17const DEFAULT_NUM_RESERVED_SPECIAL_TOKENS: u32 = 256;
20
21const KIMI_PATTERN: &str = r#"[\p{Han}]+|[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]*[\p{Ll}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]+(?i:'s|'t|'re|'ve|'m|'ll|'d)?|[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]+[\p{Ll}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]*(?i:'s|'t|'re|'ve|'m|'ll|'d)?|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+"#;
23
24pub struct TikTokenTokenizer {
25 bpe: CoreBPE,
26 special_token_ids: HashSet<u32>,
27}
28
29impl TikTokenTokenizer {
30 pub fn from_file(
37 path: &str,
38 pattern: &str,
39 special_tokens: FxHashMap<String, u32>,
40 ) -> Result<Self> {
41 let encoder = parse_tiktoken_file(path)?;
42 let special_token_ids: HashSet<u32> = special_tokens.values().copied().collect();
43
44 let bpe = CoreBPE::new(encoder, special_tokens, pattern)
45 .map_err(|err| Error::msg(format!("Error creating tiktoken BPE: {err}")))?;
46
47 Ok(Self {
48 bpe,
49 special_token_ids,
50 })
51 }
52
53 pub fn from_file_auto(path: &str) -> Result<Self> {
58 let file_path = Path::new(path);
59 let directory = file_path
60 .parent()
61 .ok_or_else(|| Error::msg("Cannot determine parent directory of tiktoken file"))?;
62
63 let pattern = detect_bpe_pattern(directory)?;
64 let encoder = parse_tiktoken_file(path)?;
65 let num_base_tokens = encoder.values().max().map_or(0, |&m| m + 1) as usize;
67 let special_tokens = load_special_tokens(directory, num_base_tokens)?;
68 let special_token_ids: HashSet<u32> = special_tokens.values().copied().collect();
69
70 let bpe = CoreBPE::new(encoder, special_tokens, pattern)
71 .map_err(|err| Error::msg(format!("Error creating tiktoken BPE: {err}")))?;
72
73 Ok(Self {
74 bpe,
75 special_token_ids,
76 })
77 }
78}
79
80impl Encoder for TikTokenTokenizer {
81 fn encode(&self, input: &str) -> Result<Encoding> {
82 let token_ids: Vec<u32> = self.bpe.encode_with_special_tokens(input);
83 Ok(Encoding::Sp(token_ids))
84 }
85
86 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
87 inputs.par_iter().map(|input| self.encode(input)).collect()
88 }
89}
90
91impl Decoder for TikTokenTokenizer {
92 fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
93 let ids: Vec<u32> = if skip_special_tokens {
94 token_ids
95 .iter()
96 .filter(|&&id| !self.special_token_ids.contains(&id))
97 .copied()
98 .collect()
99 } else {
100 token_ids.to_vec()
101 };
102
103 let bytes: Vec<u8> = self.bpe._decode_native_and_split(ids).flatten().collect();
112 match String::from_utf8(bytes) {
113 Ok(text) => Ok(DecodeResult::Complete(text)),
114 Err(e) => {
115 let text = String::from_utf8_lossy(e.as_bytes()).into_owned();
116 Ok(DecodeResult::from_decoded(text))
117 }
118 }
119 }
120}
121
122impl Tokenizer for TikTokenTokenizer {}
123
124fn parse_tiktoken_file(path: &str) -> Result<FxHashMap<Vec<u8>, u32>> {
126 let contents = std::fs::read_to_string(path)
127 .map_err(|err| Error::msg(format!("Failed to read tiktoken file '{path}': {err}")))?;
128
129 let engine = base64::engine::general_purpose::STANDARD;
130 let mut encoder = FxHashMap::default();
131
132 for line in contents.lines() {
133 let line = line.trim();
134 if line.is_empty() {
135 continue;
136 }
137 let mut parts = line.split_whitespace();
138 let token_b64 = parts
139 .next()
140 .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no token): {line}")))?;
141 let rank_str = parts
142 .next()
143 .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no rank): {line}")))?;
144
145 let token_bytes = engine
146 .decode(token_b64)
147 .map_err(|err| Error::msg(format!("Invalid base64 in tiktoken file: {err}")))?;
148 let rank: u32 = rank_str
149 .parse()
150 .map_err(|err| Error::msg(format!("Invalid rank in tiktoken file: {err}")))?;
151
152 encoder.insert(token_bytes, rank);
153 }
154
155 Ok(encoder)
156}
157
158fn detect_bpe_pattern(directory: &Path) -> Result<&'static str> {
160 let model_type: String = crate::file_json_field(&directory.join("config.json"), "model_type")
161 .map_err(|err| {
162 Error::msg(format!("Failed to read model_type from config.json: {err}"))
163 })?;
164
165 match model_type.as_str() {
166 "kimi" | "kimi_k2" | "kimi_k25" | "deepseek_v3" => Ok(KIMI_PATTERN),
172 _ => Err(Error::msg(format!(
173 "Unsupported tiktoken model_type '{model_type}'. \
174 Currently supported: kimi, kimi_k2, kimi_k25, deepseek_v3. \
175 To add a new model type, extend detect_bpe_pattern() in lib/tokenizers/src/tiktoken.rs \
176 with the appropriate BPE regex pattern. \
177 Alternatively, provide a tokenizer.json (HuggingFace format) instead."
178 ))),
179 }
180}
181
182fn load_special_tokens(directory: &Path, num_base_tokens: usize) -> Result<FxHashMap<String, u32>> {
187 let config_path = directory.join("tokenizer_config.json");
188 let mut special_tokens = FxHashMap::default();
189
190 if !config_path.exists() {
191 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
193 let id = num_base_tokens as u32 + i;
194 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
195 }
196 return Ok(special_tokens);
197 }
198
199 let contents = std::fs::read_to_string(&config_path)
200 .map_err(|err| Error::msg(format!("Failed to read tokenizer_config.json: {err}")))?;
201
202 let config: serde_json::Value = serde_json::from_str(&contents)
203 .map_err(|err| Error::msg(format!("Failed to parse tokenizer_config.json: {err}")))?;
204
205 if let Some(added_tokens) = config
206 .get("added_tokens_decoder")
207 .and_then(|v| v.as_object())
208 {
209 for (id_str, token_def) in added_tokens {
210 let id: u32 = id_str.parse().map_err(|err| {
211 Error::msg(format!(
212 "Invalid token ID '{id_str}' in added_tokens_decoder: {err}"
213 ))
214 })?;
215
216 let content = token_def
217 .get("content")
218 .and_then(|v| v.as_str())
219 .unwrap_or_else(|| {
220 tracing::warn!("Missing 'content' field for token ID {id}");
222 ""
223 });
224
225 if !content.is_empty() {
226 special_tokens.insert(content.to_string(), id);
227 }
228 }
229
230 let used_ids: HashSet<u32> = special_tokens.values().copied().collect();
232 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
233 let id = num_base_tokens as u32 + i;
234 if !used_ids.contains(&id) {
235 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
236 }
237 }
238 } else {
239 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
241 let id = num_base_tokens as u32 + i;
242 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
243 }
244 }
245
246 Ok(special_tokens)
247}
248
249#[cfg(test)]
250mod tests {
251 use super::*;
252 use crate::DecodeStream;
253 use std::io::Write;
254 use std::sync::Arc;
255
256 fn create_test_tiktoken_file(dir: &Path) -> String {
257 let engine = base64::engine::general_purpose::STANDARD;
258 let mut content = String::new();
259
260 let tokens: Vec<(&[u8], u32)> = vec![
262 (b"h", 0),
263 (b"e", 1),
264 (b"l", 2),
265 (b"o", 3),
266 (b" ", 4),
267 (b"w", 5),
268 (b"r", 6),
269 (b"d", 7),
270 (b"he", 8),
271 (b"ll", 9),
272 (b"lo", 10),
273 (b"wo", 11),
274 (b"rl", 12),
275 (b"hel", 13),
276 (b"llo", 14),
277 (b"wor", 15),
278 (b"hell", 16),
279 (b"ello", 17),
280 (b"worl", 18),
281 (b"hello", 19),
282 (b"world", 20),
283 ];
284
285 for (token, rank) in tokens {
286 let encoded = engine.encode(token);
287 content.push_str(&format!("{encoded} {rank}\n"));
288 }
289
290 let file_path = dir.join("tiktoken.model");
291 let mut file = std::fs::File::create(&file_path).unwrap();
292 file.write_all(content.as_bytes()).unwrap();
293 file_path.to_str().unwrap().to_string()
294 }
295
296 fn create_test_config(dir: &Path, model_type: &str) {
297 let config = serde_json::json!({
298 "model_type": model_type,
299 "max_position_embeddings": 32768,
300 "eos_token_id": [21]
301 });
302 let file_path = dir.join("config.json");
303 std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
304 }
305
306 fn create_test_tokenizer_config(dir: &Path, num_base_tokens: usize) {
307 let mut added_tokens = serde_json::Map::new();
308 let bos_id = num_base_tokens;
309 let eos_id = num_base_tokens + 1;
310
311 added_tokens.insert(
312 bos_id.to_string(),
313 serde_json::json!({"content": "[BOS]", "special": true}),
314 );
315 added_tokens.insert(
316 eos_id.to_string(),
317 serde_json::json!({"content": "[EOS]", "special": true}),
318 );
319
320 let config = serde_json::json!({
321 "added_tokens_decoder": added_tokens
322 });
323
324 let file_path = dir.join("tokenizer_config.json");
325 std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
326 }
327
328 #[test]
329 fn test_parse_tiktoken_file() {
330 let dir = tempfile::tempdir().unwrap();
331 let file_path = create_test_tiktoken_file(dir.path());
332 let encoder = parse_tiktoken_file(&file_path).unwrap();
333 assert_eq!(encoder.len(), 21);
334 assert_eq!(encoder[b"hello".as_slice()], 19);
335 assert_eq!(encoder[b"world".as_slice()], 20);
336 }
337
338 #[test]
339 fn test_parse_tiktoken_file_missing() {
340 let result = parse_tiktoken_file("/nonexistent/path/tiktoken.model");
341 assert!(result.is_err());
342 }
343
344 #[test]
345 fn test_tiktoken_from_file() {
346 let dir = tempfile::tempdir().unwrap();
347 let file_path = create_test_tiktoken_file(dir.path());
348
349 let mut special_tokens = FxHashMap::default();
350 special_tokens.insert("[BOS]".to_string(), 21_u32);
351 special_tokens.insert("[EOS]".to_string(), 22_u32);
352
353 let pattern = r"[\w]+|[^\w\s]+|\s+";
355
356 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
357
358 let encoding = tokenizer.encode("hello world").unwrap();
360 let ids = encoding.token_ids();
361 assert!(!ids.is_empty());
362
363 let decoded: String = tokenizer.decode(ids, false).unwrap().into();
365 assert_eq!(decoded, "hello world");
366 }
367
368 #[test]
369 fn test_tiktoken_encoding_variant() {
370 let dir = tempfile::tempdir().unwrap();
371 let file_path = create_test_tiktoken_file(dir.path());
372
373 let special_tokens = FxHashMap::default();
374 let pattern = r"[\w]+|[^\w\s]+|\s+";
375
376 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
377 let encoding = tokenizer.encode("hello").unwrap();
378
379 match &encoding {
381 Encoding::Sp(_) => {}
382 other => panic!("Expected Encoding::Sp, got {:?}", other),
383 }
384 }
385
386 #[test]
387 fn test_tiktoken_skip_special_tokens() {
388 let dir = tempfile::tempdir().unwrap();
389 let file_path = create_test_tiktoken_file(dir.path());
390
391 let mut special_tokens = FxHashMap::default();
392 special_tokens.insert("[BOS]".to_string(), 21_u32);
393 special_tokens.insert("[EOS]".to_string(), 22_u32);
394
395 let pattern = r"[\w]+|[^\w\s]+|\s+";
396
397 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
398
399 let encoding = tokenizer.encode("hello").unwrap();
401 let mut ids = vec![21u32]; ids.extend(encoding.token_ids());
403 ids.push(22); let decoded_skip: String = tokenizer.decode(&ids, true).unwrap().into();
407 assert_eq!(decoded_skip, "hello");
408
409 let decoded_all: String = tokenizer.decode(&ids, false).unwrap().into();
411 assert!(decoded_all.contains("hello"));
412 }
413
414 #[test]
415 fn test_tiktoken_from_file_auto() {
416 let dir = tempfile::tempdir().unwrap();
417 let file_path = create_test_tiktoken_file(dir.path());
418
419 create_test_config(dir.path(), "kimi");
420 create_test_tokenizer_config(dir.path(), 21);
421
422 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
423
424 let encoding = tokenizer.encode("hello world").unwrap();
426 let ids = encoding.token_ids();
427 assert!(!ids.is_empty());
428
429 let decoded: String = tokenizer.decode(ids, false).unwrap().into();
430 assert_eq!(decoded, "hello world");
431 }
432
433 #[test]
434 fn test_detect_bpe_pattern_unknown() {
435 let dir = tempfile::tempdir().unwrap();
436 create_test_config(dir.path(), "unknown_model");
437 let result = detect_bpe_pattern(dir.path());
438 assert!(result.is_err());
439 }
440
441 #[test]
442 fn test_load_special_tokens_no_config() {
443 let dir = tempfile::tempdir().unwrap();
444 let tokens = load_special_tokens(dir.path(), 100).unwrap();
445 assert_eq!(tokens.len(), 256);
446 assert_eq!(tokens["<|reserved_token_100|>"], 100);
447 assert_eq!(tokens["<|reserved_token_355|>"], 355);
448 }
449
450 #[test]
451 fn test_load_special_tokens_with_config() {
452 let dir = tempfile::tempdir().unwrap();
453 create_test_tokenizer_config(dir.path(), 100);
454 let tokens = load_special_tokens(dir.path(), 100).unwrap();
455 assert_eq!(tokens["[BOS]"], 100);
456 assert_eq!(tokens["[EOS]"], 101);
457 assert!(tokens.len() > 2);
459 }
460
461 fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
463 let engine = base64::engine::general_purpose::STANDARD;
464 let mut content = String::new();
465
466 let tokens: Vec<(&[u8], u32)> = vec![
467 (b"h", 0),
468 (b"e", 1),
469 (b"l", 2),
470 (b"o", 3),
471 (b" ", 4),
472 (b"hello", 5),
473 ];
474
475 for (token, rank) in &tokens {
476 let encoded = engine.encode(token);
477 content.push_str(&format!("{encoded} {rank}\n"));
478 }
479
480 let byte_tokens: Vec<(Vec<u8>, u32)> =
483 vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
484
485 for (token, rank) in &byte_tokens {
486 let encoded = engine.encode(token);
487 content.push_str(&format!("{encoded} {rank}\n"));
488 }
489
490 let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
492 (vec![0xF0], 200),
493 (vec![0x9F], 201),
494 (vec![0x98], 202),
495 (vec![0x80], 203),
496 ];
497
498 for (token, rank) in &emoji_tokens {
499 let encoded = engine.encode(token);
500 content.push_str(&format!("{encoded} {rank}\n"));
501 }
502
503 let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
506
507 for (token, rank) in &fffd_token {
508 let encoded = engine.encode(token);
509 content.push_str(&format!("{encoded} {rank}\n"));
510 }
511
512 let file_path = dir.join("tiktoken.model");
513 let mut file = std::fs::File::create(&file_path).unwrap();
514 file.write_all(content.as_bytes()).unwrap();
515 file_path.to_str().unwrap().to_string()
516 }
517
518 fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
519 let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
520 let special_tokens = FxHashMap::default();
521 let pattern = r"[\w]+|[^\w\s]+|\s+";
522 TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
523 }
524
525 #[test]
529 fn test_decode_single_incomplete_utf8_byte_does_not_error() {
530 let dir = tempfile::tempdir().unwrap();
531 let tokenizer = create_byte_token_tokenizer(dir.path());
532
533 let result = tokenizer.decode(&[100], false);
534 assert!(
535 result.is_ok(),
536 "decode() should not error on incomplete UTF-8 bytes"
537 );
538 let decode_result = result.unwrap();
539 assert!(
540 decode_result.is_partial(),
541 "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
542 decode_result
543 );
544 }
545
546 #[test]
548 fn test_decode_two_of_three_utf8_bytes_does_not_error() {
549 let dir = tempfile::tempdir().unwrap();
550 let tokenizer = create_byte_token_tokenizer(dir.path());
551
552 let result = tokenizer.decode(&[100, 101], false);
553 assert!(result.is_ok());
554 let decode_result = result.unwrap();
555 assert!(
556 decode_result.is_partial(),
557 "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
558 decode_result
559 );
560 }
561
562 #[test]
566 fn test_decode_complete_multibyte_utf8_produces_correct_char() {
567 let dir = tempfile::tempdir().unwrap();
568 let tokenizer = create_byte_token_tokenizer(dir.path());
569
570 let result = tokenizer.decode(&[100, 101, 102], false);
571 assert!(result.is_ok());
572 assert_eq!(String::from(result.unwrap()), "你");
573 }
574
575 #[test]
578 fn test_decode_complete_4byte_emoji_from_byte_tokens() {
579 let dir = tempfile::tempdir().unwrap();
580 let tokenizer = create_byte_token_tokenizer(dir.path());
581
582 let result = tokenizer.decode(&[200, 201, 202, 203], false);
583 assert!(result.is_ok());
584 assert_eq!(String::from(result.unwrap()), "😀");
585 }
586
587 #[test]
592 fn test_decode_legitimate_replacement_char_token_is_complete() {
593 let dir = tempfile::tempdir().unwrap();
594 let tokenizer = create_byte_token_tokenizer(dir.path());
595
596 let result = tokenizer.decode(&[300], false);
597 assert!(result.is_ok());
598 let decode_result = result.unwrap();
599 assert!(
600 decode_result.is_complete(),
601 "legitimate U+FFFD vocab token must be Complete, got: {:?}",
602 decode_result
603 );
604 assert_eq!(decode_result.as_str(), "\u{FFFD}");
605 }
606
607 #[test]
609 fn test_decode_partial_emoji_does_not_error() {
610 let dir = tempfile::tempdir().unwrap();
611 let tokenizer = create_byte_token_tokenizer(dir.path());
612
613 let result = tokenizer.decode(&[200], false);
614 assert!(result.is_ok());
615 assert!(result.unwrap().is_partial());
616 }
617
618 #[test]
620 fn test_decode_mixed_ascii_and_incomplete_bytes() {
621 let dir = tempfile::tempdir().unwrap();
622 let tokenizer = create_byte_token_tokenizer(dir.path());
623
624 let result = tokenizer.decode(&[5, 100], false);
625 assert!(result.is_ok());
626 let decode_result = result.unwrap();
627 assert!(
628 decode_result.is_partial(),
629 "trailing incomplete byte should produce DecodeResult::Partial"
630 );
631 let text: String = decode_result.into();
632 assert!(
633 text.starts_with("hello"),
634 "should start with 'hello', got: {:?}",
635 text
636 );
637 }
638
639 #[test]
643 fn test_decode_stream_incremental_multibyte_reassembly() {
644 let dir = tempfile::tempdir().unwrap();
645 let tokenizer = create_byte_token_tokenizer(dir.path());
646 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
647
648 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
649
650 let r1 = stream.step(100).unwrap();
651 assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
652
653 let r2 = stream.step(101).unwrap();
654 assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
655
656 let r3 = stream.step(102).unwrap();
657 assert!(r3.is_some(), "third byte should complete the character");
658 assert_eq!(r3.unwrap(), "你");
659 }
660
661 #[test]
663 fn test_decode_stream_incremental_emoji_reassembly() {
664 let dir = tempfile::tempdir().unwrap();
665 let tokenizer = create_byte_token_tokenizer(dir.path());
666 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
667
668 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
669
670 let r1 = stream.step(200).unwrap();
671 assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
672
673 let r2 = stream.step(201).unwrap();
674 assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
675
676 let r3 = stream.step(202).unwrap();
677 assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
678
679 let r4 = stream.step(203).unwrap();
680 assert!(r4.is_some(), "byte 4/4 should complete the emoji");
681 assert_eq!(r4.unwrap(), "😀");
682 }
683
684 #[test]
685 fn test_tiktoken_encode_batch() {
686 let dir = tempfile::tempdir().unwrap();
687 let file_path = create_test_tiktoken_file(dir.path());
688
689 let special_tokens = FxHashMap::default();
690 let pattern = r"[\w]+|[^\w\s]+|\s+";
691
692 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
693
694 let inputs = &["hello", "world"];
695 let encodings = tokenizer.encode_batch(inputs).unwrap();
696 assert_eq!(encodings.len(), 2);
697
698 for (encoding, input) in encodings.iter().zip(inputs.iter()) {
699 let decoded: String = tokenizer
700 .decode(encoding.token_ids(), false)
701 .unwrap()
702 .into();
703 assert_eq!(decoded, *input);
704 }
705 }
706
707 fn create_byte_level_tiktoken_file(dir: &Path) -> String {
710 let engine = base64::engine::general_purpose::STANDARD;
711 let mut content = String::new();
712 for byte_val in 0u16..256 {
713 let encoded = engine.encode([byte_val as u8]);
714 content.push_str(&format!("{encoded} {byte_val}\n"));
715 }
716 let file_path = dir.join("tiktoken.model");
717 std::fs::write(&file_path, &content).unwrap();
718 file_path.to_str().unwrap().to_string()
719 }
720
721 #[test]
732 fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
733 let dir = tempfile::tempdir().unwrap();
734 let file_path = create_byte_level_tiktoken_file(dir.path());
735
736 create_test_config(dir.path(), "kimi");
738
739 create_test_tokenizer_config(dir.path(), 256);
741
742 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
743
744 let single = "<|reserved_token_258|>";
751 let enc = tokenizer.encode(single).unwrap();
752 assert_eq!(
753 enc.token_ids().len(),
754 1,
755 "'{single}' should be 1 special token, got {} tokens: {:?}. \
756 This means fallback naming still uses relative offsets instead of absolute IDs.",
757 enc.token_ids().len(),
758 enc.token_ids()
759 );
760 assert_eq!(enc.token_ids()[0], 258);
761
762 let multi: String = (258u32..268)
765 .map(|id| format!("<|reserved_token_{id}|>"))
766 .collect();
767 let enc_multi = tokenizer.encode(&multi).unwrap();
768 assert_eq!(
769 enc_multi.token_ids().len(),
770 10,
771 "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
772 enc_multi.token_ids().len(),
773 enc_multi.token_ids()
774 );
775 let expected_ids: Vec<u32> = (258..268).collect();
776 assert_eq!(enc_multi.token_ids(), &expected_ids);
777 }
778
779 #[test]
783 fn test_relative_offset_naming_causes_inflation() {
784 let dir = tempfile::tempdir().unwrap();
785 let file_path = create_byte_level_tiktoken_file(dir.path());
786
787 let _encoder = parse_tiktoken_file(&file_path).unwrap();
788 let num_base_tokens = 256usize;
789
790 let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
792 bad_special_tokens.insert("[BOS]".to_string(), 256);
793 bad_special_tokens.insert("[EOS]".to_string(), 257);
794 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
795 let id = num_base_tokens as u32 + i;
796 if id != 256 && id != 257 {
797 bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
799 }
800 }
801
802 let bad_tokenizer =
803 TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
804
805 let input = "<|reserved_token_258|>";
808 let enc = bad_tokenizer.encode(input).unwrap();
809 assert!(
810 enc.token_ids().len() > 1,
811 "With buggy relative-offset naming, '{}' should NOT be recognized as a \
812 single special token. Got {} token(s): {:?}",
813 input,
814 enc.token_ids().len(),
815 enc.token_ids()
816 );
817
818 let multi: String = (258u32..268)
820 .map(|id| format!("<|reserved_token_{id}|>"))
821 .collect();
822 let enc_multi = bad_tokenizer.encode(&multi).unwrap();
823 assert!(
824 enc_multi.token_ids().len() > 10,
825 "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
826 Got {}",
827 enc_multi.token_ids().len(),
828 );
829 }
830}