use std::sync::Mutex;
use crate::audio::whisper::{
audio::{is_voice_detected, relative_energy, signal_energy},
backend::InferenceBackend,
constants::SAMPLE_RATE,
error::TranscribeError,
options::DecodingOptions,
result::{TranscriptionProgress, TranscriptionResult, TranscriptionSegment},
text::compression_ratio_of_tokens,
tokenizer::WhisperTokenizer,
transcribe::TranscribeTask,
};
pub mod agreement;
#[cfg(test)]
mod tests;
pub const ENERGY_FRAME_SAMPLES: usize = 1_600;
pub const RELATIVE_ENERGY_WINDOW: usize = 20;
#[derive(Debug, Clone, PartialEq)]
pub struct StreamState {
current_fallbacks: usize,
last_buffer_size: usize,
last_confirmed_segment_end_seconds: f32,
buffer_energy: Vec<f32>,
current_text: String,
confirmed_segments: Vec<TranscriptionSegment>,
unconfirmed_segments: Vec<TranscriptionSegment>,
unconfirmed_text: Vec<String>,
}
impl Default for StreamState {
fn default() -> Self {
Self::new()
}
}
impl StreamState {
pub const fn new() -> Self {
Self {
current_fallbacks: 0,
last_buffer_size: 0,
last_confirmed_segment_end_seconds: 0.0,
buffer_energy: Vec::new(),
current_text: String::new(),
confirmed_segments: Vec::new(),
unconfirmed_segments: Vec::new(),
unconfirmed_text: Vec::new(),
}
}
#[inline(always)]
pub const fn current_fallbacks(&self) -> usize {
self.current_fallbacks
}
#[inline(always)]
pub(crate) const fn set_current_fallbacks(&mut self, current_fallbacks: usize) -> &mut Self {
self.current_fallbacks = current_fallbacks;
self
}
#[inline(always)]
pub const fn last_buffer_size(&self) -> usize {
self.last_buffer_size
}
#[inline(always)]
pub(crate) const fn set_last_buffer_size(&mut self, last_buffer_size: usize) -> &mut Self {
self.last_buffer_size = last_buffer_size;
self
}
#[inline(always)]
pub const fn last_confirmed_segment_end_seconds(&self) -> f32 {
self.last_confirmed_segment_end_seconds
}
#[inline(always)]
pub(crate) const fn set_last_confirmed_segment_end_seconds(
&mut self,
last_confirmed_segment_end_seconds: f32,
) -> &mut Self {
self.last_confirmed_segment_end_seconds = last_confirmed_segment_end_seconds;
self
}
#[inline(always)]
pub const fn buffer_energy_slice(&self) -> &[f32] {
self.buffer_energy.as_slice()
}
#[inline(always)]
pub(crate) fn buffer_energy_mut(&mut self) -> &mut Vec<f32> {
&mut self.buffer_energy
}
#[inline(always)]
pub fn current_text(&self) -> &str {
self.current_text.as_str()
}
#[inline(always)]
pub(crate) fn set_current_text(&mut self, current_text: impl Into<String>) -> &mut Self {
self.current_text = current_text.into();
self
}
#[inline(always)]
pub const fn confirmed_segments_slice(&self) -> &[TranscriptionSegment] {
self.confirmed_segments.as_slice()
}
#[inline(always)]
pub(crate) const fn confirmed_segments_mut(&mut self) -> &mut Vec<TranscriptionSegment> {
&mut self.confirmed_segments
}
#[inline(always)]
pub const fn unconfirmed_segments_slice(&self) -> &[TranscriptionSegment] {
self.unconfirmed_segments.as_slice()
}
#[inline(always)]
pub(crate) fn set_unconfirmed_segments(
&mut self,
unconfirmed_segments: impl Into<Vec<TranscriptionSegment>>,
) -> &mut Self {
self.unconfirmed_segments = unconfirmed_segments.into();
self
}
#[inline(always)]
pub const fn unconfirmed_text_slice(&self) -> &[String] {
self.unconfirmed_text.as_slice()
}
#[inline(always)]
pub(crate) fn set_unconfirmed_text(
&mut self,
unconfirmed_text: impl Into<Vec<String>>,
) -> &mut Self {
self.unconfirmed_text = unconfirmed_text.into();
self
}
#[inline(always)]
pub(crate) const fn unconfirmed_text_mut(&mut self) -> &mut Vec<String> {
&mut self.unconfirmed_text
}
}
pub type StateChangeCallback<'a> = &'a (dyn Fn(&StreamState, &StreamState) + Sync);
pub const DEFAULT_REQUIRED_SEGMENTS_FOR_CONFIRMATION: usize = 2;
pub const DEFAULT_SILENCE_THRESHOLD: f32 = 0.3;
pub const DEFAULT_COMPRESSION_CHECK_WINDOW: usize = 60;
pub const DEFAULT_USE_VAD: bool = true;
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct AudioStreamOptions {
#[cfg_attr(
feature = "serde",
serde(default = "default_required_segments_for_confirmation")
)]
required_segments_for_confirmation: usize,
#[cfg_attr(feature = "serde", serde(default = "default_silence_threshold"))]
silence_threshold: f32,
#[cfg_attr(feature = "serde", serde(default = "default_compression_check_window"))]
compression_check_window: usize,
#[cfg_attr(feature = "serde", serde(default = "default_use_vad"))]
use_vad: bool,
}
#[cfg(feature = "serde")]
fn default_required_segments_for_confirmation() -> usize {
DEFAULT_REQUIRED_SEGMENTS_FOR_CONFIRMATION
}
#[cfg(feature = "serde")]
fn default_silence_threshold() -> f32 {
DEFAULT_SILENCE_THRESHOLD
}
#[cfg(feature = "serde")]
fn default_compression_check_window() -> usize {
DEFAULT_COMPRESSION_CHECK_WINDOW
}
#[cfg(feature = "serde")]
fn default_use_vad() -> bool {
DEFAULT_USE_VAD
}
impl Default for AudioStreamOptions {
fn default() -> Self {
Self::new()
}
}
impl AudioStreamOptions {
pub const fn new() -> Self {
Self {
required_segments_for_confirmation: DEFAULT_REQUIRED_SEGMENTS_FOR_CONFIRMATION,
silence_threshold: DEFAULT_SILENCE_THRESHOLD,
compression_check_window: DEFAULT_COMPRESSION_CHECK_WINDOW,
use_vad: DEFAULT_USE_VAD,
}
}
#[inline(always)]
pub const fn required_segments_for_confirmation(&self) -> usize {
self.required_segments_for_confirmation
}
#[must_use]
#[inline(always)]
pub const fn with_required_segments_for_confirmation(
mut self,
required_segments_for_confirmation: usize,
) -> Self {
self.set_required_segments_for_confirmation(required_segments_for_confirmation);
self
}
#[inline(always)]
pub const fn set_required_segments_for_confirmation(
&mut self,
required_segments_for_confirmation: usize,
) -> &mut Self {
self.required_segments_for_confirmation = required_segments_for_confirmation;
self
}
#[inline(always)]
pub const fn silence_threshold(&self) -> f32 {
self.silence_threshold
}
#[must_use]
#[inline(always)]
pub const fn with_silence_threshold(mut self, silence_threshold: f32) -> Self {
self.set_silence_threshold(silence_threshold);
self
}
#[inline(always)]
pub const fn set_silence_threshold(&mut self, silence_threshold: f32) -> &mut Self {
self.silence_threshold = silence_threshold;
self
}
#[inline(always)]
pub const fn compression_check_window(&self) -> usize {
self.compression_check_window
}
#[must_use]
#[inline(always)]
pub const fn with_compression_check_window(mut self, compression_check_window: usize) -> Self {
self.set_compression_check_window(compression_check_window);
self
}
#[inline(always)]
pub const fn set_compression_check_window(
&mut self,
compression_check_window: usize,
) -> &mut Self {
self.compression_check_window = compression_check_window;
self
}
#[inline(always)]
pub const fn use_vad(&self) -> bool {
self.use_vad
}
#[must_use]
#[inline(always)]
pub const fn with_use_vad(mut self) -> Self {
self.set_use_vad();
self
}
#[inline(always)]
pub const fn set_use_vad(&mut self) -> &mut Self {
self.use_vad = true;
self
}
#[must_use]
#[inline(always)]
pub const fn maybe_use_vad(mut self, use_vad: bool) -> Self {
self.update_use_vad(use_vad);
self
}
#[inline(always)]
pub const fn update_use_vad(&mut self, use_vad: bool) -> &mut Self {
self.use_vad = use_vad;
self
}
#[inline(always)]
pub const fn clear_use_vad(&mut self) -> &mut Self {
self.use_vad = false;
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display, derive_more::IsVariant)]
#[display("{}", self.as_str())]
#[non_exhaustive]
pub enum StreamUpdate {
AwaitingAudio,
AwaitingVoice,
Transcribed,
}
impl StreamUpdate {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::AwaitingAudio => "awaiting_audio",
Self::AwaitingVoice => "awaiting_voice",
Self::Transcribed => "transcribed",
}
}
}
pub fn should_stop_early(
progress: &TranscriptionProgress,
options: &DecodingOptions,
compression_check_window: usize,
) -> Option<bool> {
let tokens = progress.tokens_slice();
if tokens.len() > compression_check_window {
let window = &tokens[tokens.len() - compression_check_window..];
let compression_ratio = compression_ratio_of_tokens(window);
if compression_ratio > options.compression_ratio_threshold().unwrap_or(0.0) {
return Some(false);
}
}
if let Some(avg_logprob) = progress.avg_logprob()
&& let Some(threshold) = options.logprob_threshold()
&& avg_logprob < threshold
{
return Some(false);
}
None
}
#[derive(Debug, Clone, Default, PartialEq)]
pub(crate) struct EnergyTracker {
frames: Vec<(f32, f32)>,
consumed_samples: usize,
}
impl EnergyTracker {
fn absorb(&mut self, buffer: &[f32]) {
while self.consumed_samples + ENERGY_FRAME_SAMPLES <= buffer.len() {
let frame = &buffer[self.consumed_samples..self.consumed_samples + ENERGY_FRAME_SAMPLES];
let avg = signal_energy(frame);
let rel = if self.frames.is_empty() {
0.0
} else {
let start = self.frames.len().saturating_sub(RELATIVE_ENERGY_WINDOW);
let reference = self.frames[start..]
.iter()
.map(|&(_, avg)| avg)
.fold(f32::INFINITY, f32::min);
relative_energy(avg, reference)
};
self.frames.push((rel, avg));
self.consumed_samples += ENERGY_FRAME_SAMPLES;
}
}
fn relative_energies_from(&self, start: usize) -> Vec<f32> {
self.frames[start.min(self.frames.len())..]
.iter()
.map(|&(rel, _)| rel)
.collect()
}
}
const WAITING_FOR_SPEECH_TEXT: &str = "Waiting for speech...";
fn apply(
state: &mut StreamState,
callback: Option<StateChangeCallback<'_>>,
mutate: impl FnOnce(&mut StreamState),
) {
match callback {
Some(callback) => {
let old = state.clone();
mutate(state);
callback(&old, state);
}
None => mutate(state),
}
}
fn contains_subsequence(
haystack: &[TranscriptionSegment],
needle: &[TranscriptionSegment],
) -> bool {
!needle.is_empty()
&& haystack
.windows(needle.len())
.any(|window| window == needle)
}
fn on_progress_callback(
state: &mut StreamState,
callback: Option<StateChangeCallback<'_>>,
progress: &TranscriptionProgress,
) {
let fallbacks = progress.timings().total_decoding_fallbacks() as usize;
if progress.text().chars().count() < state.current_text().chars().count()
&& fallbacks == state.current_fallbacks()
{
let stale_text = state.current_text().to_string();
apply(state, callback, |s| {
s.unconfirmed_text_mut().push(stale_text);
});
}
let text = progress.text().to_string();
apply(state, callback, |s| {
s.set_current_text(text);
});
apply(state, callback, |s| {
s.set_current_fallbacks(fallbacks);
});
}
pub struct AudioStreamTranscriber<'ctx, B> {
backend: &'ctx B,
tokenizer: &'ctx WhisperTokenizer,
decoding_options: DecodingOptions,
stream_options: AudioStreamOptions,
state: StreamState,
state_callback: Option<StateChangeCallback<'ctx>>,
buffer: Vec<f32>,
energy: EnergyTracker,
}
impl<'ctx, B> AudioStreamTranscriber<'ctx, B> {
pub fn new(
backend: &'ctx B,
tokenizer: &'ctx WhisperTokenizer,
decoding_options: DecodingOptions,
) -> Self {
Self {
backend,
tokenizer,
decoding_options,
stream_options: AudioStreamOptions::new(),
state: StreamState::new(),
state_callback: None,
buffer: Vec::new(),
energy: EnergyTracker::default(),
}
}
#[must_use]
#[inline(always)]
pub const fn with_stream_options(mut self, stream_options: AudioStreamOptions) -> Self {
self.set_stream_options(stream_options);
self
}
#[inline(always)]
pub const fn set_stream_options(&mut self, stream_options: AudioStreamOptions) -> &mut Self {
self.stream_options = stream_options;
self
}
#[must_use]
#[inline(always)]
pub const fn with_state_callback(mut self, state_callback: StateChangeCallback<'ctx>) -> Self {
self.set_state_callback(state_callback);
self
}
#[inline(always)]
pub const fn set_state_callback(
&mut self,
state_callback: StateChangeCallback<'ctx>,
) -> &mut Self {
self.state_callback = Some(state_callback);
self
}
#[inline(always)]
pub const fn state(&self) -> &StreamState {
&self.state
}
#[inline(always)]
pub const fn buffer_len(&self) -> usize {
self.buffer.len()
}
#[inline(always)]
pub const fn stream_options(&self) -> &AudioStreamOptions {
&self.stream_options
}
}
impl<B> AudioStreamTranscriber<'_, B>
where
B: InferenceBackend,
{
pub fn push_samples(&mut self, samples: &[f32]) -> Result<StreamUpdate, TranscribeError> {
self.buffer.extend_from_slice(samples);
self.energy.absorb(&self.buffer);
let new_tail = self
.energy
.relative_energies_from(self.state.buffer_energy_slice().len());
apply(&mut self.state, self.state_callback, |s| {
s.buffer_energy_mut().extend_from_slice(&new_tail);
});
let next_buffer_samples = self
.buffer
.len()
.saturating_sub(self.state.last_buffer_size());
let next_buffer_seconds = next_buffer_samples as f32 / SAMPLE_RATE as f32;
if next_buffer_seconds <= 1.0 {
self.set_waiting_text_if_empty();
return Ok(StreamUpdate::AwaitingAudio);
}
if self.stream_options.use_vad()
&& !is_voice_detected(
self.state.buffer_energy_slice(),
next_buffer_seconds,
self.stream_options.silence_threshold(),
)
{
self.set_waiting_text_if_empty();
return Ok(StreamUpdate::AwaitingVoice);
}
let buffer_len = self.buffer.len();
apply(&mut self.state, self.state_callback, |s| {
s.set_last_buffer_size(buffer_len);
});
let transcription = self.transcribe_audio_samples()?;
apply(&mut self.state, self.state_callback, |s| {
s.set_current_text("");
});
apply(&mut self.state, self.state_callback, |s| {
s.set_unconfirmed_text(Vec::new());
});
let segments = transcription.segments_slice();
let required = self.stream_options.required_segments_for_confirmation();
if segments.len() > required {
let confirm_count = segments.len() - required;
let (confirmed, remaining) = segments.split_at(confirm_count);
if let Some(last_confirmed) = confirmed.last()
&& last_confirmed.end() > self.state.last_confirmed_segment_end_seconds()
{
let watermark = last_confirmed.end();
apply(&mut self.state, self.state_callback, |s| {
s.set_last_confirmed_segment_end_seconds(watermark);
});
if !contains_subsequence(self.state.confirmed_segments_slice(), confirmed) {
apply(&mut self.state, self.state_callback, |s| {
s.confirmed_segments_mut().extend_from_slice(confirmed);
});
}
}
apply(&mut self.state, self.state_callback, |s| {
s.set_unconfirmed_segments(remaining);
});
} else {
apply(&mut self.state, self.state_callback, |s| {
s.set_unconfirmed_segments(segments);
});
}
Ok(StreamUpdate::Transcribed)
}
fn set_waiting_text_if_empty(&mut self) {
if self.state.current_text().is_empty() {
apply(&mut self.state, self.state_callback, |s| {
s.set_current_text(WAITING_FOR_SPEECH_TEXT);
});
}
}
fn transcribe_audio_samples(&mut self) -> Result<TranscriptionResult, TranscribeError> {
let mut options = self.decoding_options.clone();
options.set_clip_timestamps(vec![self.state.last_confirmed_segment_end_seconds()]);
let compression_check_window = self.stream_options.compression_check_window();
let state_callback = self.state_callback;
let state_mutex = Mutex::new(std::mem::take(&mut self.state));
struct RestoreOnDrop<'a> {
slot: &'a mut StreamState,
mutex: &'a Mutex<StreamState>,
}
impl Drop for RestoreOnDrop<'_> {
fn drop(&mut self) {
let mut guard = self
.mutex
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*self.slot = std::mem::take(&mut guard);
}
}
let _restore = RestoreOnDrop {
slot: &mut self.state,
mutex: &state_mutex,
};
let progress_callback = |progress: &TranscriptionProgress| -> Option<bool> {
{
let mut guard = state_mutex.lock().expect("stream state mutex poisoned");
on_progress_callback(&mut guard, state_callback, progress);
}
should_stop_early(progress, &options, compression_check_window)
};
TranscribeTask::new(self.backend, self.tokenizer)
.with_progress_callback(&progress_callback)
.run(&self.buffer, &options)
}
}