//! Silero VAD v6 ONNX runtime integration (P1 VAD_002). //! //! ## Silero VAD v6 model I/O //! //! The v6 model (silero_vad.onnx from the v6.x releases) has a different //! interface from v4. Key changes: //! //! | Tensor | Shape | Dtype | Meaning | //! |---------|------------------|-------|--------------------------------------| //! | input | \[1, 576\] | f32 | 64-sample context + 512-sample frame | //! | state | \[2, 1, 128\] | f32 | LSTM state (carry across frames) | //! | sr | \[1\] | i64 | Sample rate (16000 or 8000) | //! | output | \[1, 1\] | f32 | Speech probability | //! | stateN | \[2, 1, 128\] | f32 | Updated LSTM state | //! //! Frame size: **512 samples at 16 kHz = 32 ms**. //! Context: **64 samples** prepended to each frame (last 64 samples of previous frame). //! Total input width: 512 + 64 = **576 samples**. //! //! ## Threading //! //! `SileroOnnxVad` is `Send`. The session is created once and reused — //! never re-created per callback (INV_007). //! //! ## Accumulation //! //! The capture pipeline delivers 10 ms frames (480 samples at 48 kHz → //! 160 samples at 16 kHz). Three 10 ms frames = 30 ms ≈ 32 ms. We //! accumulate 512 samples (32 ms at 16 kHz) before running inference. //! The last probability is held between inference calls so the state //! machine always has a value to work with. //! //! ## Fallback //! //! `try_new` returns `None` when the model file is missing, the ONNX //! Runtime is unavailable, or the platform is not iOS/macOS. The caller //! falls back to `WebRtcFallbackVad`. use super::{VadOutput, VoiceActivityDetector}; use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering}; use std::sync::Arc; use std::thread::JoinHandle; /// 16 kHz frame size for Silero VAD v6 (32 ms). pub const SILERO_FRAME_16K: usize = 512; /// Context size prepended to each frame (64 samples at 16 kHz). pub const SILERO_CONTEXT_16K: usize = 64; /// Total input width: context + frame. pub const SILERO_INPUT_WIDTH: usize = SILERO_CONTEXT_16K + SILERO_FRAME_16K; /// LSTM state size: 2 × 1 × 128 = 256 f32 values. pub const SILERO_STATE_SIZE: usize = 256; /// Maximum lag in 10 ms frames before the realtime callback treats /// the Silero worker as stale and falls back to the local WebRTC /// detector for that frame. pub const SILERO_MAX_STALE_FRAMES: u64 = 3; /// Silero VAD v6 ONNX backend. /// /// Operates at **16 kHz**, accumulating 32 ms frames (512 samples) /// before running inference. The caller is responsible for downsampling /// from 48 kHz before calling `process_10ms`. pub struct SileroOnnxVad { /// LSTM state [2, 1, 128] — persisted across frames. state: Box<[f32; SILERO_STATE_SIZE]>, /// Context ring: last 64 samples of the previous frame. context: Box<[f32; SILERO_CONTEXT_16K]>, /// Accumulation buffer for 16 kHz samples (fills to SILERO_FRAME_16K). accum: Vec, /// Last speech probability output (held between inference calls). last_probability: f32, /// Model path stored for diagnostics. model_path: String, /// Inner ONNX implementation (platform-specific). inner: SileroInner, } enum SileroInner { Onnx(OnnxSession), #[cfg(test)] Stub, } struct OnnxSession { session: ort::session::Session, } impl SileroOnnxVad { /// Attempt to load the Silero v6 ONNX model from `model_path`. /// /// Returns `None` when the model file is missing, the ONNX Runtime /// is unavailable, or the platform does not support ONNX. pub fn try_new(model_path: &str) -> Option { Self::try_new_onnx(model_path) } fn try_new_onnx(model_path: &str) -> Option { use tracing::{error, info}; if !std::path::Path::new(model_path).exists() { tracing::warn!( target: "chanora_audio", path = model_path, "SileroOnnxVad: model file not found; falling back to WebRtcFallbackVad" ); return None; } let session_result = std::panic::catch_unwind(|| { ort::session::Session::builder().and_then(|mut b| b.commit_from_file(model_path)) }); match session_result { Err(_) => { error!( target: "chanora_audio", path = model_path, "SileroOnnxVad: ONNX Runtime panicked during load; falling back to WebRtcFallbackVad" ); None } Ok(Ok(session)) => { info!( target: "chanora_audio", path = model_path, "SileroOnnxVad v6: model loaded" ); Some(Self { state: Box::new([0.0; SILERO_STATE_SIZE]), context: Box::new([0.0; SILERO_CONTEXT_16K]), accum: Vec::with_capacity(SILERO_FRAME_16K), last_probability: 0.0, model_path: model_path.to_owned(), inner: SileroInner::Onnx(OnnxSession { session }), }) } Ok(Err(e)) => { error!( target: "chanora_audio", path = model_path, error = %e, "SileroOnnxVad: failed to load model; falling back to WebRtcFallbackVad" ); None } } } /// Reset LSTM state and context (call on voice_leave / session restart). pub fn reset_state(&mut self) { self.state.iter_mut().for_each(|v| *v = 0.0); self.context.iter_mut().for_each(|v| *v = 0.0); self.accum.clear(); self.last_probability = 0.0; } /// Return the model path for diagnostics. pub fn model_path(&self) -> &str { &self.model_path } fn input_with_context(context: &[f32; SILERO_CONTEXT_16K], audio_frame: &[f32]) -> Vec { let mut input = Vec::with_capacity(SILERO_CONTEXT_16K + audio_frame.len()); input.extend_from_slice(context); input.extend_from_slice(audio_frame); input } fn update_context_from_frame(&mut self, audio_frame: &[f32]) { let ctx_start = audio_frame.len().saturating_sub(SILERO_CONTEXT_16K); let new_ctx = &audio_frame[ctx_start..]; let copy_len = new_ctx.len().min(SILERO_CONTEXT_16K); self.context.fill(0.0); self.context[SILERO_CONTEXT_16K - copy_len..].copy_from_slice(&new_ctx[..copy_len]); } /// Run one upstream-style `calc_level` pass over a 32 ms / 512-sample /// 16 kHz frame: concatenate prior context, pass `input/state/sr` to /// ONNX, persist `stateN`, then refresh context from the current frame. /// Updates `last_probability` and returns the new value. fn calc_level(&mut self, audio_frame: &[f32]) -> f32 { use ort::value::Value; use tracing::error; #[cfg(test)] let inner = match self.inner { SileroInner::Onnx(ref mut inner) => inner, SileroInner::Stub => return self.last_probability, }; #[cfg(not(test))] let SileroInner::Onnx(ref mut inner) = self.inner; debug_assert_eq!(audio_frame.len(), SILERO_FRAME_16K); // Build input: [1, 576] = context (64) + frame (512), matching // snakers4/silero-vad's Rust `calc_level` example. let input_vec = Self::input_with_context(self.context.as_ref(), audio_frame); // Build ndarray tensors. use ndarray::{Array, IxDyn}; let input_arr = Array::from_shape_vec(IxDyn(&[1, SILERO_INPUT_WIDTH]), input_vec); let state_arr = Array::from_shape_vec(IxDyn(&[2, 1, 128]), self.state.to_vec()); let sr_arr = Array::from_shape_vec(IxDyn(&[1]), vec![16000_i64]); let (input_arr, state_arr, sr_arr) = match (input_arr, state_arr, sr_arr) { (Ok(i), Ok(s), Ok(sr)) => (i, s, sr), _ => return self.last_probability, }; let input_val = match Value::from_array(input_arr) { Ok(v) => v, Err(e) => { error!(target: "chanora_audio", error = %e, "SileroOnnxVad: input tensor error"); return self.last_probability; } }; let state_val = match Value::from_array(state_arr) { Ok(v) => v, Err(e) => { error!(target: "chanora_audio", error = %e, "SileroOnnxVad: state tensor error"); return self.last_probability; } }; let sr_val = match Value::from_array(sr_arr) { Ok(v) => v, Err(e) => { error!(target: "chanora_audio", error = %e, "SileroOnnxVad: sr tensor error"); return self.last_probability; } }; let outputs = match inner .session .run([(&input_val).into(), (&state_val).into(), (&sr_val).into()]) { Ok(o) => o, Err(e) => { error!(target: "chanora_audio", error = %e, "SileroOnnxVad: inference failed"); return self.last_probability; } }; // Extract probability from "output". if let Ok((_, prob_data)) = outputs["output"].try_extract_tensor::() { if let Some(&p) = prob_data.first() { self.last_probability = p.clamp(0.0, 1.0); } } // Update state from "stateN". if let Ok((shape, state_data)) = outputs["stateN"].try_extract_tensor::() { let total: usize = shape.iter().map(|&d| d as usize).product(); let copy_len = total.min(SILERO_STATE_SIZE); self.state[..copy_len].copy_from_slice(&state_data[..copy_len]); } drop(outputs); // Match the upstream example: context becomes the last context_size // samples from the current frame after the model call succeeds. self.update_context_from_frame(audio_frame); self.last_probability } } impl VoiceActivityDetector for SileroOnnxVad { /// Accept one 10 ms **16 kHz** f32 mono frame (160 samples). /// /// Accumulates samples until a full 32 ms frame (512 samples) is /// ready, then runs inference. Between inference calls the last /// probability is returned unchanged. fn process_10ms(&mut self, samples: &[f32]) -> VadOutput { debug_assert_eq!( samples.len(), super::resampler::OUTPUT_FRAME_10MS, "SileroOnnxVad expects 160 samples (16 kHz 10 ms), got {}", samples.len() ); self.accum.extend_from_slice(samples); if self.accum.len() >= SILERO_FRAME_16K { let audio_frame: Vec = self.accum[..SILERO_FRAME_16K].to_vec(); if matches!(self.inner, SileroInner::Onnx(_)) { self.calc_level(&audio_frame); } else { self.update_context_from_frame(&audio_frame); } // Drain the accumulator (keep any overflow for next frame). let overflow: Vec = self.accum.drain(SILERO_FRAME_16K..).collect(); self.accum.clear(); self.accum.extend_from_slice(&overflow); } VadOutput { probability: self.last_probability, speech: self.last_probability >= 0.5, } } } // SAFETY: ONNX Runtime sessions are thread-safe for inference. // State arrays are owned by this struct and accessed only from // the single capture callback thread. unsafe impl Send for SileroOnnxVad {} struct SileroFrameMessage { seq: u64, frame: [f32; super::resampler::INPUT_FRAME_10MS], } /// Background Silero worker. The realtime callback only enqueues /// 10 ms frames and reads the latest probability atomically. pub struct SileroOnnxVadWorker { tx: Option>, latest_probability: Arc, latest_processed_seq: Arc, alive: Arc, handle: Option>, } impl SileroOnnxVadWorker { /// Start a background Silero worker if the model loads. pub fn try_new(model_path: &str) -> Option { let vad = SileroOnnxVad::try_new(model_path)?; let latest_probability = Arc::new(AtomicU32::new(0.0_f32.to_bits())); let latest_processed_seq = Arc::new(AtomicU64::new(u64::MAX)); let alive = Arc::new(AtomicBool::new(true)); let (tx, rx) = std::sync::mpsc::sync_channel::(64); let latest_probability_for_thread = latest_probability.clone(); let latest_processed_seq_for_thread = latest_processed_seq.clone(); let alive_for_thread = alive.clone(); let handle = std::thread::Builder::new() .name("chanora-silero-vad".to_string()) .spawn(move || { let mut vad = super::Resampled16kHzVad::new(vad); while alive_for_thread.load(Ordering::Relaxed) { let message = match rx.recv() { Ok(message) => message, Err(_) => break, }; let output = vad.process_10ms(&message.frame); latest_probability_for_thread.store( output.probability.clamp(0.0, 1.0).to_bits(), Ordering::Relaxed, ); latest_processed_seq_for_thread.store(message.seq, Ordering::Relaxed); } }) .ok()?; Some(Self { tx: Some(tx), latest_probability, latest_processed_seq, alive, handle: Some(handle), }) } /// Best-effort enqueue of a 10 ms frame for background inference. pub fn try_send(&self, seq: u64, frame: &[f32; super::resampler::INPUT_FRAME_10MS]) -> bool { let Some(tx) = &self.tx else { return false; }; tx.try_send(SileroFrameMessage { seq, frame: *frame }) .is_ok() } /// Latest probability published by the background worker. pub fn latest_probability(&self) -> f32 { f32::from_bits(self.latest_probability.load(Ordering::Relaxed)) } /// Number of 10 ms frames the worker is behind the capture thread. pub fn lag_frames(&self, capture_seq: u64) -> u64 { capture_seq.saturating_sub(self.latest_processed_seq.load(Ordering::Relaxed)) } /// True when the worker is too far behind to trust its latest /// probability for the current frame. pub fn is_stale(&self, capture_seq: u64) -> bool { let latest = self .latest_processed_seq .load(std::sync::atomic::Ordering::Relaxed); latest == u64::MAX || capture_seq.saturating_sub(latest) > SILERO_MAX_STALE_FRAMES } } impl Drop for SileroOnnxVadWorker { fn drop(&mut self) { self.alive.store(false, Ordering::Relaxed); let _ = self.tx.take(); let _ = self.handle.take(); } } #[cfg(test)] mod tests { use super::*; fn make_stub_vad() -> SileroOnnxVad { SileroOnnxVad { state: Box::new([0.0; SILERO_STATE_SIZE]), context: Box::new([0.0; SILERO_CONTEXT_16K]), accum: Vec::new(), last_probability: 0.0, model_path: String::new(), inner: SileroInner::Stub, } } #[test] fn try_new_returns_none_without_model_file() { let result = SileroOnnxVad::try_new("/nonexistent/silero_vad.onnx"); assert!(result.is_none()); } #[test] fn stub_accumulates_and_holds_zero_probability() { let mut vad = make_stub_vad(); let frame = vec![0.0_f32; super::super::resampler::OUTPUT_FRAME_10MS]; // Feed 3 frames (30 ms < 32 ms) — no inference yet. for _ in 0..3 { let out = VoiceActivityDetector::process_10ms(&mut vad, &frame); assert_eq!(out.probability, 0.0); } // Feed 1 more frame (40 ms > 32 ms) — accumulator drains. let out = VoiceActivityDetector::process_10ms(&mut vad, &frame); assert_eq!(out.probability, 0.0); // stub stays at 0 } #[test] fn reset_state_clears_all() { let mut vad = make_stub_vad(); vad.state[0] = 1.0; vad.context[0] = 1.0; vad.last_probability = 0.9; vad.accum.push(0.5); vad.reset_state(); assert_eq!(vad.state[0], 0.0); assert_eq!(vad.context[0], 0.0); assert_eq!(vad.last_probability, 0.0); assert!(vad.accum.is_empty()); } #[test] fn accumulates_correct_number_of_samples() { let mut vad = make_stub_vad(); let frame = vec![0.1_f32; super::super::resampler::OUTPUT_FRAME_10MS]; // 160 samples // 3 × 160 = 480 < 512 — not yet full. for _ in 0..3 { VoiceActivityDetector::process_10ms(&mut vad, &frame); } assert_eq!(vad.accum.len(), 480); // 4th frame: 640 > 512 — inference fires, 128 samples remain. VoiceActivityDetector::process_10ms(&mut vad, &frame); assert_eq!(vad.accum.len(), 128); } #[test] fn input_concatenates_context_before_frame_like_upstream_example() { let mut context = [0.0_f32; SILERO_CONTEXT_16K]; context[0] = -1.0; context[SILERO_CONTEXT_16K - 1] = 1.0; let frame = vec![0.25_f32; SILERO_FRAME_16K]; let input = SileroOnnxVad::input_with_context(&context, &frame); assert_eq!(input.len(), SILERO_INPUT_WIDTH); assert_eq!(input[0], -1.0); assert_eq!(input[SILERO_CONTEXT_16K - 1], 1.0); assert_eq!(input[SILERO_CONTEXT_16K], 0.25); assert_eq!(input[SILERO_INPUT_WIDTH - 1], 0.25); } #[test] fn context_tracks_last_64_samples_of_completed_frame() { let mut vad = make_stub_vad(); let frame = vec![0.0_f32; super::super::resampler::OUTPUT_FRAME_10MS]; for idx in 0..4 { let mut chunk = frame.clone(); let chunk_len = chunk.len(); for (sample_idx, sample) in chunk.iter_mut().enumerate() { *sample = (idx * chunk_len + sample_idx) as f32; } VoiceActivityDetector::process_10ms(&mut vad, &chunk); } let completed_frame: Vec = (0..SILERO_FRAME_16K).map(|v| v as f32).collect(); assert_eq!( vad.context.as_ref(), &completed_frame[SILERO_FRAME_16K - SILERO_CONTEXT_16K..] ); } }