Skip to main content

forge/tokenizer/
mod.rs

1//! Tokenizers: the GPT-2 byte-level BPE, and nanoGPT's character-level vocab.
2//!
3//! [`Gpt2Tokenizer`] loads the original `vocab.json` + `merges.txt`. The split
4//! pattern uses a negative lookahead, which the standard `regex` crate cannot
5//! express, so `fancy-regex` is used (see roadmap "Known Pitfalls").
6//!
7//! The generation path in [`crate::Gpt2`] is generic over the [`Tokenizer`]
8//! trait, so [`CharTokenizer`] plugs into it unchanged.
9
10use std::collections::HashMap;
11use std::path::Path;
12use std::sync::Mutex;
13
14use crate::error::{ForgeError, Result};
15
16// `self::` is required: a bare `char` path would resolve as a crate name.
17pub 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
23/// The tokenizer surface the generation path needs — exactly the four methods
24/// [`crate::Gpt2::generate`] and friends call.
25pub trait Tokenizer {
26    fn encode(&self, text: &str) -> Result<Vec<u32>>;
27    fn decode(&self, ids: &[u32]) -> String;
28    /// Raw byte-level decode. Must be append-only per token — i.e. `decode_bytes(a ++ b)`
29    /// equals `decode_bytes(a) ++ decode_bytes(b)` — because the streaming path
30    /// emits only the valid-UTF-8 prefix of the accumulated bytes.
31    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
45/// GPT-2's reversible byte -> printable-unicode mapping.
46fn 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    /// Build from the contents of `vocab.json` and `merges.txt` — the primary
71    /// form; on wasm the strings arrive from an HTTP fetch.
72    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    /// Load `vocab.json` + `merges.txt` from a directory.
106    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    /// Raw byte-level decode. Unlike [`Gpt2Tokenizer::decode`] this is exact
133    /// and append-only per token (id sequence a ++ b decodes to
134    /// bytes(a) ++ bytes(b)), which streaming builds on: a multi-byte UTF-8
135    /// character split across BPE tokens completes once the next token's
136    /// bytes arrive.
137    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    /// Standard BPE merge loop over a byte-mapped word.
150    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
186/// Delegates to the inherent methods, which stay public so existing callers
187/// compile unchanged (inherent methods win over trait methods at the call
188/// site, so this is not recursive).
189impl 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
207/// Either tokenizer, for callers that pick one at runtime (the TUI and the
208/// wasm facade both do). Generation stays generic over `impl Tokenizer`, so
209/// this dispatches without a vtable.
210pub enum AnyTokenizer {
211    /// Boxed because a `Gpt2Tokenizer` (four hash maps plus a compiled regex)
212    /// is ~6× the size of a `CharTokenizer`, and every value of this enum
213    /// would otherwise pay for the larger one.
214    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    /// Load whichever tokenizer `dir` holds: BPE when `merges.txt` is present
224    /// alongside `vocab.json`, otherwise the character vocab.
225    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}