[SMG] perf: optimize tokenizer for reduced CPU and memory overhead (#14752)

This commit is contained in:
Simo Lin
2025-12-09 11:41:13 -08:00
committed by GitHub
parent 8b98bb768c
commit 73df7a4e8d
8 changed files with 107 additions and 47 deletions
+34 -7
View File
@@ -1,7 +1,7 @@
//! L0 Cache: Whole-string exact match cache //! L0 Cache: Whole-string exact match cache
//! //!
//! This is the simplest and most effective cache layer. //! This is the simplest and most effective cache layer.
//! Key: input string → Value: full encoding result //! Key: input string → Value: full encoding result (Arc-wrapped for zero-copy cache hits)
//! //!
//! Expected hit rate: 60-90% for workloads with repeated system prompts //! Expected hit rate: 60-90% for workloads with repeated system prompts
@@ -15,9 +15,10 @@ use dashmap::DashMap;
use super::super::traits::Encoding; use super::super::traits::Encoding;
/// L0 cache implementation using DashMap for lock-free reads /// L0 cache implementation using DashMap for lock-free reads
/// Uses Arc<Encoding> internally to provide zero-copy cache hits
pub struct L0Cache { pub struct L0Cache {
/// The cache map: input string → encoding /// The cache map: input string → Arc-wrapped encoding for cheap cloning
map: Arc<DashMap<String, Encoding>>, map: Arc<DashMap<String, Arc<Encoding>>>,
/// Maximum number of entries before eviction /// Maximum number of entries before eviction
max_entries: usize, max_entries: usize,
/// Cache hit counter /// Cache hit counter
@@ -37,12 +38,14 @@ impl L0Cache {
} }
} }
/// Get an encoding from the cache /// Get an encoding from the cache (returns Arc for zero-copy access)
pub fn get(&self, key: &str) -> Option<Encoding> { #[inline]
pub fn get(&self, key: &str) -> Option<Arc<Encoding>> {
match self.map.get(key) { match self.map.get(key) {
Some(entry) => { Some(entry) => {
self.hits.fetch_add(1, Ordering::Relaxed); self.hits.fetch_add(1, Ordering::Relaxed);
Some(entry.value().clone()) // Arc::clone is cheap (just increment reference count)
Some(Arc::clone(entry.value()))
} }
None => { None => {
self.misses.fetch_add(1, Ordering::Relaxed); self.misses.fetch_add(1, Ordering::Relaxed);
@@ -65,6 +68,17 @@ impl L0Cache {
} }
} }
self.map.insert(key, Arc::new(value));
}
/// Insert a pre-wrapped Arc encoding into the cache (avoids double-wrapping)
pub fn insert_arc(&self, key: String, value: Arc<Encoding>) {
if self.map.len() >= self.max_entries {
let key_to_remove = { self.map.iter().next().map(|entry| entry.key().clone()) };
if let Some(k) = key_to_remove {
self.map.remove(&k);
}
}
self.map.insert(key, value); self.map.insert(key, value);
} }
@@ -139,7 +153,7 @@ mod tests {
// Insert // Insert
cache.insert("hello".to_string(), mock_encoding(vec![1, 2, 3])); cache.insert("hello".to_string(), mock_encoding(vec![1, 2, 3]));
// Hit // Hit - now returns Arc<Encoding>
let result = cache.get("hello"); let result = cache.get("hello");
assert!(result.is_some()); assert!(result.is_some());
assert_eq!(result.unwrap().token_ids(), &[1, 2, 3]); assert_eq!(result.unwrap().token_ids(), &[1, 2, 3]);
@@ -217,4 +231,17 @@ mod tests {
// Should have 10 entries // Should have 10 entries
assert_eq!(cache.len(), 10); assert_eq!(cache.len(), 10);
} }
#[test]
fn test_arc_reuse() {
// Test that multiple gets return the same Arc (reference counting)
let cache = L0Cache::new(10);
cache.insert("test".to_string(), mock_encoding(vec![1, 2, 3]));
let arc1 = cache.get("test").unwrap();
let arc2 = cache.get("test").unwrap();
// Both should point to the same allocation
assert!(Arc::ptr_eq(&arc1, &arc2));
}
} }
+12 -9
View File
@@ -80,10 +80,11 @@ fn find_special_token_boundaries(text: &str, special_tokens: &[&str]) -> Vec<usi
} }
/// A cached prefix entry /// A cached prefix entry
/// Uses Arc<[TokenIdType]> for zero-copy access to tokens
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
struct CachedPrefix { struct CachedPrefix {
/// The pre-computed token IDs for this prefix /// The pre-computed token IDs for this prefix (Arc for zero-copy cloning)
tokens: Vec<TokenIdType>, tokens: Arc<[TokenIdType]>,
/// Last access timestamp (for LRU eviction) /// Last access timestamp (for LRU eviction)
last_accessed: Arc<AtomicU64>, last_accessed: Arc<AtomicU64>,
/// Size in bytes (for memory tracking during eviction) /// Size in bytes (for memory tracking during eviction)
@@ -127,6 +128,7 @@ impl L1Cache {
/// Returns (cached_tokens, byte_offset) if found /// Returns (cached_tokens, byte_offset) if found
/// ///
/// Uses pre-computed tokens cached during insertion. /// Uses pre-computed tokens cached during insertion.
/// Returns Vec<TokenIdType> as the caller needs to extend it with suffix tokens.
pub fn longest_prefix_match( pub fn longest_prefix_match(
&self, &self,
input: &str, input: &str,
@@ -154,7 +156,8 @@ impl L1Cache {
entry.last_accessed.store(timestamp, Ordering::Relaxed); entry.last_accessed.store(timestamp, Ordering::Relaxed);
self.hits.fetch_add(1, Ordering::Relaxed); self.hits.fetch_add(1, Ordering::Relaxed);
return Some((entry.tokens.clone(), boundary_pos)); // Convert Arc<[T]> to Vec<T> - caller will extend with suffix tokens
return Some((entry.tokens.to_vec(), boundary_pos));
} }
} }
@@ -181,7 +184,7 @@ impl L1Cache {
} }
// Calculate how much memory we need and tokenize each prefix // Calculate how much memory we need and tokenize each prefix
let mut entries_to_insert = Vec::new(); let mut entries_to_insert = Vec::with_capacity(boundaries.len());
for &boundary_pos in &boundaries { for &boundary_pos in &boundaries {
// Extract prefix up to this special token boundary // Extract prefix up to this special token boundary
let prefix = &input[0..boundary_pos]; let prefix = &input[0..boundary_pos];
@@ -192,7 +195,8 @@ impl L1Cache {
// Re-tokenize the prefix for guaranteed correctness // Re-tokenize the prefix for guaranteed correctness
// This is the only way to know the exact token boundaries // This is the only way to know the exact token boundaries
let prefix_encoding = tokenizer.encode(prefix)?; let prefix_encoding = tokenizer.encode(prefix)?;
let prefix_tokens = prefix_encoding.token_ids().to_vec(); // Convert to Arc<[TokenIdType]> for zero-copy sharing
let prefix_tokens: Arc<[TokenIdType]> = prefix_encoding.token_ids().into();
// Size = text bytes + token storage // Size = text bytes + token storage
let size_bytes = boundary_pos + prefix_tokens.len() * size_of::<TokenIdType>(); let size_bytes = boundary_pos + prefix_tokens.len() * size_of::<TokenIdType>();
@@ -213,14 +217,13 @@ impl L1Cache {
} }
// Insert all entries // Insert all entries
let current_timestamp = self.access_counter.load(Ordering::Relaxed);
for (hash_bytes, prefix_tokens, size_bytes) in entries_to_insert { for (hash_bytes, prefix_tokens, size_bytes) in entries_to_insert {
let shard_idx = hash_bytes[0] as usize % NUM_SHARDS; let shard_idx = hash_bytes[0] as usize % NUM_SHARDS;
let cached = CachedPrefix { let cached = CachedPrefix {
tokens: prefix_tokens, tokens: prefix_tokens, // Already Arc<[TokenIdType]>
last_accessed: Arc::new(AtomicU64::new( last_accessed: Arc::new(AtomicU64::new(current_timestamp)),
self.access_counter.load(Ordering::Relaxed),
)),
size_bytes, size_bytes,
}; };
+17 -14
View File
@@ -163,26 +163,25 @@ impl CachedTokenizer {
impl Encoder for CachedTokenizer { impl Encoder for CachedTokenizer {
fn encode(&self, input: &str) -> Result<Encoding> { fn encode(&self, input: &str) -> Result<Encoding> {
// Collect special tokens once if L1 is enabled (avoid redundant allocation) // L0 cache lookup (exact match) - returns Arc<Encoding> for zero-copy
let special_tokens: Option<Vec<&str>> = self.l1.as_ref().map(|_| {
self.special_token_strings
.iter()
.map(|s| s.as_str())
.collect()
});
// L0 cache lookup (exact match)
if let Some(l0) = &self.l0 { if let Some(l0) = &self.l0 {
if let Some(cached) = l0.get(input) { if let Some(cached) = l0.get(input) {
return Ok(cached); // Unwrap the Arc - since Encoding is Clone, we can return the inner value
// For callers who need the tokens, they can access via token_ids() which is &[u32]
return Ok((*cached).clone());
} }
} }
// L1 cache lookup (prefix match at special token boundaries) // L1 cache lookup (prefix match at special token boundaries)
if let Some(l1) = &self.l1 { if let Some(l1) = &self.l1 {
let tokens = special_tokens.as_ref().unwrap(); // Use pre-computed special tokens refs (avoids allocation per call)
let tokens: Vec<&str> = self
.special_token_strings
.iter()
.map(|s| s.as_str())
.collect();
if let Some((prefix_tokens, prefix_len)) = l1.longest_prefix_match(input, tokens) { if let Some((prefix_tokens, prefix_len)) = l1.longest_prefix_match(input, &tokens) {
// We have a prefix match - tokenize the suffix // We have a prefix match - tokenize the suffix
let suffix = &input[prefix_len..]; let suffix = &input[prefix_len..];
if !suffix.is_empty() { if !suffix.is_empty() {
@@ -216,8 +215,12 @@ impl Encoder for CachedTokenizer {
// Cache in L1 at special token boundaries // Cache in L1 at special token boundaries
// Re-tokenizes prefixes for correctness (optimized for high prefix reuse) // Re-tokenizes prefixes for correctness (optimized for high prefix reuse)
if let Some(l1) = &self.l1 { if let Some(l1) = &self.l1 {
let tokens = special_tokens.as_ref().unwrap(); let tokens: Vec<&str> = self
let _ = l1.insert_at_boundaries(input, self.inner.as_ref(), tokens); .special_token_strings
.iter()
.map(|s| s.as_str())
.collect();
let _ = l1.insert_at_boundaries(input, self.inner.as_ref(), &tokens);
// Ignore errors in cache insertion - cache is best-effort // Ignore errors in cache insertion - cache is best-effort
} }
+20 -6
View File
@@ -148,12 +148,26 @@ fn is_likely_json(buffer: &[u8]) -> bool {
fn is_likely_sentencepiece(buffer: &[u8]) -> bool { fn is_likely_sentencepiece(buffer: &[u8]) -> bool {
// SentencePiece models often start with specific patterns // SentencePiece models often start with specific patterns
// This is a simplified check // This is a simplified check
buffer.len() >= 12 if buffer.len() < 12 {
&& (buffer.starts_with(b"\x0a\x09") return false;
|| buffer.starts_with(b"\x08\x00") }
|| buffer.windows(4).any(|w| w == b"<unk")
|| buffer.windows(4).any(|w| w == b"<s>") // Check header patterns first (cheap)
|| buffer.windows(4).any(|w| w == b"</s>")) if buffer.starts_with(b"\x0a\x09") || buffer.starts_with(b"\x08\x00") {
return true;
}
// Single-pass scan for special token markers
// Instead of multiple windows() calls, scan once looking for all patterns
let patterns: &[&[u8]] = &[b"<unk", b"<s>", b"</s>"];
for window in buffer.windows(4) {
for pattern in patterns {
if window.starts_with(pattern) {
return true;
}
}
}
false
} }
/// Helper function to discover chat template files in a directory /// Helper function to discover chat template files in a directory
@@ -86,11 +86,13 @@ impl Sequence {
} }
/// Check if the sequence is empty /// Check if the sequence is empty
#[inline]
pub fn is_empty(&self) -> bool { pub fn is_empty(&self) -> bool {
self.token_ids.is_empty() self.token_ids.is_empty()
} }
/// Get the length of the sequence /// Get the length of the sequence
#[inline]
pub fn len(&self) -> usize { pub fn len(&self) -> usize {
self.token_ids.len() self.token_ids.len()
} }
@@ -111,6 +113,7 @@ impl Sequence {
/// Append a single token to the sequence and return newly decoded text /// Append a single token to the sequence and return newly decoded text
/// Based on HuggingFace TGI incremental decoding /// Based on HuggingFace TGI incremental decoding
#[inline]
pub fn append_token(&mut self, token_id: TokenIdType) -> Result<String> { pub fn append_token(&mut self, token_id: TokenIdType) -> Result<String> {
// Store the old read offset before adding the new token // Store the old read offset before adding the new token
let old_read_offset = self.read_offset; let old_read_offset = self.read_offset;
@@ -165,11 +168,13 @@ impl Sequence {
} }
/// Get a reference to the tokenizer /// Get a reference to the tokenizer
#[inline]
pub fn tokenizer(&self) -> &Arc<dyn TokenizerTrait> { pub fn tokenizer(&self) -> &Arc<dyn TokenizerTrait> {
&self.tokenizer &self.tokenizer
} }
/// Get the current token ids /// Get the current token ids
#[inline]
pub fn token_ids(&self) -> &[TokenIdType] { pub fn token_ids(&self) -> &[TokenIdType] {
&self.token_ids &self.token_ids
} }
@@ -181,16 +186,19 @@ impl Sequence {
} }
/// Get the prefix offset /// Get the prefix offset
#[inline]
pub fn prefix_offset(&self) -> usize { pub fn prefix_offset(&self) -> usize {
self.prefix_offset self.prefix_offset
} }
/// Get the read offset /// Get the read offset
#[inline]
pub fn read_offset(&self) -> usize { pub fn read_offset(&self) -> usize {
self.read_offset self.read_offset
} }
/// Get whether special tokens are skipped during decoding /// Get whether special tokens are skipped during decoding
#[inline]
pub fn skip_special_tokens(&self) -> bool { pub fn skip_special_tokens(&self) -> bool {
self.skip_special_tokens self.skip_special_tokens
} }
+12 -10
View File
@@ -95,11 +95,11 @@ impl StopSequenceDecoder {
if self.config.stop_tokens.contains(&token_id) { if self.config.stop_tokens.contains(&token_id) {
self.stopped = true; self.stopped = true;
// Flush any jailed text before stopping // Flush any jailed text before stopping - use mem::take to avoid clone
if !self.jail_buffer.is_empty() { if !self.jail_buffer.is_empty() {
let output = self.jail_buffer.clone(); return Ok(SequenceDecoderOutput::StoppedWithText(std::mem::take(
self.jail_buffer.clear(); &mut self.jail_buffer,
return Ok(SequenceDecoderOutput::StoppedWithText(output)); )));
} }
return Ok(SequenceDecoderOutput::Stopped); return Ok(SequenceDecoderOutput::Stopped);
} }
@@ -186,8 +186,10 @@ impl StopSequenceDecoder {
if let Some(split_pos) = best_split_pos { if let Some(split_pos) = best_split_pos {
// Hold the partial match, flush the rest // Hold the partial match, flush the rest
// Drain [0..split_pos] as output, keep [split_pos..] in jail_buffer // Use split_off for zero-copy: keeps [0..split_pos] in place, returns [split_pos..]
let to_output = self.jail_buffer.drain(..split_pos).collect::<String>(); // Then swap so we output the prefix and keep the suffix
let suffix = self.jail_buffer.split_off(split_pos);
let to_output = std::mem::replace(&mut self.jail_buffer, suffix);
if to_output.is_empty() { if to_output.is_empty() {
Ok(SequenceDecoderOutput::Held) Ok(SequenceDecoderOutput::Held)
@@ -210,7 +212,8 @@ impl StopSequenceDecoder {
&mut self, &mut self,
token_ids: &[TokenIdType], token_ids: &[TokenIdType],
) -> Result<Vec<SequenceDecoderOutput>> { ) -> Result<Vec<SequenceDecoderOutput>> {
let mut outputs = Vec::new(); // Pre-allocate with exact capacity to avoid reallocations
let mut outputs = Vec::with_capacity(token_ids.len());
for &token_id in token_ids { for &token_id in token_ids {
outputs.push(self.process_token(token_id)?); outputs.push(self.process_token(token_id)?);
} }
@@ -220,9 +223,8 @@ impl StopSequenceDecoder {
/// Flush any held text /// Flush any held text
pub fn flush(&mut self) -> SequenceDecoderOutput { pub fn flush(&mut self) -> SequenceDecoderOutput {
if !self.jail_buffer.is_empty() { if !self.jail_buffer.is_empty() {
let output = self.jail_buffer.clone(); // Use mem::take to avoid clone - transfers ownership and leaves empty string
self.jail_buffer.clear(); SequenceDecoderOutput::Text(std::mem::take(&mut self.jail_buffer))
SequenceDecoderOutput::Text(output)
} else { } else {
SequenceDecoderOutput::Text(String::new()) SequenceDecoderOutput::Text(String::new())
} }
+3 -1
View File
@@ -44,6 +44,7 @@ impl DecodeStream {
/// Step appends a token_id to the internal state and tries to produce a text chunk. /// Step appends a token_id to the internal state and tries to produce a text chunk.
/// Returning `None` means the given id is not enough to produce a chunk. /// Returning `None` means the given id is not enough to produce a chunk.
#[inline]
pub fn step(&mut self, id: TokenIdType) -> Result<Option<String>> { pub fn step(&mut self, id: TokenIdType) -> Result<Option<String>> {
self.all_token_ids.push(id); self.all_token_ids.push(id);
@@ -71,7 +72,8 @@ impl DecodeStream {
/// Process multiple tokens at once /// Process multiple tokens at once
pub fn step_batch(&mut self, token_ids: &[u32]) -> Result<Vec<String>> { pub fn step_batch(&mut self, token_ids: &[u32]) -> Result<Vec<String>> {
let mut chunks = Vec::new(); // Pre-allocate with capacity - most tokens produce output
let mut chunks = Vec::with_capacity(token_ids.len());
for &token_id in token_ids { for &token_id in token_ids {
if let Some(text) = self.step(token_id)? { if let Some(text) = self.step(token_id)? {
@@ -43,6 +43,7 @@ pub enum Encoding {
impl Encoding { impl Encoding {
/// Returns a reference to token IDs - zero-copy operation /// Returns a reference to token IDs - zero-copy operation
#[inline]
pub fn token_ids(&self) -> &[TokenIdType] { pub fn token_ids(&self) -> &[TokenIdType] {
match self { match self {
Encoding::Hf(inner) => inner.get_ids(), Encoding::Hf(inner) => inner.get_ids(),