//! `SentencePiece` tokenization implementation //! //! This module provides a complete `SentencePiece` tokenization implementation //! with both unigram and BPE algorithms, character coverage analysis, //! and full compatibility with the original Google `SentencePiece` library. use crate::{Result, TokenId, TokenizationStats, Tokenizer}; use indexmap::IndexMap; use parking_lot::RwLock; use regex::Regex; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::sync::Arc; use unicode_normalization::UnicodeNormalization; use unicode_segmentation::UnicodeSegmentation; /// `SentencePiece` configuration #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct SentencePieceConfig { /// Vocabulary size pub vocab_size: usize, /// Model type (unigram or bpe) pub model_type: SentencePieceModelType, /// Special tokens pub special_tokens: Vec, /// Character coverage (scaled by 10000 to avoid f64) pub character_coverage_scaled: u32, /// Input sentence size limit pub input_sentence_size: usize, /// Normalize input text pub normalize_text: bool, /// Add dummy prefix to input pub add_dummy_prefix: bool, /// Minimum subword frequency pub min_frequency: u64, /// Maximum subword length pub max_subword_length: usize, /// UNK surface forms pub unk_surface: String, } impl Default for SentencePieceConfig { fn default() -> Self { Self { vocab_size: 8000, model_type: SentencePieceModelType::Unigram, special_tokens: vec![ "[UNK]".to_string(), "[BOS]".to_string(), "[EOS]".to_string(), "[PAD]".to_string(), ], character_coverage_scaled: 9995, // 0.9995 * 10000 input_sentence_size: 10000, normalize_text: true, add_dummy_prefix: true, min_frequency: 2, max_subword_length: 16, unk_surface: "▁".to_string(), } } } impl SentencePieceConfig { /// Get character coverage as f64 #[must_use] pub fn character_coverage(&self) -> f64 { f64::from(self.character_coverage_scaled) / 10000.0 } } /// `SentencePiece` model type #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub enum SentencePieceModelType { /// Unigram language model Unigram, /// Byte-Pair Encoding Bpe, } /// Subword piece with frequency and score #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct SubwordPiece { /// The subword text pub piece: String, /// Token ID pub id: TokenId, /// Frequency in training data pub frequency: u64, /// Log probability score (for unigram model) pub score: f64, /// Whether this is a prefix piece (starts with ▁) pub is_prefix: bool, } /// Character statistics for coverage analysis #[derive(Debug, Clone)] struct CharacterStats { character: char, frequency: u64, cumulative_frequency: f64, } /// `SentencePiece` tokenizer with full implementation #[derive(Debug, Serialize, Deserialize)] pub struct SentencePieceTokenizer { config: SentencePieceConfig, /// Vocabulary mapping piece -> `SubwordPiece` vocab: IndexMap, /// ID to piece mapping id_to_piece: IndexMap, /// Trie for efficient prefix matching #[serde(skip)] prefix_trie: Option, /// Precompiled normalization regex #[serde(skip)] normalization_regex: Option, /// Statistics #[serde(skip)] stats: Arc>, } /// Prefix trie for efficient tokenization #[derive(Debug)] struct PrefixTrie { children: HashMap, piece: Option, } impl PrefixTrie { fn new() -> Self { Self { children: HashMap::new(), piece: None, } } fn insert(&mut self, piece: &str, subword_piece: SubwordPiece) { let mut node = self; for ch in piece.chars() { node = node.children.entry(ch).or_insert_with(Self::new); } node.piece = Some(subword_piece); } fn find_prefixes(&self, text: &str, start: usize) -> Vec<&SubwordPiece> { let mut results = Vec::new(); let mut node = self; let chars: Vec = text.chars().collect(); for i in start..chars.len() { if let Some(next_node) = node.children.get(&chars[i]) { node = next_node; if let Some(ref piece) = node.piece { results.push(piece); } } else { break; } } results } } /// Viterbi lattice node for unigram decoding #[derive(Debug, Clone)] struct ViterbiNode { piece_id: TokenId, start: usize, end: usize, score: f64, previous: Option, } impl SentencePieceTokenizer { /// Create new `SentencePiece` tokenizer #[must_use] pub fn new(config: SentencePieceConfig) -> Self { let mut tokenizer = Self { config, vocab: IndexMap::new(), id_to_piece: IndexMap::new(), prefix_trie: None, normalization_regex: None, stats: Arc::new(RwLock::new(TokenizationStats::default())), }; tokenizer.initialize_regex(); tokenizer } /// Create from pre-trained vocabulary pub fn from_vocab( config: SentencePieceConfig, vocab: IndexMap, ) -> Result { let id_to_piece = vocab .iter() .map(|(piece, subword)| (subword.id, piece.clone())) .collect(); let mut tokenizer = Self { config, vocab, id_to_piece, prefix_trie: None, normalization_regex: None, stats: Arc::new(RwLock::new(TokenizationStats::default())), }; tokenizer.initialize_regex(); tokenizer.build_prefix_trie(); Ok(tokenizer) } /// Initialize normalization regex fn initialize_regex(&mut self) { if self.config.normalize_text { // Basic normalization regex - in production this would be more comprehensive let regex_str = r"\s+"; self.normalization_regex = Regex::new(regex_str).ok(); } } /// Build prefix trie for efficient tokenization fn build_prefix_trie(&mut self) { let mut trie = PrefixTrie::new(); for (piece, subword_piece) in &self.vocab { trie.insert(piece, subword_piece.clone()); } self.prefix_trie = Some(trie); } /// Normalize text according to `SentencePiece` rules fn normalize(&self, text: &str) -> String { let mut normalized = text.to_string(); if self.config.normalize_text { // Unicode normalization normalized = normalized.nfc().collect::(); // Whitespace normalization if let Some(ref regex) = self.normalization_regex { normalized = regex.replace_all(&normalized, " ").to_string(); } // Trim whitespace normalized = normalized.trim().to_string(); } // Add dummy prefix if configured if self.config.add_dummy_prefix { normalized = format!("▁{normalized}"); } normalized } /// Tokenize using unigram algorithm async fn tokenize_unigram(&self, text: &str) -> Result> { let normalized = self.normalize(text); if normalized.is_empty() { return Ok(vec![]); } // Build lattice using Viterbi algorithm let lattice = self.build_viterbi_lattice(&normalized)?; // Find best path through lattice let best_path = self.find_best_path(&lattice)?; Ok(best_path) } /// Build Viterbi lattice for unigram decoding fn build_viterbi_lattice(&self, text: &str) -> Result> { let chars: Vec = text.chars().collect(); let text_len = chars.len(); let mut lattice = Vec::new(); // Initialize with BOS node lattice.push(ViterbiNode { piece_id: self.get_bos_id(), start: 0, end: 0, score: 0.0, previous: None, }); // Build lattice forward for pos in 0..text_len { if let Some(ref trie) = self.prefix_trie { let prefixes = trie.find_prefixes(text, pos); for prefix in prefixes { let end_pos = pos + prefix.piece.chars().count(); if end_pos <= text_len { // Find best previous node let mut best_score = f64::NEG_INFINITY; let mut best_prev = None; for (i, node) in lattice.iter().enumerate() { if node.end == pos { let score = node.score + prefix.score; if score > best_score { best_score = score; best_prev = Some(i); } } } lattice.push(ViterbiNode { piece_id: prefix.id, start: pos, end: end_pos, score: best_score, previous: best_prev, }); } } } } // Add EOS node let mut best_score = f64::NEG_INFINITY; let mut best_prev = None; for (i, node) in lattice.iter().enumerate() { if node.end == text_len && node.score > best_score { best_score = node.score; best_prev = Some(i); } } lattice.push(ViterbiNode { piece_id: self.get_eos_id(), start: text_len, end: text_len, score: best_score, previous: best_prev, }); Ok(lattice) } /// Find best path through Viterbi lattice fn find_best_path(&self, lattice: &[ViterbiNode]) -> Result> { if lattice.is_empty() { return Ok(vec![]); } let mut path = Vec::new(); let mut current = lattice.len() - 1; // Start from EOS while let Some(prev_idx) = lattice[current].previous { if lattice[current].piece_id != self.get_bos_id() && lattice[current].piece_id != self.get_eos_id() { path.push(lattice[current].piece_id); } current = prev_idx; } path.reverse(); Ok(path) } /// Get BOS token ID fn get_bos_id(&self) -> TokenId { self.vocab .get("[BOS]") .or_else(|| self.vocab.get("")) .map_or(1, |piece| piece.id) } /// Get EOS token ID fn get_eos_id(&self) -> TokenId { self.vocab .get("[EOS]") .or_else(|| self.vocab.get("")) .map_or(2, |piece| piece.id) } /// Get UNK token ID fn get_unk_id(&self) -> TokenId { self.vocab .get("[UNK]") .or_else(|| self.vocab.get("")) .map_or(0, |piece| piece.id) } /// Tokenize using BPE algorithm (simplified) async fn tokenize_bpe(&self, text: &str) -> Result> { let normalized = self.normalize(text); let mut tokens = Vec::new(); let unk_id = self.get_unk_id(); // Simple word-level BPE tokenization for word in normalized.split_whitespace() { if let Some(piece) = self.vocab.get(word) { tokens.push(piece.id); } else { // Try to split into smaller pieces let chars: Vec = word.chars().collect(); for ch in chars { let ch_str = ch.to_string(); if let Some(piece) = self.vocab.get(&ch_str) { tokens.push(piece.id); } else { tokens.push(unk_id); } } } } Ok(tokens) } /// Analyze character coverage in training data #[must_use] pub fn analyze_character_coverage(texts: &[String], coverage: f64) -> Vec { let mut char_freq: HashMap = HashMap::new(); let mut total_chars = 0u64; // Count character frequencies for text in texts { for ch in text.chars() { *char_freq.entry(ch).or_insert(0) += 1; total_chars += 1; } } // Sort by frequency let mut char_stats: Vec = char_freq .into_iter() .map(|(ch, freq)| CharacterStats { character: ch, frequency: freq, cumulative_frequency: 0.0, }) .collect(); char_stats.sort_by(|a, b| b.frequency.cmp(&a.frequency)); // Calculate cumulative frequency let mut cumulative = 0u64; for stat in &mut char_stats { cumulative += stat.frequency; stat.cumulative_frequency = cumulative as f64 / total_chars as f64; } // Select characters up to coverage threshold char_stats .into_iter() .take_while(|stat| stat.cumulative_frequency <= coverage) .map(|stat| stat.character) .collect() } } #[async_trait::async_trait] impl Tokenizer for SentencePieceTokenizer { async fn encode(&self, text: &str) -> Result> { let start_time = std::time::Instant::now(); let token_ids = match self.config.model_type { SentencePieceModelType::Unigram => self.tokenize_unigram(text).await?, SentencePieceModelType::Bpe => self.tokenize_bpe(text).await?, }; // Update statistics { let mut stats = self.stats.write(); stats.token_count = token_ids.len(); stats.processing_time_ms = start_time.elapsed().as_millis() as u64; } Ok(token_ids) } async fn decode(&self, token_ids: &[TokenId]) -> Result { let mut pieces = Vec::new(); for &token_id in token_ids { if let Some(piece) = self.id_to_piece.get(&token_id) { pieces.push(piece.clone()); } else { pieces.push(self.config.unk_surface.clone()); } } // Join pieces and handle prefix markers let mut result = pieces.join(""); // Remove dummy prefix marker if present if result.starts_with('▁') { result = result[1..].to_string(); } // Replace remaining prefix markers with spaces result = result.replace('▁', " "); Ok(result) } fn vocab_size(&self) -> usize { self.vocab.len() } fn get_stats(&self) -> TokenizationStats { self.stats.read().clone() } fn supports(&self, _text: &str) -> bool { true // SentencePiece supports all text } } /// `SentencePiece` trainer for learning subword vocabularies #[derive(Debug)] pub struct SentencePieceTrainer { config: SentencePieceConfig, char_coverage: Vec, subword_candidates: HashMap, } impl SentencePieceTrainer { /// Create a new `SentencePiece` trainer #[must_use] pub fn new(config: SentencePieceConfig) -> Self { Self { config, char_coverage: Vec::new(), subword_candidates: HashMap::new(), } } /// Add training text pub fn add_text(&mut self, text: &str) -> Result<()> { // Extract subword candidates based on character n-grams for n in 1..=self.config.max_subword_length { for ngram in text.graphemes(true).collect::>().windows(n) { let candidate = ngram.join(""); *self.subword_candidates.entry(candidate).or_insert(0) += 1; } } Ok(()) } /// Train the `SentencePiece` model pub fn train(&mut self, training_texts: &[String]) -> Result { // Analyze character coverage self.char_coverage = SentencePieceTokenizer::analyze_character_coverage( training_texts, self.config.character_coverage(), ); // Collect and add training text for text in training_texts { self.add_text(text)?; } // Filter candidates by frequency let filtered_candidates: HashMap = self .subword_candidates .iter() .filter(|(_, freq)| **freq >= self.config.min_frequency) .map(|(k, v)| (k.clone(), *v)) .collect(); // Build vocabulary let vocab = self.build_vocabulary(filtered_candidates)?; SentencePieceTokenizer::from_vocab(self.config.clone(), vocab) } /// Build final vocabulary from candidates fn build_vocabulary( &self, candidates: HashMap, ) -> Result> { let mut vocab = IndexMap::new(); let mut next_id = 0; // Add special tokens for token in &self.config.special_tokens { let piece = SubwordPiece { piece: token.clone(), id: next_id, frequency: u64::MAX, // Special tokens have max frequency score: 0.0, is_prefix: false, }; vocab.insert(token.clone(), piece); next_id += 1; } // Sort candidates by frequency (highest first) let mut sorted_candidates: Vec<_> = candidates.into_iter().collect(); sorted_candidates.sort_by(|a, b| b.1.cmp(&a.1)); // Add top candidates up to vocab size let remaining_slots = self.config.vocab_size.saturating_sub(vocab.len()); for (piece, frequency) in sorted_candidates.into_iter().take(remaining_slots) { let score = (frequency as f64).ln(); // Simple log probability let is_prefix = piece.starts_with('▁'); let subword_piece = SubwordPiece { piece: piece.clone(), id: next_id, frequency, score, is_prefix, }; vocab.insert(piece, subword_piece); next_id += 1; } Ok(vocab) } } #[cfg(test)] mod tests { use super::*; #[test] fn sentencepiece_config_default() { let config = SentencePieceConfig::default(); assert_eq!(config.vocab_size, 8000); assert_eq!(config.character_coverage(), 0.9995); assert!(config.normalize_text); assert!(config.add_dummy_prefix); assert_eq!(config.min_frequency, 2); assert_eq!(config.max_subword_length, 16); } #[test] fn subword_piece_creation() { let piece = SubwordPiece { piece: "▁hello".to_string(), id: 42, frequency: 100, score: 4.605, // ln(100) is_prefix: true, }; assert_eq!(piece.piece, "▁hello"); assert_eq!(piece.id, 42); assert_eq!(piece.frequency, 100); assert!((piece.score - 4.605).abs() < 0.001); assert!(piece.is_prefix); } #[tokio::test] async fn sentencepiece_tokenizer_creation() { let config = SentencePieceConfig::default(); let tokenizer = SentencePieceTokenizer::new(config); assert_eq!(tokenizer.vocab_size(), 0); // Empty vocab initially assert!(tokenizer.supports("any text")); } #[test] fn character_coverage_analysis() { let texts = vec![ "hello world".to_string(), "test text".to_string(), "more examples".to_string(), ]; let coverage_chars = SentencePieceTokenizer::analyze_character_coverage(&texts, 1.0); assert!(!coverage_chars.is_empty()); assert!(coverage_chars.contains(&'e')); // Common character assert!(coverage_chars.contains(&' ')); // Space } #[test] fn sentencepiece_trainer_creation() { let config = SentencePieceConfig::default(); let trainer = SentencePieceTrainer::new(config); assert_eq!(trainer.subword_candidates.len(), 0); assert_eq!(trainer.char_coverage.len(), 0); } #[test] fn sentencepiece_trainer_add_text() { let config = SentencePieceConfig::default(); let mut trainer = SentencePieceTrainer::new(config); trainer.add_text("hello world").unwrap(); assert!(!trainer.subword_candidates.is_empty()); // Should contain character-level candidates assert!(trainer.subword_candidates.contains_key("h")); assert!(trainer.subword_candidates.contains_key("e")); assert!(trainer.subword_candidates.contains_key("l")); } #[test] fn sentencepiece_training_smoke_test() { let config = SentencePieceConfig { vocab_size: 100, min_frequency: 1, ..SentencePieceConfig::default() }; let mut trainer = SentencePieceTrainer::new(config); let training_texts = vec![ "hello world test".to_string(), "more training data".to_string(), "additional examples".to_string(), ]; let tokenizer = trainer.train(&training_texts).unwrap(); assert!(tokenizer.vocab_size() > 4); // More than just special tokens assert!(tokenizer.vocab_size() <= 100); // Respects limit } #[test] fn text_normalization() { let config = SentencePieceConfig { normalize_text: true, add_dummy_prefix: true, ..SentencePieceConfig::default() }; let tokenizer = SentencePieceTokenizer::new(config); let normalized = tokenizer.normalize(" hello world "); assert!(normalized.starts_with('▁')); assert!(!normalized.contains(" ")); // Multiple spaces should be normalized } #[tokio::test] async fn encode_decode_consistency_with_vocab() { let config = SentencePieceConfig::default(); let mut vocab = IndexMap::new(); // Add some test vocabulary vocab.insert( "[UNK]".to_string(), SubwordPiece { piece: "[UNK]".to_string(), id: 0, frequency: u64::MAX, score: 0.0, is_prefix: false, }, ); vocab.insert( "▁hello".to_string(), SubwordPiece { piece: "▁hello".to_string(), id: 1, frequency: 100, score: 4.605, is_prefix: true, }, ); vocab.insert( "▁world".to_string(), SubwordPiece { piece: "▁world".to_string(), id: 2, frequency: 50, score: 3.912, is_prefix: true, }, ); let tokenizer = SentencePieceTokenizer::from_vocab(config, vocab).unwrap(); // Test basic functionality let tokens = tokenizer.encode("hello world").await.unwrap(); assert!(!tokens.is_empty()); let decoded = tokenizer.decode(&tokens).await.unwrap(); assert!(!decoded.is_empty()); } }