tpt_tokenizer_core/
loader.rs1use 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#[derive(Debug, Clone)]
29pub enum LoadedTokenizer {
30 Bpe(BpeTokenizer),
32 WordPiece(WordPieceTokenizer),
34}
35
36impl LoadedTokenizer {
37 #[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
48pub 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#[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
87fn 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
102fn 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
124fn 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
139fn 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 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 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 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 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}