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
175#[cfg(test)]
176mod tests {
177    use std::sync::Arc;
178
179    use super::*;
180    use crate::{
181        HuggingFaceTokenizer, Tokenizer as TokenizerWrapper, traits::Tokenizer as TokenizerTrait,
182    };
183
184    const TOKENIZER_PATH: &str = concat!(
185        env!("CARGO_MANIFEST_DIR"),
186        "/tests/data/minimal-bpe/tokenizer.json"
187    );
188
189    #[test]
190    fn encode_matches_hugging_face() {
191        let baseten = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
192        let hf = HuggingFaceTokenizer::from_file(TOKENIZER_PATH).unwrap();
193
194        for text in ["Hello, world!", "Hello", " world", "He llo"] {
195            let baseten_ids = baseten.encode(text).unwrap();
196            let hf_ids = hf.encode(text).unwrap();
197            assert_eq!(
198                baseten_ids.token_ids(),
199                hf_ids.token_ids(),
200                "Baseten and Hugging Face must produce identical token IDs for '{text}'"
201            );
202        }
203    }
204
205    #[test]
206    fn batch_encode_matches_sequential_encode() {
207        let tokenizer = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
208        let inputs = ["Hello", " world", "Hello, world!"];
209        let batch = tokenizer.encode_batch(&inputs).unwrap();
210
211        for (encoding, input) in batch.iter().zip(inputs) {
212            let sequential = tokenizer.encode(input).unwrap();
213            assert_eq!(encoding.token_ids(), sequential.token_ids());
214        }
215    }
216
217    #[test]
218    fn encode_decode_roundtrip() {
219        let tokenizer = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
220        let encoding = tokenizer.encode("Hello, world!").unwrap();
221        let decoded = tokenizer.decode(encoding.token_ids(), true).unwrap();
222
223        assert_eq!(decoded.as_str(), "Hello, world!");
224    }
225
226    #[test]
227    fn works_with_decode_stream() {
228        let tokenizer = Arc::new(BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap());
229        let wrapper = TokenizerWrapper::from(tokenizer);
230        let prompt_ids = wrapper.encode("Hello").unwrap().token_ids().to_vec();
231        let continuation_ids = wrapper.encode(", world!").unwrap().token_ids().to_vec();
232        let mut stream = wrapper.decode_stream(&prompt_ids, true);
233        let mut accumulated = String::new();
234
235        for id in &continuation_ids {
236            if let Some(chunk) = stream.step(*id).unwrap() {
237                accumulated.push_str(&chunk);
238            }
239        }
240
241        let mut all_ids = prompt_ids.clone();
242        all_ids.extend_from_slice(&continuation_ids);
243        let full_text: String = wrapper.decode(&all_ids, true).unwrap().into();
244        let prompt_text: String = wrapper.decode(&prompt_ids, true).unwrap().into();
245        assert_eq!(accumulated, full_text[prompt_text.len()..]);
246    }
247
248    #[test]
249    fn prefix_cache_rejects_special_token_post_processing() {
250        let plain = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
251        assert!(plain.validate_prefix_cache().is_ok());
252
253        let with_special_tokens = BasetenTokenizer::from_file(TOKENIZER_PATH)
254            .unwrap()
255            .with_options(TokenizerOptions {
256                add_special_tokens: true,
257            });
258        assert!(with_special_tokens.validate_prefix_cache().is_err());
259    }
260
261    #[test]
262    fn segments_preserve_special_token_trust_boundary() {
263        let temp = tempfile::tempdir().unwrap();
264        let path = temp.path().join("tokenizer.json");
265        let mut json: serde_json::Value =
266            serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
267        let vocab = json["model"]["vocab"].as_object_mut().unwrap();
268        vocab.insert("<".to_string(), serde_json::json!(23));
269        vocab.insert(">".to_string(), serde_json::json!(24));
270        vocab.insert("c".to_string(), serde_json::json!(25));
271        json["added_tokens"] = serde_json::json!([{
272            "id": 26,
273            "content": "<ctl>",
274            "single_word": false,
275            "lstrip": false,
276            "rstrip": false,
277            "normalized": false,
278            "special": true
279        }]);
280        std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
281
282        let tokenizer: Arc<dyn TokenizerTrait> =
283            Arc::new(BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap());
284        let segments = [
285            EncodeSegment::new("Hello", true),
286            EncodeSegment::new("<ctl>", false),
287            EncodeSegment::new(" world!", true),
288        ];
289
290        let segmented = tokenizer.encode_segments(&segments).unwrap();
291        let flattened = tokenizer.encode("Hello<ctl> world!").unwrap();
292
293        assert_ne!(
294            segmented.token_ids(),
295            flattened.token_ids(),
296            "untrusted control-token-looking text must not become an added token"
297        );
298        assert!(flattened.token_ids().contains(&26));
299        assert!(!segmented.token_ids().contains(&26));
300    }
301
302    #[test]
303    fn segments_honor_add_special_tokens_option() {
304        let temp = tempfile::tempdir().unwrap();
305        let path = temp.path().join("tokenizer.json");
306        let mut json: serde_json::Value =
307            serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
308        json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
309        json["added_tokens"] = serde_json::json!([{
310            "id": 23,
311            "content": "<bos>",
312            "single_word": false,
313            "lstrip": false,
314            "rstrip": false,
315            "normalized": false,
316            "special": true
317        }]);
318        json["post_processor"] = serde_json::json!({
319            "type": "TemplateProcessing",
320            "single": [
321                {"SpecialToken": {"id": "<bos>", "type_id": 0}},
322                {"Sequence": {"id": "A", "type_id": 0}}
323            ],
324            "pair": [
325                {"SpecialToken": {"id": "<bos>", "type_id": 0}},
326                {"Sequence": {"id": "A", "type_id": 0}},
327                {"Sequence": {"id": "B", "type_id": 0}}
328            ],
329            "special_tokens": {
330                "<bos>": {"id": "<bos>", "ids": [23], "tokens": ["<bos>"]}
331            }
332        });
333        std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
334        let segments = [EncodeSegment::new("Hello", false)];
335
336        let plain = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
337        let plain_ids = plain
338            .encode_segments(&segments)
339            .unwrap()
340            .token_ids()
341            .to_vec();
342
343        let with_bos = BasetenTokenizer::from_file(path.to_str().unwrap())
344            .unwrap()
345            .with_options(TokenizerOptions {
346                add_special_tokens: true,
347            });
348        assert_eq!(
349            with_bos.encode_segments(&segments).unwrap().token_ids(),
350            [&[23], plain_ids.as_slice()].concat()
351        );
352    }
353
354    #[test]
355    fn merges_config_only_special_tokens() {
356        let temp = tempfile::tempdir().unwrap();
357        let tokenizer_path = temp.path().join("tokenizer.json");
358        std::fs::copy(TOKENIZER_PATH, &tokenizer_path).unwrap();
359        std::fs::write(
360            temp.path().join("tokenizer_config.json"),
361            serde_json::json!({
362                "added_tokens_decoder": {
363                    "23": {
364                        "content": "<ctl>",
365                        "special": true,
366                        "single_word": false,
367                        "lstrip": false,
368                        "rstrip": false,
369                        "normalized": false
370                    }
371                }
372            })
373            .to_string(),
374        )
375        .unwrap();
376
377        let tokenizer = BasetenTokenizer::from_file(tokenizer_path.to_str().unwrap()).unwrap();
378        let encoding = tokenizer.encode("<ctl>").unwrap();
379        assert_eq!(encoding.token_ids(), &[23]);
380        assert_eq!(tokenizer.decode(&[23], false).unwrap().as_str(), "<ctl>");
381        assert_eq!(tokenizer.decode(&[23], true).unwrap().as_str(), "");
382    }
383}