voxtral_micro/tokenizer/
encoder.rs1use 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
15pub struct TekkenEncoder {
21 bpe: CoreBPE,
22}
23
24impl TekkenEncoder {
25 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 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 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 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 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 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 for &id in &ids {
170 assert!(id >= TEXT_TOKEN_OFFSET, "Token ID {id} below offset");
171 }
172 }
173}