Skip to main content

dynamo_tokenizers/
basetenkenizer.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Baseten Tokenizer backend for high-performance BPE encoding and decoding.
5//!
6//! Some Kimi repositories ship tiktoken assets without a directly loadable
7//! `tokenizer.json`. Baseten publishes compatible tokenizer artifacts,
8//! including [`baseten/kimi-k3-tokenizer`](https://huggingface.co/baseten/kimi-k3-tokenizer).
9
10use std::path::Path;
11
12use super::{
13    EncodeSegment, Encoding, Error, Result, TokenIdType, TokenizerOptions,
14    traits::{DecodeResult, Decoder, Encoder, Tokenizer},
15};
16
17/// Tokenizer backed by the `basetenkenizer` crate.
18pub struct BasetenTokenizer {
19    tokenizer: basetenkenizer::Tokenizer,
20    options: TokenizerOptions,
21}
22
23impl BasetenTokenizer {
24    /// Load a tokenizer from a Hugging Face `tokenizer.json` file.
25    pub fn from_file(path: &str) -> Result<Self> {
26        let path = Path::new(path);
27        let raw = std::fs::read_to_string(path)
28            .map_err(|e| Error::msg(format!("Error reading Baseten tokenizer: {e}")))?;
29        let mut json: serde_json::Value = serde_json::from_str(&raw)
30            .map_err(|e| Error::msg(format!("Error parsing Baseten tokenizer: {e}")))?;
31        if let Some(parent) = path.parent() {
32            merge_special_tokens_from_config(&mut json, parent);
33        }
34        let tokenizer = basetenkenizer::Tokenizer::from_json(json)
35            .map_err(|e| Error::msg(format!("Error loading Baseten tokenizer: {e}")))?;
36        Ok(Self {
37            tokenizer,
38            options: TokenizerOptions::default(),
39        })
40    }
41}
42
43fn merge_special_tokens_from_config(json: &mut serde_json::Value, model_dir: &Path) {
44    let config_path = model_dir.join("tokenizer_config.json");
45    let Ok(raw) = std::fs::read_to_string(&config_path) else {
46        return;
47    };
48    let config: serde_json::Value = match serde_json::from_str(&raw) {
49        Ok(value) => value,
50        Err(error) => {
51            tracing::debug!(
52                target: "tokenizer",
53                path = %config_path.display(),
54                error = %error,
55                "tokenizer_config.json parse failed; skipping special-token merge"
56            );
57            return;
58        }
59    };
60    let Some(decoder) = config
61        .get("added_tokens_decoder")
62        .and_then(serde_json::Value::as_object)
63    else {
64        return;
65    };
66
67    if json.get("added_tokens").is_none() {
68        json["added_tokens"] = serde_json::json!([]);
69    }
70    let Some(added_tokens) = json
71        .get_mut("added_tokens")
72        .and_then(serde_json::Value::as_array_mut)
73    else {
74        return;
75    };
76
77    for (id, spec) in decoder {
78        let Some(id) = id.parse::<u32>().ok() else {
79            continue;
80        };
81        let Some(spec) = spec.as_object() else {
82            continue;
83        };
84        if spec.get("special").and_then(serde_json::Value::as_bool) != Some(true) {
85            continue;
86        }
87        let Some(content) = spec
88            .get("content")
89            .and_then(serde_json::Value::as_str)
90            .filter(|content| !content.is_empty())
91        else {
92            continue;
93        };
94
95        if let Some(existing) = added_tokens
96            .iter_mut()
97            .find(|token| token.get("content").and_then(serde_json::Value::as_str) == Some(content))
98        {
99            existing["special"] = serde_json::Value::Bool(true);
100            continue;
101        }
102
103        let mut token = serde_json::Map::from_iter([
104            ("id".to_string(), serde_json::json!(id)),
105            ("content".to_string(), serde_json::json!(content)),
106            ("special".to_string(), serde_json::Value::Bool(true)),
107        ]);
108        for field in ["single_word", "lstrip", "rstrip", "normalized"] {
109            token.insert(
110                field.to_string(),
111                serde_json::Value::Bool(
112                    spec.get(field)
113                        .and_then(serde_json::Value::as_bool)
114                        .unwrap_or(false),
115                ),
116            );
117        }
118        added_tokens.push(serde_json::Value::Object(token));
119    }
120}
121
122impl Encoder for BasetenTokenizer {
123    fn encode(&self, input: &str) -> Result<Encoding> {
124        let ids = self
125            .tokenizer
126            .encode_with_special_tokens(input, self.options.add_special_tokens)
127            .map_err(|e| Error::msg(format!("Baseten tokenizer encode error: {e}")))?;
128        Ok(Encoding::Sp(ids))
129    }
130
131    fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
132        self.tokenizer
133            .encode_batch(inputs, self.options.add_special_tokens)
134            .map(|ids| ids.into_iter().map(Encoding::Sp).collect())
135            .map_err(|e| Error::msg(format!("Baseten tokenizer batch encode error: {e}")))
136    }
137
138    fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
139        let segments = segments
140            .iter()
141            .map(|segment| (segment.text, segment.allow_special));
142        let ids = self
143            .tokenizer
144            .encode_segments_tiktoken_safe(segments, self.options.add_special_tokens)
145            .map_err(|e| Error::msg(format!("Baseten tokenizer segment encode error: {e}")))?;
146        Ok(Encoding::Sp(ids))
147    }
148}
149
150impl Decoder for BasetenTokenizer {
151    fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
152        self.tokenizer
153            .decode(token_ids, skip_special_tokens)
154            .map(DecodeResult::from)
155            .map_err(|e| Error::msg(format!("Baseten tokenizer decode error: {e}")))
156    }
157}
158
159impl Tokenizer for BasetenTokenizer {
160    fn validate_prefix_cache(&self) -> Result<()> {
161        if self.options.add_special_tokens {
162            return Err(Error::msg(
163                "Baseten tokenizers configured with add_special_tokens=true must remain uncached",
164            ));
165        }
166        Ok(())
167    }
168
169    fn with_options(mut self, options: TokenizerOptions) -> Self {
170        self.options = options;
171        self
172    }
173
174    fn vocab_size(&self) -> Option<usize> {
175        // `basetenkenizer::Tokenizer::vocab_size` sums the model vocab size
176        // and the added-tokens count with no dedup, but real tokenizer.json
177        // files (this crate's own TinyLlama_v1.1 fixture included) list
178        // `added_tokens` entries whose ids already exist in `model.vocab` --
179        // metadata for special tokens, not additional vocabulary. Count only
180        // added tokens whose content isn't already in the model vocab.
181        let model = self.tokenizer.model();
182        let model_size = model.vocab_size();
183        let extra = self.tokenizer.added_tokens().map_or(0, |added| {
184            added
185                .iter()
186                .filter(|info| model.token_to_id(info.content).is_none())
187                .count()
188        });
189        Some(model_size + extra)
190    }
191
192    fn token_to_id(&self, token: &str) -> Result<Option<TokenIdType>> {
193        Ok(self.tokenizer.token_to_id(token))
194    }
195
196    fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
197        let Some(added_tokens) = self.tokenizer.added_tokens() else {
198            return Ok(Vec::new());
199        };
200        let mut ids: Vec<TokenIdType> = added_tokens
201            .iter()
202            .filter_map(|info| info.special.then_some(info.id))
203            .collect();
204        ids.sort_unstable();
205        Ok(ids)
206    }
207
208    fn num_special_tokens_added(&self) -> Result<usize> {
209        // basetenkenizer exposes no direct count, but post-processing an
210        // empty sequence with add_special_tokens=true is content-length
211        // independent for every basetenkenizer::PostProcessor variant:
212        // ByteLevel is identity (adds 0); TemplateProcessing::apply_single
213        // walks a fixed template where SpecialToken pieces insert a
214        // fixed-length id vector and Sequence pieces are replaced 1:1 by
215        // whatever content came in, so the *added* length never depends on
216        // content length; Sequence(steps) folds that same guarantee across
217        // steps. So this reveals exactly what the post-processor inserts
218        // around a bare encoding, for real content of any length.
219        Ok(self.tokenizer.post_process(Vec::new(), true).len())
220    }
221}
222
223#[cfg(test)]
224mod tests {
225    use std::sync::Arc;
226
227    use super::*;
228    use crate::{
229        HuggingFaceTokenizer, Tokenizer as TokenizerWrapper, traits::Tokenizer as TokenizerTrait,
230    };
231
232    const TOKENIZER_PATH: &str = concat!(
233        env!("CARGO_MANIFEST_DIR"),
234        "/tests/data/minimal-bpe/tokenizer.json"
235    );
236
237    #[test]
238    fn encode_matches_hugging_face() {
239        let baseten = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
240        let hf = HuggingFaceTokenizer::from_file(TOKENIZER_PATH).unwrap();
241
242        for text in ["Hello, world!", "Hello", " world", "He llo"] {
243            let baseten_ids = baseten.encode(text).unwrap();
244            let hf_ids = hf.encode(text).unwrap();
245            assert_eq!(
246                baseten_ids.token_ids(),
247                hf_ids.token_ids(),
248                "Baseten and Hugging Face must produce identical token IDs for '{text}'"
249            );
250        }
251    }
252
253    #[test]
254    fn batch_encode_matches_sequential_encode() {
255        let tokenizer = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
256        let inputs = ["Hello", " world", "Hello, world!"];
257        let batch = tokenizer.encode_batch(&inputs).unwrap();
258
259        for (encoding, input) in batch.iter().zip(inputs) {
260            let sequential = tokenizer.encode(input).unwrap();
261            assert_eq!(encoding.token_ids(), sequential.token_ids());
262        }
263    }
264
265    #[test]
266    fn encode_decode_roundtrip() {
267        let tokenizer = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
268        let encoding = tokenizer.encode("Hello, world!").unwrap();
269        let decoded = tokenizer.decode(encoding.token_ids(), true).unwrap();
270
271        assert_eq!(decoded.as_str(), "Hello, world!");
272    }
273
274    #[test]
275    fn works_with_decode_stream() {
276        let tokenizer = Arc::new(BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap());
277        let wrapper = TokenizerWrapper::from(tokenizer);
278        let prompt_ids = wrapper.encode("Hello").unwrap().token_ids().to_vec();
279        let continuation_ids = wrapper.encode(", world!").unwrap().token_ids().to_vec();
280        let mut stream = wrapper.decode_stream(&prompt_ids, true);
281        let mut accumulated = String::new();
282
283        for id in &continuation_ids {
284            if let Some(chunk) = stream.step(*id).unwrap() {
285                accumulated.push_str(&chunk);
286            }
287        }
288
289        let mut all_ids = prompt_ids.clone();
290        all_ids.extend_from_slice(&continuation_ids);
291        let full_text: String = wrapper.decode(&all_ids, true).unwrap().into();
292        let prompt_text: String = wrapper.decode(&prompt_ids, true).unwrap().into();
293        assert_eq!(accumulated, full_text[prompt_text.len()..]);
294    }
295
296    #[test]
297    fn prefix_cache_rejects_special_token_post_processing() {
298        let plain = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
299        assert!(plain.validate_prefix_cache().is_ok());
300
301        let with_special_tokens = BasetenTokenizer::from_file(TOKENIZER_PATH)
302            .unwrap()
303            .with_options(TokenizerOptions {
304                add_special_tokens: true,
305            });
306        assert!(with_special_tokens.validate_prefix_cache().is_err());
307    }
308
309    #[test]
310    fn segments_preserve_special_token_trust_boundary() {
311        let temp = tempfile::tempdir().unwrap();
312        let path = temp.path().join("tokenizer.json");
313        let mut json: serde_json::Value =
314            serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
315        let vocab = json["model"]["vocab"].as_object_mut().unwrap();
316        vocab.insert("<".to_string(), serde_json::json!(23));
317        vocab.insert(">".to_string(), serde_json::json!(24));
318        vocab.insert("c".to_string(), serde_json::json!(25));
319        json["added_tokens"] = serde_json::json!([{
320            "id": 26,
321            "content": "<ctl>",
322            "single_word": false,
323            "lstrip": false,
324            "rstrip": false,
325            "normalized": false,
326            "special": true
327        }]);
328        std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
329
330        let tokenizer: Arc<dyn TokenizerTrait> =
331            Arc::new(BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap());
332        let segments = [
333            EncodeSegment::new("Hello", true),
334            EncodeSegment::new("<ctl>", false),
335            EncodeSegment::new(" world!", true),
336        ];
337
338        let segmented = tokenizer.encode_segments(&segments).unwrap();
339        let flattened = tokenizer.encode("Hello<ctl> world!").unwrap();
340
341        assert_ne!(
342            segmented.token_ids(),
343            flattened.token_ids(),
344            "untrusted control-token-looking text must not become an added token"
345        );
346        assert!(flattened.token_ids().contains(&26));
347        assert!(!segmented.token_ids().contains(&26));
348    }
349
350    #[test]
351    fn segments_honor_add_special_tokens_option() {
352        let temp = tempfile::tempdir().unwrap();
353        let path = temp.path().join("tokenizer.json");
354        let mut json: serde_json::Value =
355            serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
356        json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
357        json["added_tokens"] = serde_json::json!([{
358            "id": 23,
359            "content": "<bos>",
360            "single_word": false,
361            "lstrip": false,
362            "rstrip": false,
363            "normalized": false,
364            "special": true
365        }]);
366        json["post_processor"] = serde_json::json!({
367            "type": "TemplateProcessing",
368            "single": [
369                {"SpecialToken": {"id": "<bos>", "type_id": 0}},
370                {"Sequence": {"id": "A", "type_id": 0}}
371            ],
372            "pair": [
373                {"SpecialToken": {"id": "<bos>", "type_id": 0}},
374                {"Sequence": {"id": "A", "type_id": 0}},
375                {"Sequence": {"id": "B", "type_id": 0}}
376            ],
377            "special_tokens": {
378                "<bos>": {"id": "<bos>", "ids": [23], "tokens": ["<bos>"]}
379            }
380        });
381        std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
382        let segments = [EncodeSegment::new("Hello", false)];
383
384        let plain = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
385        let plain_ids = plain
386            .encode_segments(&segments)
387            .unwrap()
388            .token_ids()
389            .to_vec();
390
391        let with_bos = BasetenTokenizer::from_file(path.to_str().unwrap())
392            .unwrap()
393            .with_options(TokenizerOptions {
394                add_special_tokens: true,
395            });
396        assert_eq!(
397            with_bos.encode_segments(&segments).unwrap().token_ids(),
398            [&[23], plain_ids.as_slice()].concat()
399        );
400    }
401
402    #[test]
403    fn num_special_tokens_added_reflects_bos_post_processor() {
404        let temp = tempfile::tempdir().unwrap();
405        let path = temp.path().join("tokenizer.json");
406        let mut json: serde_json::Value =
407            serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
408        json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
409        json["added_tokens"] = serde_json::json!([{
410            "id": 23,
411            "content": "<bos>",
412            "single_word": false,
413            "lstrip": false,
414            "rstrip": false,
415            "normalized": false,
416            "special": true
417        }]);
418        json["post_processor"] = serde_json::json!({
419            "type": "TemplateProcessing",
420            "single": [
421                {"SpecialToken": {"id": "<bos>", "type_id": 0}},
422                {"Sequence": {"id": "A", "type_id": 0}}
423            ],
424            "pair": [
425                {"SpecialToken": {"id": "<bos>", "type_id": 0}},
426                {"Sequence": {"id": "A", "type_id": 0}},
427                {"Sequence": {"id": "B", "type_id": 0}}
428            ],
429            "special_tokens": {
430                "<bos>": {"id": "<bos>", "ids": [23], "tokens": ["<bos>"]}
431            }
432        });
433        std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
434
435        let with_bos = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
436        assert_eq!(with_bos.num_special_tokens_added().unwrap(), 1);
437
438        let plain = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
439        assert_eq!(plain.num_special_tokens_added().unwrap(), 0);
440    }
441
442    #[test]
443    fn num_special_tokens_added_is_length_independent_under_sequence_post_processor() {
444        let temp = tempfile::tempdir().unwrap();
445        let path = temp.path().join("tokenizer.json");
446        let mut json: serde_json::Value =
447            serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
448        json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
449        json["model"]["vocab"]["<eos>"] = serde_json::json!(24);
450        json["added_tokens"] = serde_json::json!([
451            {
452                "id": 23,
453                "content": "<bos>",
454                "single_word": false,
455                "lstrip": false,
456                "rstrip": false,
457                "normalized": false,
458                "special": true
459            },
460            {
461                "id": 24,
462                "content": "<eos>",
463                "single_word": false,
464                "lstrip": false,
465                "rstrip": false,
466                "normalized": false,
467                "special": true
468            }
469        ]);
470        json["post_processor"] = serde_json::json!({
471            "type": "Sequence",
472            "processors": [
473                {
474                    "type": "TemplateProcessing",
475                    "single": [
476                        {"SpecialToken": {"id": "<bos>", "type_id": 0}},
477                        {"Sequence": {"id": "A", "type_id": 0}}
478                    ],
479                    "pair": [],
480                    "special_tokens": {
481                        "<bos>": {"id": "<bos>", "ids": [23], "tokens": ["<bos>"]}
482                    }
483                },
484                {
485                    "type": "TemplateProcessing",
486                    "single": [
487                        {"Sequence": {"id": "A", "type_id": 0}},
488                        {"SpecialToken": {"id": "<eos>", "type_id": 0}}
489                    ],
490                    "pair": [],
491                    "special_tokens": {
492                        "<eos>": {"id": "<eos>", "ids": [24], "tokens": ["<eos>"]}
493                    }
494                }
495            ]
496        });
497        std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
498
499        let tokenizer = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
500        assert_eq!(tokenizer.num_special_tokens_added().unwrap(), 2);
501
502        let with_specials = tokenizer.with_options(TokenizerOptions {
503            add_special_tokens: true,
504        });
505        let plain = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
506
507        for text in ["h", "hello there world"] {
508            let without = plain.encode(text).unwrap().token_ids().len();
509            let with = with_specials.encode(text).unwrap().token_ids().len();
510            assert_eq!(
511                with - without,
512                2,
513                "'{text}' should grow by exactly num_special_tokens_added()"
514            );
515        }
516    }
517
518    #[test]
519    fn merges_config_only_special_tokens() {
520        let temp = tempfile::tempdir().unwrap();
521        let tokenizer_path = temp.path().join("tokenizer.json");
522        std::fs::copy(TOKENIZER_PATH, &tokenizer_path).unwrap();
523        std::fs::write(
524            temp.path().join("tokenizer_config.json"),
525            serde_json::json!({
526                "added_tokens_decoder": {
527                    "23": {
528                        "content": "<ctl>",
529                        "special": true,
530                        "single_word": false,
531                        "lstrip": false,
532                        "rstrip": false,
533                        "normalized": false
534                    }
535                }
536            })
537            .to_string(),
538        )
539        .unwrap();
540
541        let tokenizer = BasetenTokenizer::from_file(tokenizer_path.to_str().unwrap()).unwrap();
542        let encoding = tokenizer.encode("<ctl>").unwrap();
543        assert_eq!(encoding.token_ids(), &[23]);
544        assert_eq!(tokenizer.decode(&[23], false).unwrap().as_str(), "<ctl>");
545        assert_eq!(tokenizer.decode(&[23], true).unwrap().as_str(), "");
546    }
547
548    #[test]
549    fn vocab_introspection_accessors() {
550        let plain = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
551        assert_eq!(plain.vocab_size(), Some(23));
552        assert_eq!(plain.token_to_id("hello").unwrap(), None);
553        assert_eq!(plain.token_to_id("h").unwrap(), Some(10));
554        assert!(plain.special_token_ids().unwrap().is_empty());
555
556        let temp = tempfile::tempdir().unwrap();
557        let path = temp.path().join("tokenizer.json");
558        let mut json: serde_json::Value =
559            serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
560        json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
561        json["added_tokens"] = serde_json::json!([{
562            "id": 23,
563            "content": "<bos>",
564            "single_word": false,
565            "lstrip": false,
566            "rstrip": false,
567            "normalized": false,
568            "special": true
569        }]);
570        std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
571
572        let with_added = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
573        assert_eq!(with_added.vocab_size(), Some(24));
574        assert_eq!(with_added.token_to_id("<bos>").unwrap(), Some(23));
575        assert_eq!(with_added.special_token_ids().unwrap(), vec![23]);
576
577        let path2 = temp.path().join("tokenizer2.json");
578        let mut json2: serde_json::Value =
579            serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
580        json2["added_tokens"] = serde_json::json!([{
581            "id": 23,
582            "content": "<extra>",
583            "single_word": false,
584            "lstrip": false,
585            "rstrip": false,
586            "normalized": false,
587            "special": true
588        }]);
589        std::fs::write(&path2, serde_json::to_vec(&json2).unwrap()).unwrap();
590
591        let genuinely_added = BasetenTokenizer::from_file(path2.to_str().unwrap()).unwrap();
592        assert_eq!(genuinely_added.vocab_size(), Some(24));
593        assert_eq!(genuinely_added.token_to_id("<extra>").unwrap(), Some(23));
594    }
595}