256 lines
8.7 KiB
Rust
256 lines
8.7 KiB
Rust
//! Voice activity detection backends and helpers.
|
|
//!
|
|
//! Apple capture feeds VoiceProcessingIO/CoreAudio-processed microphone
|
|
//! frames into this module and prefers Apple CoreML Silero VAD when the
|
|
//! Swift bridge is linked. WebRTC VAD remains the realtime-safe fallback;
|
|
//! non-Apple platforms may use ONNX-backed Silero when available.
|
|
|
|
#[cfg(any(target_os = "ios", target_os = "macos"))]
|
|
pub mod apple_coreml;
|
|
pub mod resampler;
|
|
#[cfg(not(target_os = "ios"))]
|
|
pub mod silero_onnx;
|
|
|
|
use std::sync::atomic::{AtomicU64, Ordering};
|
|
use std::sync::{OnceLock, RwLock};
|
|
|
|
use crate::frame::{f32_to_i16, i16_to_f32};
|
|
use crate::AudioError;
|
|
use resampler::{Downsampler48to16, INPUT_FRAME_10MS};
|
|
|
|
#[cfg(not(target_os = "ios"))]
|
|
pub use silero_onnx::SileroOnnxVad;
|
|
|
|
/// Voice activity detector output for one 10 ms frame.
|
|
#[derive(Debug, Clone, Copy)]
|
|
pub struct VadOutput {
|
|
/// Speech confidence in the inclusive range `[0.0, 1.0]`.
|
|
pub probability: f32,
|
|
/// Immediate detector speech decision before hangover/min-duration state.
|
|
pub speech: bool,
|
|
}
|
|
|
|
/// Realtime-safe detector that consumes one 10 ms f32 mono frame.
|
|
pub trait VoiceActivityDetector: Send {
|
|
/// Process one 10 ms frame and return speech probability/state.
|
|
fn process_10ms(&mut self, samples: &[f32]) -> VadOutput;
|
|
}
|
|
|
|
/// Realtime-safe WebRTC VAD used when a model runtime is unavailable.
|
|
pub struct WebRtcFallbackVad {
|
|
vad: webrtc_vad::Vad,
|
|
frame_i16: [i16; INPUT_FRAME_10MS],
|
|
}
|
|
|
|
// `webrtc_vad::Vad` owns an FFI pointer and is only touched from the
|
|
// capture thread after construction. Moving the wrapper between threads is
|
|
// safe; sharing it concurrently is not required and not implemented.
|
|
unsafe impl Send for WebRtcFallbackVad {}
|
|
|
|
impl Default for WebRtcFallbackVad {
|
|
fn default() -> Self {
|
|
Self {
|
|
vad: webrtc_vad::Vad::new_with_rate_and_mode(
|
|
webrtc_vad::SampleRate::Rate48kHz,
|
|
webrtc_vad::VadMode::Aggressive,
|
|
),
|
|
frame_i16: [0; INPUT_FRAME_10MS],
|
|
}
|
|
}
|
|
}
|
|
|
|
impl VoiceActivityDetector for WebRtcFallbackVad {
|
|
fn process_10ms(&mut self, samples: &[f32]) -> VadOutput {
|
|
for (dst, src) in self.frame_i16.iter_mut().zip(samples.iter().copied()) {
|
|
*dst = f32_to_i16(src);
|
|
}
|
|
let speech = self.vad.is_voice_segment(&self.frame_i16).unwrap_or(false);
|
|
VadOutput {
|
|
probability: if speech { 1.0 } else { 0.0 },
|
|
speech,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Wraps any `VoiceActivityDetector` that operates at 16 kHz and
|
|
/// downsamples 48 kHz input before forwarding.
|
|
pub struct Resampled16kHzVad<D: VoiceActivityDetector> {
|
|
inner: D,
|
|
downsampler: Downsampler48to16,
|
|
}
|
|
|
|
impl<D: VoiceActivityDetector> Resampled16kHzVad<D> {
|
|
/// Wrap a 16 kHz detector so it can consume 48 kHz frames.
|
|
pub fn new(inner: D) -> Self {
|
|
Self {
|
|
inner,
|
|
downsampler: Downsampler48to16::default(),
|
|
}
|
|
}
|
|
}
|
|
|
|
impl<D: VoiceActivityDetector> VoiceActivityDetector for Resampled16kHzVad<D> {
|
|
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.inner.process_10ms(&downsampled)
|
|
}
|
|
}
|
|
|
|
/// Convert one 48 kHz i16 10 ms frame and run a detector over it.
|
|
pub fn process_i16_10ms(detector: &mut dyn VoiceActivityDetector, samples: &[i16]) -> VadOutput {
|
|
let mut frame = [0.0_f32; INPUT_FRAME_10MS];
|
|
for (dst, src) in frame.iter_mut().zip(samples.iter().copied()) {
|
|
*dst = i16_to_f32(src);
|
|
}
|
|
detector.process_10ms(&frame)
|
|
}
|
|
|
|
static SILERO_MODEL_PATH_OVERRIDE: OnceLock<RwLock<Option<String>>> = OnceLock::new();
|
|
static SILERO_MODEL_EPOCH: AtomicU64 = AtomicU64::new(0);
|
|
|
|
fn silero_model_path_override() -> &'static RwLock<Option<String>> {
|
|
SILERO_MODEL_PATH_OVERRIDE.get_or_init(|| RwLock::new(None))
|
|
}
|
|
|
|
/// Configure the preferred Silero ONNX model path on supported platforms.
|
|
///
|
|
/// The path is validated eagerly. A successful call increments the
|
|
/// model epoch so running non-iOS audio backends can reload the model
|
|
/// without an app restart.
|
|
pub fn set_silero_model_path(path: &str) -> Result<(), AudioError> {
|
|
let path = path.trim();
|
|
if path.is_empty() {
|
|
return Err(AudioError::InvalidAudioProcessingConfig(
|
|
"vad model path must not be empty".to_string(),
|
|
));
|
|
}
|
|
if !std::path::Path::new(path).is_file() {
|
|
return Err(AudioError::InvalidAudioProcessingConfig(format!(
|
|
"vad model path does not exist or is not a file: {path}"
|
|
)));
|
|
}
|
|
let mut guard = silero_model_path_override()
|
|
.write()
|
|
.map_err(|_| AudioError::Backend("vad model path lock poisoned".to_string()))?;
|
|
*guard = Some(path.to_string());
|
|
SILERO_MODEL_EPOCH.fetch_add(1, Ordering::Relaxed);
|
|
Ok(())
|
|
}
|
|
|
|
/// Monotonic counter incremented whenever the configured model path changes.
|
|
pub fn silero_model_epoch() -> u64 {
|
|
SILERO_MODEL_EPOCH.load(Ordering::Relaxed)
|
|
}
|
|
|
|
/// Return the expected path of the Silero VAD v6 ONNX model on
|
|
/// supported platforms.
|
|
/// The model is shipped as a Flutter asset and copied to the app's
|
|
/// data directory by the Dart-side asset loader.
|
|
///
|
|
/// Android and macOS may use app data/Documents locations. Desktop
|
|
/// callers can set the path explicitly via `set_silero_model_path`.
|
|
pub fn silero_model_bundle_path() -> String {
|
|
if let Ok(guard) = silero_model_path_override().read() {
|
|
if let Some(path) = guard.as_ref() {
|
|
return path.clone();
|
|
}
|
|
}
|
|
|
|
// macOS: Documents directory (written by Flutter asset loader).
|
|
// iOS keeps this fallback only for API compatibility; the ONNX
|
|
// detector is not compiled into iOS builds.
|
|
#[cfg(any(target_os = "ios", target_os = "macos"))]
|
|
{
|
|
if let Ok(home) = std::env::var("HOME") {
|
|
let docs = format!("{home}/Documents/silero_vad.onnx");
|
|
if std::path::Path::new(&docs).exists() {
|
|
return docs;
|
|
}
|
|
let bundle = format!("{home}/../Library/silero_vad.onnx");
|
|
if std::path::Path::new(&bundle).exists() {
|
|
return bundle;
|
|
}
|
|
}
|
|
"silero_vad.onnx".to_string()
|
|
}
|
|
|
|
// Android: the model is in the app's files directory, same
|
|
// Documents path pattern used by Flutter's path_provider.
|
|
#[cfg(target_os = "android")]
|
|
{
|
|
// On Android, Flutter's getApplicationDocumentsDirectory
|
|
// resolves to /data/data/<package>/app_flutter.
|
|
// The Silero model path is set explicitly via
|
|
// set_silero_model_path from Dart before voice starts,
|
|
// so this fallback is rarely needed.
|
|
"silero_vad.onnx".to_string()
|
|
}
|
|
|
|
// Desktop (Windows, Linux): rely on the override set by Dart.
|
|
#[cfg(not(any(target_os = "ios", target_os = "macos", target_os = "android")))]
|
|
{
|
|
if let Ok(cwd) = std::env::current_dir() {
|
|
let local = cwd.join("silero_vad.onnx");
|
|
if local.exists() {
|
|
return local.to_string_lossy().to_string();
|
|
}
|
|
}
|
|
"silero_vad.onnx".to_string()
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn fallback_opens_for_voiced_signal() {
|
|
let mut vad = WebRtcFallbackVad::default();
|
|
let mut frame = [0_i16; INPUT_FRAME_10MS];
|
|
for (idx, sample) in frame.iter_mut().enumerate() {
|
|
let phase = idx as f32 * 2.0 * std::f32::consts::PI * 220.0 / 48_000.0;
|
|
*sample = (phase.sin() * 12_000.0) as i16;
|
|
}
|
|
|
|
let output = process_i16_10ms(&mut vad, &frame);
|
|
|
|
assert!(output.speech);
|
|
assert!(output.probability >= 0.5);
|
|
}
|
|
|
|
#[test]
|
|
fn fallback_stays_closed_for_silence() {
|
|
let mut vad = WebRtcFallbackVad::default();
|
|
let frame = [0_i16; INPUT_FRAME_10MS];
|
|
|
|
let output = process_i16_10ms(&mut vad, &frame);
|
|
|
|
assert!(!output.speech);
|
|
assert!(output.probability < 0.5);
|
|
}
|
|
|
|
#[test]
|
|
fn set_silero_model_path_rejects_missing_file() {
|
|
let result = set_silero_model_path("/definitely/not/a/silero_vad.onnx");
|
|
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn set_silero_model_path_updates_override_and_epoch() {
|
|
let path =
|
|
std::env::temp_dir().join(format!("chanora_test_silero_{}.onnx", std::process::id()));
|
|
std::fs::write(&path, b"test").unwrap();
|
|
let before = silero_model_epoch();
|
|
|
|
set_silero_model_path(path.to_str().unwrap()).unwrap();
|
|
|
|
assert!(silero_model_epoch() > before);
|
|
assert_eq!(silero_model_bundle_path(), path.to_string_lossy());
|
|
let _ = std::fs::remove_file(path);
|
|
}
|
|
}
|