feat(voice): add iOS VAD runtime support

This commit is contained in:
Edison Jwa
2026-05-21 20:51:45 +09:00
parent 171baf6e41
commit 6af4ecab0f
73 changed files with 11529 additions and 1249 deletions
+527
View File
@@ -0,0 +1,527 @@
//! 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<f32>,
/// 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 {
#[cfg(any(target_os = "ios", target_os = "macos"))]
Onnx(OnnxSession),
#[allow(dead_code)]
Stub,
}
#[cfg(any(target_os = "ios", target_os = "macos"))]
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 is not iOS/macOS.
pub fn try_new(model_path: &str) -> Option<Self> {
#[cfg(any(target_os = "ios", target_os = "macos"))]
{
Self::try_new_onnx(model_path)
}
#[cfg(not(any(target_os = "ios", target_os = "macos")))]
{
let _ = model_path;
None
}
}
#[cfg(any(target_os = "ios", target_os = "macos"))]
fn try_new_onnx(model_path: &str) -> Option<Self> {
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;
}
#[cfg(target_os = "macos")]
if let Some(path) = bundled_onnxruntime_path() {
let _ = ort::init_from(path.to_string_lossy()).commit();
}
let session_result = std::panic::catch_unwind(|| {
ort::session::Session::builder().and_then(|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<f32> {
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.
#[cfg(any(target_os = "ios", target_os = "macos"))]
fn calc_level(&mut self, audio_frame: &[f32]) -> f32 {
use ort::value::Value;
use tracing::error;
let SileroInner::Onnx(ref mut inner) = self.inner else {
return self.last_probability;
};
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::<f32>() {
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::<f32>() {
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
}
}
#[cfg(target_os = "macos")]
fn bundled_onnxruntime_path() -> Option<std::path::PathBuf> {
let exe = std::env::current_exe().ok()?;
let app_dir = exe.parent()?;
let framework = app_dir
.join("Frameworks")
.join("onnxruntime.framework")
.join("onnxruntime");
framework.exists().then_some(framework)
}
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<f32> = self.accum[..SILERO_FRAME_16K].to_vec();
#[cfg(any(target_os = "ios", target_os = "macos"))]
{
if matches!(self.inner, SileroInner::Onnx(_)) {
self.calc_level(&audio_frame);
} else {
self.update_context_from_frame(&audio_frame);
}
}
#[cfg(not(any(target_os = "ios", target_os = "macos")))]
{
self.update_context_from_frame(&audio_frame);
}
// Drain the accumulator (keep any overflow for next frame).
let overflow: Vec<f32> = 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<std::sync::mpsc::SyncSender<SileroFrameMessage>>,
latest_probability: Arc<AtomicU32>,
latest_processed_seq: Arc<AtomicU64>,
alive: Arc<AtomicBool>,
handle: Option<JoinHandle<()>>,
}
impl SileroOnnxVadWorker {
/// Start a background Silero worker if the model loads.
pub fn try_new(model_path: &str) -> Option<Self> {
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(0));
let alive = Arc::new(AtomicBool::new(true));
let (tx, rx) = std::sync::mpsc::sync_channel::<SileroFrameMessage>(8);
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 {
self.lag_frames(capture_seq) > 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<f32> = (0..SILERO_FRAME_16K).map(|v| v as f32).collect();
assert_eq!(
vad.context.as_ref(),
&completed_frame[SILERO_FRAME_16K - SILERO_CONTEXT_16K..]
);
}
}