Skip to main content

voxtral_micro/tokenizer/
encoder.rs

1//! Tekken BPE encoder using tiktoken-rs.
2//!
3//! Builds a `CoreBPE` from tekken.json for text → token ID encoding.
4//! Token IDs are offset by 1000 (special tokens occupy 0-999).
5
6use std::path::Path;
7
8use anyhow::{Context, Result};
9use base64::prelude::*;
10use rustc_hash::FxHashMap;
11use tiktoken_rs::CoreBPE;
12
13use super::{TekkenJson, TEXT_TOKEN_OFFSET};
14
15/// Tekken BPE encoder for text → token IDs.
16///
17/// Wraps tiktoken's `CoreBPE` with the correct vocab and regex pattern
18/// from a tekken.json file. Output token IDs include the 1000 offset
19/// (matching the model's embedding table layout).
20pub struct TekkenEncoder {
21    bpe: CoreBPE,
22}
23
24impl TekkenEncoder {
25    /// Load encoder from a `tekken.json` file.
26    pub fn from_file<P: AsRef<Path>>(path: P) -> Result<Self> {
27        let path = path.as_ref();
28        let file = std::fs::File::open(path)
29            .with_context(|| format!("Failed to open tokenizer: {}", path.display()))?;
30        let reader = std::io::BufReader::new(file);
31        let tekken: TekkenJson = serde_json::from_reader(reader)
32            .with_context(|| format!("Failed to parse tekken.json: {}", path.display()))?;
33        Self::from_tekken(tekken)
34    }
35
36    /// Load encoder from a JSON string.
37    pub fn from_json(json: &str) -> Result<Self> {
38        let tekken: TekkenJson =
39            serde_json::from_str(json).context("Failed to parse tekken JSON")?;
40        Self::from_tekken(tekken)
41    }
42
43    fn from_tekken(tekken: TekkenJson) -> Result<Self> {
44        let pattern = &tekken.config.pattern;
45        let inner_vocab_size =
46            tekken.config.default_vocab_size - tekken.config.default_num_special_tokens;
47
48        // Build mergeable_ranks: bytes → rank (for non-special tokens only)
49        let mut mergeable_ranks: FxHashMap<Vec<u8>, u32> =
50            FxHashMap::with_capacity_and_hasher(inner_vocab_size, Default::default());
51
52        for entry in &tekken.vocab {
53            if entry.is_control {
54                continue;
55            }
56
57            let rank = entry.rank;
58            if rank as usize >= inner_vocab_size {
59                continue;
60            }
61
62            let bytes = if let Some(b64) = &entry.token_bytes {
63                BASE64_STANDARD
64                    .decode(b64)
65                    .with_context(|| format!("Bad base64 for rank {rank}"))?
66            } else if let Some(s) = &entry.token_str {
67                s.as_bytes().to_vec()
68            } else {
69                continue;
70            };
71
72            mergeable_ranks.insert(bytes, rank);
73        }
74
75        let special_tokens: FxHashMap<String, u32> = FxHashMap::default();
76        let bpe = CoreBPE::new(mergeable_ranks, special_tokens, pattern)
77            .map_err(|e| anyhow::anyhow!("Failed to create CoreBPE: {e}"))?;
78
79        Ok(Self { bpe })
80    }
81
82    /// Encode text to token IDs (with 1000 offset applied).
83    ///
84    /// Returns token IDs matching the model's embedding table layout:
85    /// IDs 0-999 are special tokens, 1000+ are text tokens.
86    pub fn encode(&self, text: &str) -> Vec<u32> {
87        let ranks = self.bpe.encode_ordinary(text);
88        ranks.into_iter().map(|r| r + TEXT_TOKEN_OFFSET).collect()
89    }
90}
91
92#[cfg(test)]
93mod tests {
94    use super::*;
95    use std::path::PathBuf;
96
97    fn tts_tokenizer_path() -> PathBuf {
98        PathBuf::from("models/voxtral-tts/tekken.json")
99    }
100
101    fn asr_tokenizer_path() -> PathBuf {
102        PathBuf::from("models/voxtral/tekken.json")
103    }
104
105    #[test]
106    fn test_encode_hello_world() {
107        let path = tts_tokenizer_path();
108        if !path.exists() {
109            let path2 = asr_tokenizer_path();
110            if !path2.exists() {
111                println!("Skipping: no tokenizer available");
112                return;
113            }
114            let enc = TekkenEncoder::from_file(&path2).unwrap();
115            let ids = enc.encode("Hello world");
116            // Both ASR and TTS tokenizers use the same Tekken vocab
117            assert_eq!(ids, vec![22177, 4304]);
118            return;
119        }
120        let enc = TekkenEncoder::from_file(&path).unwrap();
121        let ids = enc.encode("Hello world");
122        assert_eq!(ids, vec![22177, 4304]);
123    }
124
125    #[test]
126    fn test_encode_various_texts() {
127        let path = tts_tokenizer_path();
128        let path = if path.exists() {
129            path
130        } else {
131            let p = asr_tokenizer_path();
132            if !p.exists() {
133                println!("Skipping: no tokenizer available");
134                return;
135            }
136            p
137        };
138
139        let enc = TekkenEncoder::from_file(&path).unwrap();
140
141        // Verified against Python: Tekkenizer.encode(text, bos=False, eos=False)
142        assert_eq!(
143            enc.encode("Mary had a little lamb"),
144            vec![48650, 1880, 1261, 4945, 56914]
145        );
146        assert_eq!(
147            enc.encode("The quick brown fox jumps over the lazy dog."),
148            vec![1784, 7586, 22980, 94137, 72993, 2136, 1278, 42757, 10575, 1046]
149        );
150    }
151
152    #[test]
153    fn test_encode_offset() {
154        let path = tts_tokenizer_path();
155        let path = if path.exists() {
156            path
157        } else {
158            let p = asr_tokenizer_path();
159            if !p.exists() {
160                println!("Skipping: no tokenizer available");
161                return;
162            }
163            p
164        };
165
166        let enc = TekkenEncoder::from_file(&path).unwrap();
167        let ids = enc.encode("a");
168        // All IDs should be >= 1000 (TEXT_TOKEN_OFFSET)
169        for &id in &ids {
170            assert!(id >= TEXT_TOKEN_OFFSET, "Token ID {id} below offset");
171        }
172    }
173}