1use 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
20pub 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 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 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 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 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 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 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 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}