use std::{
cell::Cell,
sync::atomic::{AtomicBool, Ordering},
time::Instant,
};
use crate::audio::whisper::{
backend::InferenceBackend,
constants::{DEFAULT_LANGUAGE_CODE, MAX_TOKEN_CONTEXT, SECONDS_PER_TIME_TOKEN, language_code},
decode::{
filter::{
LanguageLogitsFilter, LogitsFilter, SuppressBlankFilter, SuppressTokensFilter,
TimestampRulesFilter,
},
sampler::GreedyTokenSampler,
},
error::DecodeError,
options::DecodingOptions,
result::{DecodingResult, TranscriptionProgress, TranscriptionTimings},
text,
tokenizer::{SpecialTokens, WhisperTokenizer},
};
pub mod filter;
pub mod sampler;
#[cfg(test)]
mod tests;
fn log_sum_exp(logits: &[f32]) -> f32 {
let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
if !max.is_finite() {
return f32::NEG_INFINITY;
}
max + logits.iter().map(|&v| (v - max).exp()).sum::<f32>().ln()
}
pub type TranscriptionProgressCallback<'a> =
&'a (dyn Fn(&TranscriptionProgress) -> Option<bool> + Sync);
pub fn prefill_tokens(
options: &DecodingOptions,
tokenizer: &WhisperTokenizer,
is_multilingual: bool,
) -> Vec<u32> {
let special = tokenizer.special_tokens();
let mut tokens: Vec<u32> = vec![special.start_of_transcript_token()];
if is_multilingual {
let lang = if options.language().is_empty() {
DEFAULT_LANGUAGE_CODE
} else {
options.language()
};
let language_token = tokenizer
.token_to_id(&format!("<|{lang}|>"))
.unwrap_or_else(|| special.english_token());
tokens.push(language_token);
let task_token = tokenizer
.token_to_id(&format!("<|{}|>", options.task().as_str()))
.unwrap_or_else(|| special.transcribe_token());
tokens.push(task_token);
}
let timestamps_token = if options.without_timestamps() {
special.no_timestamps_token()
} else {
special.time_token_begin()
};
tokens.push(timestamps_token);
let prompt_tokens = options.prompt_tokens_slice();
if !prompt_tokens.is_empty() {
let max_prompt_len = MAX_TOKEN_CONTEXT / 2 - 1;
let start = prompt_tokens.len().saturating_sub(max_prompt_len);
let trimmed = prompt_tokens[start..]
.iter()
.copied()
.filter(|&t| t < special.special_token_begin());
let mut prefixed = Vec::with_capacity(1 + (prompt_tokens.len() - start) + tokens.len());
prefixed.push(special.start_of_previous_token());
prefixed.extend(trimmed);
prefixed.extend(tokens);
tokens = prefixed;
}
let prefix_tokens = options.prefix_tokens_slice();
if !prefix_tokens.is_empty() {
let start = prefix_tokens.len().saturating_sub(MAX_TOKEN_CONTEXT / 2);
tokens.extend(
prefix_tokens[start..]
.iter()
.copied()
.filter(|&t| t < special.special_token_begin()),
);
}
tokens
}
pub(crate) fn create_logits_filters(
options: &DecodingOptions,
sample_begin_prefilled: usize,
initial_prompt_len: usize,
special: &SpecialTokens,
is_multilingual: bool,
) -> Vec<Box<dyn LogitsFilter>> {
let mut filters: Vec<Box<dyn LogitsFilter>> = Vec::new();
if options.suppress_blank() {
filters.push(Box::new(SuppressBlankFilter::new(
special,
sample_begin_prefilled,
)));
}
if !options.suppress_tokens_slice().is_empty() {
let filtered: Vec<u32> = options
.suppress_tokens_slice()
.iter()
.copied()
.filter(|&t| t < special.special_token_begin())
.collect();
filters.push(Box::new(SuppressTokensFilter::new(filtered)));
}
if !options.without_timestamps() {
let max_initial_timestamp_index = options
.max_initial_timestamp()
.map(|seconds| (seconds / SECONDS_PER_TIME_TOKEN) as usize);
filters.push(Box::new(TimestampRulesFilter::new(
special,
initial_prompt_len,
max_initial_timestamp_index,
is_multilingual,
)));
}
filters
}
#[allow(clippy::too_many_arguments)] pub fn decode_text<B>(
backend: &B,
encoder_output: &B::EncoderOutput,
state: &mut B::DecoderState,
initial_prompt: &[u32],
sampler: &mut GreedyTokenSampler,
options: &DecodingOptions,
tokenizer: &WhisperTokenizer,
timings: &mut TranscriptionTimings,
early_stop: &AtomicBool,
observed_language_token: &Cell<Option<u32>>,
callback: Option<TranscriptionProgressCallback<'_>>,
) -> Result<DecodingResult, DecodeError>
where
B: InferenceBackend,
{
let special = *tokenizer.special_tokens();
let dims = backend.dims();
let prefilled_index = 0usize; let initial_prompt_index = initial_prompt.len();
let mut current_tokens: Vec<u32> = initial_prompt.to_vec();
let mut log_probs: Vec<f32> = vec![0.0; current_tokens.len()];
let mut next_token = *initial_prompt
.last()
.expect("initial_prompt must contain at least the start-of-transcript token");
let filters = create_logits_filters(
options,
prefilled_index,
initial_prompt_index,
&special,
dims.is_multilingual(),
);
let loop_count = options.sample_length().min(MAX_TOKEN_CONTEXT - 1);
let mut logits: Vec<f32> = Vec::with_capacity(dims.vocab());
let mut is_first_token_log_prob_too_low = false;
let mut first_token_log_prob = 0.0f32;
for token_index in prefilled_index..loop_count {
let is_prefill = token_index + 1 < initial_prompt_index; let is_last_prefill_token = token_index + 1 == initial_prompt_index; let is_first_token = token_index == prefilled_index;
if token_index < initial_prompt_index {
let prompt_is_timestamp = current_tokens[token_index] >= special.time_token_begin();
let model_predicted_timestamp = next_token >= special.time_token_begin();
if !(is_last_prefill_token && prompt_is_timestamp && model_predicted_timestamp) {
next_token = current_tokens[token_index];
} else {
current_tokens[token_index] = next_token;
}
}
let step_start = Instant::now();
backend.decode_step(next_token, token_index, encoder_output, state, &mut logits)?;
timings.set_decoding_predictions(
timings.decoding_predictions() + step_start.elapsed().as_secs_f64(),
);
let filter_start = Instant::now();
for filter in &filters {
filter
.filter(&mut logits, ¤t_tokens)
.map_err(DecodeError::UnmaskableToken)?; }
timings
.set_decoding_filtering(timings.decoding_filtering() + filter_start.elapsed().as_secs_f64());
let sample_start = Instant::now();
let sample = sampler.sample(&logits); timings
.set_decoding_sampling(timings.decoding_sampling() + sample_start.elapsed().as_secs_f64());
next_token = sample.token();
let next_token_log_prob = sample.logprob();
if is_first_token {
first_token_log_prob = next_token_log_prob;
if let Some(threshold) = options.first_token_logprob_threshold() {
is_first_token_log_prob_too_low = next_token_log_prob < threshold; }
}
if !is_prefill
&& observed_language_token.get().is_none()
&& tokenizer.all_language_tokens().contains(&next_token)
{
observed_language_token.set(Some(next_token));
}
let is_segment_completed = sample.completed()
|| current_tokens.len() >= MAX_TOKEN_CONTEXT - 1
|| is_first_token_log_prob_too_low;
if is_segment_completed {
timings.set_total_decoding_loops(timings.total_decoding_loops() + 1.0);
break;
}
if !is_prefill {
current_tokens.push(next_token); log_probs.push(next_token_log_prob);
}
backend.commit_alignment_row(state);
if let Some(callback) = callback {
let word_tokens: Vec<u32> = current_tokens
.iter()
.copied()
.filter(|&t| t < special.special_token_begin())
.collect();
let text_tokens = if options.skip_special_tokens() {
&word_tokens
} else {
¤t_tokens
};
let progress = TranscriptionProgress::new(
timings.clone(),
tokenizer.decode(text_tokens, false)?,
current_tokens.clone(),
)
.with_avg_logprob(log_probs.iter().sum::<f32>() / log_probs.len() as f32)
.with_compression_ratio(text::compression_ratio_of_tokens(¤t_tokens));
if callback(&progress) == Some(false) && !is_prefill {
early_stop.store(true, Ordering::Relaxed);
}
}
timings.set_total_decoding_loops(timings.total_decoding_loops() + 1.0);
if early_stop.load(Ordering::Relaxed) {
break; }
}
let observed_language = observed_language_token
.get()
.map(|token| {
let decoded = tokenizer.decode(&[token], false)?;
Ok::<String, DecodeError>(text::trim_special_token_chars(&decoded).to_string())
})
.transpose()?;
Ok(
finalize_decoding_result(
current_tokens,
log_probs,
first_token_log_prob,
sampler,
options,
tokenizer,
)?
.maybe_observed_language(observed_language)
.maybe_early_stopped(early_stop.load(Ordering::Relaxed)),
)
}
fn finalize_decoding_result(
mut current_tokens: Vec<u32>,
mut log_probs: Vec<f32>,
first_token_log_prob: f32,
sampler: &GreedyTokenSampler,
options: &DecodingOptions,
tokenizer: &WhisperTokenizer,
) -> Result<DecodingResult, DecodeError> {
let special = tokenizer.special_tokens();
sampler.finalize(&mut current_tokens, &mut log_probs);
let start_index = current_tokens
.iter()
.position(|&t| t == special.start_of_transcript_token())
.unwrap_or(0);
let end_index = current_tokens
.iter()
.position(|&t| t == special.end_token())
.unwrap_or(current_tokens.len());
let filtered_tokens = ¤t_tokens[start_index..=end_index];
let filtered_log_probs = &log_probs[start_index..=end_index];
let sum_log_probs: f32 = filtered_log_probs.iter().sum();
let avg_log_probs = sum_log_probs / filtered_log_probs.len() as f32;
let token_log_probs: Vec<(u32, f32)> = filtered_tokens
.iter()
.copied()
.zip(filtered_log_probs.iter().copied())
.collect();
let word_tokens: Vec<u32> = filtered_tokens
.iter()
.copied()
.filter(|&t| t < special.special_token_begin())
.collect();
let final_compression_ratio = text::compression_ratio_of_tokens(&word_tokens);
let temperature = (sampler.temperature() * 1000.0).round() / 1000.0;
let no_speech_prob = 0.0;
let (language, language_probs) = if !options.language().is_empty() {
(
options.language().to_string(),
vec![(options.language().to_string(), 0.0)],
)
} else {
match filtered_tokens
.iter()
.position(|&t| tokenizer.all_language_tokens().contains(&t))
{
Some(index) => {
let decoded = tokenizer.decode(&filtered_tokens[index..=index], false)?;
let lang = text::trim_special_token_chars(&decoded).to_string();
let prob = filtered_log_probs[index];
(lang.clone(), vec![(lang, prob)])
}
None => {
let lang = DEFAULT_LANGUAGE_CODE.to_string();
(lang.clone(), vec![(lang, 0.0)])
}
}
};
let text = tokenizer.decode(filtered_tokens, false)?;
Ok(
DecodingResult::new()
.with_language(language)
.with_language_probs(language_probs)
.with_tokens(filtered_tokens.to_vec())
.with_token_log_probs(token_log_probs)
.with_text(text)
.with_avg_logprob(avg_log_probs)
.with_no_speech_prob(no_speech_prob)
.with_temperature(temperature)
.with_compression_ratio(final_compression_ratio)
.with_first_token_log_prob(first_token_log_prob),
)
}
pub fn detect_language<B>(
backend: &B,
encoder_output: &B::EncoderOutput,
state: &mut B::DecoderState,
tokenizer: &WhisperTokenizer,
sampler: &mut GreedyTokenSampler,
timings: &mut TranscriptionTimings,
) -> Result<DecodingResult, DecodeError>
where
B: InferenceBackend,
{
let result = detect_language_probe(backend, encoder_output, state, tokenizer, sampler, timings);
backend.reset_decoder_state(state);
result
}
fn detect_language_probe<B>(
backend: &B,
encoder_output: &B::EncoderOutput,
state: &mut B::DecoderState,
tokenizer: &WhisperTokenizer,
sampler: &mut GreedyTokenSampler,
timings: &mut TranscriptionTimings,
) -> Result<DecodingResult, DecodeError>
where
B: InferenceBackend,
{
let special = *tokenizer.special_tokens();
let filter = LanguageLogitsFilter::new(tokenizer.all_language_tokens(), 0);
let mut logits: Vec<f32> = Vec::with_capacity(backend.dims().vocab());
let step_start = Instant::now();
backend.decode_step(
special.start_of_transcript_token(),
0,
encoder_output,
state,
&mut logits,
)?;
timings
.set_decoding_predictions(timings.decoding_predictions() + step_start.elapsed().as_secs_f64());
let prompt = [special.start_of_transcript_token()];
filter
.filter(&mut logits, &prompt)
.map_err(DecodeError::UnmaskableToken)?;
let sample_start = Instant::now();
let sample = sampler.sample(&logits);
timings.set_decoding_sampling(timings.decoding_sampling() + sample_start.elapsed().as_secs_f64());
let decoded = tokenizer.decode(&[sample.token()], false)?;
let trimmed = text::trim_special_token_chars(&decoded).to_string();
let mut language_probs: Vec<(String, f32)> = Vec::new();
if tokenizer.all_language_tokens().contains(&sample.token()) {
language_probs.push((trimmed.clone(), sample.logprob()));
}
let language = if language_code(&trimmed).is_some() {
trimmed
} else {
DEFAULT_LANGUAGE_CODE.to_string()
};
Ok(
DecodingResult::new()
.with_language(language)
.with_language_probs(language_probs),
)
}