Skip to main content

dynamo_tokenizers/
fastokens.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Fastokens backend using the `fastokens` crate for high-performance BPE encoding.
5//!
6//! This module preserves the existing hybrid behavior: `fastokens` handles encoding and
7//! `HuggingFaceTokenizer` handles decoding. Both are loaded from the same `tokenizer.json`
8//! file.
9
10use std::path::Path;
11
12use rayon::prelude::*;
13
14use super::{
15    EncodeSegment, Encoding, Error, Result, TokenIdType,
16    hf::HuggingFaceTokenizer,
17    traits::{DecodeResult, Decoder, Encoder, Tokenizer},
18};
19
20/// Hybrid tokenizer: fast BPE encoding via `fastokens`, decoding via HuggingFace.
21///
22/// Both backends are loaded from the same `tokenizer.json` file.
23pub struct FastTokenizer {
24    fast_encoder: fastokens::Tokenizer,
25    hf_decoder: HuggingFaceTokenizer,
26}
27
28impl FastTokenizer {
29    pub fn from_file(path: &str) -> Result<Self> {
30        let fast_encoder = fastokens::Tokenizer::from_file(Path::new(path))
31            .map_err(|e| Error::msg(format!("Error loading fastokens tokenizer: {e}")))?;
32        let hf_decoder = HuggingFaceTokenizer::from_file(path)?;
33        Ok(Self {
34            fast_encoder,
35            hf_decoder,
36        })
37    }
38}
39
40impl Encoder for FastTokenizer {
41    fn encode(&self, input: &str) -> Result<Encoding> {
42        let ids = self
43            .fast_encoder
44            .encode(input)
45            .map_err(|e| Error::msg(format!("Fastokens encode error: {e}")))?;
46        Ok(Encoding::Sp(ids))
47    }
48
49    fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
50        inputs.par_iter().map(|input| self.encode(input)).collect()
51    }
52
53    fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
54        let segments: Vec<fastokens::EncodeSegment<'_>> = segments
55            .iter()
56            .map(|segment| fastokens::EncodeSegment {
57                text: segment.text,
58                allow_special: segment.allow_special,
59            })
60            .collect();
61        let ids = self
62            .fast_encoder
63            .encode_segments(&segments)
64            .map_err(|e| Error::msg(format!("Fastokens segmented encode error: {e}")))?;
65        Ok(Encoding::Sp(ids))
66    }
67}
68
69impl Decoder for FastTokenizer {
70    fn has_unstable_suffix(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> bool {
71        self.hf_decoder
72            .has_unstable_suffix(token_ids, skip_special_tokens)
73    }
74
75    fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
76        self.hf_decoder.decode(token_ids, skip_special_tokens)
77    }
78}
79
80impl Tokenizer for FastTokenizer {
81    fn validate_prefix_cache(&self) -> Result<()> {
82        Ok(())
83    }
84
85    // `fast_encoder` and `hf_decoder` are loaded from the same tokenizer.json,
86    // so the HF side's vocabulary introspection applies to both.
87    fn vocab_size(&self) -> Option<usize> {
88        self.hf_decoder.vocab_size()
89    }
90
91    fn token_to_id(&self, token: &str) -> Result<Option<TokenIdType>> {
92        self.hf_decoder.token_to_id(token)
93    }
94
95    fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
96        self.hf_decoder.special_token_ids()
97    }
98
99    fn num_special_tokens_added(&self) -> Result<usize> {
100        Ok(0)
101    }
102}
103
104#[cfg(test)]
105mod tests {
106    use super::*;
107    use crate::{HuggingFaceTokenizer, TokenizerOptions};
108
109    // Minimal synthetic BPE tokenizer with no normalizer or post-processor --
110    // compatible with fastokens. Vocab covers: H,T,a,d,e,h,i,l,o,r,s,t,w + punctuation.
111    const TOKENIZER_PATH: &str = concat!(
112        env!("CARGO_MANIFEST_DIR"),
113        "/tests/data/minimal-bpe/tokenizer.json"
114    );
115    const SEGMENTED_TOKENIZER_PATH: &str = concat!(
116        env!("CARGO_MANIFEST_DIR"),
117        "/tests/data/sample-models/TinyLlama_v1.1/tokenizer.json"
118    );
119
120    #[test]
121    fn byte_fallback_stream_matches_full_decode() {
122        let tokenizer: crate::Tokenizer =
123            std::sync::Arc::new(FastTokenizer::from_file(SEGMENTED_TOKENIZER_PATH).unwrap()).into();
124        for (pieces, skip) in [
125            (vec!["<0x61>", "<0xF5>"], false),
126            (vec!["<0x61>", "</s>", "<0xF5>"], false),
127            (vec!["<0x61>", "</s>", "<0xF5>"], true),
128        ] {
129            let ids: Vec<_> = pieces
130                .iter()
131                .map(|piece| tokenizer.token_to_id(piece).unwrap().unwrap())
132                .collect();
133            let expected: String = tokenizer.decode(&ids, skip).unwrap().into();
134            let mut stream = tokenizer.decode_stream(&[], skip);
135            let mut actual = String::new();
136            for id in ids {
137                actual.push_str(&stream.step(id).unwrap().unwrap_or_default());
138            }
139            actual.push_str(&stream.finish().unwrap().unwrap_or_default());
140            assert_eq!(actual, expected, "{pieces:?}, skip={skip}");
141        }
142    }
143
144    #[test]
145    fn test_fast_encode_decode_roundtrip() {
146        let tokenizer = FastTokenizer::from_file(TOKENIZER_PATH).unwrap();
147        // Encode then decode: verifies both paths execute without error.
148        // With a null decoder, HF inserts spaces between tokens so exact equality
149        // is not expected here -- we just verify the operations succeed and produce
150        // non-empty results.
151        let text = "Hello, world!";
152        let encoding = tokenizer.encode(text).unwrap();
153        assert!(!encoding.token_ids().is_empty());
154        let decoded: String = tokenizer.decode(encoding.token_ids(), true).unwrap().into();
155        assert!(!decoded.is_empty());
156        // The decoded text should contain the same non-space characters
157        let enc_chars: String = text.chars().filter(|c| !c.is_whitespace()).collect();
158        let dec_chars: String = decoded.chars().filter(|c| !c.is_whitespace()).collect();
159        assert_eq!(
160            enc_chars, dec_chars,
161            "non-space characters must be preserved"
162        );
163    }
164
165    #[test]
166    fn test_fast_matches_hf_encoding() {
167        let fast = FastTokenizer::from_file(TOKENIZER_PATH).unwrap();
168        let hf = HuggingFaceTokenizer::from_file(TOKENIZER_PATH).unwrap();
169
170        for text in &["Hello, world!", "Hello", " world", "He llo"] {
171            let fast_ids = fast.encode(text).unwrap();
172            let hf_ids = hf.encode(text).unwrap();
173            assert_eq!(
174                fast_ids.token_ids(),
175                hf_ids.token_ids(),
176                "fastokens and HuggingFace must produce identical token IDs for '{text}'"
177            );
178        }
179    }
180
181    #[test]
182    fn test_fast_batch_encode() {
183        let tokenizer = FastTokenizer::from_file(TOKENIZER_PATH).unwrap();
184        let inputs = &["Hello", " world", "Hello, world!"];
185        let encodings = tokenizer.encode_batch(inputs).unwrap();
186        assert_eq!(encodings.len(), inputs.len());
187        for (enc, input) in encodings.iter().zip(inputs.iter()) {
188            assert!(
189                !enc.token_ids().is_empty(),
190                "encoding for '{input}' must be non-empty"
191            );
192        }
193    }
194
195    #[test]
196    fn test_fast_segmented_encoding_preserves_trust_boundaries() {
197        let tokenizer = FastTokenizer::from_file(SEGMENTED_TOKENIZER_PATH).unwrap();
198        let upstream =
199            fastokens::Tokenizer::from_file(std::path::Path::new(SEGMENTED_TOKENIZER_PATH))
200                .unwrap();
201        let marker = "<s>";
202
203        let trusted = tokenizer
204            .encode_segments(&[EncodeSegment::control(marker)])
205            .unwrap();
206        assert_eq!(
207            trusted.token_ids(),
208            &[upstream.token_to_id(marker).unwrap()],
209            "trusted renderer output must recognize the control token"
210        );
211
212        let ordinary = tokenizer
213            .encode_segments(&[EncodeSegment::ordinary(marker)])
214            .unwrap();
215        assert_ne!(
216            ordinary.token_ids(),
217            trusted.token_ids(),
218            "untrusted content must encode the control-token spelling as ordinary text"
219        );
220
221        let segments = [
222            EncodeSegment::ordinary("hello "),
223            EncodeSegment::control(marker),
224            EncodeSegment::ordinary(marker),
225        ];
226        let upstream_segments = [
227            fastokens::EncodeSegment::ordinary("hello "),
228            fastokens::EncodeSegment::special(marker),
229            fastokens::EncodeSegment::ordinary(marker),
230        ];
231        let actual = tokenizer.encode_segments(&segments).unwrap();
232        let expected = upstream.encode_segments(&upstream_segments).unwrap();
233        assert_eq!(actual.token_ids(), expected);
234
235        assert!(
236            tokenizer
237                .encode_segments(&[])
238                .unwrap()
239                .token_ids()
240                .is_empty()
241        );
242    }
243
244    #[test]
245    fn test_fast_with_decode_stream() {
246        use crate::Tokenizer as TokenizerWrapper;
247        use std::sync::Arc;
248
249        let tokenizer = Arc::new(FastTokenizer::from_file(TOKENIZER_PATH).unwrap());
250        let wrapper = TokenizerWrapper::from(tokenizer);
251
252        // Encode a prompt and a continuation, then step through the decode stream
253        let prompt_ids = wrapper.encode("Hello").unwrap().token_ids().to_vec();
254        let continuation = ", world!";
255        let cont_ids = wrapper.encode(continuation).unwrap().token_ids().to_vec();
256
257        let mut stream = wrapper.decode_stream(&prompt_ids, true);
258        // Accumulate incremental chunks from decode_stream
259        let mut accumulated = String::new();
260        for id in &cont_ids {
261            if let Some(chunk) = stream.step(*id).unwrap() {
262                accumulated.push_str(&chunk);
263            }
264        }
265
266        // DecodeStream uses prompt tokens as context, so the expected text is
267        // decode(prompt + continuation) minus decode(prompt) -- not a bare
268        // decode(continuation) which lacks the surrounding context.
269        let mut all_ids = prompt_ids.clone();
270        all_ids.extend_from_slice(&cont_ids);
271        let full_text: String = wrapper.decode(&all_ids, true).unwrap().into();
272        let prompt_text: String = wrapper.decode(&prompt_ids, true).unwrap().into();
273        let expected = &full_text[prompt_text.len()..];
274        assert_eq!(
275            accumulated, expected,
276            "streamed chunks must equal context-aware decoded continuation"
277        );
278    }
279
280    #[test]
281    fn vocabulary_metadata_forwards_to_hf_decoder() {
282        let fast = FastTokenizer::from_file(TOKENIZER_PATH).unwrap();
283        let hf = HuggingFaceTokenizer::from_file(TOKENIZER_PATH).unwrap();
284        assert_eq!(fast.vocab_size(), hf.vocab_size());
285        assert_eq!(
286            fast.token_to_id("Hello").unwrap(),
287            hf.token_to_id("Hello").unwrap()
288        );
289        assert_eq!(
290            fast.special_token_ids().unwrap(),
291            hf.special_token_ids().unwrap()
292        );
293    }
294
295    #[test]
296    fn special_token_accounting_matches_fast_encoder() {
297        let fast = FastTokenizer::from_file(SEGMENTED_TOKENIZER_PATH).unwrap();
298        let upstream =
299            fastokens::Tokenizer::from_file(std::path::Path::new(SEGMENTED_TOKENIZER_PATH))
300                .unwrap();
301        let hf = HuggingFaceTokenizer::from_file(SEGMENTED_TOKENIZER_PATH).unwrap();
302
303        assert_eq!(hf.num_special_tokens_added().unwrap(), 1);
304        assert_eq!(fast.num_special_tokens_added().unwrap(), 0);
305        let hf_with_special_tokens = hf.with_options(TokenizerOptions {
306            add_special_tokens: true,
307        });
308
309        for text in ["hello", "hello there"] {
310            let fast_ids = fast.encode(text).unwrap();
311            assert_eq!(
312                fast_ids.token_ids(),
313                upstream.encode(text).unwrap(),
314                "FastTokenizer must match the encoder that omits the HF post-processor"
315            );
316            assert_eq!(
317                hf_with_special_tokens
318                    .encode(text)
319                    .unwrap()
320                    .token_ids()
321                    .len(),
322                fast_ids.token_ids().len() + 1,
323                "the HF post-processor must add the BOS token FastTokenizer omits"
324            );
325        }
326    }
327}