use std::path::Path;
use unicode_categories::UnicodeCategories;
use crate::audio::whisper::{constants, error::TokenizerError, options::WordGrouping};
#[cfg(feature = "nl-recognizer")]
pub mod nl_recognizer;
const DEFAULT_WHITESPACE_TOKEN: u32 = 220;
const DEFAULT_SPECIAL_TOKEN_BEGIN: u32 = 50_257;
const DEFAULT_END_TOKEN: u32 = 50_257;
const DEFAULT_START_OF_PREVIOUS_TOKEN: u32 = 50_361;
const DEFAULT_START_OF_TRANSCRIPT_TOKEN: u32 = 50_258;
const DEFAULT_ENGLISH_TOKEN: u32 = 50_259;
const DEFAULT_TRANSCRIBE_TOKEN: u32 = 50_359;
const DEFAULT_TRANSLATE_TOKEN: u32 = 50_358;
const DEFAULT_NO_SPEECH_TOKEN: u32 = 50_362;
const DEFAULT_NO_TIMESTAMPS_TOKEN: u32 = 50_363;
const DEFAULT_TIME_TOKEN_BEGIN: u32 = 50_364;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SpecialTokens {
end_token: u32,
english_token: u32,
no_speech_token: u32,
no_timestamps_token: u32,
special_token_begin: u32,
start_of_previous_token: u32,
start_of_transcript_token: u32,
time_token_begin: u32,
transcribe_token: u32,
translate_token: u32,
whitespace_token: u32,
}
impl SpecialTokens {
fn probe(tokenizer: &tokenizers::Tokenizer) -> Self {
let end_token = tokenizer
.token_to_id("<|endoftext|>")
.unwrap_or(DEFAULT_END_TOKEN);
let english_token = tokenizer
.token_to_id("<|en|>")
.unwrap_or(DEFAULT_ENGLISH_TOKEN);
let no_speech_token = tokenizer
.token_to_id("<|nospeech|>")
.unwrap_or(DEFAULT_NO_SPEECH_TOKEN);
let no_timestamps_token = tokenizer
.token_to_id("<|notimestamps|>")
.unwrap_or(DEFAULT_NO_TIMESTAMPS_TOKEN);
let special_token_begin = tokenizer
.token_to_id("<|endoftext|>")
.unwrap_or(DEFAULT_SPECIAL_TOKEN_BEGIN);
let start_of_previous_token = tokenizer
.token_to_id("<|startofprev|>")
.unwrap_or(DEFAULT_START_OF_PREVIOUS_TOKEN);
let start_of_transcript_token = tokenizer
.token_to_id("<|startoftranscript|>")
.unwrap_or(DEFAULT_START_OF_TRANSCRIPT_TOKEN);
let time_token_begin = tokenizer
.token_to_id("<|0.00|>")
.unwrap_or(DEFAULT_TIME_TOKEN_BEGIN);
let transcribe_token = tokenizer
.token_to_id("<|transcribe|>")
.unwrap_or(DEFAULT_TRANSCRIBE_TOKEN);
let translate_token = tokenizer
.token_to_id("<|translate|>")
.unwrap_or(DEFAULT_TRANSLATE_TOKEN);
let whitespace_token = tokenizer
.token_to_id(" ")
.unwrap_or(DEFAULT_WHITESPACE_TOKEN);
Self {
end_token,
english_token,
no_speech_token,
no_timestamps_token,
special_token_begin,
start_of_previous_token,
start_of_transcript_token,
time_token_begin,
transcribe_token,
translate_token,
whitespace_token,
}
}
#[inline(always)]
pub const fn whisper_defaults() -> Self {
Self {
end_token: DEFAULT_END_TOKEN,
english_token: DEFAULT_ENGLISH_TOKEN,
no_speech_token: DEFAULT_NO_SPEECH_TOKEN,
no_timestamps_token: DEFAULT_NO_TIMESTAMPS_TOKEN,
special_token_begin: DEFAULT_SPECIAL_TOKEN_BEGIN,
start_of_previous_token: DEFAULT_START_OF_PREVIOUS_TOKEN,
start_of_transcript_token: DEFAULT_START_OF_TRANSCRIPT_TOKEN,
time_token_begin: DEFAULT_TIME_TOKEN_BEGIN,
transcribe_token: DEFAULT_TRANSCRIBE_TOKEN,
translate_token: DEFAULT_TRANSLATE_TOKEN,
whitespace_token: DEFAULT_WHITESPACE_TOKEN,
}
}
#[inline(always)]
pub const fn end_token(&self) -> u32 {
self.end_token
}
#[inline(always)]
pub const fn english_token(&self) -> u32 {
self.english_token
}
#[inline(always)]
pub const fn no_speech_token(&self) -> u32 {
self.no_speech_token
}
#[inline(always)]
pub const fn no_timestamps_token(&self) -> u32 {
self.no_timestamps_token
}
#[inline(always)]
pub const fn special_token_begin(&self) -> u32 {
self.special_token_begin
}
#[inline(always)]
pub const fn start_of_previous_token(&self) -> u32 {
self.start_of_previous_token
}
#[inline(always)]
pub const fn start_of_transcript_token(&self) -> u32 {
self.start_of_transcript_token
}
#[inline(always)]
pub const fn time_token_begin(&self) -> u32 {
self.time_token_begin
}
#[inline(always)]
pub const fn transcribe_token(&self) -> u32 {
self.transcribe_token
}
#[inline(always)]
pub const fn translate_token(&self) -> u32 {
self.translate_token
}
#[inline(always)]
pub const fn whitespace_token(&self) -> u32 {
self.whitespace_token
}
}
fn is_single_punctuation_scalar(s: &str) -> bool {
let trimmed = s.trim_matches(|c: char| c.is_separator_space() || c == '\u{0009}');
let mut chars = trimmed.chars();
match (chars.next(), chars.next()) {
(Some(c), None) => c.is_punctuation(),
_ => false,
}
}
fn config_boolean(value: &serde_json::Value) -> Option<bool> {
match value {
serde_json::Value::Bool(b) => Some(*b),
serde_json::Value::Number(n) => match n.as_i64() {
Some(int) => Some(int == 1),
None => {
const INT_MIN: f64 = -9_223_372_036_854_775_808.0; const PAST_INT_MAX: f64 = 9_223_372_036_854_775_808.0;
let float = n.as_f64()?;
(float > INT_MIN && float < PAST_INT_MAX && float.fract() == 0.0).then_some(float == 1.0)
}
},
serde_json::Value::String(s) => match s.to_lowercase().as_str() {
"true" | "t" | "1" => Some(true),
"false" | "f" | "0" => Some(false),
_ => None,
},
serde_json::Value::Null | serde_json::Value::Array(_) | serde_json::Value::Object(_) => None,
}
}
fn clean_up_tokenization_spaces_from(folder: &Path) -> bool {
let Some(config) = std::fs::read(folder.join("tokenizer_config.json"))
.ok()
.and_then(|bytes| serde_json::from_slice::<serde_json::Value>(&bytes).ok())
else {
return true;
};
config
.get("cleanUpTokenizationSpaces")
.or_else(|| config.get("clean_up_tokenization_spaces"))
.and_then(config_boolean)
.unwrap_or(true)
}
fn clean_up_tokenization(text: String) -> String {
let can_match = text
.as_bytes()
.windows(2)
.any(|pair| pair[0] == b' ' && matches!(pair[1], b'.' | b'?' | b'!' | b',' | b'\'' | b'n'));
if !can_match {
return text;
}
text
.replace(" .", ".")
.replace(" ?", "?")
.replace(" !", "!")
.replace(" ,", ",")
.replace(" ' ", "'")
.replace(" n't", "n't")
.replace(" 'm", "'m")
.replace(" 's", "'s")
.replace(" 've", "'ve")
.replace(" 're", "'re")
}
fn indexed_id_domain(tokenizer: &tokenizers::Tokenizer) -> usize {
tokenizer
.get_vocab(true)
.into_values()
.max()
.map_or(0, |largest| largest as usize + 1)
}
#[derive(Debug)]
pub struct WhisperTokenizer {
tokenizer: tokenizers::Tokenizer,
special_tokens: SpecialTokens,
vocab_size: usize,
clean_up_tokenization_spaces: bool,
language_table: Vec<(u32, &'static str)>,
language_ids: Vec<u32>,
}
impl WhisperTokenizer {
pub fn from_folder(folder: impl AsRef<Path>) -> Result<Self, TokenizerError> {
let folder = folder.as_ref();
let path = folder.join("tokenizer.json");
if !path.is_file() {
return Err(TokenizerError::FileNotFound(vec![path]));
}
let tokenizer = tokenizers::Tokenizer::from_file(&path)?;
let special_tokens = SpecialTokens::probe(&tokenizer);
let vocab_size = indexed_id_domain(&tokenizer);
let clean_up_tokenization_spaces = clean_up_tokenization_spaces_from(folder);
let mut language_table: Vec<(u32, &'static str)> = Vec::new();
for &(_, code) in constants::languages() {
let Some(id) = tokenizer.token_to_id(&format!("<|{code}|>")) else {
continue;
};
if id > special_tokens.special_token_begin
&& !language_table.iter().any(|&(existing, _)| existing == id)
{
language_table.push((id, code));
}
}
let language_ids: Vec<u32> = language_table.iter().map(|&(id, _)| id).collect();
Ok(Self {
tokenizer,
special_tokens,
vocab_size,
clean_up_tokenization_spaces,
language_table,
language_ids,
})
}
pub fn encode(&self, text: &str) -> Result<Vec<u32>, TokenizerError> {
Ok(self.tokenizer.encode(text, false)?.get_ids().to_vec())
}
pub fn decode(&self, ids: &[u32], skip_special: bool) -> Result<String, TokenizerError> {
let decoded = self.tokenizer.decode(ids, skip_special)?;
Ok(if self.clean_up_tokenization_spaces {
clean_up_tokenization(decoded)
} else {
decoded
})
}
#[inline(always)]
pub fn token_to_id(&self, token: &str) -> Option<u32> {
self.tokenizer.token_to_id(token)
}
#[inline(always)]
pub fn id_to_token(&self, id: u32) -> Option<String> {
self.tokenizer.id_to_token(id)
}
#[inline(always)]
pub const fn special_tokens(&self) -> &SpecialTokens {
&self.special_tokens
}
#[inline(always)]
pub const fn vocab_size(&self) -> usize {
self.vocab_size
}
#[inline(always)]
pub fn all_language_tokens(&self) -> &[u32] {
self.language_ids.as_slice()
}
pub fn language_for_token(&self, id: u32) -> Option<&'static str> {
self
.language_table
.iter()
.find(|&&(tid, _)| tid == id)
.map(|&(_, code)| code)
}
pub fn split_to_word_tokens(
&self,
tokens: &[u32],
language_code: &str,
grouping: WordGrouping,
) -> Result<Vec<(String, Vec<u32>)>, TokenizerError> {
let unicode_split = match grouping {
WordGrouping::FineGrained => {
matches!(language_code, "zh" | "ja" | "th" | "lo" | "my" | "yue")
}
WordGrouping::SwiftParity => matches!(language_code, "ja" | "th" | "lo" | "my"),
};
if unicode_split {
self.split_tokens_on_unicode(tokens)
} else {
self.split_tokens_on_spaces(tokens)
}
}
fn split_tokens_on_unicode(
&self,
tokens: &[u32],
) -> Result<Vec<(String, Vec<u32>)>, TokenizerError> {
let decoded_full = self.decode(tokens, false)?;
let mut words: Vec<(String, Vec<u32>)> = Vec::new();
let mut current_tokens: Vec<u32> = Vec::new();
for &token in tokens {
current_tokens.push(token);
let decoded = self.decode(¤t_tokens, false)?;
let has_unicode_in_full_string = decoded.find('\u{FFFD}').is_some_and(|offset| {
decoded_full
.get(offset..)
.and_then(|rest| rest.chars().next())
== Some('\u{FFFD}')
});
if !decoded.contains('\u{FFFD}') || has_unicode_in_full_string {
words.push((decoded, std::mem::take(&mut current_tokens)));
}
}
Ok(words)
}
fn split_tokens_on_spaces(
&self,
tokens: &[u32],
) -> Result<Vec<(String, Vec<u32>)>, TokenizerError> {
let subwords = self.split_tokens_on_unicode(tokens)?;
let mut words: Vec<(String, Vec<u32>)> = Vec::new();
for (subword, subword_tokens) in subwords {
let is_special = subword_tokens
.first()
.is_some_and(|&id| id >= self.special_tokens.special_token_begin);
let starts_with_space = subword.starts_with(' ');
let is_punctuation = is_single_punctuation_scalar(&subword);
if is_special || starts_with_space || is_punctuation || words.is_empty() {
words.push((subword, subword_tokens));
} else {
let last = words.len() - 1;
words[last].0.push_str(&subword);
words[last].1.extend(subword_tokens);
}
}
Ok(words)
}
}
#[cfg(test)]
mod tests;