1use std::sync::atomic::Ordering;
4
5use crate::utils::gguf_metadata::ContentMetadata;
6use crate::DEBUG;
7use ahash::AHashMap;
8use anyhow::Result;
9use hanzo_ml::quantized::gguf_file::Value;
10use itertools::Itertools;
11use tokenizers::pre_tokenizers::{
12 sequence::Sequence,
13 split::{Split, SplitPattern},
14 PreTokenizerWrapper,
15};
16use tokenizers::tokenizer::normalizer::SplitDelimiterBehavior;
17use tokenizers::{
18 decoders::{
19 self, byte_fallback::ByteFallback, byte_level::ByteLevel, fuse::Fuse, strip::Strip,
20 },
21 models::{bpe::BpeBuilder, unigram::Unigram},
22 normalizers::{self, Prepend, Replace},
23 processors, AddedToken, DecoderWrapper, ModelWrapper, NormalizerWrapper, Tokenizer,
24};
25use tracing::info;
26
27use super::Content;
28
29pub(crate) struct GgufTokenizerConversion {
30 pub tokenizer: Tokenizer,
31 pub bos: Option<String>,
32 pub eos: Option<String>,
33 pub unk: Option<String>,
34}
35
36struct PropsGGUF {
37 model: String,
38 tokens: Vec<String>,
39 added_tokens: Option<Vec<String>>,
40 scores: Option<Vec<f32>>,
41 merges: Option<Vec<String>>,
42 unk: Option<u32>,
43 bos: Option<u32>,
44 eos: u32,
45}
46
47impl TryFrom<ContentMetadata<'_>> for PropsGGUF {
48 type Error = anyhow::Error;
49
50 fn try_from(c: ContentMetadata) -> Result<Self, Self::Error> {
51 let required = ["model", "tokens", "eos_token_id"];
52 c.has_required_keys(&required)?;
53
54 let props = Self {
55 model: c.get_value("model")?,
56 tokens: c.get_value("tokens")?,
57 added_tokens: c.get_value("added_tokens").ok(),
58 scores: c.get_value("scores").ok(),
59 merges: c.get_value("merges").ok(),
60 unk: c.get_value("unknown_token_id").ok(),
61 eos: c.get_value("eos_token_id")?,
62 bos: c.get_value("bos_token_id").ok(),
63 };
64
65 Ok(props)
66 }
67}
68
69pub fn convert_gguf_to_hf_tokenizer<R: std::io::Seek + std::io::Read>(
70 content: &Content<'_, R>,
71) -> Result<GgufTokenizerConversion> {
72 let metadata = ContentMetadata {
73 path_prefix: "tokenizer.ggml",
74 metadata: content.get_metadata(),
75 };
76
77 let md_get = |s: &str| match metadata.metadata.get(s) {
78 None => hanzo_ml::bail!("cannot find {s} in metadata"),
79 Some(v) => Ok(v),
80 };
81
82 let mut token_types = Vec::<i32>::new();
83 if metadata.metadata.contains_key("tokenizer.ggml.token_type") {
84 let vtypes: &Vec<Value> = md_get("tokenizer.ggml.token_type")
85 .unwrap()
86 .to_vec()
87 .unwrap();
88 let v: Vec<i32> = vtypes.iter().map(|v| v.to_i32().unwrap()).collect();
89 token_types.extend(v);
90 }
91
92 let props = PropsGGUF::try_from(metadata)?;
93
94 let (mut tokenizer, kind) = match props.model.as_str() {
95 "llama" | "replit" => unigram_tokenizer(&props)?,
96 "gpt2" => bpe_tokenizer(&props)?,
97 other => {
98 anyhow::bail!("Tokenizer model `{other}` not supported.");
99 }
100 };
101
102 let mut num_special_tokens = 0;
104 #[allow(clippy::needless_range_loop)]
105 if token_types.len() == props.tokens.len() {
106 for i in 0..props.tokens.len() {
107 if token_types[i] != 1i32 {
108 let tk = props.tokens[i].clone();
109 tokenizer.add_special_tokens(&[AddedToken::from(tk.to_string(), true)]);
110 num_special_tokens += 1;
111 }
112 }
113 }
114
115 info!(
116 "GGUF tokenizer model is `{model}`, kind: `{kind:?}`, num tokens: {}, num special tokens {}, num added tokens: {}, num merges: {}, num scores: {}",
117 tokenizer.get_vocab_size(true),
118 num_special_tokens,
119 props.added_tokens.as_ref().map(|x| x.len()).unwrap_or(0),
120 props.merges.as_ref().map(|x| x.len()).unwrap_or(0),
121 props.scores.as_ref().map(|x| x.len()).unwrap_or(0),
122 model = props.model,
123 );
124 if DEBUG.load(Ordering::Relaxed) {
125 info!("Tokenizer: {tokenizer:?}");
126 }
127
128 let unk = match props.unk {
129 Some(u) => Some(props.tokens[u as usize].clone()),
130 _ => None,
131 };
132
133 let bos = match props.bos {
134 Some(b) => Some(props.tokens[b as usize].clone()),
135 None => None,
136 };
137
138 Ok(GgufTokenizerConversion {
139 tokenizer,
140 bos,
141 eos: Some(props.tokens[props.eos as usize].clone()),
142 unk,
143 })
144}
145
146#[derive(Debug)]
149enum TokenizerKind {
150 Unigram,
151 Bpe,
152}
153
154fn unigram_tokenizer(p: &PropsGGUF) -> Result<(Tokenizer, TokenizerKind)> {
155 let PropsGGUF { unk, eos, bos, .. } = *p;
156 let unk = unk.unwrap_or(0);
158
159 let model = {
161 let vocab: Vec<(String, f64)> = {
162 let Some(s) = p.scores.as_ref() else {
163 anyhow::bail!(
164 "`llama` unigram tokenizer is missing required metadata `tokenizer.ggml.scores`"
165 );
166 };
167 let scores = s.iter().cloned().map(|f_32| f_32 as f64);
168
169 p.tokens.iter().cloned().zip(scores).collect()
170 };
171
172 Unigram::from(vocab, Some(unk as usize), true).map_err(anyhow::Error::msg)?
173 };
174
175 let decoder = Decoder::Sequence(vec![
178 Decoder::Replace("▁", " "),
179 Decoder::ByteFallback,
180 Decoder::Fuse,
181 Decoder::Strip(' ', 1, 0),
182 ]);
183
184 let normalizer = Normalizer::Sequence(vec![
185 Normalizer::Prepend("▁"),
186 Normalizer::Replace(" ", "▁"),
187 ]);
188
189 let mut tokenizer: Tokenizer = TokenizerX::new(
190 ModelWrapper::Unigram(model),
191 Some(decoder),
192 Some(normalizer),
193 )?;
194
195 for v in [bos, Some(eos), Some(unk)].iter().flatten() {
197 let tk = p.tokens[*v as usize].clone();
198 tokenizer.add_special_tokens(&[AddedToken::from(tk.to_string(), true)]);
199 }
200 Ok((tokenizer, TokenizerKind::Unigram))
201}
202
203fn bpe_tokenizer(p: &PropsGGUF) -> Result<(Tokenizer, TokenizerKind)> {
204 let merges = p
207 .merges
208 .as_ref()
209 .ok_or(anyhow::Error::msg("BPE tokenizer must include merges"))?
210 .iter()
211 .map(|merge| {
212 let split: (&str, &str) = merge
213 .splitn(2, ' ')
214 .collect_tuple()
215 .expect("Failed to convert split into 2-tuple");
216 (split.0.to_string(), split.1.to_string())
217 })
218 .collect::<Vec<_>>();
219
220 let mut vocab = AHashMap::new();
221 for (i, token) in p.tokens.iter().enumerate() {
222 #[allow(clippy::cast_possible_truncation)]
223 vocab.insert(token.clone(), i as u32);
224 }
225
226 let PropsGGUF { bos, eos, unk, .. } = *p;
227
228 let mut bpe = BpeBuilder::new().vocab_and_merges(vocab, merges);
229 if let Some(unk) = unk {
230 bpe = bpe.unk_token(p.tokens[unk as usize].to_string());
231 };
232
233 let bpe = bpe.build().map_err(anyhow::Error::msg)?;
234
235 let mut tokenizer = TokenizerX::new(
236 ModelWrapper::BPE(bpe),
237 Some(Decoder::ByteLevel(true, true, true)),
238 None,
239 )?;
240
241 let split = Split::new(
242 SplitPattern::Regex("(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+".to_string()),
243 SplitDelimiterBehavior::Isolated,
244 false,
245 ).unwrap();
246
247 let pre_tokenizer = Sequence::new(vec![
253 PreTokenizerWrapper::Split(split),
254 PreTokenizerWrapper::ByteLevel(ByteLevel::new(false, false, false)),
255 ]);
256
257 tokenizer.with_pre_tokenizer(Some(pre_tokenizer));
258
259 tokenizer.with_decoder(Some(decoders::byte_level::ByteLevel::new(
260 false, false, false,
261 )));
262 tokenizer.with_post_processor(Some(processors::byte_level::ByteLevel::new(
263 false, false, false,
264 )));
265
266 for v in [bos, Some(eos), unk].iter().flatten() {
267 let tk = p.tokens[*v as usize].clone();
268 tokenizer.add_special_tokens(&[AddedToken::from(tk.to_string(), true)]);
269 }
270
271 Ok((tokenizer, TokenizerKind::Bpe))
272}
273
274struct TokenizerX;
278
279impl TokenizerX {
280 #[allow(clippy::new_ret_no_self)]
281 fn new<'a>(
282 model: ModelWrapper,
283 decoder: Option<Decoder<'a>>,
284 normalizer: Option<Normalizer<'a>>,
285 ) -> Result<Tokenizer> {
286 let mut tokenizer = Tokenizer::new(model);
287
288 if let Some(decoder) = decoder {
290 let d = DecoderWrapper::try_from(decoder)?;
291 tokenizer.with_decoder(Some(d));
292 }
293 if let Some(normalizer) = normalizer {
294 let n: NormalizerWrapper = NormalizerWrapper::try_from(normalizer)?;
295 tokenizer.with_normalizer(Some(n));
296 }
297
298 Ok(tokenizer)
299 }
300}
301
302enum Decoder<'a> {
305 ByteFallback,
306 Fuse,
307 Replace(&'a str, &'a str),
308 Strip(char, usize, usize),
309 Sequence(Vec<Self>),
310 ByteLevel(bool, bool, bool),
311}
312
313impl TryFrom<Decoder<'_>> for DecoderWrapper {
315 type Error = anyhow::Error;
316
317 fn try_from(variant: Decoder) -> Result<Self, Self::Error> {
318 let value: DecoderWrapper = match variant {
319 Decoder::ByteFallback => ByteFallback::default().into(),
320 Decoder::Fuse => Fuse::default().into(),
321 Decoder::Replace(pattern, content) => Replace::new(pattern, content)
322 .map_err(anyhow::Error::msg)?
323 .into(),
324 Decoder::Strip(content, start, stop) => Strip::new(content, start, stop).into(),
325 Decoder::Sequence(decoders) => {
326 let seq = decoders
327 .into_iter()
328 .map(DecoderWrapper::try_from)
329 .collect::<Result<Vec<DecoderWrapper>>>()?;
330
331 decoders::sequence::Sequence::new(seq).into()
332 }
333 Decoder::ByteLevel(add_prefix_space, trim_offsets, use_regex) => {
334 ByteLevel::new(add_prefix_space, trim_offsets, use_regex).into()
335 }
336 };
337
338 Ok(value)
339 }
340}
341
342enum Normalizer<'a> {
345 Prepend(&'a str),
346 Replace(&'a str, &'a str),
347 Sequence(Vec<Self>),
348}
349
350impl TryFrom<Normalizer<'_>> for NormalizerWrapper {
351 type Error = anyhow::Error;
352
353 fn try_from(variant: Normalizer) -> Result<Self, Self::Error> {
354 let value: NormalizerWrapper = match variant {
355 Normalizer::Prepend(prepend) => Prepend::new(prepend.to_owned()).into(),
356 Normalizer::Replace(pattern, content) => Replace::new(pattern, content)
357 .map_err(anyhow::Error::msg)?
358 .into(),
359 Normalizer::Sequence(decoders) => {
360 let seq = decoders
361 .into_iter()
362 .map(NormalizerWrapper::try_from)
363 .collect::<Result<Vec<NormalizerWrapper>>>()?;
364
365 normalizers::Sequence::new(seq).into()
366 }
367 };
368
369 Ok(value)
370 }
371}
372
373#[cfg(test)]
374mod tests {
375 use anyhow::Result;
376 use hf_hub::{api::sync::ApiBuilder, Repo, RepoType};
377 use tokenizers::Tokenizer;
378
379 #[allow(dead_code)]
380 #[derive(Debug)]
381 enum TokenizerType {
382 Llama,
384 Replit,
385 Gpt2,
386 Rwkv,
387 }
388
389 fn get_gguf_tokenizer(tokenizer: TokenizerType) -> Result<Tokenizer> {
390 match tokenizer {
391 TokenizerType::Llama => {
392 let api = ApiBuilder::new().with_progress(true).build().unwrap();
393 let api = api.repo(Repo::with_revision(
394 "hanzoai/mistralrs_tests".to_string(),
395 RepoType::Model,
396 "main".to_string(),
397 ));
398
399 let filename = api.get("llama_gguf_tokenizer.json").unwrap();
400 let tokenizer = Tokenizer::from_file(filename).expect("Valid tokenizer");
401 Ok(tokenizer)
402 }
403 TokenizerType::Gpt2 => {
404 let api = ApiBuilder::new().with_progress(true).build().unwrap();
405 let api = api.repo(Repo::with_revision(
406 "hanzoai/mistralrs_tests".to_string(),
407 RepoType::Model,
408 "main".to_string(),
409 ));
410
411 let filename = api.get("gpt2_gguf_tokenizer.json").unwrap();
412 let tokenizer = Tokenizer::from_file(filename).expect("Valid tokenizer");
413 Ok(tokenizer)
414 }
415 other => anyhow::bail!("Cannot get testing HF tokenizer for type {other:?}"),
416 }
417 }
418
419 fn get_hf_tokenizer(tokenizer: TokenizerType) -> Result<Tokenizer> {
420 match tokenizer {
421 TokenizerType::Llama => {
422 let api = ApiBuilder::new().with_progress(true).build().unwrap();
423 let api = api.repo(Repo::with_revision(
424 "hanzoai/mistralrs_tests".to_string(),
425 RepoType::Model,
426 "main".to_string(),
427 ));
428
429 let tokenizer_filename = api.get("tokenizer.json").unwrap();
430 Ok(Tokenizer::from_file(tokenizer_filename).unwrap())
431 }
432 TokenizerType::Gpt2 => {
433 let api = ApiBuilder::new().with_progress(true).build().unwrap();
434 let api = api.repo(Repo::with_revision(
435 "hanzoai/mistralrs_tests".to_string(),
436 RepoType::Model,
437 "main".to_string(),
438 ));
439
440 let tokenizer_filename = api.get("tokenizer_gpt2.json").unwrap();
441 Ok(Tokenizer::from_file(tokenizer_filename).unwrap())
442 }
443 other => anyhow::bail!("Cannot get testing HF tokenizer for type {other:?}"),
444 }
445 }
446
447 fn get_test_passage() -> String {
449 let passage = "Hello, world! \n🚀 (normal) 😶🌫️ (compound emoji, zwj sequence) ✅ (emoji as single token)\n你好世界!\nNǐ hǎo shìjiè!";
450
451 passage.to_owned()
452 }
453
454 fn codec_roundtrip(
456 tokenizer: &Tokenizer,
457 passage: &str,
458 add_special_tokens: bool,
459 ) -> Result<String> {
460 let tokenized = tokenizer
461 .encode_fast(passage, add_special_tokens)
462 .map_err(anyhow::Error::msg)?;
463
464 decode(tokenizer, tokenized.get_ids(), !add_special_tokens)
466 }
467
468 fn decode(
469 tokenizer: &Tokenizer,
470 token_ids: &[u32],
471 skip_special_tokens: bool,
472 ) -> Result<String> {
473 tokenizer
474 .decode(token_ids, skip_special_tokens)
475 .map_err(anyhow::Error::msg)
476 }
477
478 #[test]
479 fn test_encode_decode_llama() -> Result<()> {
480 use rand::rng;
481 use rand::seq::SliceRandom;
482
483 let passage = get_test_passage();
484 let hf_tokenizer = get_hf_tokenizer(TokenizerType::Llama)?;
485 let gguf_tokenizer = get_gguf_tokenizer(TokenizerType::Llama)?;
486
487 let hf_decoded = codec_roundtrip(&hf_tokenizer, passage.as_str(), false)?;
489 let gguf_decoded = codec_roundtrip(&gguf_tokenizer, passage.as_str(), false)?;
490 assert_eq!(hf_decoded, gguf_decoded);
491 assert_eq!(passage, gguf_decoded);
492
493 #[allow(clippy::cast_possible_truncation)]
504 let mut tokens = (0..hf_tokenizer.get_vocab_size(false) as u32).collect::<Vec<_>>();
505 tokens.shuffle(&mut rng());
506
507 let hf_decoded = decode(&hf_tokenizer, &tokens, false)?;
509 let gguf_decoded = decode(&gguf_tokenizer, &tokens, false)?;
510 assert_eq!(hf_decoded, gguf_decoded);
511
512 let hf_decoded = decode(&hf_tokenizer, &tokens, true)?;
514 let gguf_decoded = decode(&gguf_tokenizer, &tokens, true)?;
515 assert_eq!(hf_decoded, gguf_decoded);
516
517 Ok(())
518 }
519
520 #[test]
521 fn test_encode_decode_gpt2() -> Result<()> {
522 use rand::rng;
523 use rand::seq::SliceRandom;
524
525 let passage = get_test_passage();
526 let hf_tokenizer = get_hf_tokenizer(TokenizerType::Gpt2)?;
527 let gguf_tokenizer = get_gguf_tokenizer(TokenizerType::Gpt2)?;
528
529 let hf_decoded = codec_roundtrip(&hf_tokenizer, passage.as_str(), false)?;
531 let gguf_decoded = codec_roundtrip(&gguf_tokenizer, passage.as_str(), false)?;
532 assert_eq!(hf_decoded, gguf_decoded);
533 assert_eq!(passage, gguf_decoded);
534
535 #[allow(clippy::cast_possible_truncation)]
546 let mut tokens = (0..hf_tokenizer.get_vocab_size(false) as u32).collect::<Vec<_>>();
547 tokens.shuffle(&mut rng());
548
549 let hf_decoded = decode(&hf_tokenizer, &tokens, false)?;
551 let gguf_decoded = decode(&gguf_tokenizer, &tokens, false)?;
552 assert_eq!(hf_decoded, gguf_decoded);
553
554 let hf_decoded = decode(&hf_tokenizer, &tokens, true)?;
556 let gguf_decoded = decode(&gguf_tokenizer, &tokens, true)?;
557 assert_eq!(hf_decoded, gguf_decoded);
558
559 Ok(())
560 }
561}