dynamo_bench/coding/
tokenizer.rs1use anyhow::{Context, Result, anyhow, bail};
5use std::path::{Path, PathBuf};
6use tokenizers::Tokenizer;
7
8#[derive(Clone, Debug)]
9pub struct HfTokenizerFactory {
10 tokenizer_path: PathBuf,
11}
12
13impl HfTokenizerFactory {
14 pub fn resolve(model_or_path: &str) -> Result<Self> {
15 Ok(Self {
16 tokenizer_path: resolve_tokenizer_path(model_or_path)?,
17 })
18 }
19}
20
21pub trait TokenizerWorker: Send + 'static {
22 fn encode(&mut self, text: &str) -> Result<Vec<u32>>;
23
24 fn encode_with_word_overlap(
25 &mut self,
26 text: &str,
27 previous_text: Option<&str>,
28 previous_tokens: Option<&[u32]>,
29 overlap_words: usize,
30 ) -> Result<Vec<u32>> {
31 if overlap_words == 0 {
32 return self.encode(text);
33 }
34 let Some(previous_text) = previous_text else {
35 return self.encode(text);
36 };
37 let Some(previous_tokens) = previous_tokens else {
38 return self.encode(text);
39 };
40 if !text.starts_with(previous_text) {
41 return self.encode(text);
42 }
43
44 let overlap_start = last_word_overlap_start(previous_text, overlap_words);
45 let previous_suffix_tokens = self.encode(&previous_text[overlap_start..])?;
46 let prefix_token_count = previous_tokens
47 .len()
48 .saturating_sub(previous_suffix_tokens.len());
49 let suffix_tokens = self.encode(&text[overlap_start..])?;
50
51 let mut merged = Vec::with_capacity(prefix_token_count + suffix_tokens.len());
52 merged.extend_from_slice(&previous_tokens[..prefix_token_count]);
53 merged.extend(suffix_tokens);
54 Ok(merged)
55 }
56}
57
58pub trait TokenizerFactory: Clone + Send + Sync + 'static {
59 type Worker: TokenizerWorker;
60
61 fn create_worker(&self) -> Result<Self::Worker>;
62}
63
64impl TokenizerFactory for HfTokenizerFactory {
65 type Worker = HfTokenizerWorker;
66
67 fn create_worker(&self) -> Result<Self::Worker> {
68 HfTokenizerWorker::from_file(&self.tokenizer_path)
69 }
70}
71
72pub struct HfTokenizerWorker {
73 tokenizer: Tokenizer,
74}
75
76impl HfTokenizerWorker {
77 pub fn from_file(path: &Path) -> Result<Self> {
78 let tokenizer = Tokenizer::from_file(path).map_err(|error| {
79 anyhow!("failed to load tokenizer from {}: {error}", path.display())
80 })?;
81 Ok(Self { tokenizer })
82 }
83}
84
85impl TokenizerWorker for HfTokenizerWorker {
86 fn encode(&mut self, text: &str) -> Result<Vec<u32>> {
87 let encoding = self
88 .tokenizer
89 .encode(text, false)
90 .map_err(|error| anyhow!("failed to tokenize input: {error}"))?;
91 Ok(encoding.get_ids().to_vec())
92 }
93}
94
95fn resolve_tokenizer_path(model_or_path: &str) -> Result<PathBuf> {
96 let path = Path::new(model_or_path);
97 if path.is_file() {
98 return Ok(path.to_path_buf());
99 }
100
101 if path.is_dir() {
102 let tokenizer_path = path.join("tokenizer.json");
103 if tokenizer_path.exists() {
104 return Ok(tokenizer_path);
105 }
106 bail!(
107 "directory '{}' does not contain tokenizer.json",
108 path.display()
109 );
110 }
111
112 let cache = hf_hub::Cache::default();
113 let api = hf_hub::api::sync::ApiBuilder::from_cache(cache)
114 .with_progress(true)
115 .build()
116 .context("failed to create HuggingFace API client")?;
117 let repo = api.model(model_or_path.to_string());
118 let tokenizer_path = repo
119 .get("tokenizer.json")
120 .with_context(|| format!("failed to download tokenizer.json from '{model_or_path}'"))?;
121 Ok(tokenizer_path)
122}
123
124pub(crate) fn last_word_overlap_start(text: &str, overlap_words: usize) -> usize {
125 if overlap_words == 0 || text.is_empty() {
126 return text.len();
127 }
128
129 let mut starts = Vec::new();
130 let mut in_word = false;
131 for (index, ch) in text.char_indices() {
132 if ch.is_whitespace() {
133 in_word = false;
134 continue;
135 }
136 if !in_word {
137 starts.push(index);
138 in_word = true;
139 }
140 }
141
142 if starts.len() <= overlap_words {
143 return 0;
144 }
145 starts[starts.len() - overlap_words]
146}
147
148#[cfg(test)]
149mod tests {
150 use super::{TokenizerWorker, last_word_overlap_start};
151 use anyhow::Result;
152 use std::thread;
153 use std::time::Duration;
154
155 struct StubWorker;
156
157 impl TokenizerWorker for StubWorker {
158 fn encode(&mut self, text: &str) -> Result<Vec<u32>> {
159 if text.contains("slow") {
160 thread::sleep(Duration::from_millis(5));
161 }
162 Ok(text
163 .split_whitespace()
164 .map(|word| word.len() as u32)
165 .collect())
166 }
167 }
168
169 #[test]
170 fn overlap_start_returns_zero_for_short_text() {
171 assert_eq!(last_word_overlap_start("one two", 10), 0);
172 }
173
174 #[test]
175 fn overlap_encoding_falls_back_on_prefix_break() {
176 let mut worker = StubWorker;
177 let tokens = worker
178 .encode_with_word_overlap("a b c", Some("z y"), Some(&[1, 2]), 2)
179 .unwrap();
180 assert_eq!(tokens, vec![1, 1, 1]);
181 }
182}