Skip to main content

dynamo_bench/coding/
tokenizer.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use 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}