Skip to main content

dynamo_tokenizers/
tiktoken.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use 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    EncodeSegment, Encoding, Error, Result, TokenIdType,
14    traits::{DecodeResult, Decoder, Encoder, Tokenizer},
15};
16
17/// Number of reserved special-token slots to generate when filling gaps in the vocabulary.
18/// Most tiktoken-based models reserve 256 IDs above the base vocabulary for special tokens.
19const DEFAULT_NUM_RESERVED_SPECIAL_TOKENS: u32 = 256;
20
21/// Kimi BPE pattern from moonshotai/Kimi-K2-Instruct/tokenization_kimi.py
22const 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    special_tokens: Vec<String>,
28}
29
30fn sorted_special_token_strings(special_tokens: &FxHashMap<String, u32>) -> Vec<String> {
31    let mut strings: Vec<String> = special_tokens.keys().cloned().collect();
32    strings.sort();
33    strings.dedup();
34    strings
35}
36
37impl TikTokenTokenizer {
38    /// Create a TikTokenTokenizer from a tiktoken model file.
39    ///
40    /// # Arguments
41    /// * `path` - Path to the `.model` or `.tiktoken` file (base64 rank-per-line format)
42    /// * `pattern` - BPE regex pattern string
43    /// * `special_tokens` - Map of special token strings to their IDs
44    pub fn from_file(
45        path: &str,
46        pattern: &str,
47        special_tokens: FxHashMap<String, u32>,
48    ) -> Result<Self> {
49        let encoder = parse_tiktoken_file(path)?;
50        let special_token_ids: HashSet<u32> = special_tokens.values().copied().collect();
51        let special_token_strings = sorted_special_token_strings(&special_tokens);
52
53        let bpe = CoreBPE::new(encoder, special_tokens, pattern)
54            .map_err(|err| Error::msg(format!("Error creating tiktoken BPE: {err}")))?;
55
56        Ok(Self {
57            bpe,
58            special_token_ids,
59            special_tokens: special_token_strings,
60        })
61    }
62
63    /// Create a TikTokenTokenizer from a tiktoken model file, auto-detecting
64    /// the BPE pattern from `config.json` and special tokens from `tokenizer_config.json`.
65    ///
66    /// The tiktoken file and config files must be in the same directory.
67    pub fn from_file_auto(path: &str) -> Result<Self> {
68        let file_path = Path::new(path);
69        let directory = file_path
70            .parent()
71            .ok_or_else(|| Error::msg("Cannot determine parent directory of tiktoken file"))?;
72
73        let pattern = detect_bpe_pattern(directory)?;
74        let encoder = parse_tiktoken_file(path)?;
75        // Use max rank + 1 (not len) to avoid ID collisions with sparse/non-contiguous ranks
76        let num_base_tokens = encoder.values().max().map_or(0, |&m| m + 1) as usize;
77        let special_tokens = load_special_tokens(directory, num_base_tokens)?;
78        let special_token_ids: HashSet<u32> = special_tokens.values().copied().collect();
79        let special_token_strings = sorted_special_token_strings(&special_tokens);
80
81        let bpe = CoreBPE::new(encoder, special_tokens, pattern)
82            .map_err(|err| Error::msg(format!("Error creating tiktoken BPE: {err}")))?;
83
84        Ok(Self {
85            bpe,
86            special_token_ids,
87            special_tokens: special_token_strings,
88        })
89    }
90
91    /// Atomic special-token strings registered with the underlying TikToken BPE.
92    ///
93    /// [`CoreBPE::encode_with_special_tokens`] recognizes these strings outside ordinary
94    /// BPE. For non-overlapping token spellings, splitting immediately after a match
95    /// preserves tokenization.
96    pub fn special_tokens(&self) -> &[String] {
97        &self.special_tokens
98    }
99}
100
101impl Encoder for TikTokenTokenizer {
102    fn encode(&self, input: &str) -> Result<Encoding> {
103        let token_ids: Vec<u32> = self.bpe.encode_with_special_tokens(input);
104        Ok(Encoding::Sp(token_ids))
105    }
106
107    fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
108        inputs.par_iter().map(|input| self.encode(input)).collect()
109    }
110
111    fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
112        let mut token_ids = Vec::new();
113        for segment in segments {
114            if segment.allow_special {
115                token_ids.extend(self.bpe.encode_with_special_tokens(segment.text));
116            } else {
117                token_ids.extend(self.bpe.encode_ordinary(segment.text));
118            }
119        }
120        Ok(Encoding::Sp(token_ids))
121    }
122}
123
124impl Decoder for TikTokenTokenizer {
125    fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
126        let ids: Vec<u32> = if skip_special_tokens {
127            token_ids
128                .iter()
129                .filter(|&&id| !self.special_token_ids.contains(&id))
130                .copied()
131                .collect()
132        } else {
133            token_ids.to_vec()
134        };
135
136        // Try strict UTF-8 first: valid bytes get `Complete` with zero extra allocation
137        // (takes ownership of the Vec). This correctly handles vocabulary tokens whose
138        // raw bytes are EF BF BD (legitimate U+FFFD) -- they are valid UTF-8 and must
139        // not be confused with incomplete multi-byte sequences.
140        //
141        // On failure, fall back to lossy conversion so partial multi-byte sequences
142        // become U+FFFD, then classify via the trailing-FFFD heuristic. This path is
143        // only hit during incremental detokenization of byte-fallback tokens.
144        let bytes: Vec<u8> = self.bpe._decode_native_and_split(ids).flatten().collect();
145        match String::from_utf8(bytes) {
146            Ok(text) => Ok(DecodeResult::Complete(text)),
147            Err(e) => {
148                let text = String::from_utf8_lossy(e.as_bytes()).into_owned();
149                Ok(DecodeResult::from_decoded(text))
150            }
151        }
152    }
153}
154
155impl Tokenizer for TikTokenTokenizer {
156    fn validate_prefix_cache(&self) -> Result<()> {
157        Ok(())
158    }
159}
160
161/// Parse a tiktoken model file (base64-encoded token + rank per line).
162fn parse_tiktoken_file(path: &str) -> Result<FxHashMap<Vec<u8>, u32>> {
163    let contents = std::fs::read_to_string(path)
164        .map_err(|err| Error::msg(format!("Failed to read tiktoken file '{path}': {err}")))?;
165
166    let engine = base64::engine::general_purpose::STANDARD;
167    let mut encoder = FxHashMap::default();
168
169    for line in contents.lines() {
170        let line = line.trim();
171        if line.is_empty() {
172            continue;
173        }
174        let mut parts = line.split_whitespace();
175        let token_b64 = parts
176            .next()
177            .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no token): {line}")))?;
178        let rank_str = parts
179            .next()
180            .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no rank): {line}")))?;
181
182        let token_bytes = engine
183            .decode(token_b64)
184            .map_err(|err| Error::msg(format!("Invalid base64 in tiktoken file: {err}")))?;
185        let rank: u32 = rank_str
186            .parse()
187            .map_err(|err| Error::msg(format!("Invalid rank in tiktoken file: {err}")))?;
188
189        encoder.insert(token_bytes, rank);
190    }
191
192    Ok(encoder)
193}
194
195/// Detect the BPE pattern for a model by reading `model_type` from `config.json`.
196fn detect_bpe_pattern(directory: &Path) -> Result<&'static str> {
197    let model_type: String = crate::file_json_field(&directory.join("config.json"), "model_type")
198        .map_err(|err| {
199        Error::msg(format!("Failed to read model_type from config.json: {err}"))
200    })?;
201
202    match model_type.as_str() {
203        // baseten-admin/Kimi-2.5-text-nvfp4-v3 model has model_type: "deepseek_v3" in its config.json
204        // because Kimi K2.5 is built on the DeepSeek V3 architecture.
205        // it still ships the Kimi tiktoken tokenizer file, so the KIMI_PATTERN BPE regex is the
206        // correct pattern to use.  No pure DeepSeek V3 model uses tiktoken.model files
207        // (they use tokenizer.json instead) so this match is safe.
208        "kimi" | "kimi_k2" | "kimi_k25" | "kimi_k3" | "kimi_linear" | "deepseek_v3" => {
209            Ok(KIMI_PATTERN)
210        }
211        _ => Err(Error::msg(format!(
212            "Unsupported tiktoken model_type '{model_type}'. \
213             Currently supported: kimi, kimi_k2, kimi_k25, kimi_k3, kimi_linear, deepseek_v3. \
214             To add a new model type, extend detect_bpe_pattern() in lib/tokenizers/src/tiktoken.rs \
215             with the appropriate BPE regex pattern. \
216             Alternatively, provide a tokenizer.json (HuggingFace format) instead."
217        ))),
218    }
219}
220
221/// Load special tokens from `tokenizer_config.json` in the model directory.
222///
223/// Reads the `added_tokens_decoder` field which maps string token IDs to token definitions.
224/// Falls back to generating `<|reserved_token_{id}|>` names for unmapped IDs.
225fn load_special_tokens(directory: &Path, num_base_tokens: usize) -> Result<FxHashMap<String, u32>> {
226    let config_path = directory.join("tokenizer_config.json");
227    let mut special_tokens = FxHashMap::default();
228
229    if !config_path.exists() {
230        // No tokenizer_config.json β€” generate default reserved tokens
231        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
232            let id = num_base_tokens as u32 + i;
233            special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
234        }
235        return Ok(special_tokens);
236    }
237
238    let contents = std::fs::read_to_string(&config_path)
239        .map_err(|err| Error::msg(format!("Failed to read tokenizer_config.json: {err}")))?;
240
241    let config: serde_json::Value = serde_json::from_str(&contents)
242        .map_err(|err| Error::msg(format!("Failed to parse tokenizer_config.json: {err}")))?;
243
244    if let Some(added_tokens) = config
245        .get("added_tokens_decoder")
246        .and_then(|v| v.as_object())
247    {
248        for (id_str, token_def) in added_tokens {
249            let id: u32 = id_str.parse().map_err(|err| {
250                Error::msg(format!(
251                    "Invalid token ID '{id_str}' in added_tokens_decoder: {err}"
252                ))
253            })?;
254
255            let content = token_def
256                .get("content")
257                .and_then(|v| v.as_str())
258                .unwrap_or_else(|| {
259                    // This shouldn't happen in well-formed configs, but handle gracefully
260                    tracing::warn!("Missing 'content' field for token ID {id}");
261                    ""
262                });
263
264            if !content.is_empty() {
265                special_tokens.insert(content.to_string(), id);
266            }
267        }
268
269        // Fill in any gaps with reserved tokens for the expected range
270        let used_ids: HashSet<u32> = special_tokens.values().copied().collect();
271        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
272            let id = num_base_tokens as u32 + i;
273            if !used_ids.contains(&id) {
274                special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
275            }
276        }
277    } else {
278        // No added_tokens_decoder β€” generate default reserved tokens
279        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
280            let id = num_base_tokens as u32 + i;
281            special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
282        }
283    }
284
285    Ok(special_tokens)
286}
287
288#[cfg(test)]
289mod tests {
290    use super::*;
291    use crate::DecodeStream;
292    use std::io::Write;
293    use std::sync::Arc;
294
295    fn create_test_tiktoken_file(dir: &Path) -> String {
296        let engine = base64::engine::general_purpose::STANDARD;
297        let mut content = String::new();
298
299        // Create some simple token entries: single bytes with sequential ranks
300        let tokens: Vec<(&[u8], u32)> = vec![
301            (b"h", 0),
302            (b"e", 1),
303            (b"l", 2),
304            (b"o", 3),
305            (b" ", 4),
306            (b"w", 5),
307            (b"r", 6),
308            (b"d", 7),
309            (b"he", 8),
310            (b"ll", 9),
311            (b"lo", 10),
312            (b"wo", 11),
313            (b"rl", 12),
314            (b"hel", 13),
315            (b"llo", 14),
316            (b"wor", 15),
317            (b"hell", 16),
318            (b"ello", 17),
319            (b"worl", 18),
320            (b"hello", 19),
321            (b"world", 20),
322        ];
323
324        for (token, rank) in tokens {
325            let encoded = engine.encode(token);
326            content.push_str(&format!("{encoded} {rank}\n"));
327        }
328
329        let file_path = dir.join("tiktoken.model");
330        let mut file = std::fs::File::create(&file_path).unwrap();
331        file.write_all(content.as_bytes()).unwrap();
332        file_path.to_str().unwrap().to_string()
333    }
334
335    fn create_test_config(dir: &Path, model_type: &str) {
336        let config = serde_json::json!({
337            "model_type": model_type,
338            "max_position_embeddings": 32768,
339            "eos_token_id": [21]
340        });
341        let file_path = dir.join("config.json");
342        std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
343    }
344
345    fn create_test_tokenizer_config(dir: &Path, num_base_tokens: usize) {
346        let mut added_tokens = serde_json::Map::new();
347        let bos_id = num_base_tokens;
348        let eos_id = num_base_tokens + 1;
349
350        added_tokens.insert(
351            bos_id.to_string(),
352            serde_json::json!({"content": "[BOS]", "special": true}),
353        );
354        added_tokens.insert(
355            eos_id.to_string(),
356            serde_json::json!({"content": "[EOS]", "special": true}),
357        );
358
359        let config = serde_json::json!({
360            "added_tokens_decoder": added_tokens
361        });
362
363        let file_path = dir.join("tokenizer_config.json");
364        std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
365    }
366
367    #[test]
368    fn test_parse_tiktoken_file() {
369        let dir = tempfile::tempdir().unwrap();
370        let file_path = create_test_tiktoken_file(dir.path());
371        let encoder = parse_tiktoken_file(&file_path).unwrap();
372        assert_eq!(encoder.len(), 21);
373        assert_eq!(encoder[b"hello".as_slice()], 19);
374        assert_eq!(encoder[b"world".as_slice()], 20);
375    }
376
377    #[test]
378    fn test_parse_tiktoken_file_missing() {
379        let result = parse_tiktoken_file("/nonexistent/path/tiktoken.model");
380        assert!(result.is_err());
381    }
382
383    #[test]
384    fn test_tiktoken_from_file() {
385        let dir = tempfile::tempdir().unwrap();
386        let file_path = create_test_tiktoken_file(dir.path());
387
388        let mut special_tokens = FxHashMap::default();
389        special_tokens.insert("[BOS]".to_string(), 21_u32);
390        special_tokens.insert("[EOS]".to_string(), 22_u32);
391
392        // Use a simple pattern for testing
393        let pattern = r"[\w]+|[^\w\s]+|\s+";
394
395        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
396
397        assert_eq!(
398            tokenizer.special_tokens(),
399            &["[BOS]".to_string(), "[EOS]".to_string()]
400        );
401
402        // Test encode
403        let encoding = tokenizer.encode("hello world").unwrap();
404        let ids = encoding.token_ids();
405        assert!(!ids.is_empty());
406
407        // Test decode roundtrip
408        let decoded: String = tokenizer.decode(ids, false).unwrap().into();
409        assert_eq!(decoded, "hello world");
410    }
411
412    #[test]
413    fn test_tiktoken_encoding_variant() {
414        let dir = tempfile::tempdir().unwrap();
415        let file_path = create_test_tiktoken_file(dir.path());
416
417        let special_tokens = FxHashMap::default();
418        let pattern = r"[\w]+|[^\w\s]+|\s+";
419
420        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
421        assert!(tokenizer.special_tokens().is_empty());
422        let encoding = tokenizer.encode("hello").unwrap();
423
424        // Verify it produces the Sp variant
425        match &encoding {
426            Encoding::Sp(_) => {}
427            other => panic!("Expected Encoding::Sp, got {:?}", other),
428        }
429    }
430
431    #[test]
432    fn test_tiktoken_skip_special_tokens() {
433        let dir = tempfile::tempdir().unwrap();
434        let file_path = create_test_tiktoken_file(dir.path());
435
436        let mut special_tokens = FxHashMap::default();
437        special_tokens.insert("[BOS]".to_string(), 21_u32);
438        special_tokens.insert("[EOS]".to_string(), 22_u32);
439
440        let pattern = r"[\w]+|[^\w\s]+|\s+";
441
442        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
443
444        // Encode hello and prepend/append special tokens
445        let encoding = tokenizer.encode("hello").unwrap();
446        let mut ids = vec![21u32]; // [BOS]
447        ids.extend(encoding.token_ids());
448        ids.push(22); // [EOS]
449
450        // Decode with skip_special_tokens=true should strip special tokens
451        let decoded_skip: String = tokenizer.decode(&ids, true).unwrap().into();
452        assert_eq!(decoded_skip, "hello");
453
454        // Decode with skip_special_tokens=false should include them
455        let decoded_all: String = tokenizer.decode(&ids, false).unwrap().into();
456        assert!(decoded_all.contains("hello"));
457    }
458
459    #[test]
460    fn test_tiktoken_from_file_auto() {
461        let dir = tempfile::tempdir().unwrap();
462        let file_path = create_test_tiktoken_file(dir.path());
463
464        create_test_config(dir.path(), "kimi");
465        create_test_tokenizer_config(dir.path(), 21);
466
467        let mut expected_specials: Vec<String> = load_special_tokens(dir.path(), 21)
468            .unwrap()
469            .into_keys()
470            .collect();
471        expected_specials.sort();
472        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
473        assert_eq!(tokenizer.special_tokens(), expected_specials);
474
475        // Basic encode/decode roundtrip
476        let encoding = tokenizer.encode("hello world").unwrap();
477        let ids = encoding.token_ids();
478        assert!(!ids.is_empty());
479
480        let decoded: String = tokenizer.decode(ids, false).unwrap().into();
481        assert_eq!(decoded, "hello world");
482    }
483
484    #[test]
485    fn test_detect_bpe_pattern_unknown() {
486        let dir = tempfile::tempdir().unwrap();
487        create_test_config(dir.path(), "unknown_model");
488        let result = detect_bpe_pattern(dir.path());
489        assert!(result.is_err());
490    }
491
492    #[test]
493    fn test_kimi_k3_uses_kimi_bpe_pattern() {
494        let dir = tempfile::tempdir().unwrap();
495        create_test_config(dir.path(), "kimi_k3");
496        assert_eq!(detect_bpe_pattern(dir.path()).unwrap(), KIMI_PATTERN);
497    }
498
499    #[test]
500    fn test_detect_bpe_pattern_kimi_linear() {
501        let dir = tempfile::tempdir().unwrap();
502        create_test_config(dir.path(), "kimi_linear");
503        assert_eq!(detect_bpe_pattern(dir.path()).unwrap(), KIMI_PATTERN);
504    }
505
506    #[test]
507    fn test_load_special_tokens_no_config() {
508        let dir = tempfile::tempdir().unwrap();
509        let tokens = load_special_tokens(dir.path(), 100).unwrap();
510        assert_eq!(tokens.len(), 256);
511        assert_eq!(tokens["<|reserved_token_100|>"], 100);
512        assert_eq!(tokens["<|reserved_token_355|>"], 355);
513    }
514
515    #[test]
516    fn test_load_special_tokens_with_config() {
517        let dir = tempfile::tempdir().unwrap();
518        create_test_tokenizer_config(dir.path(), 100);
519        let tokens = load_special_tokens(dir.path(), 100).unwrap();
520        assert_eq!(tokens["[BOS]"], 100);
521        assert_eq!(tokens["[EOS]"], 101);
522        // Should also have reserved tokens filling gaps
523        assert!(tokens.len() > 2);
524    }
525
526    /// Helper: create a tiktoken file that includes raw byte tokens (byte fallback tokens).
527    fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
528        let engine = base64::engine::general_purpose::STANDARD;
529        let mut content = String::new();
530
531        let tokens: Vec<(&[u8], u32)> = vec![
532            (b"h", 0),
533            (b"e", 1),
534            (b"l", 2),
535            (b"o", 3),
536            (b" ", 4),
537            (b"hello", 5),
538        ];
539
540        for (token, rank) in &tokens {
541            let encoded = engine.encode(token);
542            content.push_str(&format!("{encoded} {rank}\n"));
543        }
544
545        // Byte-fallback tokens: individual bytes that form CJK character "δ½ " (U+4F60)
546        // UTF-8 encoding: 0xE4 0xBD 0xA0
547        let byte_tokens: Vec<(Vec<u8>, u32)> =
548            vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
549
550        for (token, rank) in &byte_tokens {
551            let encoded = engine.encode(token);
552            content.push_str(&format!("{encoded} {rank}\n"));
553        }
554
555        // Bytes for emoji "πŸ˜€" (U+1F600) β€” 4-byte UTF-8: 0xF0 0x9F 0x98 0x80
556        let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
557            (vec![0xF0], 200),
558            (vec![0x9F], 201),
559            (vec![0x98], 202),
560            (vec![0x80], 203),
561        ];
562
563        for (token, rank) in &emoji_tokens {
564            let encoded = engine.encode(token);
565            content.push_str(&format!("{encoded} {rank}\n"));
566        }
567
568        // Legitimate U+FFFD token: valid UTF-8 bytes EF BF BD (replacement character
569        // as an actual vocabulary entry, not an artifact of lossy conversion)
570        let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
571
572        for (token, rank) in &fffd_token {
573            let encoded = engine.encode(token);
574            content.push_str(&format!("{encoded} {rank}\n"));
575        }
576
577        let file_path = dir.join("tiktoken.model");
578        let mut file = std::fs::File::create(&file_path).unwrap();
579        file.write_all(content.as_bytes()).unwrap();
580        file_path.to_str().unwrap().to_string()
581    }
582
583    fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
584        let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
585        let special_tokens = FxHashMap::default();
586        let pattern = r"[\w]+|[^\w\s]+|\s+";
587        TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
588    }
589
590    /// Reproduces the original panic: decoding a single byte-fallback token that is
591    /// part of a multi-byte UTF-8 character. Before the fix, CoreBPE::decode() would
592    /// call String::from_utf8() on [0xE4] and error with "incomplete utf-8 byte sequence".
593    #[test]
594    fn test_decode_single_incomplete_utf8_byte_does_not_error() {
595        let dir = tempfile::tempdir().unwrap();
596        let tokenizer = create_byte_token_tokenizer(dir.path());
597
598        let result = tokenizer.decode(&[100], false);
599        assert!(
600            result.is_ok(),
601            "decode() should not error on incomplete UTF-8 bytes"
602        );
603        let decode_result = result.unwrap();
604        assert!(
605            decode_result.is_partial(),
606            "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
607            decode_result
608        );
609    }
610
611    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
612    #[test]
613    fn test_decode_two_of_three_utf8_bytes_does_not_error() {
614        let dir = tempfile::tempdir().unwrap();
615        let tokenizer = create_byte_token_tokenizer(dir.path());
616
617        let result = tokenizer.decode(&[100, 101], false);
618        assert!(result.is_ok());
619        let decode_result = result.unwrap();
620        assert!(
621            decode_result.is_partial(),
622            "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
623            decode_result
624        );
625    }
626
627    /// When all bytes of a multi-byte character are present, the concatenated bytes form
628    /// valid UTF-8, so this test passes both before and after the fix. It serves as a
629    /// correctness check that the lossy conversion doesn't corrupt complete characters.
630    #[test]
631    fn test_decode_complete_multibyte_utf8_produces_correct_char() {
632        let dir = tempfile::tempdir().unwrap();
633        let tokenizer = create_byte_token_tokenizer(dir.path());
634
635        let result = tokenizer.decode(&[100, 101, 102], false);
636        assert!(result.is_ok());
637        assert_eq!(String::from(result.unwrap()), "δ½ ");
638    }
639
640    /// All 4 emoji bytes together form valid UTF-8, so this passes both before and after
641    /// the fix. Validates that lossy conversion doesn't alter complete multi-byte sequences.
642    #[test]
643    fn test_decode_complete_4byte_emoji_from_byte_tokens() {
644        let dir = tempfile::tempdir().unwrap();
645        let tokenizer = create_byte_token_tokenizer(dir.path());
646
647        let result = tokenizer.decode(&[200, 201, 202, 203], false);
648        assert!(result.is_ok());
649        assert_eq!(String::from(result.unwrap()), "πŸ˜€");
650    }
651
652    /// Regression test: a vocabulary token whose raw bytes are EF BF BD (the valid
653    /// UTF-8 encoding of U+FFFD) must decode as `Complete`, not `Partial`. Before the
654    /// from_utf8 fast-path fix, from_utf8_lossy + the trailing-FFFD heuristic would
655    /// misclassify this as Partial, causing the incremental decoder to suppress it.
656    #[test]
657    fn test_decode_legitimate_replacement_char_token_is_complete() {
658        let dir = tempfile::tempdir().unwrap();
659        let tokenizer = create_byte_token_tokenizer(dir.path());
660
661        let result = tokenizer.decode(&[300], false);
662        assert!(result.is_ok());
663        let decode_result = result.unwrap();
664        assert!(
665            decode_result.is_complete(),
666            "legitimate U+FFFD vocab token must be Complete, got: {:?}",
667            decode_result
668        );
669        assert_eq!(decode_result.as_str(), "\u{FFFD}");
670    }
671
672    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
673    #[test]
674    fn test_decode_partial_emoji_does_not_error() {
675        let dir = tempfile::tempdir().unwrap();
676        let tokenizer = create_byte_token_tokenizer(dir.path());
677
678        let result = tokenizer.decode(&[200], false);
679        assert!(result.is_ok());
680        assert!(result.unwrap().is_partial());
681    }
682
683    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
684    #[test]
685    fn test_decode_mixed_ascii_and_incomplete_bytes() {
686        let dir = tempfile::tempdir().unwrap();
687        let tokenizer = create_byte_token_tokenizer(dir.path());
688
689        let result = tokenizer.decode(&[5, 100], false);
690        assert!(result.is_ok());
691        let decode_result = result.unwrap();
692        assert!(
693            decode_result.is_partial(),
694            "trailing incomplete byte should produce DecodeResult::Partial"
695        );
696        let text: String = decode_result.into();
697        assert!(
698            text.starts_with("hello"),
699            "should start with 'hello', got: {:?}",
700            text
701        );
702    }
703
704    /// End-to-end incremental detokenization: DecodeStream buffers partial bytes,
705    /// emits the complete character once all bytes arrive.
706    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
707    #[test]
708    fn test_decode_stream_incremental_multibyte_reassembly() {
709        let dir = tempfile::tempdir().unwrap();
710        let tokenizer = create_byte_token_tokenizer(dir.path());
711        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
712
713        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
714
715        let r1 = stream.step(100).unwrap();
716        assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
717
718        let r2 = stream.step(101).unwrap();
719        assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
720
721        let r3 = stream.step(102).unwrap();
722        assert!(r3.is_some(), "third byte should complete the character");
723        assert_eq!(r3.unwrap(), "δ½ ");
724    }
725
726    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
727    #[test]
728    fn test_decode_stream_incremental_emoji_reassembly() {
729        let dir = tempfile::tempdir().unwrap();
730        let tokenizer = create_byte_token_tokenizer(dir.path());
731        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
732
733        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
734
735        let r1 = stream.step(200).unwrap();
736        assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
737
738        let r2 = stream.step(201).unwrap();
739        assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
740
741        let r3 = stream.step(202).unwrap();
742        assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
743
744        let r4 = stream.step(203).unwrap();
745        assert!(r4.is_some(), "byte 4/4 should complete the emoji");
746        assert_eq!(r4.unwrap(), "πŸ˜€");
747    }
748
749    #[test]
750    fn test_tiktoken_encode_batch() {
751        let dir = tempfile::tempdir().unwrap();
752        let file_path = create_test_tiktoken_file(dir.path());
753
754        let special_tokens = FxHashMap::default();
755        let pattern = r"[\w]+|[^\w\s]+|\s+";
756
757        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
758
759        let inputs = &["hello", "world"];
760        let encodings = tokenizer.encode_batch(inputs).unwrap();
761        assert_eq!(encodings.len(), 2);
762
763        for (encoding, input) in encodings.iter().zip(inputs.iter()) {
764            let decoded: String = tokenizer
765                .decode(encoding.token_ids(), false)
766                .unwrap()
767                .into();
768            assert_eq!(decoded, *input);
769        }
770    }
771
772    /// Helper: create a tiktoken file containing all 256 single-byte tokens (ranks 0..255).
773    /// This gives a complete byte-level base vocabulary so any ASCII string can be encoded.
774    fn create_byte_level_tiktoken_file(dir: &Path) -> String {
775        let engine = base64::engine::general_purpose::STANDARD;
776        let mut content = String::new();
777        for byte_val in 0u16..256 {
778            let encoded = engine.encode([byte_val as u8]);
779            content.push_str(&format!("{encoded} {byte_val}\n"));
780        }
781        let file_path = dir.join("tiktoken.model");
782        std::fs::write(&file_path, &content).unwrap();
783        file_path.to_str().unwrap().to_string()
784    }
785
786    #[test]
787    fn test_segment_encoding_preserves_special_vs_ordinary_boundary() {
788        let dir = tempfile::tempdir().unwrap();
789        let file_path = create_byte_level_tiktoken_file(dir.path());
790        let marker = "<|open|>";
791        let marker_id = 256;
792        let mut special_tokens = FxHashMap::default();
793        special_tokens.insert(marker.to_string(), marker_id);
794        let tokenizer =
795            TikTokenTokenizer::from_file(&file_path, r"(?s).+", special_tokens).unwrap();
796
797        let structural = tokenizer
798            .encode_segments(&[EncodeSegment::control(marker)])
799            .unwrap();
800        let user_text = tokenizer
801            .encode_segments(&[EncodeSegment::ordinary(marker)])
802            .unwrap();
803
804        assert_eq!(structural.token_ids(), &[marker_id]);
805        assert_ne!(user_text.token_ids(), &[marker_id]);
806        let decoded: String = tokenizer
807            .decode(user_text.token_ids(), false)
808            .unwrap()
809            .into();
810        assert_eq!(decoded, marker);
811    }
812
813    fn has_nontrivial_self_overlap(token: &str) -> bool {
814        let bytes = token.as_bytes();
815        (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
816    }
817
818    fn have_ambiguous_overlap(a: &str, b: &str) -> bool {
819        let a = a.as_bytes();
820        let b = b.as_bytes();
821
822        if a.windows(b.len()).any(|window| window == b)
823            || b.windows(a.len()).any(|window| window == a)
824        {
825            return true;
826        }
827
828        let max_overlap = a.len().min(b.len());
829        (1..max_overlap).any(|overlap| {
830            a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
831        })
832    }
833
834    fn assert_special_split_invariant(
835        tokenizer: &TikTokenTokenizer,
836        left: &str,
837        special: &str,
838        right: &str,
839    ) {
840        let whole = tokenizer
841            .encode(&format!("{left}{special}{right}"))
842            .unwrap()
843            .token_ids()
844            .to_vec();
845        let mut split = tokenizer
846            .encode(&format!("{left}{special}"))
847            .unwrap()
848            .token_ids()
849            .to_vec();
850        split.extend_from_slice(tokenizer.encode(right).unwrap().token_ids());
851        assert_eq!(
852            whole, split,
853            "splitting after registered special token {special:?} must preserve tokenization"
854        );
855    }
856
857    #[test]
858    fn test_registered_special_tokens_are_cache_safe_boundaries() {
859        let dir = tempfile::tempdir().unwrap();
860        let file_path = create_byte_level_tiktoken_file(dir.path());
861        create_test_config(dir.path(), "kimi");
862        create_test_tokenizer_config(dir.path(), 256);
863
864        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
865        let specials = tokenizer.special_tokens();
866        assert!(!specials.is_empty());
867
868        for (index, special) in specials.iter().enumerate() {
869            assert!(
870                !has_nontrivial_self_overlap(special),
871                "registered special token has an ambiguous self-overlap: {special:?}"
872            );
873            for other in &specials[index + 1..] {
874                assert!(
875                    !have_ambiguous_overlap(special, other),
876                    "registered special tokens overlap ambiguously: {special:?} and {other:?}"
877                );
878            }
879
880            for (left, right) in [
881                ("ordinary prefix ", " ordinary suffix"),
882                ("line before\n", "\tUnicode after: εŒ—δΊ¬ πŸ˜€"),
883                ("", " trailing text"),
884            ] {
885                assert_special_split_invariant(&tokenizer, left, special, right);
886            }
887        }
888
889        let bos = specials
890            .iter()
891            .find(|token| token.as_str() == "[BOS]")
892            .unwrap();
893        let eos = specials
894            .iter()
895            .find(|token| token.as_str() == "[EOS]")
896            .unwrap();
897        assert_special_split_invariant(
898            &tokenizer,
899            "prefix ",
900            bos,
901            &format!("{eos}<|reserved_token_258|>Unicode 尾部"),
902        );
903    }
904
905    /// Regression test for Kimi K2.5 tokenizer inflation.
906    ///
907    /// Python's tokenization_kimi.py names unnamed reserved tokens by absolute ID:
908    ///   `<|reserved_token_{absolute_id}|>`
909    ///
910    /// The Rust code previously used relative offsets (0..255) as the naming index,
911    /// so when a prompt contained `<|reserved_token_258|>` the Rust tokenizer did NOT
912    /// recognize it as a special token. Each occurrence was encoded as multiple BPE
913    /// tokens instead of 1, inflating an 8192-token prompt to 9038 tokens and causing
914    /// TRT-LLM to reject the request.
915    #[test]
916    fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
917        let dir = tempfile::tempdir().unwrap();
918        let file_path = create_byte_level_tiktoken_file(dir.path());
919
920        // config.json with kimi model type -> triggers KIMI_PATTERN
921        create_test_config(dir.path(), "kimi");
922
923        // tokenizer_config.json: BOS at 256, EOS at 257. Base vocab is IDs 0..255.
924        create_test_tokenizer_config(dir.path(), 256);
925
926        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
927
928        // ID 256 = [BOS], ID 257 = [EOS].
929        // ID 258 = first UNNAMED reserved token.
930        //   With fix:   named <|reserved_token_258|>  (absolute ID)
931        //   Before fix: named <|reserved_token_2|>    (relative offset)
932
933        // Single unnamed reserved token should be recognized as 1 special token.
934        let single = "<|reserved_token_258|>";
935        let enc = tokenizer.encode(single).unwrap();
936        assert_eq!(
937            enc.token_ids().len(),
938            1,
939            "'{single}' should be 1 special token, got {} tokens: {:?}. \
940             This means fallback naming still uses relative offsets instead of absolute IDs.",
941            enc.token_ids().len(),
942            enc.token_ids()
943        );
944        assert_eq!(enc.token_ids()[0], 258);
945
946        // Multiple unnamed reserved tokens in sequence (mini version of the benchmark).
947        // IDs 258..268 are all unnamed; with the fix they're <|reserved_token_258|>..267.
948        let multi: String = (258u32..268)
949            .map(|id| format!("<|reserved_token_{id}|>"))
950            .collect();
951        let enc_multi = tokenizer.encode(&multi).unwrap();
952        assert_eq!(
953            enc_multi.token_ids().len(),
954            10,
955            "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
956            enc_multi.token_ids().len(),
957            enc_multi.token_ids()
958        );
959        let expected_ids: Vec<u32> = (258..268).collect();
960        assert_eq!(enc_multi.token_ids(), &expected_ids);
961    }
962
963    /// Confirm that the old relative-offset naming would cause token inflation.
964    /// Manually builds a tokenizer whose special-token map uses the WRONG names
965    /// (relative offsets), then shows the same string encodes as many tokens.
966    #[test]
967    fn test_relative_offset_naming_causes_inflation() {
968        let dir = tempfile::tempdir().unwrap();
969        let file_path = create_byte_level_tiktoken_file(dir.path());
970
971        let _encoder = parse_tiktoken_file(&file_path).unwrap();
972        let num_base_tokens = 256usize;
973
974        // Build special tokens with the OLD (buggy) relative-offset naming
975        let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
976        bad_special_tokens.insert("[BOS]".to_string(), 256);
977        bad_special_tokens.insert("[EOS]".to_string(), 257);
978        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
979            let id = num_base_tokens as u32 + i;
980            if id != 256 && id != 257 {
981                // OLD naming: relative offset i, not absolute id
982                bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
983            }
984        }
985
986        let bad_tokenizer =
987            TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
988
989        // With the wrong naming, <|reserved_token_258|> is NOT recognized as special.
990        // It gets split into byte-level BPE tokens -> many more than 1.
991        let input = "<|reserved_token_258|>";
992        let enc = bad_tokenizer.encode(input).unwrap();
993        assert!(
994            enc.token_ids().len() > 1,
995            "With buggy relative-offset naming, '{}' should NOT be recognized as a \
996             single special token. Got {} token(s): {:?}",
997            input,
998            enc.token_ids().len(),
999            enc.token_ids()
1000        );
1001
1002        // Show the inflation: 10 reserved tokens produce far more than 10 BPE tokens.
1003        let multi: String = (258u32..268)
1004            .map(|id| format!("<|reserved_token_{id}|>"))
1005            .collect();
1006        let enc_multi = bad_tokenizer.encode(&multi).unwrap();
1007        assert!(
1008            enc_multi.token_ids().len() > 10,
1009            "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
1010             Got {}",
1011            enc_multi.token_ids().len(),
1012        );
1013    }
1014}