Skip to main content

toktrie_tiktoken/
lib.rs

1//! This crate integrates the [`tiktoken`](tiktoken_rs) BPE tokenizer (used by OpenAI models)
2//! with [`toktrie`], providing a [`TokenizerEnv`] implementation backed by tiktoken's [`CoreBPE`].
3
4use anyhow::{bail, Result};
5use std::sync::Arc;
6use tiktoken_rs::{CoreBPE, Rank};
7use toktrie::{TokEnv, TokRxInfo, TokTrie, TokenId, TokenizerEnv};
8
9/// A tiktoken BPE tokenizer paired with a [`TokTrie`] for efficient
10/// constrained-decoding support. Implements [`TokenizerEnv`].
11pub struct TikTokenBPE {
12    /// The underlying tiktoken [`CoreBPE`] encoder.
13    pub bpe: CoreBPE,
14    tok_trie: TokTrie,
15}
16
17impl TikTokenBPE {
18    /// Creates a new `TikTokenBPE` from a BPE encoder vocabulary, special tokens,
19    /// a regex pattern, an optional vocabulary size override, and an EOS token ID.
20    ///
21    /// Empty token slots are filled with placeholder special tokens.
22    /// Returns an error if `n_vocab_override` is smaller than the actual vocabulary.
23    pub fn new(
24        encoder: Vec<(Vec<u8>, Rank)>,
25        special_tokens_encoder: Vec<(String, Rank)>,
26        pattern: &str,
27        n_vocab_override: Option<usize>,
28        eos_token: u32,
29    ) -> Result<TikTokenBPE> {
30        let mut n_vocab = encoder.len() + special_tokens_encoder.len();
31        let mut tokens = vec![vec![]; n_vocab];
32
33        for (bytes, idx) in encoder.iter() {
34            while tokens.len() <= *idx as usize {
35                tokens.push(vec![]);
36            }
37            tokens[*idx as usize] = bytes.clone();
38        }
39
40        for (name, idx) in special_tokens_encoder.iter() {
41            while tokens.len() <= *idx as usize {
42                tokens.push(vec![]);
43            }
44            let mut spec_bytes = Vec::with_capacity(name.len() + 1);
45            spec_bytes.push(TokTrie::SPECIAL_TOKEN_MARKER);
46            spec_bytes.extend_from_slice(name.as_bytes());
47            tokens[*idx as usize] = spec_bytes;
48        }
49
50        n_vocab = tokens.len();
51
52        if let Some(n_vocab_override) = n_vocab_override {
53            if n_vocab_override < n_vocab {
54                bail!("vocab size too small; {} vs {}", n_vocab_override, n_vocab);
55            }
56            n_vocab = n_vocab_override;
57            tokens.resize(n_vocab, vec![]);
58        }
59
60        for (i, token) in tokens.iter_mut().enumerate() {
61            if token.is_empty() {
62                let mut name = format!(".<[{i}]>").into_bytes();
63                name[0] = TokTrie::SPECIAL_TOKEN_MARKER;
64                *token = name;
65            }
66        }
67
68        let tok_trie = TokTrie::from(
69            &TokRxInfo {
70                vocab_size: n_vocab as u32,
71                tok_eos: eos_token,
72                tok_end_of_turn: None,
73                tok_unk: None,
74                tok_pad: None,
75                tok_bos: None,
76            },
77            &tokens,
78        );
79
80        let bpe = CoreBPE::new(
81            encoder.into_iter().collect(),
82            special_tokens_encoder.into_iter().collect(),
83            pattern,
84        )?;
85
86        Ok(TikTokenBPE { bpe, tok_trie })
87    }
88
89    /// Returns the [`TokRxInfo`] metadata for this tokenizer.
90    pub fn tokrx_info(&self) -> TokRxInfo {
91        *self.tok_trie.info()
92    }
93
94    /// Replaces the set of end-of-sequence tokens recognized by the trie.
95    pub fn set_eos_tokens(&mut self, tokens: &[TokenId]) {
96        self.tok_trie = self.tok_trie.with_eos_tokens(tokens);
97    }
98
99    /// Wraps this tokenizer in an `Arc`, returning a [`TokEnv`].
100    pub fn to_env(self) -> TokEnv {
101        Arc::new(self)
102    }
103}
104
105impl TokenizerEnv for TikTokenBPE {
106    fn tok_trie(&self) -> &TokTrie {
107        &self.tok_trie
108    }
109
110    /// Tokenizes raw bytes using trie-based greedy fallback to tiktoken BPE encoding.
111    fn tokenize_bytes(&self, s: &[u8]) -> Vec<TokenId> {
112        self.tok_trie
113            .tokenize_with_greedy_fallback(s, |s| self.bpe.encode_ordinary(s))
114    }
115
116    /// Like [`tokenize_bytes`](Self::tokenize_bytes), but also recognizes special tokens
117    /// registered in the trie.
118    fn tokenize_bytes_special(&self, s: &[u8]) -> Vec<TokenId> {
119        self.tok_trie.tokenize_with_greedy_fallback(s, |s| {
120            self.tok_trie
121                .tokenize_with_special(s, |s| self.bpe.encode_ordinary(s))
122        })
123    }
124}