Skip to main content

tpt_tokenizer_core/
loader.rs

1//! Loader for the modern unified Hugging Face `tokenizer.json` format.
2//!
3//! Most current Hub repositories ship a single `tokenizer.json` instead of the
4//! legacy `vocab.txt` + `merges.txt` pair. This module parses that file with the
5//! crate's internal [JSON parser](crate::json) (no `serde`) and produces the
6//! appropriate concrete tokenizer.
7//!
8//! Supported `model.type` values: `"BPE"` and `"WordPiece"`. The loader also
9//! honours `added_tokens` (registered as atomic special tokens), detects
10//! GPT-2 byte-level BPE from the pre-tokenizer, and detects lowercasing from a
11//! BERT normalizer.
12
13use alloc::collections::BTreeMap;
14use alloc::string::{String, ToString};
15use alloc::vec::Vec;
16
17use crate::bpe::BpeTokenizer;
18use crate::error::TokenizerError;
19use crate::json::{self, JsonValue};
20use crate::tokenizer::TokenId;
21use crate::wordpiece::WordPieceTokenizer;
22
23/// A tokenizer loaded from a `tokenizer.json`, tagged by its underlying scheme.
24///
25/// Both variants implement [`Tokenizer`](crate::Tokenizer); match on this to
26/// recover the concrete type, or call [`LoadedTokenizer::as_tokenizer`] for a
27/// trait object.
28#[derive(Debug, Clone)]
29pub enum LoadedTokenizer {
30    /// A Byte-Pair Encoding tokenizer.
31    Bpe(BpeTokenizer),
32    /// A WordPiece tokenizer.
33    WordPiece(WordPieceTokenizer),
34}
35
36impl LoadedTokenizer {
37    /// Borrow the loaded tokenizer as a [`Tokenizer`](crate::Tokenizer) trait
38    /// object.
39    #[must_use]
40    pub fn as_tokenizer(&self) -> &dyn crate::Tokenizer {
41        match self {
42            LoadedTokenizer::Bpe(t) => t,
43            LoadedTokenizer::WordPiece(t) => t,
44        }
45    }
46}
47
48/// Parse a `tokenizer.json` document from a string.
49///
50/// # Errors
51/// Returns [`TokenizerError::MalformedFile`] if the JSON is invalid, the model
52/// type is unsupported, or a required field is missing.
53pub fn from_tokenizer_json_str(text: &str) -> Result<LoadedTokenizer, TokenizerError> {
54    let root = json::parse(text).map_err(TokenizerError::MalformedFile)?;
55    let model = root
56        .get("model")
57        .ok_or_else(|| malformed("missing \"model\" object"))?;
58    let model_type = model
59        .get("type")
60        .and_then(JsonValue::as_str)
61        .ok_or_else(|| malformed("missing \"model.type\""))?;
62
63    match model_type {
64        "BPE" => load_bpe(&root, model).map(LoadedTokenizer::Bpe),
65        "WordPiece" => load_wordpiece(&root, model).map(LoadedTokenizer::WordPiece),
66        other => Err(malformed(&alloc::format!(
67            "unsupported model.type {other:?} (only BPE and WordPiece are supported)"
68        ))),
69    }
70}
71
72/// Load a `tokenizer.json` from disk.
73///
74/// # Errors
75/// Returns [`TokenizerError::Io`] on a read failure, or the same errors as
76/// [`from_tokenizer_json_str`] on a parse failure.
77#[cfg(feature = "std")]
78pub fn from_tokenizer_json_file(path: &str) -> Result<LoadedTokenizer, TokenizerError> {
79    let text = std::fs::read_to_string(path)?;
80    from_tokenizer_json_str(&text)
81}
82
83fn malformed(msg: &str) -> TokenizerError {
84    TokenizerError::MalformedFile(msg.to_string())
85}
86
87/// Extract a `token -> id` map from a JSON object.
88fn parse_vocab(value: &JsonValue) -> Result<BTreeMap<String, TokenId>, TokenizerError> {
89    let obj = value
90        .as_object()
91        .ok_or_else(|| malformed("\"model.vocab\" must be an object"))?;
92    let mut vocab = BTreeMap::new();
93    for (token, id) in obj {
94        let id = id
95            .as_u32()
96            .ok_or_else(|| malformed("vocab id is not a non-negative integer"))?;
97        vocab.insert(token.clone(), id);
98    }
99    Ok(vocab)
100}
101
102/// Collect `(content, id, is_special)` for every entry in a top-level
103/// `added_tokens` array.
104fn parse_added_tokens(root: &JsonValue) -> Vec<(String, TokenId, bool)> {
105    let Some(added) = root.get("added_tokens").and_then(JsonValue::as_array) else {
106        return Vec::new();
107    };
108    let mut out = Vec::new();
109    for entry in added {
110        let (Some(content), Some(id)) = (
111            entry.get("content").and_then(JsonValue::as_str),
112            entry.get("id").and_then(JsonValue::as_u32),
113        ) else {
114            continue;
115        };
116        let special = entry
117            .get("special")
118            .is_some_and(|v| matches!(v, JsonValue::Bool(true)));
119        out.push((content.to_string(), id, special));
120    }
121    out
122}
123
124/// Detect whether a pre-tokenizer (possibly a `Sequence`) uses GPT-2 byte-level
125/// splitting.
126fn detect_byte_level(root: &JsonValue) -> bool {
127    fn contains_byte_level(v: &JsonValue) -> bool {
128        if v.get("type").and_then(JsonValue::as_str) == Some("ByteLevel") {
129            return true;
130        }
131        if let Some(list) = v.get("pretokenizers").and_then(JsonValue::as_array) {
132            return list.iter().any(contains_byte_level);
133        }
134        false
135    }
136    root.get("pre_tokenizer").is_some_and(contains_byte_level)
137}
138
139/// Detect a lowercasing normalizer (BERT `lowercase: true`, or a `Lowercase`
140/// normalizer, possibly inside a `Sequence`).
141fn detect_lowercase(root: &JsonValue) -> bool {
142    fn is_lower(v: &JsonValue) -> bool {
143        match v.get("type").and_then(JsonValue::as_str) {
144            Some("Lowercase") => return true,
145            Some("BertNormalizer") => {
146                if matches!(v.get("lowercase"), Some(JsonValue::Bool(true))) {
147                    return true;
148                }
149            }
150            _ => {}
151        }
152        if let Some(list) = v.get("normalizers").and_then(JsonValue::as_array) {
153            return list.iter().any(is_lower);
154        }
155        false
156    }
157    root.get("normalizer").is_some_and(is_lower)
158}
159
160fn load_bpe(root: &JsonValue, model: &JsonValue) -> Result<BpeTokenizer, TokenizerError> {
161    let mut vocab = parse_vocab(
162        model
163            .get("vocab")
164            .ok_or_else(|| malformed("missing \"model.vocab\""))?,
165    )?;
166
167    let merges_val = model
168        .get("merges")
169        .and_then(JsonValue::as_array)
170        .ok_or_else(|| malformed("missing \"model.merges\" array"))?;
171    let mut merges = Vec::with_capacity(merges_val.len());
172    for entry in merges_val {
173        let pair = match entry {
174            // Newer format: ["a", "b"].
175            JsonValue::Array(parts) if parts.len() == 2 => {
176                let a = parts[0]
177                    .as_str()
178                    .ok_or_else(|| malformed("merge entry element is not a string"))?;
179                let b = parts[1]
180                    .as_str()
181                    .ok_or_else(|| malformed("merge entry element is not a string"))?;
182                (a.to_string(), b.to_string())
183            }
184            // Legacy format: "a b".
185            JsonValue::String(s) => {
186                let mut it = s.splitn(2, ' ');
187                match (it.next(), it.next()) {
188                    (Some(a), Some(b)) => (a.to_string(), b.to_string()),
189                    _ => return Err(malformed("merge string is not a space-separated pair")),
190                }
191            }
192            _ => return Err(malformed("unrecognised merge entry")),
193        };
194        merges.push(pair);
195    }
196
197    // Fold added tokens into the vocab and collect the special ones.
198    let mut specials = BTreeMap::new();
199    for (content, id, special) in parse_added_tokens(root) {
200        vocab.entry(content.clone()).or_insert(id);
201        if special {
202            specials.insert(content, id);
203        }
204    }
205
206    let mut tok = BpeTokenizer::from_vocab_merges(vocab, merges);
207    if detect_byte_level(root) {
208        tok = tok.with_byte_level();
209    }
210    if !specials.is_empty() {
211        tok = tok.with_special_tokens(specials);
212    }
213    Ok(tok)
214}
215
216fn load_wordpiece(
217    root: &JsonValue,
218    model: &JsonValue,
219) -> Result<WordPieceTokenizer, TokenizerError> {
220    let mut vocab = parse_vocab(
221        model
222            .get("vocab")
223            .ok_or_else(|| malformed("missing \"model.vocab\""))?,
224    )?;
225
226    // Fold in any added tokens so their ids are decodable.
227    for (content, id, _special) in parse_added_tokens(root) {
228        vocab.entry(content).or_insert(id);
229    }
230
231    let unk = model
232        .get("unk_token")
233        .and_then(JsonValue::as_str)
234        .unwrap_or("[UNK]")
235        .to_string();
236
237    let mut tok = WordPieceTokenizer::from_vocab(vocab, &unk)?;
238    if detect_lowercase(root) {
239        tok = tok.with_lowercase();
240    }
241    Ok(tok)
242}