1use std::collections::HashSet;
14use std::path::Path;
15
16use rayon::prelude::*;
17use rustc_hash::FxHashMap;
18
19use super::{
20 EncodeSegment, Encoding, Error, Result, TokenIdType,
21 hf::HuggingFaceTokenizer,
22 tiktoken,
23 traits::{DecodeResult, Decoder, Encoder, Tokenizer},
24};
25
26fn fast_encode(encoder: &fastokens::Tokenizer, input: &str) -> Result<Encoding> {
27 let ids = encoder
28 .encode(input)
29 .map_err(|e| Error::msg(format!("Fastokens encode error: {e}")))?;
30 Ok(Encoding::Sp(ids))
31}
32
33fn fast_encode_segments(
34 encoder: &fastokens::Tokenizer,
35 segments: &[EncodeSegment<'_>],
36) -> Result<Encoding> {
37 let segments: Vec<fastokens::EncodeSegment<'_>> = segments
38 .iter()
39 .map(|segment| fastokens::EncodeSegment {
40 text: segment.text,
41 allow_special: segment.allow_special,
42 })
43 .collect();
44 let ids = encoder
45 .encode_segments(&segments)
46 .map_err(|e| Error::msg(format!("Fastokens segmented encode error: {e}")))?;
47 Ok(Encoding::Sp(ids))
48}
49
50pub struct FastTokenizer {
54 fast_encoder: fastokens::Tokenizer,
55 hf_decoder: HuggingFaceTokenizer,
56}
57
58impl FastTokenizer {
59 pub fn from_file(path: &str) -> Result<Self> {
60 let fast_encoder = fastokens::Tokenizer::from_file(Path::new(path))
61 .map_err(|e| Error::msg(format!("Error loading fastokens tokenizer: {e}")))?;
62 let hf_decoder = HuggingFaceTokenizer::from_file(path)?;
63 Ok(Self {
64 fast_encoder,
65 hf_decoder,
66 })
67 }
68}
69
70impl Encoder for FastTokenizer {
71 fn encode(&self, input: &str) -> Result<Encoding> {
72 fast_encode(&self.fast_encoder, input)
73 }
74
75 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
76 inputs.par_iter().map(|input| self.encode(input)).collect()
77 }
78
79 fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
80 fast_encode_segments(&self.fast_encoder, segments)
81 }
82}
83
84impl Decoder for FastTokenizer {
85 fn has_unstable_suffix(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> bool {
86 self.hf_decoder
87 .has_unstable_suffix(token_ids, skip_special_tokens)
88 }
89
90 fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
91 self.hf_decoder.decode(token_ids, skip_special_tokens)
92 }
93}
94
95impl Tokenizer for FastTokenizer {
96 fn validate_prefix_cache(&self) -> Result<()> {
97 Ok(())
98 }
99
100 fn vocab_size(&self) -> Option<usize> {
103 self.hf_decoder.vocab_size()
104 }
105
106 fn token_to_id(&self, token: &str) -> Result<Option<TokenIdType>> {
107 self.hf_decoder.token_to_id(token)
108 }
109
110 fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
111 self.hf_decoder.special_token_ids()
112 }
113
114 fn num_special_tokens_added(&self) -> Result<usize> {
115 Ok(0)
116 }
117}
118
119pub struct FastTikTokenTokenizer {
130 inner: fastokens::Tokenizer,
131 id_to_bytes: FxHashMap<u32, Vec<u8>>,
134 special_token_ids: HashSet<u32>,
135 special_tokens: Vec<String>,
136}
137
138impl FastTikTokenTokenizer {
139 pub fn from_file_auto(path: &str) -> Result<Self> {
142 let directory = Path::new(path)
143 .parent()
144 .ok_or_else(|| Error::msg("Cannot determine parent directory of tiktoken file"))?;
145 let pattern = tiktoken::detect_bpe_pattern(directory)?;
146 let encoder = tiktoken::parse_tiktoken_file(path)?;
147 let num_base_tokens = encoder.values().max().map_or(0, |&m| m + 1) as usize;
148 let special_tokens = tiktoken::load_special_tokens(directory, num_base_tokens)?;
149
150 let ranks: Vec<(Vec<u8>, u32)> = encoder.into_iter().collect();
151 let mut id_to_bytes: FxHashMap<u32, Vec<u8>> = ranks
152 .iter()
153 .map(|(bytes, rank)| (*rank, bytes.clone()))
154 .collect();
155 id_to_bytes.extend(
156 special_tokens
157 .iter()
158 .map(|(content, &id)| (id, content.as_bytes().to_vec())),
159 );
160 let special_token_ids = special_tokens.values().copied().collect();
161 let special_token_strings = tiktoken::sorted_special_token_strings(&special_tokens);
162
163 let config =
164 fastokens::tiktoken::TiktokenConfig::new(pattern, special_tokens.into_iter().collect());
165 let inner = fastokens::Tokenizer::from_tiktoken_ranks(&ranks, config).map_err(|e| {
166 Error::msg(format!(
167 "Error loading fastokens tiktoken tokenizer from {path}: {e}"
168 ))
169 })?;
170 Ok(Self {
171 inner,
172 id_to_bytes,
173 special_token_ids,
174 special_tokens: special_token_strings,
175 })
176 }
177
178 pub fn special_tokens(&self) -> &[String] {
181 &self.special_tokens
182 }
183}
184
185impl Encoder for FastTikTokenTokenizer {
186 fn encode(&self, input: &str) -> Result<Encoding> {
187 fast_encode(&self.inner, input)
188 }
189
190 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
191 inputs.par_iter().map(|input| self.encode(input)).collect()
192 }
193
194 fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
195 fast_encode_segments(&self.inner, segments)
196 }
197}
198
199impl Decoder for FastTikTokenTokenizer {
200 fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
201 let mut bytes = Vec::new();
202 for id in token_ids {
203 if skip_special_tokens && self.special_token_ids.contains(id) {
204 continue;
205 }
206 if let Some(token) = self.id_to_bytes.get(id) {
207 bytes.extend_from_slice(token);
208 }
209 }
210 match String::from_utf8(bytes) {
211 Ok(text) => Ok(DecodeResult::Complete(text)),
212 Err(e) => Ok(DecodeResult::from_decoded(
213 String::from_utf8_lossy(e.as_bytes()).into_owned(),
214 )),
215 }
216 }
217}
218
219impl Tokenizer for FastTikTokenTokenizer {
220 fn validate_prefix_cache(&self) -> Result<()> {
221 Err(Error::msg(
222 "fastokens over tiktoken.model does not satisfy the prefix-cache invariant: its \
223 scanner (plain text) and regex path (text containing special tokens) disagree on \
224 the case folding of (?i:'s), so encode(prefix) + encode(suffix) can differ from \
225 encode(prefix + suffix) across a special-token boundary",
226 ))
227 }
228
229 fn token_to_id(&self, token: &str) -> Result<Option<TokenIdType>> {
230 Ok(self.inner.token_to_id(token))
231 }
232
233 fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
234 let mut ids: Vec<TokenIdType> = self.special_token_ids.iter().copied().collect();
235 ids.sort_unstable();
236 Ok(ids)
237 }
238
239 fn num_special_tokens_added(&self) -> Result<usize> {
240 Ok(0)
241 }
242}
243
244#[cfg(test)]
245mod tests {
246 use super::*;
247 use crate::{HuggingFaceTokenizer, TokenizerOptions};
248
249 const TOKENIZER_PATH: &str = concat!(
252 env!("CARGO_MANIFEST_DIR"),
253 "/tests/data/minimal-bpe/tokenizer.json"
254 );
255 const SEGMENTED_TOKENIZER_PATH: &str = concat!(
256 env!("CARGO_MANIFEST_DIR"),
257 "/tests/data/sample-models/TinyLlama_v1.1/tokenizer.json"
258 );
259
260 #[test]
261 fn byte_fallback_stream_matches_full_decode() {
262 let tokenizer: crate::Tokenizer =
263 std::sync::Arc::new(FastTokenizer::from_file(SEGMENTED_TOKENIZER_PATH).unwrap()).into();
264 for (pieces, skip) in [
265 (vec!["<0x61>", "<0xF5>"], false),
266 (vec!["<0x61>", "</s>", "<0xF5>"], false),
267 (vec!["<0x61>", "</s>", "<0xF5>"], true),
268 ] {
269 let ids: Vec<_> = pieces
270 .iter()
271 .map(|piece| tokenizer.token_to_id(piece).unwrap().unwrap())
272 .collect();
273 let expected: String = tokenizer.decode(&ids, skip).unwrap().into();
274 let mut stream = tokenizer.decode_stream(&[], skip);
275 let mut actual = String::new();
276 for id in ids {
277 actual.push_str(&stream.step(id).unwrap().unwrap_or_default());
278 }
279 actual.push_str(&stream.finish().unwrap().unwrap_or_default());
280 assert_eq!(actual, expected, "{pieces:?}, skip={skip}");
281 }
282 }
283
284 #[test]
285 fn test_fast_encode_decode_roundtrip() {
286 let tokenizer = FastTokenizer::from_file(TOKENIZER_PATH).unwrap();
287 let text = "Hello, world!";
292 let encoding = tokenizer.encode(text).unwrap();
293 assert!(!encoding.token_ids().is_empty());
294 let decoded: String = tokenizer.decode(encoding.token_ids(), true).unwrap().into();
295 assert!(!decoded.is_empty());
296 let enc_chars: String = text.chars().filter(|c| !c.is_whitespace()).collect();
298 let dec_chars: String = decoded.chars().filter(|c| !c.is_whitespace()).collect();
299 assert_eq!(
300 enc_chars, dec_chars,
301 "non-space characters must be preserved"
302 );
303 }
304
305 #[test]
306 fn test_fast_matches_hf_encoding() {
307 let fast = FastTokenizer::from_file(TOKENIZER_PATH).unwrap();
308 let hf = HuggingFaceTokenizer::from_file(TOKENIZER_PATH).unwrap();
309
310 for text in &["Hello, world!", "Hello", " world", "He llo"] {
311 let fast_ids = fast.encode(text).unwrap();
312 let hf_ids = hf.encode(text).unwrap();
313 assert_eq!(
314 fast_ids.token_ids(),
315 hf_ids.token_ids(),
316 "fastokens and HuggingFace must produce identical token IDs for '{text}'"
317 );
318 }
319 }
320
321 #[test]
322 fn test_fast_batch_encode() {
323 let tokenizer = FastTokenizer::from_file(TOKENIZER_PATH).unwrap();
324 let inputs = &["Hello", " world", "Hello, world!"];
325 let encodings = tokenizer.encode_batch(inputs).unwrap();
326 assert_eq!(encodings.len(), inputs.len());
327 for (enc, input) in encodings.iter().zip(inputs.iter()) {
328 assert!(
329 !enc.token_ids().is_empty(),
330 "encoding for '{input}' must be non-empty"
331 );
332 }
333 }
334
335 #[test]
336 fn test_fast_segmented_encoding_preserves_trust_boundaries() {
337 let tokenizer = FastTokenizer::from_file(SEGMENTED_TOKENIZER_PATH).unwrap();
338 let upstream =
339 fastokens::Tokenizer::from_file(std::path::Path::new(SEGMENTED_TOKENIZER_PATH))
340 .unwrap();
341 let marker = "<s>";
342
343 let trusted = tokenizer
344 .encode_segments(&[EncodeSegment::control(marker)])
345 .unwrap();
346 assert_eq!(
347 trusted.token_ids(),
348 &[upstream.token_to_id(marker).unwrap()],
349 "trusted renderer output must recognize the control token"
350 );
351
352 let ordinary = tokenizer
353 .encode_segments(&[EncodeSegment::ordinary(marker)])
354 .unwrap();
355 assert_ne!(
356 ordinary.token_ids(),
357 trusted.token_ids(),
358 "untrusted content must encode the control-token spelling as ordinary text"
359 );
360
361 let segments = [
362 EncodeSegment::ordinary("hello "),
363 EncodeSegment::control(marker),
364 EncodeSegment::ordinary(marker),
365 ];
366 let upstream_segments = [
367 fastokens::EncodeSegment::ordinary("hello "),
368 fastokens::EncodeSegment::special(marker),
369 fastokens::EncodeSegment::ordinary(marker),
370 ];
371 let actual = tokenizer.encode_segments(&segments).unwrap();
372 let expected = upstream.encode_segments(&upstream_segments).unwrap();
373 assert_eq!(actual.token_ids(), expected);
374
375 assert!(
376 tokenizer
377 .encode_segments(&[])
378 .unwrap()
379 .token_ids()
380 .is_empty()
381 );
382 }
383
384 #[test]
385 fn test_fast_with_decode_stream() {
386 use crate::Tokenizer as TokenizerWrapper;
387 use std::sync::Arc;
388
389 let tokenizer = Arc::new(FastTokenizer::from_file(TOKENIZER_PATH).unwrap());
390 let wrapper = TokenizerWrapper::from(tokenizer);
391
392 let prompt_ids = wrapper.encode("Hello").unwrap().token_ids().to_vec();
394 let continuation = ", world!";
395 let cont_ids = wrapper.encode(continuation).unwrap().token_ids().to_vec();
396
397 let mut stream = wrapper.decode_stream(&prompt_ids, true);
398 let mut accumulated = String::new();
400 for id in &cont_ids {
401 if let Some(chunk) = stream.step(*id).unwrap() {
402 accumulated.push_str(&chunk);
403 }
404 }
405
406 let mut all_ids = prompt_ids.clone();
410 all_ids.extend_from_slice(&cont_ids);
411 let full_text: String = wrapper.decode(&all_ids, true).unwrap().into();
412 let prompt_text: String = wrapper.decode(&prompt_ids, true).unwrap().into();
413 let expected = &full_text[prompt_text.len()..];
414 assert_eq!(
415 accumulated, expected,
416 "streamed chunks must equal context-aware decoded continuation"
417 );
418 }
419
420 #[test]
421 fn vocabulary_metadata_forwards_to_hf_decoder() {
422 let fast = FastTokenizer::from_file(TOKENIZER_PATH).unwrap();
423 let hf = HuggingFaceTokenizer::from_file(TOKENIZER_PATH).unwrap();
424 assert_eq!(fast.vocab_size(), hf.vocab_size());
425 assert_eq!(
426 fast.token_to_id("Hello").unwrap(),
427 hf.token_to_id("Hello").unwrap()
428 );
429 assert_eq!(
430 fast.special_token_ids().unwrap(),
431 hf.special_token_ids().unwrap()
432 );
433 }
434
435 #[test]
436 fn special_token_accounting_matches_fast_encoder() {
437 let fast = FastTokenizer::from_file(SEGMENTED_TOKENIZER_PATH).unwrap();
438 let upstream =
439 fastokens::Tokenizer::from_file(std::path::Path::new(SEGMENTED_TOKENIZER_PATH))
440 .unwrap();
441 let hf = HuggingFaceTokenizer::from_file(SEGMENTED_TOKENIZER_PATH).unwrap();
442
443 assert_eq!(hf.num_special_tokens_added().unwrap(), 1);
444 assert_eq!(fast.num_special_tokens_added().unwrap(), 0);
445 let hf_with_special_tokens = hf.with_options(TokenizerOptions {
446 add_special_tokens: true,
447 });
448
449 for text in ["hello", "hello there"] {
450 let fast_ids = fast.encode(text).unwrap();
451 assert_eq!(
452 fast_ids.token_ids(),
453 upstream.encode(text).unwrap(),
454 "FastTokenizer must match the encoder that omits the HF post-processor"
455 );
456 assert_eq!(
457 hf_with_special_tokens
458 .encode(text)
459 .unwrap()
460 .token_ids()
461 .len(),
462 fast_ids.token_ids().len() + 1,
463 "the HF post-processor must add the BOS token FastTokenizer omits"
464 );
465 }
466 }
467}
468
469#[cfg(test)]
470mod tiktoken_parity_tests {
471 use super::*;
472 use crate::TikTokenTokenizer;
473
474 const TIKTOKEN_PATH: &str = concat!(
475 env!("CARGO_MANIFEST_DIR"),
476 "/tests/data/sample-models/mock-tiktoken-bpe/tiktoken.model"
477 );
478 const CONTRACTION_PATH: &str = concat!(
481 env!("CARGO_MANIFEST_DIR"),
482 "/tests/data/sample-models/mock-tiktoken-contraction/tiktoken.model"
483 );
484
485 fn pair() -> (TikTokenTokenizer, FastTikTokenTokenizer) {
486 let reference = TikTokenTokenizer::from_file_auto(TIKTOKEN_PATH).unwrap();
487 let fast = FastTikTokenTokenizer::from_file_auto(TIKTOKEN_PATH).unwrap();
488 (reference, fast)
489 }
490
491 fn corpus() -> Vec<String> {
492 vec![
493 String::new(),
494 "hello world".into(),
495 "Hello, World! 123 4567 89".into(),
496 " leading and trailing ".into(),
497 "tabs\tand\nnewlines\r\nmixed spacing".into(),
498 "<|im_start|>user\nhi there<|im_end|><|im_start|>assistant\n".into(),
499 "a literal <|im_start|> inside plain text".into(),
500 "emoji 😀🚀 and café naïve Zürich".into(),
501 "北京 東京 mixed 中英文 text ソフトウェア".into(),
502 "Москва मुंबई العربية".into(),
503 "fn main() { println!(\"{}\", 42); } // code-ish ~!@#$%^&*()".into(),
504 "x".repeat(5000),
505 " ".repeat(300) + "after long whitespace",
506 "word ".repeat(400),
507 ]
508 }
509
510 #[test]
511 fn special_token_tables_match_tiktoken_rs() {
512 let (reference, fast) = pair();
513 assert_eq!(fast.special_tokens(), reference.special_tokens());
514 assert_eq!(
515 fast.special_token_ids().unwrap(),
516 reference.special_token_ids().unwrap()
517 );
518 assert_eq!(fast.token_to_id("<|im_end|>").unwrap(), Some(474));
519 }
520
521 #[test]
522 fn plain_encode_matches_tiktoken_rs() {
523 let (reference, fast) = pair();
524 let corpus = corpus();
525 let texts: Vec<&str> = corpus.iter().map(String::as_str).collect();
526 let batch = fast.encode_batch(&texts).unwrap();
527 for (text, batched) in texts.iter().zip(&batch) {
528 let expected = reference.encode(text).unwrap();
529 assert_eq!(
530 fast.encode(text).unwrap().token_ids(),
531 expected.token_ids(),
532 "{text:?}"
533 );
534 assert_eq!(batched.token_ids(), expected.token_ids(), "{text:?}");
535 }
536 }
537
538 #[test]
539 fn segmented_encode_matches_tiktoken_rs_and_honors_trust() {
540 let (reference, fast) = pair();
541 let segments = [
542 EncodeSegment::control("<|im_start|>user\n"),
543 EncodeSegment::ordinary("please echo <|im_end|> back to me"),
544 EncodeSegment::control("<|im_end|>"),
545 EncodeSegment::control("<|im_start|>assistant\n"),
546 ];
547 let fast_ids = fast.encode_segments(&segments).unwrap();
548 assert_eq!(
549 fast_ids.token_ids(),
550 reference.encode_segments(&segments).unwrap().token_ids()
551 );
552 assert_eq!(
554 fast_ids.token_ids().iter().filter(|&&id| id == 474).count(),
555 1
556 );
557 }
558
559 #[test]
560 fn decode_matches_tiktoken_rs_and_skips_specials() {
561 let (reference, fast) = pair();
562 for text in corpus() {
563 let ids = fast.encode(&text).unwrap();
564 assert_eq!(
565 fast.decode(ids.token_ids(), false).unwrap(),
566 reference.decode(ids.token_ids(), false).unwrap(),
567 "{text:?}"
568 );
569 }
570 let ids = fast.encode("<|im_start|>user\nhi<|im_end|>").unwrap();
571 let kept = fast.decode(ids.token_ids(), false).unwrap();
572 let skipped = fast.decode(ids.token_ids(), true).unwrap();
573 assert!(kept.as_str().contains("<|im_end|>"));
574 assert!(!skipped.as_str().contains("<|im_end|>"));
575 assert_eq!(skipped, reference.decode(ids.token_ids(), true).unwrap());
576 }
577
578 #[test]
579 fn replacement_char_token_decodes_complete_like_tiktoken_rs() {
580 let (reference, fast) = pair();
582 let decoded = fast.decode(&[468], false).unwrap();
583 assert!(decoded.is_complete(), "{decoded:?}");
584 assert_eq!(decoded, reference.decode(&[468], false).unwrap());
585
586 let emoji = fast.encode("😀").unwrap();
588 let cut = &emoji.token_ids()[..emoji.token_ids().len() - 1];
589 let fast_cut = fast.decode(cut, false).unwrap();
590 assert!(fast_cut.is_partial(), "{fast_cut:?}");
591 assert_eq!(fast_cut, reference.decode(cut, false).unwrap());
592
593 assert_eq!(fast.decode(&[468, 9_999_999], false).unwrap(), decoded);
595 }
596
597 #[test]
598 fn plain_text_prefix_cache_is_rejected() {
599 let fast = FastTikTokenTokenizer::from_file_auto(CONTRACTION_PATH).unwrap();
600 let specials = fast.special_tokens().to_vec();
601 let error = crate::CachedTokenizer::new(std::sync::Arc::new(fast), specials, 1 << 20)
602 .err()
603 .expect("fastokens over tiktoken.model must not be prefix-cached")
604 .to_string();
605 assert!(error.contains("fastokens"), "{error}");
606 }
607
608 #[test]
615 fn canary_fastokens_scanner_disagrees_with_its_regex_path() {
616 let reference = TikTokenTokenizer::from_file_auto(CONTRACTION_PATH).unwrap();
617 let fast = FastTikTokenTokenizer::from_file_auto(CONTRACTION_PATH).unwrap();
618 let special = "<|end_of_msg|>";
619 let suffix = " I'\u{17f}";
620 let full = format!("{special}{suffix}");
621
622 let ref_full = reference.encode(&full).unwrap().token_ids().to_vec();
623 let ref_suffix = reference.encode(suffix).unwrap().token_ids().to_vec();
624 assert_eq!(
625 ref_full[1..],
626 ref_suffix[..],
627 "tiktoken-rs must be self-consistent"
628 );
629
630 let fast_full = fast.encode(&full).unwrap().token_ids().to_vec();
631 let fast_suffix = fast.encode(suffix).unwrap().token_ids().to_vec();
632 assert_eq!(fast_full, ref_full, "the regex path matches tiktoken-rs");
633 assert_ne!(
634 fast_full[1..],
635 fast_suffix[..],
636 "fastokens' scanner now agrees with its regex path; revisit validate_prefix_cache"
637 );
638 }
639}