use std::sync::Mutex;
use crate::audio::vad::{
CHUNK_SAMPLES, InferError, ModelError, VadModel, VadModelOptions, VadState,
};
use crate::audio::whisper::audio::vad::VoiceActivityDetector;
pub const DEFAULT_SPEECH_THRESHOLD: f32 = 0.5;
#[derive(Debug, Default)]
struct DetectionLatch {
generation: u64,
last_error: Option<InferError>,
}
#[derive(Debug)]
pub struct SileroVad {
model: Mutex<VadModel>,
threshold: f32,
latch: Mutex<DetectionLatch>,
}
impl SileroVad {
pub const FRAME_LENGTH_SAMPLES: usize = CHUNK_SAMPLES;
pub fn load(path: impl AsRef<std::path::Path>) -> Result<Self, ModelError> {
Ok(Self::from_model(VadModel::load(path)?))
}
pub fn load_with(
path: impl AsRef<std::path::Path>,
options: VadModelOptions,
) -> Result<Self, ModelError> {
Ok(Self::from_model(VadModel::load_with(path, options)?))
}
pub fn from_model(model: VadModel) -> Self {
Self {
model: Mutex::new(model),
threshold: DEFAULT_SPEECH_THRESHOLD,
latch: Mutex::new(DetectionLatch::default()),
}
}
#[inline(always)]
pub const fn threshold(&self) -> f32 {
self.threshold
}
#[inline(always)]
pub const fn set_threshold(&mut self, threshold: f32) -> &mut Self {
self.threshold = threshold;
self
}
#[must_use]
#[inline(always)]
pub fn with_threshold(mut self, threshold: f32) -> Self {
self.threshold = threshold;
self
}
}
impl VoiceActivityDetector for SileroVad {
fn voice_activity(&self, samples: &[f32]) -> Vec<bool> {
let model = self.model.lock().expect("SileroVad model mutex poisoned");
let mut state = VadState::initial();
samples
.chunks(CHUNK_SAMPLES)
.map(
|chunk| match model.predict_chunk_with_state(chunk, &state) {
Ok((probability, next)) => {
state = next;
probability >= self.threshold
}
Err(error) => {
let mut latch = self
.latch
.lock()
.expect("SileroVad detection latch poisoned");
latch.generation = latch.generation.wrapping_add(1);
latch.last_error = Some(error);
false
}
},
)
.collect()
}
fn detection_generation(&self) -> u64 {
self
.latch
.lock()
.expect("SileroVad detection latch poisoned")
.generation
}
fn last_detection_error(&self) -> Option<Box<dyn std::error::Error + Send + Sync + 'static>> {
self
.latch
.lock()
.expect("SileroVad detection latch poisoned")
.last_error
.clone()
.map(|error| Box::new(error) as Box<dyn std::error::Error + Send + Sync + 'static>)
}
#[inline(always)]
fn frame_length_samples(&self) -> usize {
Self::FRAME_LENGTH_SAMPLES
}
}