pub mod sentencepiece;
mod vocab;
pub use sentencepiece::SentencePieceTokenizer;
pub use vocab::{special_tokens, MergeRule, Vocabulary};
use crate::error::{WhisperError, WhisperResult};
fn bytes_to_singleton_tokens(bytes: &[u8]) -> Vec<Vec<u8>> {
bytes.iter().map(|&b| vec![b]).collect()
}
#[derive(Debug, Clone)]
pub struct BpeTokenizer {
vocab: Vocabulary,
}
impl BpeTokenizer {
#[must_use]
pub fn new(vocab: Vocabulary) -> Self {
Self { vocab }
}
#[must_use]
pub fn from_vocabulary(vocab: Vocabulary) -> Self {
Self::new(vocab)
}
#[must_use]
pub fn with_base_tokens() -> Self {
Self {
vocab: Vocabulary::with_base_tokens(),
}
}
pub fn encode(&self, text: &str) -> WhisperResult<Vec<u32>> {
if text.is_empty() {
return Ok(vec![]);
}
let tokens = self.apply_bpe_merges(text.as_bytes());
self.tokens_to_ids(&tokens)
}
fn apply_bpe_merges(&self, bytes: &[u8]) -> Vec<Vec<u8>> {
let mut tokens: Vec<Vec<u8>> = bytes_to_singleton_tokens(bytes);
while let Some((idx, merged_bytes)) = self.find_best_merge(&tokens) {
let second = tokens.remove(idx + 1);
tokens[idx].extend_from_slice(&second);
debug_assert_eq!(tokens[idx], merged_bytes);
}
tokens
}
fn tokens_to_ids(&self, tokens: &[Vec<u8>]) -> WhisperResult<Vec<u32>> {
tokens
.iter()
.map(|token_bytes| {
self.vocab.get_id(token_bytes).ok_or_else(|| {
WhisperError::Tokenizer(format!(
"unknown token: {:?}",
String::from_utf8_lossy(token_bytes)
))
})
})
.collect()
}
fn find_best_merge(&self, tokens: &[Vec<u8>]) -> Option<(usize, Vec<u8>)> {
if tokens.len() < 2 {
return None;
}
let mut best_priority = usize::MAX;
let mut best_idx = None;
let mut best_merged = None;
for i in 0..tokens.len() - 1 {
let first = &tokens[i];
let second = &tokens[i + 1];
if let Some(priority) = self.vocab.merge_priority(first, second) {
if priority < best_priority {
best_priority = priority;
best_idx = Some(i);
let mut merged = first.clone();
merged.extend_from_slice(second);
best_merged = Some(merged);
}
}
}
best_idx.and_then(|idx| best_merged.map(|merged| (idx, merged)))
}
pub fn decode(&self, tokens: &[u32]) -> WhisperResult<String> {
self.decode_with_options(tokens, false)
}
pub fn decode_with_options(&self, tokens: &[u32], skip_special: bool) -> WhisperResult<String> {
if tokens.is_empty() {
return Ok(String::new());
}
let filtered: Vec<u32> = if skip_special {
tokens
.iter()
.filter(|&&t| t < special_tokens::EOT)
.copied()
.collect()
} else {
tokens.to_vec()
};
self.vocab
.decode(&filtered)
.ok_or_else(|| WhisperError::Tokenizer("invalid token ID".into()))
}
#[must_use]
pub const fn vocab(&self) -> &Vocabulary {
&self.vocab
}
#[must_use]
pub fn vocab_size(&self) -> usize {
self.vocab.len()
}
}
impl Default for BpeTokenizer {
fn default() -> Self {
Self::with_base_tokens()
}
}
#[derive(Debug, Clone)]
pub enum Tokenizer {
Bpe(BpeTokenizer),
SentencePiece(SentencePieceTokenizer),
}
impl Tokenizer {
pub fn encode(&self, text: &str) -> WhisperResult<Vec<u32>> {
match self {
Self::Bpe(t) => t.encode(text),
Self::SentencePiece(t) => t.encode(text),
}
}
pub fn decode(&self, tokens: &[u32]) -> WhisperResult<String> {
match self {
Self::Bpe(t) => t.decode(tokens),
Self::SentencePiece(t) => t.decode(tokens),
}
}
#[must_use]
pub fn vocab_size(&self) -> usize {
match self {
Self::Bpe(t) => t.vocab_size(),
Self::SentencePiece(t) => t.vocab_size(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tokenizer_new() {
let vocab = Vocabulary::new();
let tokenizer = BpeTokenizer::new(vocab);
assert_eq!(tokenizer.vocab_size(), 0);
}
#[test]
fn test_tokenizer_with_base_tokens() {
let tokenizer = BpeTokenizer::with_base_tokens();
assert_eq!(tokenizer.vocab_size(), 256);
}
#[test]
fn test_tokenizer_default() {
let tokenizer = BpeTokenizer::default();
assert_eq!(tokenizer.vocab_size(), 256);
}
#[test]
fn test_tokenizer_encode_empty() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.encode("");
assert!(matches!(result, Ok(ref v) if v.is_empty()));
}
#[test]
fn test_tokenizer_encode_single_char() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.encode("a").expect("encode should succeed");
assert_eq!(result, vec![97]); }
#[test]
fn test_tokenizer_encode_ascii() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.encode("hi").expect("encode should succeed");
assert_eq!(result, vec![104, 105]); }
#[test]
fn test_tokenizer_encode_hello() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.encode("hello").expect("encode should succeed");
assert_eq!(result, vec![104, 101, 108, 108, 111]);
}
#[test]
fn test_tokenizer_encode_with_merges() {
let mut vocab = Vocabulary::with_base_tokens();
vocab.add_merge(vec![104], vec![105]);
let tokenizer = BpeTokenizer::new(vocab);
let result = tokenizer.encode("hi").expect("encode should succeed");
assert_eq!(result, vec![256]); }
#[test]
fn test_tokenizer_encode_multiple_merges() {
let mut vocab = Vocabulary::with_base_tokens();
vocab.add_merge(vec![104], vec![101]); vocab.add_merge(vec![108], vec![108]); vocab.add_merge(vec![108, 108], vec![111]);
let tokenizer = BpeTokenizer::new(vocab);
let result = tokenizer.encode("hello").expect("encode should succeed");
assert_eq!(result, vec![256, 258]);
}
#[test]
fn test_tokenizer_encode_space() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.encode(" ").expect("encode should succeed");
assert_eq!(result, vec![32]); }
#[test]
fn test_tokenizer_encode_newline() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.encode("\n").expect("encode should succeed");
assert_eq!(result, vec![10]); }
#[test]
fn test_tokenizer_encode_utf8_emoji() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.encode("😀").expect("encode should succeed");
assert_eq!(result, vec![240, 159, 152, 128]);
}
#[test]
fn test_tokenizer_encode_utf8_japanese() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.encode("こ").expect("encode should succeed");
assert_eq!(result, vec![227, 129, 147]);
}
#[test]
fn test_tokenizer_decode_empty() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.decode(&[]);
assert!(matches!(result, Ok(ref s) if s.is_empty()));
}
#[test]
fn test_tokenizer_decode_single_token() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.decode(&[97]).expect("decode should succeed");
assert_eq!(result, "a");
}
#[test]
fn test_tokenizer_decode_hello() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer
.decode(&[104, 101, 108, 108, 111])
.expect("decode should succeed");
assert_eq!(result, "hello");
}
#[test]
fn test_tokenizer_decode_with_merged_tokens() {
let mut vocab = Vocabulary::with_base_tokens();
vocab.add_merge(vec![104], vec![105]);
let tokenizer = BpeTokenizer::new(vocab);
let result = tokenizer.decode(&[256]).expect("decode should succeed");
assert_eq!(result, "hi");
}
#[test]
fn test_tokenizer_decode_invalid_token() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer.decode(&[50000]); assert!(result.is_err());
}
#[test]
fn test_tokenizer_decode_skips_special() {
let tokenizer = BpeTokenizer::with_base_tokens();
let result = tokenizer
.decode(&[104, 105, special_tokens::EOT])
.expect("decode should succeed");
assert_eq!(result, "hi");
}
#[test]
fn test_tokenizer_roundtrip_ascii() {
let tokenizer = BpeTokenizer::with_base_tokens();
let text = "Hello, World!";
let tokens = tokenizer.encode(text).expect("encode should succeed");
let decoded = tokenizer.decode(&tokens).expect("decode should succeed");
assert_eq!(decoded, text);
}
#[test]
fn test_tokenizer_roundtrip_with_merges() {
let mut vocab = Vocabulary::with_base_tokens();
vocab.add_merge(vec![116], vec![104]); vocab.add_merge(vec![116, 104], vec![101]);
let tokenizer = BpeTokenizer::new(vocab);
let text = "the";
let tokens = tokenizer.encode(text).expect("encode should succeed");
let decoded = tokenizer.decode(&tokens).expect("decode should succeed");
assert_eq!(decoded, text);
}
#[test]
fn test_tokenizer_roundtrip_utf8() {
let tokenizer = BpeTokenizer::with_base_tokens();
let text = "Hello 世界 🌍";
let tokens = tokenizer.encode(text).expect("encode should succeed");
let decoded = tokenizer.decode(&tokens).expect("decode should succeed");
assert_eq!(decoded, text);
}
#[test]
fn test_decode_with_options_skip_special() {
let tokenizer = BpeTokenizer::with_base_tokens();
let tokens = vec![104, 105, special_tokens::SOT, special_tokens::EOT];
let result = tokenizer
.decode_with_options(&tokens, true)
.expect("decode should succeed");
assert_eq!(result, "hi");
}
#[test]
fn test_decode_with_options_keep_special() {
let tokenizer = BpeTokenizer::with_base_tokens();
let tokens = vec![104, 105]; let result = tokenizer
.decode_with_options(&tokens, false)
.expect("decode should succeed");
assert_eq!(result, "hi");
}
#[test]
fn test_encode_never_empty_for_nonempty_input() {
let tokenizer = BpeTokenizer::with_base_tokens();
let inputs = ["a", " ", "\n", "hello", "🎉"];
for input in &inputs {
let tokens = tokenizer.encode(input).expect("encode should succeed");
assert!(
!tokens.is_empty(),
"encode('{}') should produce non-empty output",
input
);
}
}
#[test]
fn test_encode_produces_valid_tokens() {
let tokenizer = BpeTokenizer::with_base_tokens();
let text = "The quick brown fox jumps over the lazy dog.";
let tokens = tokenizer.encode(text).expect("encode should succeed");
for &token_id in &tokens {
assert!(
tokenizer.vocab().get_bytes(token_id).is_some(),
"token {} should be valid",
token_id
);
}
}
fn test_token_ids(count: usize) -> Vec<u32> {
(0..count).map(|i| (i % 128) as u32 + 32).collect()
}
mod property_tests {
use super::*;
use proptest::prelude::*;
proptest! {
#![proptest_config(ProptestConfig::with_cases(50))]
#[test]
fn property_encode_produces_valid_tokens(s in "[a-zA-Z0-9 ]{1,50}") {
let tokenizer = BpeTokenizer::with_base_tokens();
if let Ok(tokens) = tokenizer.encode(&s) {
for token in &tokens {
prop_assert!(
tokenizer.vocab().get_bytes(*token).is_some(),
"token {} from encoding '{}' should be valid",
token,
s
);
}
}
}
#[test]
fn property_decode_produces_string(token_count in 1usize..20) {
let tokenizer = BpeTokenizer::with_base_tokens();
let tokens = test_token_ids(token_count);
if let Ok(result) = tokenizer.decode(&tokens) {
prop_assert!(result.len() <= token_count * 4); }
}
#[test]
fn property_vocab_size_positive(_dummy in 0..1i32) {
let tokenizer = BpeTokenizer::with_base_tokens();
let size = tokenizer.vocab_size();
prop_assert!(size > 0, "vocab size should be positive");
let vocab = tokenizer.vocab();
prop_assert_eq!(size, vocab.len());
}
#[test]
fn property_special_tokens_reasonable(_dummy in 0..1i32) {
let tokenizer = BpeTokenizer::with_base_tokens();
let special_tokens = [
special_tokens::SOT,
special_tokens::EOT,
special_tokens::TRANSCRIBE,
special_tokens::TRANSLATE,
special_tokens::NO_TIMESTAMPS,
];
for token in special_tokens {
prop_assert!(
(token as usize) >= 50256 && (token as usize) < 60000,
"special token {} should be in Whisper's special token range (50256-60000)",
token
);
prop_assert!(tokenizer.vocab_size() > 0, "tokenizer should have vocabulary");
}
}
}
}
#[test]
fn test_from_vocabulary_alias() {
let vocab = Vocabulary::new();
let tokenizer = BpeTokenizer::from_vocabulary(vocab);
assert_eq!(tokenizer.vocab_size(), 0);
}
#[test]
fn test_tokenizer_enum_bpe_encode() {
let tokenizer = Tokenizer::Bpe(BpeTokenizer::with_base_tokens());
let result = tokenizer.encode("hi");
assert!(result.is_ok());
assert!(!result.expect("encode should succeed").is_empty());
}
#[test]
fn test_tokenizer_enum_bpe_decode() {
let tokenizer = Tokenizer::Bpe(BpeTokenizer::with_base_tokens());
let encoded = tokenizer.encode("ab").expect("encode should succeed");
let decoded = tokenizer.decode(&encoded);
assert!(decoded.is_ok());
assert_eq!(decoded.expect("decode should succeed"), "ab");
}
#[test]
fn test_tokenizer_enum_bpe_vocab_size() {
let tokenizer = Tokenizer::Bpe(BpeTokenizer::with_base_tokens());
assert_eq!(tokenizer.vocab_size(), 256);
}
#[test]
fn test_tokenizer_enum_sentencepiece_encode() {
let tokenizer = Tokenizer::SentencePiece(SentencePieceTokenizer::new(32768));
let result = tokenizer.encode("hello");
assert!(result.is_ok());
}
#[test]
fn test_tokenizer_enum_sentencepiece_decode() {
let tokenizer = Tokenizer::SentencePiece(SentencePieceTokenizer::new(32768));
let result = tokenizer.decode(&[]);
assert!(result.is_ok());
}
#[test]
fn test_tokenizer_enum_sentencepiece_vocab_size() {
let tokenizer = Tokenizer::SentencePiece(SentencePieceTokenizer::new(32768));
assert!(tokenizer.vocab_size() > 0);
}
#[test]
fn test_tokenizer_enum_debug() {
let tokenizer = Tokenizer::Bpe(BpeTokenizer::with_base_tokens());
let debug = format!("{tokenizer:?}");
assert!(debug.contains("Bpe"));
}
#[test]
fn test_tokenizer_enum_clone() {
let tokenizer = Tokenizer::Bpe(BpeTokenizer::with_base_tokens());
let cloned = tokenizer.clone();
assert_eq!(tokenizer.vocab_size(), cloned.vocab_size());
}
}