1use std::collections::HashMap;
11use std::path::Path;
12use std::sync::Mutex;
13
14use crate::error::{ForgeError, Result};
15
16pub mod char;
18pub use self::char::CharTokenizer;
19
20const SPLIT_PATTERN: &str =
21 r"'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+";
22
23pub trait Tokenizer {
26 fn encode(&self, text: &str) -> Result<Vec<u32>>;
27 fn decode(&self, ids: &[u32]) -> String;
28 fn decode_bytes(&self, ids: &[u32]) -> Vec<u8>;
32 fn vocab_size(&self) -> usize;
33}
34
35pub struct Gpt2Tokenizer {
36 encoder: HashMap<String, u32>,
37 decoder: HashMap<u32, String>,
38 bpe_ranks: HashMap<(String, String), usize>,
39 byte_encoder: HashMap<u8, char>,
40 byte_decoder: HashMap<char, u8>,
41 pattern: fancy_regex::Regex,
42 cache: Mutex<HashMap<String, Vec<String>>>,
43}
44
45fn bytes_to_unicode() -> HashMap<u8, char> {
47 let mut bs: Vec<u16> = (b'!'..=b'~').map(u16::from).collect();
48 bs.extend((0xA1u16..=0xACu16).chain(0xAEu16..=0xFFu16));
49 let mut map = HashMap::new();
50 for &b in &bs {
51 map.insert(b as u8, char::from_u32(b as u32).unwrap());
52 }
53 let mut n = 0u32;
54 for b in 0u16..256 {
55 if !bs.contains(&b) {
56 map.insert(b as u8, char::from_u32(256 + n).unwrap());
57 n += 1;
58 }
59 }
60 map
61}
62
63impl Gpt2Tokenizer {
64 pub fn from_files(vocab_path: impl AsRef<Path>, merges_path: impl AsRef<Path>) -> Result<Self> {
65 let vocab_json = std::fs::read_to_string(vocab_path)?;
66 let merges = std::fs::read_to_string(merges_path)?;
67 Self::from_strs(&vocab_json, &merges)
68 }
69
70 pub fn from_strs(vocab_json: &str, merges: &str) -> Result<Self> {
73 let encoder: HashMap<String, u32> = serde_json::from_str(vocab_json)?;
74 let decoder = encoder.iter().map(|(k, v)| (*v, k.clone())).collect();
75
76 let mut bpe_ranks = HashMap::new();
77 for (i, line) in merges
78 .lines()
79 .filter(|l| !l.starts_with("#version") && !l.trim().is_empty())
80 .enumerate()
81 {
82 let mut parts = line.split(' ');
83 let (a, b) = (
84 parts.next().ok_or_else(|| bad_merge(line))?,
85 parts.next().ok_or_else(|| bad_merge(line))?,
86 );
87 bpe_ranks.insert((a.to_string(), b.to_string()), i);
88 }
89
90 let byte_encoder = bytes_to_unicode();
91 let byte_decoder = byte_encoder.iter().map(|(&b, &c)| (c, b)).collect();
92 let pattern = fancy_regex::Regex::new(SPLIT_PATTERN)
93 .map_err(|e| ForgeError::Tokenizer(format!("pattern: {e}")))?;
94 Ok(Gpt2Tokenizer {
95 encoder,
96 decoder,
97 bpe_ranks,
98 byte_encoder,
99 byte_decoder,
100 pattern,
101 cache: Mutex::new(HashMap::new()),
102 })
103 }
104
105 pub fn from_dir(dir: impl AsRef<Path>) -> Result<Self> {
107 let dir = dir.as_ref();
108 Self::from_files(dir.join("vocab.json"), dir.join("merges.txt"))
109 }
110
111 pub fn encode(&self, text: &str) -> Result<Vec<u32>> {
112 let mut ids = Vec::new();
113 for m in self.pattern.find_iter(text) {
114 let piece = m
115 .map_err(|e| ForgeError::Tokenizer(format!("regex: {e}")))?
116 .as_str();
117 let mapped: String = piece.bytes().map(|b| self.byte_encoder[&b]).collect();
118 for token in self.bpe(&mapped) {
119 let id = self.encoder.get(&token).ok_or_else(|| {
120 ForgeError::Tokenizer(format!("token {token:?} not in vocab"))
121 })?;
122 ids.push(*id);
123 }
124 }
125 Ok(ids)
126 }
127
128 pub fn decode(&self, ids: &[u32]) -> String {
129 String::from_utf8_lossy(&self.decode_bytes(ids)).into_owned()
130 }
131
132 pub fn decode_bytes(&self, ids: &[u32]) -> Vec<u8> {
138 ids.iter()
139 .filter_map(|id| self.decoder.get(id))
140 .flat_map(|s| s.chars())
141 .filter_map(|c| self.byte_decoder.get(&c).copied())
142 .collect()
143 }
144
145 pub fn vocab_size(&self) -> usize {
146 self.encoder.len()
147 }
148
149 fn bpe(&self, word: &str) -> Vec<String> {
151 if let Some(hit) = self.cache.lock().unwrap().get(word) {
152 return hit.clone();
153 }
154 let mut parts: Vec<String> = word.chars().map(|c| c.to_string()).collect();
155 while parts.len() > 1 {
156 let best = parts
157 .windows(2)
158 .filter_map(|w| {
159 self.bpe_ranks
160 .get(&(w[0].clone(), w[1].clone()))
161 .map(|&r| (r, (w[0].clone(), w[1].clone())))
162 })
163 .min_by_key(|(r, _)| *r);
164 let Some((_, (a, b))) = best else { break };
165 let mut merged = Vec::with_capacity(parts.len());
166 let mut i = 0;
167 while i < parts.len() {
168 if i + 1 < parts.len() && parts[i] == a && parts[i + 1] == b {
169 merged.push(format!("{a}{b}"));
170 i += 2;
171 } else {
172 merged.push(parts[i].clone());
173 i += 1;
174 }
175 }
176 parts = merged;
177 }
178 self.cache
179 .lock()
180 .unwrap()
181 .insert(word.to_string(), parts.clone());
182 parts
183 }
184}
185
186impl Tokenizer for Gpt2Tokenizer {
190 fn encode(&self, text: &str) -> Result<Vec<u32>> {
191 Gpt2Tokenizer::encode(self, text)
192 }
193
194 fn decode(&self, ids: &[u32]) -> String {
195 Gpt2Tokenizer::decode(self, ids)
196 }
197
198 fn decode_bytes(&self, ids: &[u32]) -> Vec<u8> {
199 Gpt2Tokenizer::decode_bytes(self, ids)
200 }
201
202 fn vocab_size(&self) -> usize {
203 Gpt2Tokenizer::vocab_size(self)
204 }
205}
206
207pub enum AnyTokenizer {
211 Bpe(Box<Gpt2Tokenizer>),
215 Char(CharTokenizer),
216}
217
218impl AnyTokenizer {
219 pub fn bpe(t: Gpt2Tokenizer) -> Self {
220 AnyTokenizer::Bpe(Box::new(t))
221 }
222
223 pub fn from_dir(dir: impl AsRef<Path>) -> Result<Self> {
226 let dir = dir.as_ref();
227 if dir.join("merges.txt").exists() {
228 Ok(AnyTokenizer::bpe(Gpt2Tokenizer::from_dir(dir)?))
229 } else {
230 Ok(AnyTokenizer::Char(CharTokenizer::from_json(
231 &std::fs::read_to_string(dir.join("vocab.json"))?,
232 )?))
233 }
234 }
235
236 pub fn kind(&self) -> &'static str {
237 match self {
238 AnyTokenizer::Bpe(_) => "bpe",
239 AnyTokenizer::Char(_) => "char",
240 }
241 }
242}
243
244impl Tokenizer for AnyTokenizer {
245 fn encode(&self, text: &str) -> Result<Vec<u32>> {
246 match self {
247 AnyTokenizer::Bpe(t) => t.encode(text),
248 AnyTokenizer::Char(t) => t.encode(text),
249 }
250 }
251
252 fn decode(&self, ids: &[u32]) -> String {
253 match self {
254 AnyTokenizer::Bpe(t) => t.decode(ids),
255 AnyTokenizer::Char(t) => t.decode(ids),
256 }
257 }
258
259 fn decode_bytes(&self, ids: &[u32]) -> Vec<u8> {
260 match self {
261 AnyTokenizer::Bpe(t) => t.decode_bytes(ids),
262 AnyTokenizer::Char(t) => t.decode_bytes(ids),
263 }
264 }
265
266 fn vocab_size(&self) -> usize {
267 match self {
268 AnyTokenizer::Bpe(t) => t.vocab_size(),
269 AnyTokenizer::Char(t) => t.vocab_size(),
270 }
271 }
272}
273
274fn bad_merge(line: &str) -> ForgeError {
275 ForgeError::Tokenizer(format!("malformed merges line: {line:?}"))
276}