[model-gateway] fix tokenizer to match transformers special token handling (#16087)

This commit is contained in:
Simo Lin
2025-12-29 08:13:03 -08:00
committed by GitHub
parent b5d9fc873b
commit c31f62722c
18 changed files with 265 additions and 132 deletions
@@ -120,7 +120,7 @@ fn bench_encode_throughput(c: &mut Criterion) {
let tokenizer_clone = tokenizer.clone();
// Get token count once
let encoding = tokenizer.encode(prompt).unwrap();
let encoding = tokenizer.encode(prompt, false).unwrap();
let token_count = encoding.token_ids().len();
// Track if metrics have been printed for this test case
@@ -134,7 +134,7 @@ fn bench_encode_throughput(c: &mut Criterion) {
b.iter_custom(|iters| {
let start = Instant::now();
for _ in 0..iters {
black_box(tokenizer.encode(prompt).unwrap());
black_box(tokenizer.encode(prompt, false).unwrap());
}
let duration = start.elapsed();
@@ -178,7 +178,7 @@ fn bench_batch_encode(c: &mut Criterion) {
let batch_sizes = vec![1, 8, 16, 32, 64, 128];
let prompt = MEDIUM_PROMPT;
let prompt_len = prompt.len();
let encoding = tokenizer.encode(prompt).unwrap();
let encoding = tokenizer.encode(prompt, false).unwrap();
let token_count = encoding.token_ids().len();
let mut group = c.benchmark_group("batch_encode");
@@ -199,7 +199,7 @@ fn bench_batch_encode(c: &mut Criterion) {
b.iter_custom(|iters| {
let start = Instant::now();
for _ in 0..iters {
black_box(tokenizer.encode_batch(&prompts).unwrap());
black_box(tokenizer.encode_batch(&prompts, false).unwrap());
}
let duration = start.elapsed();
@@ -272,7 +272,7 @@ fn bench_concurrent_encode(c: &mut Criterion) {
let mut local_chars = 0u64;
while start.elapsed() < Duration::from_millis(500) {
let _ = tokenizer.encode(prompt).unwrap();
let _ = tokenizer.encode(prompt, false).unwrap();
local_ops += 1;
local_chars += prompt.len() as u64;
}
@@ -325,7 +325,7 @@ fn bench_decode_performance(c: &mut Criterion) {
);
let test_text = "The quick brown fox jumps over the lazy dog. ".repeat(10);
let encoding = tokenizer.encode(&test_text).unwrap();
let encoding = tokenizer.encode(&test_text, false).unwrap();
let tokens = encoding.token_ids();
let num_tokens = tokens.len();
@@ -444,7 +444,7 @@ fn bench_streaming_decode_100k(c: &mut Criterion) {
);
let sample_text = "The quick brown fox jumps over the lazy dog. ".repeat(1000);
let encoding = tokenizer.encode(&sample_text).unwrap();
let encoding = tokenizer.encode(&sample_text, false).unwrap();
let all_tokens = encoding.token_ids();
let mut group = c.benchmark_group("streaming_100k");
@@ -587,13 +587,13 @@ fn bench_latency_distribution(c: &mut Criterion) {
// Warm up
for _ in 0..100 {
let _ = tokenizer.encode(prompt).unwrap();
let _ = tokenizer.encode(prompt, false).unwrap();
}
// Measure for statistics
for _ in 0..1000 {
let start = Instant::now();
let _ = tokenizer.encode(prompt).unwrap();
let _ = tokenizer.encode(prompt, false).unwrap();
let latency = start.elapsed();
latencies.push(latency);
}
@@ -623,7 +623,7 @@ fn bench_latency_distribution(c: &mut Criterion) {
// Regular benchmark iterations
let start = Instant::now();
for _ in 0..iters {
black_box(tokenizer.encode(prompt).unwrap());
black_box(tokenizer.encode(prompt, false).unwrap());
}
start.elapsed()
};
@@ -712,7 +712,7 @@ fn bench_concurrent_streaming(c: &mut Criterion) {
let tokens_per_sequence = 10_000;
let sample_text = "The quick brown fox jumps over the lazy dog. ".repeat(100);
let encoding = tokenizer.encode(&sample_text).unwrap();
let encoding = tokenizer.encode(&sample_text, false).unwrap();
let token_batch: Vec<u32> = encoding.token_ids().to_vec();
let mut group = c.benchmark_group("concurrent_streaming");
@@ -795,7 +795,7 @@ fn bench_stop_sequences(c: &mut Criterion) {
.with_stop_token(2);
let sample_text = "Hello world! This is a test. ### Stop here. Continue after.".repeat(100);
let encoding = tokenizer.encode(&sample_text).unwrap();
let encoding = tokenizer.encode(&sample_text, false).unwrap();
let tokens = encoding.token_ids();
let mut group = c.benchmark_group("stop_sequences");
@@ -937,7 +937,7 @@ fn bench_multithreaded_encode(c: &mut Criterion) {
thread::spawn(move || {
for _ in 0..operations_per_thread {
let encoding = tokenizer.encode(test_prompt).unwrap();
let encoding = tokenizer.encode(test_prompt, false).unwrap();
total_tok.fetch_add(
encoding.token_ids().len() as u64,
Ordering::Relaxed,
@@ -1005,7 +1005,7 @@ fn bench_multithreaded_decode(c: &mut Criterion) {
// Generate tokens for decoding
let test_text = "The quick brown fox jumps over the lazy dog. ".repeat(100);
let encoding = tokenizer.encode(&test_text).unwrap();
let encoding = tokenizer.encode(&test_text, false).unwrap();
let test_tokens: Vec<u32> = encoding.token_ids().to_vec();
let mut group = c.benchmark_group("multithreaded_decode");
@@ -1102,7 +1102,7 @@ fn bench_memory_efficiency(c: &mut Criterion) {
);
let large_text = "The quick brown fox jumps over the lazy dog. ".repeat(1000);
let encoding = tokenizer.encode(&large_text).unwrap();
let encoding = tokenizer.encode(&large_text, false).unwrap();
let mut group = c.benchmark_group("memory");
@@ -1324,7 +1324,7 @@ fn bench_l1_cache_chat_template(c: &mut Criterion) {
for _ in 0..iters {
// Simulate 100 requests with different queries (realistic workload)
for prompt in &test_prompts {
black_box(tokenizer.encode(prompt).unwrap());
black_box(tokenizer.encode(prompt, false).unwrap());
}
}
let duration = start.elapsed();
@@ -1372,7 +1372,7 @@ fn bench_l1_cache_chat_template(c: &mut Criterion) {
for _ in 0..iters {
for prompt in &test_prompts {
black_box(cached.encode(prompt).unwrap());
black_box(cached.encode(prompt, false).unwrap());
}
}
let duration = start.elapsed();
@@ -1423,13 +1423,13 @@ fn bench_l1_cache_chat_template(c: &mut Criterion) {
cached.clear_cache(); // Start fresh
// Prime with first request to populate L1 with system prefix
cached.encode(&test_prompts[0]).unwrap();
cached.encode(&test_prompts[0], false).unwrap();
let start = Instant::now();
for _ in 0..iters {
// All subsequent requests benefit from L1 prefix cache
for prompt in &test_prompts {
black_box(cached.encode(prompt).unwrap());
black_box(cached.encode(prompt, false).unwrap());
}
}
let duration = start.elapsed();
@@ -1507,12 +1507,12 @@ fn bench_l1_cache_chat_template(c: &mut Criterion) {
b.iter_custom(|iters| {
cached.clear_cache();
cached.encode(&test_prompts[0]).unwrap(); // Prime cache
cached.encode(&test_prompts[0], false).unwrap(); // Prime cache
let start = Instant::now();
for _ in 0..iters {
for prompt in &test_prompts {
black_box(cached.encode(prompt).unwrap());
black_box(cached.encode(prompt, false).unwrap());
}
}
let duration = start.elapsed();
@@ -1555,12 +1555,12 @@ fn bench_l1_cache_chat_template(c: &mut Criterion) {
b.iter_custom(|iters| {
cached.clear_cache();
cached.encode(&test_prompts[0]).unwrap(); // Prime cache
cached.encode(&test_prompts[0], false).unwrap(); // Prime cache
let start = Instant::now();
for _ in 0..iters {
for prompt in &test_prompts {
black_box(cached.encode(prompt).unwrap());
black_box(cached.encode(prompt, false).unwrap());
}
}
let duration = start.elapsed();
@@ -1623,7 +1623,7 @@ fn bench_l1_cache_chat_template(c: &mut Criterion) {
for _ in 0..iters {
// Simulate progressive conversation (each turn shares prefix with previous)
for turn in &test_turns {
black_box(cached.encode(turn).unwrap());
black_box(cached.encode(turn, false).unwrap());
}
}
let duration = start.elapsed();
@@ -1692,12 +1692,12 @@ fn bench_l1_cache_chat_template(c: &mut Criterion) {
b.iter_custom(|iters| {
cached.clear_cache();
cached.encode(&test_prompts[0]).unwrap(); // Prime cache
cached.encode(&test_prompts[0], false).unwrap(); // Prime cache
let start = Instant::now();
for _ in 0..iters {
for prompt in &test_prompts {
black_box(cached.encode(prompt).unwrap());
black_box(cached.encode(prompt, false).unwrap());
}
}
let duration = start.elapsed();
+10 -3
View File
@@ -344,12 +344,19 @@ mod tests {
}
impl Encoder for TestTokenizer {
fn encode(&self, _input: &str) -> anyhow::Result<Encoding> {
fn encode(&self, _input: &str, _add_special_tokens: bool) -> anyhow::Result<Encoding> {
Ok(Encoding::Sp(Vec::new()))
}
fn encode_batch(&self, inputs: &[&str]) -> anyhow::Result<Vec<Encoding>> {
inputs.iter().map(|_| self.encode("")).collect()
fn encode_batch(
&self,
inputs: &[&str],
add_special_tokens: bool,
) -> anyhow::Result<Vec<Encoding>> {
inputs
.iter()
.map(|_| self.encode("", add_special_tokens))
.collect()
}
}
@@ -59,8 +59,8 @@ impl ChatPreparationStage {
}
};
// Step 3: Tokenize the processed text
let encoding = match tokenizer.encode(&processed_messages.text) {
// Step 3: Tokenize the processed text (no special tokens - chat template already handles them)
let encoding = match tokenizer.encode(&processed_messages.text, false) {
Ok(encoding) => encoding,
Err(e) => {
error!(function = "ChatPreparationStage::execute", error = %e, "Tokenization failed");
@@ -60,9 +60,10 @@ impl PipelineStage for EmbeddingPreparationStage {
let tokenizer =
utils::resolve_tokenizer(ctx, "EmbeddingPreparationStage::execute").map_err(|e| *e)?;
// Tokenize
// Tokenize with special tokens (BOS/EOS) for embeddings
// This matches Python's transformers behavior which reads add_bos_token/add_eos_token from tokenizer_config.json
let token_ids = tokenizer
.encode(&text)
.encode(&text, true)
.map_err(|e| {
error!(
function = "EmbeddingPreparationStage::execute",
@@ -119,8 +119,9 @@ impl GeneratePreparationStage {
tokenizer: &Arc<dyn Tokenizer>,
text: &str,
) -> Result<(String, Vec<u32>), String> {
// Don't add special tokens - raw text generation uses text as-is
let encoding = tokenizer
.encode(text)
.encode(text, false)
.map_err(|e| format!("Tokenization failed: {}", e))?;
Ok((text.to_string(), encoding.token_ids().to_vec()))
}
@@ -96,7 +96,8 @@ pub async fn tokenize(registry: &Arc<TokenizerRegistry>, request: TokenizeReques
let mut all_char_counts: Vec<i32> = Vec::with_capacity(texts.len());
for text in texts {
let encoding = match tokenizer.encode(text) {
// Don't add special tokens for tokenize API (matches Python behavior)
let encoding = match tokenizer.encode(text, false) {
Ok(enc) => enc,
Err(e) => {
error!("Tokenization failed: {}", e);
+10 -9
View File
@@ -176,6 +176,7 @@ impl L1Cache {
input: &str,
tokenizer: &E,
special_tokens: &[&str],
add_special_tokens: bool,
) -> anyhow::Result<()> {
let boundaries = find_special_token_boundaries(input, special_tokens);
@@ -194,7 +195,7 @@ impl L1Cache {
// Re-tokenize the prefix for guaranteed correctness
// This is the only way to know the exact token boundaries
let prefix_encoding = tokenizer.encode(prefix)?;
let prefix_encoding = tokenizer.encode(prefix, add_special_tokens)?;
// Convert to Arc<[TokenIdType]> for zero-copy sharing
let prefix_tokens: Arc<[TokenIdType]> = prefix_encoding.token_ids().into();
@@ -357,7 +358,7 @@ mod tests {
// Insert at special token boundaries (re-tokenizes prefixes)
cache
.insert_at_boundaries(input1, &tokenizer, special_tokens)
.insert_at_boundaries(input1, &tokenizer, special_tokens, false)
.unwrap();
// Should have cached at special token boundaries
@@ -384,7 +385,7 @@ mod tests {
let input = "<|im_start|>user\nHi<|im_end|>";
cache
.insert_at_boundaries(input, &tokenizer, special_tokens)
.insert_at_boundaries(input, &tokenizer, special_tokens, false)
.unwrap();
// Should cache at <|im_start|> boundary (has suffix left)
@@ -405,7 +406,7 @@ mod tests {
let input = "<|im_start|>system\nYou are a helpful AI assistant that provides detailed and accurate responses.<|im_end|><|im_start|>user\nHello there! How are you today? Can you help me understand how tokenization works in language models?<|im_end|><|im_start|>assistant\nI'm doing well, thank you! I'd be happy to explain tokenization. Tokenization is the process of breaking text into smaller units called tokens.<|im_end|>";
cache
.insert_at_boundaries(input, &tokenizer, special_tokens)
.insert_at_boundaries(input, &tokenizer, special_tokens, false)
.unwrap();
// Should have multiple entries at special token boundaries
@@ -432,7 +433,7 @@ mod tests {
let input = "<|im_start|>system\nYou are a helpful assistant that provides detailed answers.<|im_end|><|im_start|>user\nHello there! How are you today?<|im_end|>";
cache
.insert_at_boundaries(input, &tokenizer, special_tokens)
.insert_at_boundaries(input, &tokenizer, special_tokens, false)
.unwrap();
// Try to find match
@@ -454,7 +455,7 @@ mod tests {
let input = "<|im_start|>system\nYou are a helpful assistant that provides clear and detailed responses.<|im_end|><|im_start|>user\nHello there!<|im_end|>";
cache
.insert_at_boundaries(input, &tokenizer, special_tokens)
.insert_at_boundaries(input, &tokenizer, special_tokens, false)
.unwrap();
assert!(!cache.is_empty());
@@ -476,7 +477,7 @@ mod tests {
// Insert first conversation
let input1 = "<|im_start|>system\nYou are a helpful assistant specialized in mathematics.<|im_end|><|im_start|>user\nCan you explain calculus to me?<|im_end|><|im_start|>assistant\nCertainly! Calculus is a branch of mathematics that studies continuous change.<|im_end|><|eot_id|>";
cache
.insert_at_boundaries(input1, &tokenizer, special_tokens)
.insert_at_boundaries(input1, &tokenizer, special_tokens, false)
.unwrap();
// Access the first entry to update its timestamp
@@ -486,7 +487,7 @@ mod tests {
// Insert second conversation
let input2 = "<|im_start|>system\nYou are a helpful assistant specialized in physics.<|im_end|><|im_start|>user\nWhat is quantum mechanics?<|im_end|><|im_start|>assistant\nQuantum mechanics is the fundamental theory describing nature at atomic and subatomic scales.<|im_end|><|eot_id|>";
cache
.insert_at_boundaries(input2, &tokenizer, special_tokens)
.insert_at_boundaries(input2, &tokenizer, special_tokens, false)
.unwrap();
// Access the second entry to make it more recent
@@ -496,7 +497,7 @@ mod tests {
// Insert third conversation (should trigger eviction of oldest)
let input3 = "<|im_start|>system\nYou are a helpful assistant specialized in chemistry.<|im_end|><|im_start|>user\nExplain the periodic table to me please.<|im_end|><|im_start|>assistant\nThe periodic table is a tabular arrangement of chemical elements organized by atomic number and electron configuration.<|im_end|><|eot_id|>";
cache
.insert_at_boundaries(input3, &tokenizer, special_tokens)
.insert_at_boundaries(input3, &tokenizer, special_tokens, false)
.unwrap();
// Verify cache didn't exceed max memory
+19 -12
View File
@@ -162,8 +162,11 @@ impl CachedTokenizer {
}
impl Encoder for CachedTokenizer {
fn encode(&self, input: &str) -> Result<Encoding> {
fn encode(&self, input: &str, add_special_tokens: bool) -> Result<Encoding> {
// L0 cache lookup (exact match) - returns Arc<Encoding> for zero-copy
// Note: L0 cache doesn't distinguish by add_special_tokens flag
// This is acceptable for the current use case where embeddings always use true
// and chat always uses false with different input content
if let Some(l0) = &self.l0 {
if let Some(cached) = l0.get(input) {
// Unwrap the Arc - since Encoding is Clone, we can return the inner value
@@ -185,7 +188,7 @@ impl Encoder for CachedTokenizer {
// We have a prefix match - tokenize the suffix
let suffix = &input[prefix_len..];
if !suffix.is_empty() {
let suffix_encoding = self.inner.encode(suffix)?;
let suffix_encoding = self.inner.encode(suffix, add_special_tokens)?;
// Merge prefix tokens + suffix tokens
// Safe because we're splitting at special token boundaries
@@ -205,7 +208,7 @@ impl Encoder for CachedTokenizer {
}
// Full tokenization (both L0 and L1 miss)
let encoding = self.inner.encode(input)?;
let encoding = self.inner.encode(input, add_special_tokens)?;
// Cache in L0
if let Some(l0) = &self.l0 {
@@ -220,17 +223,21 @@ impl Encoder for CachedTokenizer {
.iter()
.map(|s| s.as_str())
.collect();
let _ = l1.insert_at_boundaries(input, self.inner.as_ref(), &tokens);
let _ =
l1.insert_at_boundaries(input, self.inner.as_ref(), &tokens, add_special_tokens);
// Ignore errors in cache insertion - cache is best-effort
}
Ok(encoding)
}
fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
fn encode_batch(&self, inputs: &[&str], add_special_tokens: bool) -> Result<Vec<Encoding>> {
// Process each input in parallel, leveraging thread-safe caches
// This maintains the parallelism from the underlying HuggingFaceTokenizer
inputs.par_iter().map(|&input| self.encode(input)).collect()
inputs
.par_iter()
.map(|&input| self.encode(input, add_special_tokens))
.collect()
}
}
@@ -276,10 +283,10 @@ mod tests {
let input = "Hello world";
// First call - miss
let result1 = cached.encode(input).unwrap();
let result1 = cached.encode(input, false).unwrap();
// Second call - hit
let result2 = cached.encode(input).unwrap();
let result2 = cached.encode(input, false).unwrap();
// Results should be identical
assert_eq!(result1.token_ids(), result2.token_ids());
@@ -304,8 +311,8 @@ mod tests {
let input = "Hello world";
// Both calls should work even without cache
let result1 = cached.encode(input).unwrap();
let result2 = cached.encode(input).unwrap();
let result1 = cached.encode(input, false).unwrap();
let result2 = cached.encode(input, false).unwrap();
assert_eq!(result1.token_ids(), result2.token_ids());
@@ -320,7 +327,7 @@ mod tests {
let inputs = vec!["Hello", "world", "Hello"]; // "Hello" repeated
let results = cached.encode_batch(&inputs).unwrap();
let results = cached.encode_batch(&inputs, false).unwrap();
assert_eq!(results.len(), 3);
@@ -330,7 +337,7 @@ mod tests {
// After batch processing, cache should be populated
// Subsequent calls should hit the cache
let _ = cached.encode("Hello").unwrap();
let _ = cached.encode("Hello", false).unwrap();
let stats = cached.cache_stats().unwrap();
// Should have at least 1 hit from the call above (cache was populated by batch)
+1 -1
View File
@@ -452,7 +452,7 @@ mod tests {
assert!(tokenizer.vocab_size() > 0);
let text = "Hello, world!";
let encoding = tokenizer.encode(text).unwrap();
let encoding = tokenizer.encode(text, false).unwrap();
let decoded = tokenizer.decode(encoding.token_ids(), false).unwrap();
assert_eq!(decoded, text);
}
+120 -35
View File
@@ -1,7 +1,8 @@
use std::collections::HashMap;
use anyhow::{Error, Result};
use tokenizers::tokenizer::Tokenizer as HfTokenizer;
use tokenizers::{processors::template::TemplateProcessing, tokenizer::Tokenizer as HfTokenizer};
use tracing::debug;
use super::{
chat_template::{
@@ -38,7 +39,7 @@ impl HuggingFaceTokenizer {
file_path: &str,
chat_template_path: Option<&str>,
) -> Result<Self> {
let tokenizer = HfTokenizer::from_file(file_path)
let mut tokenizer = HfTokenizer::from_file(file_path)
.map_err(|e| Error::msg(format!("Failed to load tokenizer: {}", e)))?;
// Extract special tokens
@@ -51,14 +52,19 @@ impl HuggingFaceTokenizer {
.map(|(token, &id)| (id, token.clone()))
.collect();
// Load chat template
let chat_template = if let Some(template_path) = chat_template_path {
// Load from specified .jinja file
Self::load_chat_template_from_file(template_path)?
} else {
// Try to load from tokenizer_config.json
Self::load_chat_template(file_path)
};
// Load chat template and tokenizer config
let (chat_template, add_bos_token, add_eos_token) =
if let Some(template_path) = chat_template_path {
// Load from specified .jinja file
(
Self::load_chat_template_from_file(template_path)?,
None,
None,
)
} else {
// Try to load from tokenizer_config.json
Self::load_chat_template_and_config(file_path)
};
// Detect content format once at initialization
let content_format = if let Some(ref template) = chat_template {
@@ -67,6 +73,25 @@ impl HuggingFaceTokenizer {
ChatTemplateContentFormat::String // Default if no template
};
// Configure post_processor based on tokenizer_config.json (matches Python transformers)
// Only modify when at least one setting is explicitly true
let needs_eos = add_eos_token == Some(true);
let needs_bos = match add_bos_token {
Some(true) => true,
Some(false) => false,
// Not set: preserve existing behavior from tokenizer.json
None => needs_eos && Self::tokenizer_adds_special_tokens(&tokenizer),
};
if needs_bos || needs_eos {
if let Some(post_processor) =
Self::build_post_processor(needs_bos, needs_eos, &special_tokens, &vocab)
{
debug!(needs_bos, needs_eos, "Configured post_processor");
tokenizer.with_post_processor(Some(post_processor));
}
}
Ok(HuggingFaceTokenizer {
tokenizer,
special_tokens,
@@ -77,6 +102,54 @@ impl HuggingFaceTokenizer {
})
}
/// Check if the tokenizer's post_processor adds special tokens (e.g., BOS)
fn tokenizer_adds_special_tokens(tokenizer: &HfTokenizer) -> bool {
tokenizer
.encode("", true)
.map(|enc| !enc.get_ids().is_empty())
.unwrap_or(false)
}
/// Build a TemplateProcessing post_processor (matches Python transformers' update_post_processor)
/// Template format: "{bos}:0 $A:0 {eos}:0" with optional BOS/EOS based on config
fn build_post_processor(
add_bos_token: bool,
add_eos_token: bool,
special_tokens: &SpecialTokens,
vocab: &HashMap<String, TokenIdType>,
) -> Option<TemplateProcessing> {
// Build template string exactly like Python:
// single = f"{(bos + ':0 ') if add_bos_token else ''}$A:0{(' ' + eos + ':0') if add_eos_token else ''}"
let mut template = String::with_capacity(32);
let mut tokens = Vec::with_capacity(2);
if add_bos_token {
let bos = special_tokens.bos_token.as_ref()?;
let bos_id = vocab.get(bos).copied()?;
template.push_str(bos);
template.push_str(":0 ");
tokens.push((bos.clone(), bos_id));
}
template.push_str("$A:0");
if add_eos_token {
let eos = special_tokens.eos_token.as_ref()?;
let eos_id = vocab.get(eos).copied()?;
template.push(' ');
template.push_str(eos);
template.push_str(":0");
tokens.push((eos.clone(), eos_id));
}
TemplateProcessing::builder()
.try_single(template.as_str())
.ok()?
.special_tokens(tokens)
.build()
.ok()
}
/// Create from an existing HuggingFace tokenizer
pub fn from_tokenizer(tokenizer: HfTokenizer) -> Self {
let special_tokens = Self::extract_special_tokens(&tokenizer);
@@ -130,21 +203,33 @@ impl HuggingFaceTokenizer {
}
}
/// Try to load chat template from tokenizer_config.json
fn load_chat_template(tokenizer_path: &str) -> Option<String> {
// Try to find tokenizer_config.json in the same directory
let path = std::path::Path::new(tokenizer_path);
let dir = path.parent()?;
let config_path = dir.join("tokenizer_config.json");
/// Load chat template and special token settings from tokenizer_config.json
/// Returns Option<bool> to distinguish between explicit false vs not set
fn load_chat_template_and_config(
tokenizer_path: &str,
) -> (Option<String>, Option<bool>, Option<bool>) {
(|| {
let path = std::path::Path::new(tokenizer_path);
let config_path = path.parent()?.join("tokenizer_config.json");
if config_path.exists() {
if let Ok(template) =
super::chat_template::load_chat_template_from_config(config_path.to_str()?)
{
return template;
if !config_path.exists() {
return None;
}
}
None
let config_str = config_path.to_str()?;
let content = std::fs::read_to_string(&config_path).ok()?;
let config: serde_json::Value = serde_json::from_str(&content).ok()?;
let chat_template = super::chat_template::load_chat_template_from_config(config_str)
.ok()
.flatten();
let add_bos_token = config.get("add_bos_token").and_then(|v| v.as_bool());
let add_eos_token = config.get("add_eos_token").and_then(|v| v.as_bool());
Some((chat_template, add_bos_token, add_eos_token))
})()
.unwrap_or((None, None, None))
}
/// Load chat template from a file (.jinja or .json containing Jinja)
@@ -217,23 +302,23 @@ impl HuggingFaceTokenizer {
}
impl Encoder for HuggingFaceTokenizer {
fn encode(&self, input: &str) -> Result<Encoding> {
fn encode(&self, input: &str, add_special_tokens: bool) -> Result<Encoding> {
self.tokenizer
.encode(input, false)
.encode(input, add_special_tokens)
.map_err(|e| Error::msg(format!("Encoding failed: {}", e)))
.map(|encoding| Encoding::Hf(Box::new(encoding)))
}
fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
let encodings = self
.tokenizer
.encode_batch(inputs.to_vec(), false)
.map_err(|e| Error::msg(format!("Batch encoding failed: {}", e)))?;
Ok(encodings
.into_iter()
.map(|e| Encoding::Hf(Box::new(e)))
.collect())
fn encode_batch(&self, inputs: &[&str], add_special_tokens: bool) -> Result<Vec<Encoding>> {
self.tokenizer
.encode_batch(inputs.to_vec(), add_special_tokens)
.map_err(|e| Error::msg(format!("Batch encoding failed: {}", e)))
.map(|encodings| {
encodings
.into_iter()
.map(|e| Encoding::Hf(Box::new(e)))
.collect()
})
}
}
+6 -3
View File
@@ -67,7 +67,7 @@ impl MockTokenizer {
}
impl Encoder for MockTokenizer {
fn encode(&self, input: &str) -> Result<Encoding> {
fn encode(&self, input: &str, _add_special_tokens: bool) -> Result<Encoding> {
// Simple word-based tokenization using the vocab
// Split by whitespace and look up each word (decoder adds spaces back)
let tokens: Vec<u32> = input
@@ -78,8 +78,11 @@ impl Encoder for MockTokenizer {
Ok(Encoding::Sp(tokens))
}
fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
inputs.iter().map(|input| self.encode(input)).collect()
fn encode_batch(&self, inputs: &[&str], add_special_tokens: bool) -> Result<Vec<Encoding>> {
inputs
.iter()
.map(|input| self.encode(input, add_special_tokens))
.collect()
}
}
+10 -4
View File
@@ -74,13 +74,19 @@ impl Tokenizer {
}
/// Direct encode method
pub fn encode(&self, input: &str) -> Result<Encoding> {
self.0.encode(input)
///
/// Set `add_special_tokens` to `true` for embeddings (to add BOS/EOS tokens configured in tokenizer_config.json),
/// or `false` for chat completion (where the chat template handles special tokens).
pub fn encode(&self, input: &str, add_special_tokens: bool) -> Result<Encoding> {
self.0.encode(input, add_special_tokens)
}
/// Direct batch encode method
pub fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
self.0.encode_batch(inputs)
///
/// Set `add_special_tokens` to `true` for embeddings (to add BOS/EOS tokens configured in tokenizer_config.json),
/// or `false` for chat completion (where the chat template handles special tokens).
pub fn encode_batch(&self, inputs: &[&str], add_special_tokens: bool) -> Result<Vec<Encoding>> {
self.0.encode_batch(inputs, add_special_tokens)
}
/// Direct decode method
+8 -5
View File
@@ -105,8 +105,11 @@ impl Sequence {
}
/// Append text to the sequence by encoding it
pub fn append_text(&mut self, input: &str) -> Result<()> {
let encoding = self.tokenizer.encode(input)?;
///
/// Set `add_special_tokens` to `true` for embeddings, or `false` for chat completion
/// where the chat template already handles special tokens.
pub fn append_text(&mut self, input: &str, add_special_tokens: bool) -> Result<()> {
let encoding = self.tokenizer.encode(input, add_special_tokens)?;
self.token_ids.extend(encoding.token_ids());
Ok(())
}
@@ -222,7 +225,7 @@ mod tests {
let tokenizer = Arc::new(MockTokenizer::new());
let mut seq = Sequence::new(tokenizer);
seq.append_text("Hello").unwrap();
seq.append_text("Hello", false).unwrap();
assert!(!seq.is_empty());
assert!(!seq.is_empty());
@@ -253,7 +256,7 @@ mod tests {
let tokenizer = Arc::new(MockTokenizer::new());
let mut seq = Sequence::new(tokenizer);
seq.append_text("Hello world").unwrap();
seq.append_text("Hello world", false).unwrap();
assert!(!seq.is_empty());
seq.clear();
@@ -268,7 +271,7 @@ mod tests {
let tokenizer = Arc::new(MockTokenizer::new());
let mut seq = Sequence::new(tokenizer);
seq.append_text("Test").unwrap();
seq.append_text("Test", false).unwrap();
let debug_str = format!("{:?}", seq);
assert!(debug_str.contains("Sequence"));
assert!(debug_str.contains("token count"));
+4 -4
View File
@@ -7,7 +7,7 @@ use super::*;
#[test]
fn test_mock_tokenizer_encode() {
let tokenizer = mock::MockTokenizer::new();
let encoding = tokenizer.encode("Hello world").unwrap();
let encoding = tokenizer.encode("Hello world", false).unwrap();
let token_ids = encoding.token_ids();
assert_eq!(token_ids, &[1, 2]); // "Hello" -> 1, "world" -> 2
}
@@ -37,7 +37,7 @@ fn test_tokenizer_wrapper() {
let mock_tokenizer = Arc::new(mock::MockTokenizer::new());
let tokenizer = Tokenizer::from_arc(mock_tokenizer);
let encoding = tokenizer.encode("Hello world").unwrap();
let encoding = tokenizer.encode("Hello world", false).unwrap();
assert_eq!(encoding.token_ids(), &[1, 2]);
let text = tokenizer.decode(&[1, 2], false).unwrap();
@@ -103,7 +103,7 @@ fn test_special_tokens() {
fn test_batch_encode() {
let tokenizer = mock::MockTokenizer::new();
let inputs = vec!["Hello", "world", "test"];
let encodings = tokenizer.encode_batch(&inputs).unwrap();
let encodings = tokenizer.encode_batch(&inputs, false).unwrap();
assert_eq!(encodings.len(), 3);
assert_eq!(encodings[0].token_ids(), &[1]); // "Hello" -> 1
@@ -124,7 +124,7 @@ fn test_thread_safety() {
let tokenizer_clone = tokenizer.clone();
thread::spawn(move || {
let text = "Hello test".to_string();
let encoding = tokenizer_clone.encode(&text).unwrap();
let encoding = tokenizer_clone.encode(&text, false).unwrap();
let decoded = tokenizer_clone.decode(encoding.token_ids(), false).unwrap();
assert!(decoded.contains("Hello") || decoded.contains("test"));
i
+9 -5
View File
@@ -132,13 +132,17 @@ impl TiktokenTokenizer {
}
impl Encoder for TiktokenTokenizer {
fn encode(&self, input: &str) -> Result<Encoding> {
fn encode(&self, input: &str, _add_special_tokens: bool) -> Result<Encoding> {
// tiktoken uses encode_ordinary which doesn't add special tokens
let tokens = self.tokenizer.encode_ordinary(input);
Ok(Encoding::Tiktoken(tokens))
}
fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
inputs.iter().map(|input| self.encode(input)).collect()
fn encode_batch(&self, inputs: &[&str], add_special_tokens: bool) -> Result<Vec<Encoding>> {
inputs
.iter()
.map(|input| self.encode(input, add_special_tokens))
.collect()
}
}
@@ -216,7 +220,7 @@ mod tests {
let tokenizer = TiktokenTokenizer::new(TiktokenModel::Cl100kBase).unwrap();
let text = "Hello, world!";
let encoding = tokenizer.encode(text).unwrap();
let encoding = tokenizer.encode(text, false).unwrap();
let decoded = tokenizer.decode(encoding.token_ids(), false).unwrap();
assert_eq!(decoded, text);
@@ -227,7 +231,7 @@ mod tests {
let tokenizer = TiktokenTokenizer::new(TiktokenModel::Cl100kBase).unwrap();
let texts = vec!["Hello", "World", "Test"];
let encodings = tokenizer.encode_batch(&texts).unwrap();
let encodings = tokenizer.encode_batch(&texts, false).unwrap();
assert_eq!(encodings.len(), 3);
for (i, encoding) in encodings.iter().enumerate() {
+2 -2
View File
@@ -10,8 +10,8 @@ pub type TokenIdType = u32;
/// Core encoding trait - separate from decoding for modularity
pub trait Encoder: Send + Sync {
fn encode(&self, input: &str) -> Result<Encoding>;
fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>>;
fn encode(&self, input: &str, add_special_tokens: bool) -> Result<Encoding>;
fn encode_batch(&self, inputs: &[&str], add_special_tokens: bool) -> Result<Vec<Encoding>>;
}
/// Core decoding trait - can be implemented independently
@@ -174,21 +174,25 @@ async fn test_cache_produces_identical_tokens() {
// Tokenize with base (no cache)
let base_encoding = base_tokenizer
.encode(turn)
.encode(turn, false)
.expect("Base tokenization failed");
let base_tokens = base_encoding.token_ids();
// Tokenize with L0-only
let l0_encoding = l0_tokenizer.encode(turn).expect("L0 tokenization failed");
let l0_encoding = l0_tokenizer
.encode(turn, false)
.expect("L0 tokenization failed");
let l0_tokens = l0_encoding.token_ids();
// Tokenize with L1-only
let l1_encoding = l1_tokenizer.encode(turn).expect("L1 tokenization failed");
let l1_encoding = l1_tokenizer
.encode(turn, false)
.expect("L1 tokenization failed");
let l1_tokens = l1_encoding.token_ids();
// Tokenize with L0+L1
let l0_l1_encoding = l0_l1_tokenizer
.encode(turn)
.encode(turn, false)
.expect("L0+L1 tokenization failed");
let l0_l1_tokens = l0_l1_encoding.token_ids();
@@ -397,13 +401,13 @@ async fn test_cache_correctness_with_edge_cases() {
test_count += 1;
let base_tokens = base_tokenizer
.encode(query)
.encode(query, false)
.expect("Base encoding failed")
.token_ids()
.to_vec();
let cached_tokens = cached_tokenizer
.encode(query)
.encode(query, false)
.expect("Cached encoding failed")
.token_ids()
.to_vec();
@@ -33,7 +33,7 @@ fn compute_hashes_for_tokenizer<E: Encoder>(tokenizer: &E, prompts: &[&str]) ->
.iter()
.map(|&prompt| {
tokenizer
.encode(prompt)
.encode(prompt, false)
.expect("Failed to encode prompt")
.get_hash()
})
@@ -63,7 +63,9 @@ fn test_tokenizer_encode_decode_lifecycle() {
.expect("Failed to load HuggingFace tokenizer");
for prompt in TEST_PROMPTS.iter() {
let encoding = tokenizer.encode(prompt).expect("Failed to encode prompt");
let encoding = tokenizer
.encode(prompt, false)
.expect("Failed to encode prompt");
let decoded = tokenizer
.decode(encoding.token_ids(), false)
@@ -82,10 +84,14 @@ fn test_sequence_operations() {
);
for prompt in TEST_PROMPTS.iter() {
let encoding = tokenizer.encode(prompt).expect("Failed to encode prompt");
let encoding = tokenizer
.encode(prompt, false)
.expect("Failed to encode prompt");
let mut sequence = Sequence::new(tokenizer.clone());
sequence.append_text(prompt).expect("Failed to append text");
sequence
.append_text(prompt, false)
.expect("Failed to append text");
assert_eq!(
sequence.len(),
@@ -123,7 +129,9 @@ fn test_decode_stream() {
);
for prompt in TEST_PROMPTS.iter() {
let encoding = tokenizer.encode(prompt).expect("Failed to encode prompt");
let encoding = tokenizer
.encode(prompt, false)
.expect("Failed to encode prompt");
let mut decoder = DecodeStream::new(tokenizer.clone(), &[], false);
let mut output = String::new();
@@ -148,11 +156,11 @@ fn test_long_sequence_incremental_decode_with_prefill() {
for (input_text, output_text) in LONG_TEST_PROMPTS.iter() {
let input_encoding = tokenizer
.encode(input_text)
.encode(input_text, false)
.expect("Failed to encode input");
let output_encoding = tokenizer
.encode(output_text)
.encode(output_text, false)
.expect("Failed to encode output");
let mut decoder = DecodeStream::new(tokenizer.clone(), input_encoding.token_ids(), false);
@@ -191,7 +199,7 @@ fn test_stop_sequence_decoder() {
let mut decoder = StopSequenceDecoder::new(tokenizer.clone(), config, false);
let encoding = tokenizer.encode(input).expect("Failed to encode");
let encoding = tokenizer.encode(input, false).expect("Failed to encode");
let mut output = String::new();
let mut stopped = false;
@@ -238,7 +246,9 @@ fn test_factory_creation() {
let tokenizer = factory::create_tokenizer(tokenizer_path.to_str().unwrap())
.expect("Failed to create tokenizer via factory");
let encoding = tokenizer.encode(TEST_PROMPTS[0]).expect("Failed to encode");
let encoding = tokenizer
.encode(TEST_PROMPTS[0], false)
.expect("Failed to encode");
let decoded = tokenizer
.decode(encoding.token_ids(), false)
@@ -254,7 +264,7 @@ fn test_batch_encoding() {
.expect("Failed to load tokenizer");
let encodings = tokenizer
.encode_batch(&TEST_PROMPTS)
.encode_batch(&TEST_PROMPTS, false)
.expect("Failed to batch encode");
assert_eq!(encodings.len(), TEST_PROMPTS.len());
@@ -300,7 +310,7 @@ fn test_thread_safety() {
let tokenizer_clone = tokenizer.clone();
thread::spawn(move || {
let encoding = tokenizer_clone
.encode(prompt)
.encode(prompt, false)
.expect("Failed to encode in thread");
let decoded = tokenizer_clone
.decode(encoding.token_ids(), false)