aria_inference/
tokenizer.rs1use aria_kernel::EngineError;
7use std::path::Path;
8use std::sync::Arc;
9use tokenizers::Tokenizer;
10
11const STOP_TOKEN_STRINGS: &[&str] = &[
12 "<|im_end|>",
13 "<|endoftext|>",
14 "<|eot_id|>",
15 "</s>",
16 "<end_of_turn>",
17 "<turn|>",
18 "<eos>",
19 "<|end|>",
20 "<pad>",
21];
22
23#[derive(Clone)]
24pub struct BundleTokenizer {
25 inner: Arc<Tokenizer>,
26 stop_ids: Vec<u32>,
27}
28
29impl std::fmt::Debug for BundleTokenizer {
30 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31 f.debug_struct("BundleTokenizer")
32 .field("vocab_size", &self.inner.get_vocab_size(false))
33 .field("stop_ids", &self.stop_ids)
34 .finish()
35 }
36}
37
38impl BundleTokenizer {
39 fn wrap(tok: Tokenizer) -> Self {
40 let stop_ids = collect_stop_ids(&tok);
41 Self {
42 inner: Arc::new(tok),
43 stop_ids,
44 }
45 }
46
47 pub fn try_load(dir: &Path) -> Result<Option<Self>, EngineError> {
49 let path = dir.join("tokenizer.json");
50 if !path.is_file() {
51 return Ok(None);
52 }
53 let tok = Tokenizer::from_file(&path).map_err(|e| {
54 EngineError::Format(format!(
55 "tokenizer.json load failed ({}): {e}",
56 path.display()
57 ))
58 })?;
59 Ok(Some(Self::wrap(tok)))
60 }
61
62 pub fn from_tokenizer_json(raw: &str) -> Result<Self, EngineError> {
63 let tok = Tokenizer::from_bytes(raw.as_bytes())
64 .map_err(|e| EngineError::Format(format!("tokenizer.json parse failed: {e}")))?;
65 Ok(Self::wrap(tok))
66 }
67
68 pub fn encode(&self, text: &str) -> Result<Vec<u32>, EngineError> {
70 let enc = self
71 .inner
72 .encode(text, false)
73 .map_err(|e| EngineError::InvalidParam(format!("tokenizer encode failed: {e}")))?;
74 Ok(enc.get_ids().to_vec())
75 }
76
77 pub fn decode(&self, ids: &[u32]) -> String {
79 self.decode_opts(ids, true)
80 }
81
82 pub fn decode_opts(&self, ids: &[u32], skip_special: bool) -> String {
83 match self.inner.decode(ids, skip_special) {
84 Ok(s) => s,
85 Err(_) => decode_placeholders(ids),
86 }
87 }
88
89 pub fn is_stop(&self, id: u32) -> bool {
90 self.stop_ids.contains(&id)
91 }
92
93 pub fn stop_ids(&self) -> &[u32] {
94 &self.stop_ids
95 }
96
97 pub fn has_token(&self, token: &str) -> bool {
98 self.inner.token_to_id(token).is_some()
99 }
100
101 pub fn chat_family_hint(&self) -> Option<&'static str> {
104 if self.has_token("<|im_start|>") {
105 if self.has_token("<think>") {
106 Some("qwen/qwen3-0.6b")
107 } else {
108 Some("chatml")
109 }
110 } else if self.has_token("<|turn>") {
111 Some("gemma/gemma-4-e2b-it")
112 } else if self.has_token("<start_of_turn>") {
113 Some("gemma/gemma-3-1b-it")
114 } else if self.has_token("<|eot_id|>") {
115 Some("llama")
116 } else {
117 None
118 }
119 }
120}
121
122fn collect_stop_ids(tok: &Tokenizer) -> Vec<u32> {
123 let mut ids = Vec::new();
124 let mut push = |id: u32| {
125 if !ids.contains(&id) {
126 ids.push(id);
127 }
128 };
129 for name in STOP_TOKEN_STRINGS {
130 if let Some(id) = tok.token_to_id(name) {
131 push(id);
132 }
133 if let Ok(enc) = tok.encode(*name, false) {
135 let got = enc.get_ids();
136 if got.len() == 1 {
137 push(got[0]);
138 }
139 }
140 }
141 for (id, added) in tok.get_added_tokens_decoder() {
142 let content = added.content;
143 if STOP_TOKEN_STRINGS.contains(&content.as_str()) {
144 push(id);
145 }
146 }
147 ids
148}
149
150pub fn decode_placeholders(ids: &[u32]) -> String {
152 ids.iter().map(|t| format!("<{t}>")).collect()
153}
154
155pub fn encode_naive(text: &str, vocab_size: u32) -> Vec<u32> {
157 let vocab = vocab_size.max(1);
158 if text.is_empty() {
159 return vec![1 % vocab];
160 }
161 text.bytes().map(|b| (b as u32) % vocab).collect()
162}
163
164#[cfg(test)]
165mod tests {
166 use super::*;
167
168 fn word_level_json() -> String {
169 serde_json::json!({
170 "version": "1.0",
171 "truncation": null,
172 "padding": null,
173 "added_tokens": [
174 {
175 "id": 2,
176 "content": "[UNK]",
177 "single_word": false,
178 "lstrip": false,
179 "rstrip": false,
180 "normalized": false,
181 "special": true
182 }
183 ],
184 "normalizer": null,
185 "pre_tokenizer": { "type": "Whitespace" },
186 "post_processor": null,
187 "decoder": null,
188 "model": {
189 "type": "WordLevel",
190 "vocab": {
191 "Hello": 0,
192 "world": 1,
193 "[UNK]": 2
194 },
195 "unk_token": "[UNK]"
196 }
197 })
198 .to_string()
199 }
200
201 #[test]
202 fn encode_decode_roundtrip_word_level() {
203 let tok = BundleTokenizer::from_tokenizer_json(&word_level_json()).unwrap();
204 let ids = tok.encode("Hello world").unwrap();
205 assert_eq!(ids, vec![0, 1]);
206 assert_eq!(tok.decode(&ids), "Hello world");
207 }
208
209 #[test]
210 fn decode_skips_special() {
211 let tok = BundleTokenizer::from_tokenizer_json(&word_level_json()).unwrap();
212 assert_eq!(tok.decode(&[0, 2, 1]), "Hello world");
213 assert!(tok.decode_opts(&[0, 2, 1], false).contains("[UNK]"));
214 }
215
216 #[test]
217 fn try_load_from_dir() {
218 let dir = tempfile::tempdir().unwrap();
219 std::fs::write(dir.path().join("tokenizer.json"), word_level_json()).unwrap();
220 let tok = BundleTokenizer::try_load(dir.path())
221 .unwrap()
222 .expect("loaded");
223 assert_eq!(tok.encode("Hello").unwrap(), vec![0]);
224 }
225
226 #[test]
227 fn try_load_missing_is_none() {
228 let dir = tempfile::tempdir().unwrap();
229 assert!(BundleTokenizer::try_load(dir.path()).unwrap().is_none());
230 }
231
232 #[test]
233 fn naive_encode_fallback() {
234 assert_eq!(encode_naive("AB", 256), vec![65, 66]);
235 assert_eq!(encode_naive("", 16), vec![1]);
236 }
237
238 #[test]
239 fn stop_ids_from_added_tokens() {
240 let mut v: serde_json::Value = serde_json::from_str(&word_level_json()).unwrap();
241 v["added_tokens"]
242 .as_array_mut()
243 .unwrap()
244 .push(serde_json::json!({
245 "id": 3,
246 "content": "<|im_end|>",
247 "single_word": false,
248 "lstrip": false,
249 "rstrip": false,
250 "normalized": false,
251 "special": true
252 }));
253 v["model"]["vocab"]["<|im_end|>"] = serde_json::json!(3);
254 let tok = BundleTokenizer::from_tokenizer_json(&v.to_string()).unwrap();
255 assert!(tok.is_stop(3));
256 assert!(!tok.is_stop(0));
257 }
258}