Skip to main content

hanzo_engine/gguf/
gguf_tokenizer.rs

1// https://github.com/huggingface/transformers/blob/8685b3c5d2dd2550527773d2a02499495a759e31/src/transformers/convert_slow_tokenizer.py
2
3use 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    //token type other than 1 treated as special token
103    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// TODO: Add support for additional tokenizer models: WordPiece, WordLevel
147// https://docs.rs/tokenizers/latest/tokenizers/models/enum.ModelWrapper.html
148#[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    // Unigram (SentencePiece) default UNK is 0
157    let unk = unk.unwrap_or(0);
158
159    // Create the Tokenizer model:
160    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    // Decoder + Normalizer config reference:
176    // https://github.com/hanzoai/engine/pull/389#discussion_r1630620763
177    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    // Add special tokens (bos, eos, unk):
196    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    // BPE merges have each string item as a space-delimited pair:
205    // https://github.com/hanzoai/engine/pull/397#discussion_r1631988370
206    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    // example:
248    // "type": "ByteLevel",
249    // "add_prefix_space": false,
250    // "trim_offsets": false,
251    // "use_regex": false
252    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
274// This is a workaround to have a better builder API.
275// Upstream `TokenizerBuilder` is difficult to work with:
276// https://github.com/huggingface/tokenizers/issues/1549
277struct 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        // Handle local enum to remote enum type:
289        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
302// Convenient alternative to upstream:
303// https://docs.rs/tokenizers/latest/tokenizers/decoders/enum.DecoderWrapper.html
304enum 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
313// Convert into upstream type wrapped enum variants:
314impl 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
342// Convenient alternative to upstream:
343// https://docs.rs/tokenizers/latest/tokenizers/normalizers/enum.NormalizerWrapper.html
344enum 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        /// Mistral v0.1 tokenizer
383        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    // Content based upon https://github.com/ggerganov/llama.cpp/blob/master/tests/test-tokenizer-random.py#L99-L161
448    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    // The provided passage should encode and decode back into the same passage string:
455    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        // NOTE: The special tokens bool param meaning differs between encode() / decode():
465        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        // Without adding special tokens
488        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        // With special tokens added
494        // SKIPPED:
495        // - Bugged the GGUF tokenizer does not prepend `<s> `
496        // - Due to HF tokenizer using BPE (tokenizer.json) while GGUF tokenizer uses Unigram (metadata)?
497        /*
498        let hf_decoded = codec_roundtrip(&hf_tokenizer, passage.as_str(), true)?;
499        let gguf_decoded = codec_roundtrip(&gguf_tokenizer, passage.as_str(), true)?;
500        assert_eq!(hf_decoded, gguf_decoded);
501        */
502
503        #[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        // Without skipping special tokens
508        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        // With skipping special tokens
513        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        // Without adding special tokens
530        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        // With special tokens added
536        // SKIPPED:
537        // - Bugged the GGUF tokenizer does not prepend `<s> `
538        // - Due to HF tokenizer using BPE (tokenizer.json) while GGUF tokenizer uses Unigram (metadata)?
539        /*
540        let hf_decoded = codec_roundtrip(&hf_tokenizer, passage.as_str(), true)?;
541        let gguf_decoded = codec_roundtrip(&gguf_tokenizer, passage.as_str(), true)?;
542        assert_eq!(hf_decoded, gguf_decoded);
543        */
544
545        #[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        // Without skipping special tokens
550        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        // With skipping special tokens
555        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}