Skip to main content

aria_inference/
tokenizer.rs

1//! Bundle tokenizer sidecar via HuggingFace `tokenizers`.
2//!
3//! Loads `tokenizer.json` for real encode + decode. Without a sidecar, Session falls
4//! back to naive byte encode / `<id>` decode placeholders (tiny fixtures).
5
6use 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    /// Load from a bundle directory. Returns `Ok(None)` if no `tokenizer.json`.
48    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    /// Encode text → token ids (no extra specials; chat template is already in `text`).
69    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    /// Decode ids → UTF-8, skipping special tokens by default.
78    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    /// Prefer tokenizer specials over a possibly-wrong Session family (serve used to
102    /// hardcode gemma even for Qwen bundles).
103    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        // Added tokens sometimes miss token_to_id; encode as a whole piece.
134        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
150/// Fallback when no tokenizer sidecar: stable `<id>` placeholders (legacy demos).
151pub fn decode_placeholders(ids: &[u32]) -> String {
152    ids.iter().map(|t| format!("<{t}>")).collect()
153}
154
155/// Fallback encode when no sidecar: map UTF-8 bytes into `[0, vocab)`.
156pub 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}