Skip to main content

llama_tokenizer/
lib.rs

1//! # llama-tokenizer
2//!
3//! Deterministic tokenization for llama.rs.
4//!
5//! This crate provides:
6//! - A `Tokenizer` trait for pluggable tokenization backends
7//! - A reference whitespace tokenizer for testing
8//! - Streaming decoding with UTF-8 handling
9//! - Chat template support (future)
10
11use std::collections::HashMap;
12use std::sync::RwLock;
13
14/// Error type for tokenization operations.
15#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
16pub enum TokenizerError {
17    #[error("Invalid token ID: {0}")]
18    InvalidToken(i32),
19    #[error("Encoding error: {0}")]
20    EncodingError(String),
21    #[error("Decoding error: {0}")]
22    DecodingError(String),
23}
24
25pub type TokenizerResult<T> = std::result::Result<T, TokenizerError>;
26
27/// Core tokenizer trait. Implementations can be swapped without changing app code.
28pub trait Tokenizer: Send + Sync {
29    /// Encode text into a sequence of token IDs.
30    fn encode(&self, text: &str) -> TokenizerResult<Vec<i32>>;
31
32    /// Decode a complete sequence of tokens into text.
33    fn decode(&self, tokens: &[i32]) -> TokenizerResult<String>;
34
35    /// Decode a single token and accumulate with partial UTF-8 state.
36    /// For streaming decoding, this allows emitting printable characters immediately.
37    fn decode_token(&self, token: i32, state: &mut DecodingState) -> TokenizerResult<String>;
38
39    /// Get vocabulary size.
40    fn vocab_size(&self) -> usize;
41}
42
43/// Streaming decoding state for handling partial UTF-8 sequences.
44#[derive(Debug, Clone, Default)]
45pub struct DecodingState {
46    buffer: String,
47    pending_utf8: Vec<u8>,
48    emitted_any: bool,
49}
50
51impl DecodingState {
52    pub fn new() -> Self {
53        Self::default()
54    }
55
56    pub fn buffer(&self) -> &str {
57        &self.buffer
58    }
59
60    pub fn clear(&mut self) {
61        self.buffer.clear();
62        self.pending_utf8.clear();
63        self.emitted_any = false;
64    }
65}
66
67/// Reference whitespace tokenizer for Milestone A testing.
68///
69/// - Splits on whitespace
70/// - Bidirectional (encode/decode)
71/// - Deterministic
72/// - Used for golden tests before real tokenizer.json loading
73pub struct WhitespaceTokenizer {
74    state: RwLock<VocabState>,
75}
76
77#[derive(Debug, Default)]
78struct VocabState {
79    vocab: HashMap<i32, String>,
80    reverse_vocab: HashMap<String, i32>,
81    next_id: i32,
82}
83
84impl WhitespaceTokenizer {
85    pub fn new() -> Self {
86        Self {
87            state: RwLock::new(VocabState::default()),
88        }
89    }
90
91    fn decode_id(&self, token: i32) -> TokenizerResult<String> {
92        let state = self
93            .state
94            .read()
95            .map_err(|_| TokenizerError::DecodingError("tokenizer lock poisoned".to_string()))?;
96
97        state
98            .vocab
99            .get(&token)
100            .cloned()
101            .ok_or(TokenizerError::InvalidToken(token))
102    }
103}
104
105impl Default for WhitespaceTokenizer {
106    fn default() -> Self {
107        Self::new()
108    }
109}
110
111impl Tokenizer for WhitespaceTokenizer {
112    fn encode(&self, text: &str) -> TokenizerResult<Vec<i32>> {
113        let mut state = self
114            .state
115            .write()
116            .map_err(|_| TokenizerError::EncodingError("tokenizer lock poisoned".to_string()))?;
117
118        let mut ids = Vec::new();
119        for word in text.split_whitespace() {
120            let id = if let Some(id) = state.reverse_vocab.get(word) {
121                *id
122            } else {
123                let id = state.next_id;
124                state.next_id += 1;
125                state.reverse_vocab.insert(word.to_string(), id);
126                state.vocab.insert(id, word.to_string());
127                id
128            };
129            ids.push(id);
130        }
131
132        Ok(ids)
133    }
134
135    fn decode(&self, tokens: &[i32]) -> TokenizerResult<String> {
136        let mut words = Vec::with_capacity(tokens.len());
137        for &id in tokens {
138            words.push(self.decode_id(id)?);
139        }
140        Ok(words.join(" "))
141    }
142
143    fn decode_token(&self, token: i32, state: &mut DecodingState) -> TokenizerResult<String> {
144        let word = self.decode_id(token)?;
145        let emitted = if state.emitted_any {
146            format!(" {}", word)
147        } else {
148            word
149        };
150        state.buffer.push_str(&emitted);
151        state.emitted_any = true;
152        Ok(emitted)
153    }
154
155    fn vocab_size(&self) -> usize {
156        self.state.read().map(|s| s.vocab.len()).unwrap_or(0)
157    }
158}
159
160#[cfg(test)]
161mod tests {
162    use super::*;
163
164    #[test]
165    fn encode_whitespace_simple() {
166        let tok = WhitespaceTokenizer::new();
167        let ids = tok.encode("hello world").unwrap();
168        assert_eq!(ids.len(), 2);
169    }
170
171    #[test]
172    fn encode_empty_string() {
173        let tok = WhitespaceTokenizer::new();
174        let ids = tok.encode("").unwrap();
175        assert!(ids.is_empty());
176    }
177
178    #[test]
179    fn encode_single_word() {
180        let tok = WhitespaceTokenizer::new();
181        let ids = tok.encode("hello").unwrap();
182        assert_eq!(ids.len(), 1);
183    }
184
185    #[test]
186    fn encode_multiple_spaces() {
187        let tok = WhitespaceTokenizer::new();
188        let ids = tok.encode("hello    world").unwrap();
189        assert_eq!(ids.len(), 2);
190    }
191
192    #[test]
193    fn decode_roundtrip() {
194        let tok = WhitespaceTokenizer::new();
195        let original = "hello world test";
196        let encoded = tok.encode(original).unwrap();
197        let decoded = tok.decode(&encoded).unwrap();
198        assert_eq!(decoded, original);
199    }
200
201    #[test]
202    fn decode_empty_tokens() {
203        let tok = WhitespaceTokenizer::new();
204        let decoded = tok.decode(&[]).unwrap();
205        assert_eq!(decoded, "");
206    }
207
208    #[test]
209    fn streaming_decode_state() {
210        let tok: &dyn Tokenizer = &WhitespaceTokenizer::new();
211        let encoded = tok.encode("hello world").unwrap();
212        let mut state = DecodingState::new();
213
214        assert_eq!(state.buffer(), "");
215        assert_eq!(tok.decode_token(encoded[0], &mut state).unwrap(), "hello");
216        assert_eq!(state.buffer(), "hello");
217        assert_eq!(tok.decode_token(encoded[1], &mut state).unwrap(), " world");
218        assert_eq!(state.buffer(), "hello world");
219
220        state.clear();
221        assert_eq!(state.buffer(), "");
222    }
223
224    #[test]
225    fn decode_invalid_token_errors() {
226        let tok = WhitespaceTokenizer::new();
227        tok.encode("hello").unwrap();
228        assert_eq!(
229            tok.decode(&[999]).unwrap_err(),
230            TokenizerError::InvalidToken(999)
231        );
232    }
233
234    #[test]
235    fn vocab_size_reflects_built_vocab() {
236        let tok = WhitespaceTokenizer::new();
237        assert_eq!(tok.vocab_size(), 0);
238        tok.encode("hello world hello").unwrap();
239        assert_eq!(tok.vocab_size(), 2);
240    }
241}