1use std::{collections::HashSet, path::Path};
5
6use tokenizers::tokenizer::{AddedToken, PostProcessor as _, Tokenizer as HfTokenizer};
7
8use super::{
9 Encoding, Error, Result, TokenIdType, TokenizerOptions,
10 traits::{DecodeResult, Decoder, Encoder, Tokenizer},
11};
12
13pub struct HuggingFaceTokenizer {
14 tokenizer: HfTokenizer,
15 byte_fallback_ids: HashSet<u32>,
16 special_ids: HashSet<u32>,
17 options: TokenizerOptions,
21}
22
23impl HuggingFaceTokenizer {
24 pub fn from_file(model_name: &str) -> Result<Self> {
29 let mut tokenizer = HfTokenizer::from_file(model_name)
30 .map_err(|err| Error::msg(format!("Error loading tokenizer: {}", err)))?;
31
32 if let Some(parent) = Path::new(model_name).parent() {
33 merge_special_tokens_from_config(&mut tokenizer, parent);
34 }
35
36 Ok(Self::from_tokenizer(tokenizer))
37 }
38
39 pub fn from_tokenizer(tokenizer: HfTokenizer) -> Self {
40 let has_byte_fallback = tokenizer.get_decoder().is_some_and(|decoder| {
44 fn contains(value: &serde_json::Value) -> bool {
45 value["type"] == "ByteFallback"
46 || value["decoders"]
47 .as_array()
48 .is_some_and(|items| items.iter().any(contains))
49 }
50 serde_json::to_value(decoder).is_ok_and(|value| contains(&value))
51 });
52 let byte_fallback_ids = if has_byte_fallback {
53 tokenizer
54 .get_vocab(true)
55 .into_iter()
56 .filter_map(|(token, id)| {
57 (token.len() == 6
58 && token.starts_with("<0x")
59 && token.ends_with('>')
60 && u8::from_str_radix(&token[3..5], 16).is_ok())
61 .then_some(id)
62 })
63 .collect()
64 } else {
65 HashSet::new()
66 };
67 let special_ids = tokenizer
68 .get_added_tokens_decoder()
69 .into_iter()
70 .filter_map(|(id, token)| token.special.then_some(id))
71 .collect();
72 HuggingFaceTokenizer {
73 tokenizer,
74 byte_fallback_ids,
75 special_ids,
76 options: TokenizerOptions::default(),
77 }
78 }
79
80 pub fn from_tokenizer_with_model_dir(tokenizer: HfTokenizer, model_dir: &Path) -> Self {
83 let mut tokenizer = tokenizer;
84 merge_special_tokens_from_config(&mut tokenizer, model_dir);
85 Self::from_tokenizer(tokenizer)
86 }
87}
88
89pub fn merge_special_tokens_from_config(tokenizer: &mut HfTokenizer, model_dir: &Path) {
98 let cfg_path = model_dir.join("tokenizer_config.json");
99 let Ok(raw) = std::fs::read_to_string(&cfg_path) else {
100 return;
101 };
102 let cfg: serde_json::Value = match serde_json::from_str(&raw) {
103 Ok(v) => v,
104 Err(e) => {
105 tracing::debug!(
106 target: "tokenizer",
107 path = %cfg_path.display(),
108 error = %e,
109 "tokenizer_config.json parse failed; skipping special-token merge"
110 );
111 return;
112 }
113 };
114 let Some(decoder) = cfg.get("added_tokens_decoder").and_then(|v| v.as_object()) else {
115 return;
116 };
117
118 let mut to_add: Vec<AddedToken> = Vec::new();
119 for (_id, spec) in decoder {
120 let obj = match spec.as_object() {
121 Some(o) => o,
122 None => continue,
123 };
124 if obj.get("special").and_then(|v| v.as_bool()) != Some(true) {
127 continue;
128 }
129 let Some(content) = obj.get("content").and_then(|v| v.as_str()) else {
130 continue;
131 };
132 if content.is_empty() {
133 continue;
134 }
135 let single_word = obj
136 .get("single_word")
137 .and_then(|v| v.as_bool())
138 .unwrap_or(false);
139 let lstrip = obj.get("lstrip").and_then(|v| v.as_bool()).unwrap_or(false);
140 let rstrip = obj.get("rstrip").and_then(|v| v.as_bool()).unwrap_or(false);
141 let normalized = obj
142 .get("normalized")
143 .and_then(|v| v.as_bool())
144 .unwrap_or(false);
145 let token = AddedToken::from(content.to_string(), true)
146 .single_word(single_word)
147 .lstrip(lstrip)
148 .rstrip(rstrip)
149 .normalized(normalized);
150 to_add.push(token);
151 }
152
153 if to_add.is_empty() {
154 return;
155 }
156 let added = tokenizer.add_special_tokens(&to_add);
159 if added > 0 {
160 let promoted: Vec<&str> = to_add.iter().map(|t| t.content.as_str()).collect();
166 tracing::warn!(
167 target: "tokenizer",
168 path = %cfg_path.display(),
169 added,
170 candidates = to_add.len(),
171 promoted = ?promoted,
172 "merged additional special tokens from tokenizer_config.json"
173 );
174 }
175}
176
177impl Encoder for HuggingFaceTokenizer {
178 fn encode(&self, input: &str) -> Result<Encoding> {
179 let encoding = self
181 .tokenizer
182 .encode(input, self.options.add_special_tokens)
183 .map_err(|err| Error::msg(format!("Error tokenizing input: {err}")))?;
184
185 Ok(Encoding::Hf(Box::new(encoding)))
186 }
187
188 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
189 let hf_encodings = self
190 .tokenizer
191 .encode_batch(inputs.to_vec(), self.options.add_special_tokens)
192 .map_err(|err| Error::msg(format!("Error batch tokenizing input: {err}")))?;
193
194 let encodings = hf_encodings
195 .into_iter()
196 .map(|enc| Encoding::Hf(Box::new(enc)))
197 .collect();
198
199 Ok(encodings)
200 }
201}
202
203impl Decoder for HuggingFaceTokenizer {
204 fn has_unstable_suffix(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> bool {
205 token_ids
206 .iter()
207 .rev()
208 .find(|id| !skip_special_tokens || !self.special_ids.contains(id))
209 .is_some_and(|id| self.byte_fallback_ids.contains(id))
210 }
211
212 fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
213 let text = self
215 .tokenizer
216 .decode(token_ids, skip_special_tokens)
217 .map_err(|err| Error::msg(format!("Error de-tokenizing input: {err}")))?;
218
219 Ok(text.into())
220 }
221}
222
223impl Tokenizer for HuggingFaceTokenizer {
224 fn validate_prefix_cache(&self) -> Result<()> {
225 if self.options.add_special_tokens {
226 return Err(Error::msg(
227 "HuggingFace tokenizers configured with add_special_tokens=true must remain uncached",
228 ));
229 }
230 Ok(())
231 }
232
233 fn with_options(mut self, options: TokenizerOptions) -> Self {
236 self.options = options;
237 self
238 }
239
240 fn vocab_size(&self) -> Option<usize> {
241 Some(self.tokenizer.get_vocab_size(true))
242 }
243
244 fn token_to_id(&self, token: &str) -> Result<Option<TokenIdType>> {
245 Ok(self.tokenizer.token_to_id(token))
246 }
247
248 fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
249 let mut ids: Vec<TokenIdType> = self
250 .tokenizer
251 .get_added_tokens_decoder()
252 .into_iter()
253 .filter_map(|(id, token)| token.special.then_some(id))
254 .collect();
255 ids.sort_unstable();
256 Ok(ids)
257 }
258
259 fn num_special_tokens_added(&self) -> Result<usize> {
260 Ok(self
261 .tokenizer
262 .get_post_processor()
263 .map_or(0, |processor| processor.added_tokens(false)))
264 }
265}
266
267impl From<HfTokenizer> for HuggingFaceTokenizer {
268 fn from(tokenizer: HfTokenizer) -> Self {
269 Self::from_tokenizer(tokenizer)
270 }
271}
272
273#[cfg(test)]
274mod byte_fallback_stream_tests {
275 use super::*;
276
277 fn tokenizer() -> super::super::Tokenizer {
278 let hf: HfTokenizer = serde_json::from_value(serde_json::json!({
279 "version": "1.0", "truncation": null, "padding": null,
280 "added_tokens": [{"id": 6, "content": "<eos>", "special": true,
281 "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}],
282 "normalizer": null, "pre_tokenizer": null, "post_processor": null,
283 "decoder": {"type": "Sequence", "decoders": [
284 {"type": "ByteFallback"}, {"type": "Fuse"}]},
285 "model": {"type": "BPE", "vocab": {
286 "<0x61>": 0, "<0xF5>": 1, "<0xC3>": 2, "<0xA9>": 3,
287 " hello": 4, "!": 5, "<eos>": 6}, "merges": [], "byte_fallback": true}
288 }))
289 .unwrap();
290 std::sync::Arc::new(HuggingFaceTokenizer::from_tokenizer(hf)).into()
291 }
292
293 #[test]
294 fn byte_runs_match_full_decode() {
295 let tokenizer = tokenizer();
296 for (ids, skip) in [
300 (vec![0, 1, 4], false),
301 (vec![2, 3, 1, 4], false),
302 (vec![2, 3, 5], false),
303 (vec![0], false),
304 (vec![2, 3], false),
305 (vec![2], false),
306 (vec![0, 1], false),
307 (vec![0, 6, 1, 4], false),
308 (vec![0, 6, 1, 4], true),
309 ] {
310 let expected: String = tokenizer.decode(&ids, skip).unwrap().into();
311 let mut stream = tokenizer.decode_stream(&[], skip);
312 let mut actual = String::new();
313 for id in &ids {
314 actual.push_str(&stream.step(*id).unwrap().unwrap_or_default());
315 }
316 actual.push_str(&stream.finish().unwrap().unwrap_or_default());
317 assert_eq!(actual, expected, "ids={ids:?}, skip={skip}");
318 assert_eq!(stream.finish().unwrap(), None);
319 }
320 }
321
322 #[test]
323 fn only_byte_suffix_is_delayed_and_prompt_is_not_emitted() {
324 let tokenizer = tokenizer();
325 let mut stream = tokenizer.decode_stream(&[4], true);
326 assert_eq!(stream.step(5).unwrap().as_deref(), Some("!"));
327 assert_eq!(stream.step(0).unwrap(), None);
328 assert_eq!(stream.step(1).unwrap(), None);
329 assert_eq!(stream.step(4).unwrap().as_deref(), Some("�� hello"));
330 assert_eq!(stream.finish().unwrap(), None);
331 }
332}
333
334#[cfg(test)]
335mod tests {
336 use super::*;
345 use std::fs;
346 use tempfile::TempDir;
347
348 #[test]
349 fn merge_gate_round_trips_through_decode() {
350 const TOKENIZER_JSON: &str = r#"{
356 "version": "1.0",
357 "truncation": null,
358 "padding": null,
359 "added_tokens": [
360 {"id": 0, "content": "<unk>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}
361 ],
362 "normalizer": null,
363 "pre_tokenizer": null,
364 "post_processor": null,
365 "decoder": null,
366 "model": {
367 "type": "WordLevel",
368 "vocab": {"<unk>": 0, "hello": 1, "world": 2, "<|special_kept|>": 3, "<|special_dropped|>": 4},
369 "unk_token": "<unk>"
370 }
371 }"#;
372
373 const TOKENIZER_CONFIG_JSON: &str = r#"{
382 "added_tokens_decoder": {
383 "3": {"content": "<|special_kept|>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
384 "4": {"content": "<|special_dropped|>", "special": false, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}
385 }
386 }"#;
387
388 let dir = TempDir::new().unwrap();
389 fs::write(dir.path().join("tokenizer.json"), TOKENIZER_JSON).unwrap();
390 fs::write(
391 dir.path().join("tokenizer_config.json"),
392 TOKENIZER_CONFIG_JSON,
393 )
394 .unwrap();
395
396 let mut tokenizer = HfTokenizer::from_file(dir.path().join("tokenizer.json")).unwrap();
397 merge_special_tokens_from_config(&mut tokenizer, dir.path());
398
399 let specials: Vec<String> = {
401 let mut v: Vec<String> = tokenizer
402 .get_added_tokens_decoder()
403 .values()
404 .filter(|t| t.special)
405 .map(|t| t.content.clone())
406 .collect();
407 v.sort();
408 v
409 };
410 assert_eq!(
411 specials,
412 vec!["<unk>".to_string(), "<|special_kept|>".to_string()],
413 "<|special_kept|> promoted; <|special_dropped|> stayed non-special"
414 );
415
416 let enc_kept = tokenizer.encode("<|special_kept|>", false).unwrap();
420 let decoded_strip = tokenizer.decode(enc_kept.get_ids(), true).unwrap();
421 assert!(
422 !decoded_strip.contains("<|special_kept|>"),
423 "promoted special:true token must be stripped under skip_special_tokens=true; got {decoded_strip:?}"
424 );
425
426 let enc_drop = tokenizer.encode("<|special_dropped|>", false).unwrap();
427 let decoded_keep = tokenizer.decode(enc_drop.get_ids(), true).unwrap();
428 assert!(
429 decoded_keep.contains("<|special_dropped|>"),
430 "non-promoted special:false token must survive skip_special_tokens=true; got {decoded_keep:?}"
431 );
432 }
433
434 #[test]
435 fn add_special_tokens_flag_controls_encode() {
436 const TOKENIZER_JSON: &str = r#"{
441 "version": "1.0",
442 "truncation": null,
443 "padding": null,
444 "added_tokens": [
445 {"id": 0, "content": "<unk>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
446 {"id": 3, "content": "<bos>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}
447 ],
448 "normalizer": null,
449 "pre_tokenizer": null,
450 "post_processor": {
451 "type": "TemplateProcessing",
452 "single": [
453 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
454 {"Sequence": {"id": "A", "type_id": 0}}
455 ],
456 "pair": [
457 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
458 {"Sequence": {"id": "A", "type_id": 0}},
459 {"Sequence": {"id": "B", "type_id": 0}}
460 ],
461 "special_tokens": {
462 "<bos>": {"id": "<bos>", "ids": [3], "tokens": ["<bos>"]}
463 }
464 },
465 "decoder": null,
466 "model": {
467 "type": "WordLevel",
468 "vocab": {"<unk>": 0, "hello": 1, "world": 2, "<bos>": 3},
469 "unk_token": "<unk>"
470 }
471 }"#;
472
473 let dir = TempDir::new().unwrap();
474 fs::write(dir.path().join("tokenizer.json"), TOKENIZER_JSON).unwrap();
475 let path = dir.path().join("tokenizer.json");
476 let path = path.to_str().unwrap();
477
478 let ids = |enc: &Encoding| match enc {
479 Encoding::Hf(e) => e.get_ids().to_vec(),
480 _ => panic!("expected Hf encoding"),
481 };
482
483 let plain = HuggingFaceTokenizer::from_file(path).unwrap();
485 assert_eq!(ids(&plain.encode("hello").unwrap()), vec![1]);
486
487 let with_bos =
488 HuggingFaceTokenizer::from_file(path)
489 .unwrap()
490 .with_options(TokenizerOptions {
491 add_special_tokens: true,
492 });
493 assert_eq!(ids(&with_bos.encode("hello").unwrap()), vec![3, 1]);
494 let batch = with_bos.encode_batch(&["hello", "world"]).unwrap();
495 assert_eq!(ids(&batch[0]), vec![3, 1]);
496 assert_eq!(ids(&batch[1]), vec![3, 2]);
497
498 use crate::Tokenizer as TokenizerWrapper;
502 let wrapper_plain = TokenizerWrapper::from_file(path).unwrap();
503 assert_eq!(ids(&wrapper_plain.encode("hello").unwrap()), vec![1]);
504
505 let wrapper_bos = TokenizerWrapper::from_file_with_options(
506 path,
507 TokenizerOptions {
508 add_special_tokens: true,
509 },
510 )
511 .unwrap();
512 assert_eq!(ids(&wrapper_bos.encode("hello").unwrap()), vec![3, 1]);
513 }
514
515 #[test]
516 fn vocab_introspection_accessors() {
517 const TOKENIZER_JSON: &str = r#"{
518 "version": "1.0",
519 "truncation": null,
520 "padding": null,
521 "added_tokens": [
522 {"id": 0, "content": "<unk>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}
523 ],
524 "normalizer": null,
525 "pre_tokenizer": null,
526 "post_processor": null,
527 "decoder": null,
528 "model": {
529 "type": "WordLevel",
530 "vocab": {"<unk>": 0, "hello": 1, "world": 2},
531 "unk_token": "<unk>"
532 }
533 }"#;
534
535 let dir = TempDir::new().unwrap();
536 fs::write(dir.path().join("tokenizer.json"), TOKENIZER_JSON).unwrap();
537 let path = dir.path().join("tokenizer.json");
538
539 let tokenizer = HuggingFaceTokenizer::from_file(path.to_str().unwrap()).unwrap();
540 assert_eq!(tokenizer.vocab_size(), Some(3));
541 assert_eq!(tokenizer.token_to_id("hello").unwrap(), Some(1));
542 assert_eq!(tokenizer.special_token_ids().unwrap(), vec![0]);
543 assert_eq!(tokenizer.num_special_tokens_added().unwrap(), 0);
544 }
545
546 #[test]
547 fn num_special_tokens_added_reflects_post_processor_additions() {
548 const TOKENIZER_JSON: &str = r#"{
549 "version": "1.0",
550 "truncation": null,
551 "padding": null,
552 "added_tokens": [
553 {"id": 0, "content": "<unk>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
554 {"id": 3, "content": "<bos>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}
555 ],
556 "normalizer": null,
557 "pre_tokenizer": null,
558 "post_processor": {
559 "type": "TemplateProcessing",
560 "single": [
561 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
562 {"Sequence": {"id": "A", "type_id": 0}}
563 ],
564 "pair": [
565 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
566 {"Sequence": {"id": "A", "type_id": 0}},
567 {"Sequence": {"id": "B", "type_id": 0}}
568 ],
569 "special_tokens": {
570 "<bos>": {"id": "<bos>", "ids": [3], "tokens": ["<bos>"]}
571 }
572 },
573 "decoder": null,
574 "model": {
575 "type": "WordLevel",
576 "vocab": {"<unk>": 0, "hello": 1, "world": 2, "<bos>": 3},
577 "unk_token": "<unk>"
578 }
579 }"#;
580
581 let dir = TempDir::new().unwrap();
582 fs::write(dir.path().join("tokenizer.json"), TOKENIZER_JSON).unwrap();
583 let path = dir.path().join("tokenizer.json");
584
585 let tokenizer = HuggingFaceTokenizer::from_file(path.to_str().unwrap()).unwrap();
586 assert_eq!(tokenizer.num_special_tokens_added().unwrap(), 1);
587 }
588}