//! TEN VAD ONNX backend. //! //! TEN's ONNX graph does not accept raw PCM. It expects the same feature //! stack produced by TEN's `AUP_Aed_aivad_proc`: three context frames of //! 40 log-mel powers plus one pitch feature, followed by four recurrent //! state tensors. This module ports that preprocessing path to Rust and //! keeps ONNX Runtime off the realtime callback where possible. use crate::frame::f32_to_i16; use rustfft::{num_complex::Complex32, FftPlanner}; use super::resampler::{Downsampler48to16, INPUT_FRAME_10MS}; use super::{VadOutput, VoiceActivityDetector}; const SAMPLE_RATE_16K: f32 = 16_000.0; const HOP_16K: usize = 256; const WINDOW_16K: usize = 768; const FFT_SIZE: usize = 1024; const N_BINS: usize = FFT_SIZE / 2 + 1; const MEL_BANDS: usize = 40; const FEATURE_LEN: usize = 41; const CONTEXT: usize = 3; const HIDDEN: usize = 64; const POWER_NORMALIZER: f32 = 32768.0 * 32768.0; const EPS: f32 = 1.0e-20; const FEATURE_MEANS: [f32; FEATURE_LEN] = [ -8.198236, -6.2657166, -5.4838185, -4.7586913, -4.417089, -4.142893, -3.9128504, -3.845928, -3.6570904, -3.7234187, -3.8761342, -3.843891, -3.6904051, -3.7560658, -3.6986961, -3.650463, -3.7004688, -3.5673213, -3.4989002, -3.477807, -3.458816, -3.4449239, -3.4013286, -3.3062613, -3.2785568, -3.2332509, -3.198616, -3.2045264, -3.2087986, -3.257838, -3.3813767, -3.5340214, -3.640868, -3.7268589, -3.773731, -3.8046672, -3.832901, -3.8711205, -3.990593, -4.4802895, 92.3569, ]; const FEATURE_STDS: [f32; FEATURE_LEN] = [ 5.166064, 4.9772096, 4.698896, 4.6306214, 4.634348, 4.641156, 4.6406765, 4.666367, 4.6505346, 4.640021, 4.6374, 4.620099, 4.5963163, 4.562655, 4.5543604, 4.5669107, 4.56249, 4.5624127, 4.5852995, 4.6001797, 4.592846, 4.5859227, 4.5834966, 4.626093, 4.626958, 4.6262894, 4.637006, 4.683016, 4.726814, 4.7342896, 4.753227, 4.849723, 4.869435, 4.884483, 4.921327, 4.9592123, 4.996619, 5.0448236, 5.072217, 5.0964394, 115.21369, ]; /// TEN VAD using ONNX Runtime and Rust-ported TEN feature preprocessing. pub struct TenOnnxVad { session: ort::session::Session, downsampler: Downsampler48to16, hop_accum: Vec, sample_fifo: Vec, feature_stack: [[f32; FEATURE_LEN]; CONTEXT], states: [[f32; HIDDEN]; 4], mel_filters: Vec<[f32; N_BINS]>, fft: std::sync::Arc>, fft_buffer: Vec, last_probability: f32, last_speech: bool, } unsafe impl Send for TenOnnxVad {} impl TenOnnxVad { /// Load TEN VAD ONNX model. pub fn try_new(model_path: &str) -> Option { if !std::path::Path::new(model_path).exists() { tracing::warn!(target: "chanora_audio", path = model_path, "TEN VAD ONNX model not found"); return None; } let session = match std::panic::catch_unwind(|| { ort::session::Session::builder().and_then(|mut b| b.commit_from_file(model_path)) }) { Ok(Ok(session)) => session, Ok(Err(error)) => { tracing::warn!(target: "chanora_audio", %error, path = model_path, "TEN VAD ONNX model load failed"); return None; } Err(_) => { tracing::warn!(target: "chanora_audio", path = model_path, "TEN VAD ONNX Runtime panicked during load"); return None; } }; let mut fft_planner = FftPlanner::::new(); let fft = fft_planner.plan_fft_forward(FFT_SIZE); tracing::info!(target: "chanora_audio", path = model_path, "TEN VAD ONNX model loaded"); Some(Self { session, downsampler: Downsampler48to16::default(), hop_accum: Vec::with_capacity(HOP_16K + super::resampler::OUTPUT_FRAME_10MS), sample_fifo: Vec::with_capacity(WINDOW_16K + HOP_16K), feature_stack: [[0.0; FEATURE_LEN]; CONTEXT], states: [[0.0; HIDDEN]; 4], mel_filters: build_mel_filters(), fft, fft_buffer: vec![Complex32::ZERO; FFT_SIZE], last_probability: 0.0, last_speech: false, }) } fn process_hop(&mut self, hop: &[f32]) { self.sample_fifo.extend_from_slice(hop); let frame = if self.sample_fifo.len() >= WINDOW_16K { let start = self.sample_fifo.len() - WINDOW_16K; self.sample_fifo[start..].to_vec() } else { let mut padded = vec![0.0; WINDOW_16K - self.sample_fifo.len()]; padded.extend_from_slice(&self.sample_fifo); padded }; if self.sample_fifo.len() > WINDOW_16K { let excess = self.sample_fifo.len() - WINDOW_16K; self.sample_fifo.drain(..excess); } let feature = compute_feature( &self.mel_filters, self.fft.as_ref(), &mut self.fft_buffer, &frame, ); self.feature_stack.copy_within(1..CONTEXT, 0); self.feature_stack[CONTEXT - 1] = feature; self.run_onnx(); } fn run_onnx(&mut self) { use ndarray::{Array, IxDyn}; use ort::value::Value; let input: Vec = self.feature_stack.iter().flatten().copied().collect(); let input_arr = match Array::from_shape_vec(IxDyn(&[1, CONTEXT, FEATURE_LEN]), input) { Ok(v) => v, Err(_) => return, }; let state_arrs = [0, 1, 2, 3] .map(|idx| Array::from_shape_vec(IxDyn(&[1, HIDDEN]), self.states[idx].to_vec())); let input_val = match Value::from_array(input_arr) { Ok(v) => v, Err(error) => { tracing::warn!(target: "chanora_audio", %error, "TEN VAD input tensor error"); return; } }; let state_vals = match state_arrs { [Ok(a), Ok(b), Ok(c), Ok(d)] => [a, b, c, d], _ => return, }; let state_vals = match state_vals.map(Value::from_array) { [Ok(a), Ok(b), Ok(c), Ok(d)] => [a, b, c, d], _ => return, }; let outputs = match self.session.run([ (&input_val).into(), (&state_vals[0]).into(), (&state_vals[1]).into(), (&state_vals[2]).into(), (&state_vals[3]).into(), ]) { Ok(outputs) => outputs, Err(error) => { tracing::warn!(target: "chanora_audio", %error, "TEN VAD ONNX inference failed"); return; } }; if let Ok((_, prob)) = outputs["output_1"].try_extract_tensor::() { if let Some(&p) = prob.first() { self.last_probability = p.clamp(0.0, 1.0); self.last_speech = self.last_probability >= 0.5; } } for (idx, name) in ["output_2", "output_3", "output_6", "output_7"] .iter() .enumerate() { if let Ok((_, state)) = outputs[*name].try_extract_tensor::() { let copy_len = state.len().min(HIDDEN); self.states[idx][..copy_len].copy_from_slice(&state[..copy_len]); } } } } fn compute_feature( mel_filters: &[[f32; N_BINS]], fft: &dyn rustfft::Fft, fft_buffer: &mut [Complex32], frame: &[f32], ) -> [f32; FEATURE_LEN] { let power = power_spectrum(fft, fft_buffer, frame); let mut feature = [0.0; FEATURE_LEN]; for band in 0..MEL_BANDS { let energy = mel_filters[band] .iter() .zip(power.iter()) .map(|(w, p)| w * p) .sum::() / POWER_NORMALIZER; let log_energy = (energy + EPS).ln(); feature[band] = (log_energy - FEATURE_MEANS[band]) / (FEATURE_STDS[band] + EPS); } let pitch_hz = estimate_pitch_hz(frame); feature[MEL_BANDS] = (pitch_hz - FEATURE_MEANS[MEL_BANDS]) / (FEATURE_STDS[MEL_BANDS] + EPS); feature } impl VoiceActivityDetector for TenOnnxVad { fn process_10ms(&mut self, samples: &[f32]) -> VadOutput { debug_assert_eq!(samples.len(), INPUT_FRAME_10MS); let mut input = [0.0_f32; INPUT_FRAME_10MS]; input.copy_from_slice(samples); let downsampled = self.downsampler.process_frame_10ms(&input); self.hop_accum.extend_from_slice(&downsampled); while self.hop_accum.len() >= HOP_16K { let hop: Vec = self.hop_accum[..HOP_16K].to_vec(); self.hop_accum.drain(..HOP_16K); self.process_hop(&hop); } VadOutput { probability: self.last_probability, speech: self.last_speech, } } } fn hz_to_mel(hz: f32) -> f32 { 2595.0 * (1.0 + hz / 700.0).log10() } fn mel_to_hz(mel: f32) -> f32 { 700.0 * (10.0_f32.powf(mel / 2595.0) - 1.0) } fn build_mel_filters() -> Vec<[f32; N_BINS]> { let low_mel = hz_to_mel(0.0); let high_mel = hz_to_mel(8000.0); let mut bins = [0_usize; MEL_BANDS + 2]; for idx in 0..bins.len() { let mel = idx as f32 * (high_mel - low_mel) / (MEL_BANDS as f32 + 1.0) + low_mel; let hz = mel_to_hz(mel); let mut bin = ((FFT_SIZE as f32 + 1.0) * hz / SAMPLE_RATE_16K).floor() as usize; bin = bin.min(N_BINS - 1); if idx > 0 && bin == bins[idx - 1] { bin = (bin + 1).min(N_BINS - 1); } bins[idx] = bin; } let mut filters = vec![[0.0_f32; N_BINS]; MEL_BANDS]; for band in 0..MEL_BANDS { let left = bins[band]; let center = bins[band + 1].max(left + 1); let right = bins[band + 2].max(center + 1).min(N_BINS - 1); for (i, weight) in filters[band] .iter_mut() .enumerate() .take(center.min(N_BINS)) .skip(left) { *weight = (i - left) as f32 / (center - left) as f32; } for (i, weight) in filters[band] .iter_mut() .enumerate() .take(right + 1) .skip(center) { *weight = (right - i) as f32 / (right - center).max(1) as f32; } } filters } fn power_spectrum( fft: &dyn rustfft::Fft, fft_buffer: &mut [Complex32], frame: &[f32], ) -> [f32; N_BINS] { debug_assert_eq!(fft_buffer.len(), FFT_SIZE); fft_buffer.fill(Complex32::ZERO); for (idx, sample) in frame.iter().take(WINDOW_16K).enumerate() { let hann = 0.5 - 0.5 * (2.0 * std::f32::consts::PI * idx as f32 / WINDOW_16K as f32).cos(); fft_buffer[idx].re = f32_to_i16(*sample) as f32 * hann; } fft.process(fft_buffer); let mut out = [0.0_f32; N_BINS]; for (dst, bin) in out.iter_mut().zip(fft_buffer.iter()) { *dst = bin.norm_sqr(); } out } fn estimate_pitch_hz(frame: &[f32]) -> f32 { let min_lag = (SAMPLE_RATE_16K / 400.0) as usize; let max_lag = (SAMPLE_RATE_16K / 60.0) as usize; let mut best_lag = 0_usize; let mut best_corr = 0.0_f32; for lag in min_lag..=max_lag.min(frame.len().saturating_sub(1)) { let mut corr = 0.0_f32; let mut energy = 0.0_f32; for i in lag..frame.len() { corr += frame[i] * frame[i - lag]; energy += frame[i - lag] * frame[i - lag]; } let norm = if energy > 1.0e-8 { corr / energy.sqrt() } else { 0.0 }; if norm > best_corr { best_corr = norm; best_lag = lag; } } if best_lag == 0 || best_corr < 0.01 { 0.0 } else { SAMPLE_RATE_16K / best_lag as f32 } } // --------------------------------------------------------------------------- // Background worker — same pattern as SileroOnnxVadWorker so the realtime // callback never blocks on STFT / pitch / ONNX inference. // --------------------------------------------------------------------------- use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64}; use std::sync::Arc; use std::thread::JoinHandle; /// Maximum number of 10 ms frames the worker may lag before the callback /// treats its output as stale and uses WebRTC fallback instead. const TEN_MAX_STALE_FRAMES: u64 = 8; struct TenFrameMessage { seq: u64, frame: [f32; INPUT_FRAME_10MS], } /// Background TEN VAD worker. The realtime callback only enqueues 10 ms /// frames and reads the latest probability atomically. pub struct TenOnnxVadWorker { tx: Option>, latest_probability: Arc, latest_processed_seq: Arc, alive: Arc, handle: Option>, } impl TenOnnxVadWorker { /// Start a background TEN worker if the model loads. pub fn try_new(model_path: &str) -> Option { let vad = TenOnnxVad::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::(128); let prob_arc = latest_probability.clone(); let seq_arc = latest_processed_seq.clone(); let alive_arc = alive.clone(); let handle = std::thread::Builder::new() .name("chanora-ten-vad".to_string()) .spawn(move || { let mut vad = vad; while alive_arc.load(std::sync::atomic::Ordering::Relaxed) { let msg = match rx.recv() { Ok(m) => m, Err(_) => break, }; let mut frame_f32 = [0.0_f32; INPUT_FRAME_10MS]; frame_f32.copy_from_slice(&msg.frame); let out = VoiceActivityDetector::process_10ms(&mut vad, &frame_f32); prob_arc.store( out.probability.clamp(0.0, 1.0).to_bits(), std::sync::atomic::Ordering::Relaxed, ); seq_arc.store(msg.seq, std::sync::atomic::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; INPUT_FRAME_10MS]) -> bool { let Some(tx) = &self.tx else { return false; }; tx.try_send(TenFrameMessage { 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(std::sync::atomic::Ordering::Relaxed), ) } /// True when the worker is too far behind to trust its output. 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) > TEN_MAX_STALE_FRAMES } } impl Drop for TenOnnxVadWorker { fn drop(&mut self) { self.alive .store(false, std::sync::atomic::Ordering::Relaxed); drop(self.tx.take()); if let Some(h) = self.handle.take() { let _ = h.join(); } } } #[cfg(test)] mod tests { use super::*; #[test] fn mel_filter_bank_has_expected_shape() { let filters = build_mel_filters(); assert_eq!(filters.len(), MEL_BANDS); assert!(filters.iter().all(|f| f.iter().any(|&v| v > 0.0))); } #[test] fn preprocessing_produces_finite_features() { let filters = build_mel_filters(); let mut planner = FftPlanner::::new(); let fft = planner.plan_fft_forward(FFT_SIZE); let mut fft_buffer = vec![Complex32::ZERO; FFT_SIZE]; let frame = vec![0.0_f32; WINDOW_16K]; let feature = compute_feature(&filters, fft.as_ref(), &mut fft_buffer, &frame); assert!(feature.iter().all(|v| v.is_finite())); } }