1use anyhow::{bail, Result};
5use std::sync::Arc;
6use tiktoken_rs::{CoreBPE, Rank};
7use toktrie::{TokEnv, TokRxInfo, TokTrie, TokenId, TokenizerEnv};
8
9pub struct TikTokenBPE {
12 pub bpe: CoreBPE,
14 tok_trie: TokTrie,
15}
16
17impl TikTokenBPE {
18 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 pub fn tokrx_info(&self) -> TokRxInfo {
91 *self.tok_trie.info()
92 }
93
94 pub fn set_eos_tokens(&mut self, tokens: &[TokenId]) {
96 self.tok_trie = self.tok_trie.with_eos_tokens(tokens);
97 }
98
99 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 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 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}