use std::path::Path;
use anyhow::{Context, Result};
use parking_lot::Mutex;
use crate::runtime::{
factory::RuntimeFactory,
session::RuntimeSession,
tensor::{Shape, Tensor, TensorData},
};
pub const VAD_MODEL_FILE: &str = "silero_vad.onnx";
pub const VAD_SAMPLE_RATE: i64 = 16000;
pub const VAD_FRAME_SAMPLES: usize = 512;
const VAD_STATE_LEN: usize = 2 * 128;
#[derive(Debug, Clone, Copy)]
pub struct VadConfig {
pub threshold: f32,
pub min_silence_ms: u32,
pub min_speech_ms: u32,
pub speech_pad_ms: u32,
}
impl Default for VadConfig {
fn default() -> Self {
Self {
threshold: 0.5,
min_silence_ms: 500,
min_speech_ms: 250,
speech_pad_ms: 100,
}
}
}
impl VadConfig {
fn ms_to_samples(ms: u32) -> usize {
(VAD_SAMPLE_RATE as usize * ms as usize) / 1000
}
}
pub struct SileroVad {
session: Mutex<Box<dyn RuntimeSession>>,
input_tensors: Mutex<Vec<Tensor>>,
}
impl SileroVad {
pub fn load(model_path: &Path) -> Result<Self> {
let factory = crate::runtime::cpu_factory();
Self::load_with_factory(model_path, factory.as_ref())
}
pub fn load_with_factory(model_path: &Path, factory: &dyn RuntimeFactory) -> Result<Self> {
tracing::debug!("Loading VAD model from {}", model_path.display());
let runtime = factory
.cpu_fallback()
.create(1)
.map_err(|e| anyhow::anyhow!(e))
.context("Failed to create runtime for VAD model")?;
let session = runtime
.load_session(model_path, false)
.map_err(|e| anyhow::anyhow!(e))
.context("Failed to load VAD model")?;
tracing::info!("VAD model loaded from {}", model_path.display());
Ok(Self {
session: Mutex::new(session),
input_tensors: Mutex::new(vec![
Tensor::new_checked(
Shape::new(vec![1, VAD_FRAME_SAMPLES]),
TensorData::F32(vec![0.0; VAD_FRAME_SAMPLES]),
),
Tensor::new_checked(
Shape::new(vec![2, 1, 128]),
TensorData::F32(vec![0.0; VAD_STATE_LEN]),
),
Tensor::new_checked(Shape::new(vec![1]), TensorData::I64(vec![VAD_SAMPLE_RATE])),
]),
})
}
fn run_frame(&self, frame: &[f32], state: &mut [f32; VAD_STATE_LEN]) -> Result<f32> {
let mut input = [0.0f32; VAD_FRAME_SAMPLES];
let n = frame.len().min(VAD_FRAME_SAMPLES);
input[..n].copy_from_slice(&frame[..n]);
let outputs = {
let mut inputs = self.input_tensors.lock();
inputs[0]
.as_f32_mut()
.context("VAD frame tensor is not f32")?
.copy_from_slice(&input);
inputs[1]
.as_f32_mut()
.context("VAD state tensor is not f32")?
.copy_from_slice(state);
let session = self.session.lock();
session.run(&inputs).context("VAD model inference failed")?
};
let mut prob = 0.0f32;
let mut new_state = [0.0f32; VAD_STATE_LEN];
for output in outputs {
let view = output.view();
if let Some(data) = view.data().as_f32() {
if data.len() == VAD_STATE_LEN {
new_state.copy_from_slice(data);
} else if data.len() == 1 {
prob = data[0];
}
}
}
state.copy_from_slice(&new_state);
Ok(prob)
}
pub fn frame_probs(&self, samples: &[f32]) -> Result<Vec<f32>> {
self.frame_probs_with_abort(samples, None)
}
pub(crate) fn frame_probs_with_abort(
&self,
samples: &[f32],
abort: Option<&dyn Fn() -> bool>,
) -> Result<Vec<f32>> {
let mut state = [0.0f32; VAD_STATE_LEN];
let mut probs = Vec::with_capacity(samples.len() / VAD_FRAME_SAMPLES + 1);
let mut i = 0;
let mut since_check = 0usize;
while i < samples.len() {
if let Some(abort) = abort {
since_check += 1;
if since_check >= 64 {
since_check = 0;
if abort() {
anyhow::bail!("cancelled");
}
}
}
let end = (i + VAD_FRAME_SAMPLES).min(samples.len());
probs.push(self.run_frame(&samples[i..end], &mut state)?);
i = end;
}
Ok(probs)
}
pub fn speech_regions(&self, samples: &[f32], cfg: &VadConfig) -> Result<Vec<(usize, usize)>> {
self.speech_regions_with_abort(samples, cfg, None)
}
pub(crate) fn speech_regions_with_abort(
&self,
samples: &[f32],
cfg: &VadConfig,
abort: Option<&dyn Fn() -> bool>,
) -> Result<Vec<(usize, usize)>> {
let probs = self.frame_probs_with_abort(samples, abort)?;
Ok(regions_from_probs(
&probs,
VAD_FRAME_SAMPLES,
samples.len(),
cfg,
))
}
}
pub fn regions_from_probs(
probs: &[f32],
frame_samples: usize,
total_samples: usize,
cfg: &VadConfig,
) -> Vec<(usize, usize)> {
if probs.is_empty() || total_samples == 0 {
return Vec::new();
}
let min_silence = VadConfig::ms_to_samples(cfg.min_silence_ms);
let min_speech = VadConfig::ms_to_samples(cfg.min_speech_ms);
let pad = VadConfig::ms_to_samples(cfg.speech_pad_ms);
let mut regions: Vec<(usize, usize)> = Vec::new();
let mut run_start: Option<usize> = None;
for (i, &p) in probs.iter().enumerate() {
let speech = p >= cfg.threshold;
if speech && run_start.is_none() {
run_start = Some(i * frame_samples);
} else if !speech && let Some(s) = run_start.take() {
regions.push((s, i * frame_samples));
}
}
if let Some(s) = run_start.take() {
regions.push((s, total_samples));
}
if regions.is_empty() {
return regions;
}
let mut merged: Vec<(usize, usize)> = Vec::with_capacity(regions.len());
for (s, e) in regions {
match merged.last_mut() {
Some(last) if s.saturating_sub(last.1) < min_silence => last.1 = e,
_ => merged.push((s, e)),
}
}
merged.retain(|(s, e)| e - s >= min_speech);
if merged.is_empty() {
return merged;
}
let mut padded: Vec<(usize, usize)> = Vec::with_capacity(merged.len());
for (s, e) in merged {
let ps = s.saturating_sub(pad);
let pe = (e + pad).min(total_samples);
match padded.last_mut() {
Some(last) if ps <= last.1 => last.1 = last.1.max(pe),
_ => padded.push((ps, pe)),
}
}
padded
}
pub fn remap_compressed_seconds(
t_compressed_s: f64,
regions: &[(usize, usize)],
sample_rate: f64,
) -> f64 {
if regions.is_empty() {
return t_compressed_s;
}
let target = (t_compressed_s * sample_rate).max(0.0);
let mut acc = 0.0f64; for &(s, e) in regions {
let len = (e - s) as f64;
if target <= acc + len {
let into = (target - acc).max(0.0);
return (s as f64 + into) / sample_rate;
}
acc += len;
}
let &(_, end) = regions.last().expect("non-empty checked above");
end as f64 / sample_rate
}
#[cfg(feature = "file-decode")]
pub(crate) struct VadSegmenter {
threshold: f32,
min_silence: usize,
min_speech: usize,
pad: usize,
state: [f32; VAD_STATE_LEN],
raw: Vec<f32>,
raw_start: usize,
pos: usize,
in_run: bool,
merged: Option<(usize, usize)>,
open_idx: Option<usize>,
regions: Vec<(usize, usize)>,
out_idx: usize,
copied_to: usize,
}
#[cfg(feature = "file-decode")]
impl VadSegmenter {
pub(crate) fn new(cfg: &VadConfig) -> Self {
Self {
threshold: cfg.threshold,
min_silence: VadConfig::ms_to_samples(cfg.min_silence_ms),
min_speech: VadConfig::ms_to_samples(cfg.min_speech_ms),
pad: VadConfig::ms_to_samples(cfg.speech_pad_ms),
state: [0.0; VAD_STATE_LEN],
raw: Vec::new(),
raw_start: 0,
pos: 0,
in_run: false,
merged: None,
open_idx: None,
regions: Vec::new(),
out_idx: 0,
copied_to: 0,
}
}
pub(crate) fn regions(&self) -> &[(usize, usize)] {
&self.regions
}
pub(crate) fn push(
&mut self,
vad: &SileroVad,
samples: &[f32],
out: &mut Vec<f32>,
) -> Result<()> {
self.push_with(samples, out, |frame, state| vad.run_frame(frame, state))
}
pub(crate) fn finish(
&mut self,
vad: &SileroVad,
total: usize,
out: &mut Vec<f32>,
) -> Result<()> {
self.finish_with(total, out, |frame, state| vad.run_frame(frame, state))
}
fn push_with<F>(&mut self, samples: &[f32], out: &mut Vec<f32>, mut score: F) -> Result<()>
where
F: FnMut(&[f32], &mut [f32; VAD_STATE_LEN]) -> Result<f32>,
{
self.raw.extend_from_slice(samples);
let avail = self.raw_start + self.raw.len();
while self.pos + VAD_FRAME_SAMPLES <= avail {
let off = self.pos - self.raw_start;
let prob = score(&self.raw[off..off + VAD_FRAME_SAMPLES], &mut self.state)?;
self.step(prob, self.pos + VAD_FRAME_SAMPLES);
}
self.flush(out);
self.trim();
Ok(())
}
fn finish_with<F>(&mut self, total: usize, out: &mut Vec<f32>, mut score: F) -> Result<()>
where
F: FnMut(&[f32], &mut [f32; VAD_STATE_LEN]) -> Result<f32>,
{
while self.pos < total {
let off = self.pos - self.raw_start;
let end = (off + VAD_FRAME_SAMPLES).min(self.raw.len());
let prob = score(&self.raw[off..end], &mut self.state)?;
self.step(prob, (self.pos + VAD_FRAME_SAMPLES).min(total));
}
self.close(total);
self.flush(out);
Ok(())
}
fn close(&mut self, total: usize) {
if self.in_run
&& let Some(m) = self.merged.as_mut()
{
m.1 = total;
self.in_run = false;
}
if let Some((ms, me)) = self.merged.take() {
self.finalize(ms, me);
}
for r in &mut self.regions {
r.1 = r.1.min(total);
}
self.pos = self.pos.max(total);
}
fn step(&mut self, prob: f32, end: usize) {
let start = self.pos;
if prob >= self.threshold {
if !self.in_run {
match self.merged {
Some((_, me)) if start.saturating_sub(me) < self.min_silence => {}
Some((ms, me)) => {
self.finalize(ms, me);
self.merged = None;
}
None => {}
}
if self.merged.is_none() {
self.merged = Some((start, start));
}
self.in_run = true;
}
if let Some(m) = self.merged.as_mut() {
m.1 = end;
}
if let Some((ms, me)) = self.merged {
match self.open_idx {
Some(i) => self.regions[i].1 = self.regions[i].1.max(me),
None if me - ms >= self.min_speech => {
self.open_idx = Some(self.push_padded(ms.saturating_sub(self.pad), me));
}
None => {}
}
}
} else {
self.in_run = false;
if let Some((ms, me)) = self.merged
&& end.saturating_sub(me) >= self.min_silence
{
self.finalize(ms, me);
self.merged = None;
}
}
self.pos = end;
}
fn push_padded(&mut self, ps: usize, pe: usize) -> usize {
match self.regions.last_mut() {
Some(last) if ps <= last.1 => last.1 = last.1.max(pe),
_ => self.regions.push((ps, pe)),
}
self.regions.len() - 1
}
fn finalize(&mut self, ms: usize, me: usize) {
let idx = self.open_idx.take();
if me - ms < self.min_speech {
return;
}
let pe = me + self.pad;
match idx {
Some(i) => self.regions[i].1 = self.regions[i].1.max(pe),
None => {
self.push_padded(ms.saturating_sub(self.pad), pe);
}
}
}
fn decided_to(&self) -> usize {
let base = self.regions.last().map_or(0, |r| r.1);
let bound = match self.merged {
Some(_) if self.open_idx.is_some() => base,
Some((ms, _)) => base.max(ms.saturating_sub(self.pad)),
None => base.max(self.pos.saturating_sub(self.pad)),
};
bound.min(self.pos)
}
fn flush(&mut self, out: &mut Vec<f32>) {
let decided = self.decided_to().min(self.raw_start + self.raw.len());
while self.out_idx < self.regions.len() {
let (s, e) = self.regions[self.out_idx];
let from = self.copied_to.max(s);
let to = e.min(decided);
if to > from {
debug_assert!(
from >= self.raw_start,
"released PCM at {from} below the retained start {}",
self.raw_start
);
out.extend_from_slice(&self.raw[from - self.raw_start..to - self.raw_start]);
self.copied_to = to;
}
if self.out_idx + 1 < self.regions.len() {
self.out_idx += 1;
} else {
break;
}
}
}
fn trim(&mut self) {
let decided = self.decided_to();
let keep_from = match self.merged {
Some((ms, _)) if self.open_idx.is_none() => decided.min(ms.saturating_sub(self.pad)),
_ => decided,
};
let drop = keep_from.saturating_sub(self.raw_start).min(self.raw.len());
if drop > 0 {
self.raw.drain(..drop);
self.raw_start += drop;
}
}
#[cfg(test)]
fn retained(&self) -> usize {
self.raw.len()
}
}
pub struct VadEndpointer {
state: [f32; VAD_STATE_LEN],
leftover: Vec<f32>,
hangover: Hangover,
}
impl VadEndpointer {
pub fn new(cfg: &VadConfig) -> Self {
Self {
state: [0.0f32; VAD_STATE_LEN],
leftover: Vec::with_capacity(VAD_FRAME_SAMPLES),
hangover: Hangover::new(cfg),
}
}
pub fn push(&mut self, vad: &SileroVad, samples: &[f32]) -> Result<bool> {
self.leftover.extend_from_slice(samples);
let mut endpoint = false;
let mut off = 0;
while off + VAD_FRAME_SAMPLES <= self.leftover.len() {
let prob = vad.run_frame(
&self.leftover[off..off + VAD_FRAME_SAMPLES],
&mut self.state,
)?;
off += VAD_FRAME_SAMPLES;
if self.hangover.update(prob, VAD_FRAME_SAMPLES) {
endpoint = true;
}
}
if off > 0 {
self.leftover.drain(..off);
}
Ok(endpoint)
}
}
#[derive(Debug)]
pub struct Hangover {
threshold: f32,
min_silence_samples: usize,
seen_speech: bool,
trailing_silence: usize,
armed: bool,
}
impl Hangover {
fn new(cfg: &VadConfig) -> Self {
Self {
threshold: cfg.threshold,
min_silence_samples: VadConfig::ms_to_samples(cfg.min_silence_ms),
seen_speech: false,
trailing_silence: 0,
armed: false,
}
}
fn update(&mut self, prob: f32, frame_samples: usize) -> bool {
if prob >= self.threshold {
self.seen_speech = true;
self.armed = true;
self.trailing_silence = 0;
return false;
}
if !self.seen_speech {
return false;
}
self.trailing_silence += frame_samples;
if self.armed && self.trailing_silence >= self.min_silence_samples {
self.armed = false; return true;
}
false
}
}
#[cfg(test)]
mod tests {
use super::*;
fn cfg(
threshold: f32,
min_silence_ms: u32,
min_speech_ms: u32,
speech_pad_ms: u32,
) -> VadConfig {
VadConfig {
threshold,
min_silence_ms,
min_speech_ms,
speech_pad_ms,
}
}
#[test]
fn test_ms_to_samples_16khz() {
assert_eq!(VadConfig::ms_to_samples(1000), 16000);
assert_eq!(VadConfig::ms_to_samples(500), 8000);
assert_eq!(VadConfig::ms_to_samples(0), 0);
}
#[test]
fn test_regions_empty_probs_is_empty() {
let c = VadConfig::default();
assert!(regions_from_probs(&[], 512, 0, &c).is_empty());
assert!(regions_from_probs(&[0.9, 0.9], 512, 0, &c).is_empty());
}
#[test]
fn test_regions_all_silence_is_empty() {
let c = cfg(0.5, 0, 0, 0);
let probs = vec![0.1f32; 10];
assert!(regions_from_probs(&probs, 512, 10 * 512, &c).is_empty());
}
#[test]
fn test_regions_single_block_no_pad_no_mins() {
let c = cfg(0.5, 0, 0, 0);
let probs = [0.1, 0.9, 0.9, 0.1];
let r = regions_from_probs(&probs, 100, 400, &c);
assert_eq!(r, vec![(100, 300)]);
}
#[test]
fn test_regions_trailing_speech_clamps_to_total() {
let c = cfg(0.5, 0, 0, 0);
let probs = [0.1, 0.9, 0.9];
let r = regions_from_probs(&probs, 100, 250, &c);
assert_eq!(r, vec![(100, 250)]);
}
#[test]
fn test_regions_min_silence_merges_short_gap() {
let c = cfg(0.5, 100, 0, 0); let probs = [0.9, 0.1, 0.9];
let r = regions_from_probs(&probs, 100, 300, &c);
assert_eq!(r, vec![(0, 300)]);
}
#[test]
fn test_regions_long_gap_keeps_two_regions() {
let c = cfg(0.5, 0, 0, 0);
let probs = [0.9, 0.1, 0.1, 0.9];
let r = regions_from_probs(&probs, 100, 400, &c);
assert_eq!(r, vec![(0, 100), (300, 400)]);
}
#[test]
fn test_regions_min_speech_drops_short_blip() {
let c = cfg(0.5, 0, 100, 0);
let probs = [0.1, 0.9, 0.1];
assert!(regions_from_probs(&probs, 100, 300, &c).is_empty());
}
#[test]
fn test_regions_padding_extends_and_clamps() {
let c = cfg(0.5, 0, 0, 10); let probs = [0.1, 0.9, 0.1];
let r = regions_from_probs(&probs, 100, 1000, &c);
assert_eq!(r, vec![(0, 360)]);
}
#[test]
fn test_regions_padding_merges_overlapping_neighbours() {
let c = cfg(0.5, 0, 0, 50); let probs = [0.9, 0.1, 0.1, 0.9, 0.1];
let r = regions_from_probs(&probs, 100, 2000, &c);
assert_eq!(r, vec![(0, 1200)]);
}
#[test]
fn test_hangover_fires_once_after_min_silence() {
let c = cfg(0.5, 100, 0, 0); let mut h = Hangover::new(&c);
assert!(!h.update(0.9, 512));
assert!(!h.update(0.1, 512)); assert!(!h.update(0.1, 512)); assert!(!h.update(0.1, 512)); assert!(h.update(0.1, 512)); assert!(!h.update(0.1, 512));
}
#[test]
fn test_hangover_no_fire_before_any_speech() {
let c = cfg(0.5, 0, 0, 0);
let mut h = Hangover::new(&c);
for _ in 0..10 {
assert!(!h.update(0.1, 512));
}
}
#[test]
fn test_hangover_rearms_for_next_utterance() {
let c = cfg(0.5, 50, 0, 0); let mut h = Hangover::new(&c);
h.update(0.9, 512); assert!(!h.update(0.1, 512)); assert!(h.update(0.1, 512)); assert!(!h.update(0.9, 512));
assert!(!h.update(0.1, 512)); assert!(h.update(0.1, 512)); }
#[cfg(feature = "file-decode")]
mod segmenter {
use super::*;
fn stream(
probs: &[f32],
total: usize,
cfg: &VadConfig,
) -> (Vec<(usize, usize)>, Vec<f32>, usize) {
let raw: Vec<f32> = (0..total).map(|i| i as f32).collect();
let mut seg = VadSegmenter::new(cfg);
let mut out = Vec::new();
let mut it = probs.iter().copied();
let mut peak = 0usize;
let mut i = 0usize;
let mut chunk = 1usize;
while i < total {
let end = (i + chunk).min(total);
seg.push_with(&raw[i..end], &mut out, |_, _| Ok(it.next().unwrap_or(0.0)))
.expect("push");
peak = peak.max(seg.retained());
i = end;
chunk = chunk % 977 + 1;
}
seg.finish_with(total, &mut out, |_, _| Ok(it.next().unwrap_or(0.0)))
.expect("finish");
(seg.regions().to_vec(), out, peak)
}
fn batch(probs: &[f32], total: usize, cfg: &VadConfig) -> (Vec<(usize, usize)>, Vec<f32>) {
let regions = regions_from_probs(probs, VAD_FRAME_SAMPLES, total, cfg);
let out = regions
.iter()
.flat_map(|&(s, e)| (s..e).map(|i| i as f32))
.collect();
(regions, out)
}
fn assert_stream_matches_batch(probs: &[f32], total: usize, cfg: &VadConfig) {
let (got_regions, got_samples, _) = stream(probs, total, cfg);
let (want_regions, want_samples) = batch(probs, total, cfg);
assert_eq!(
got_regions, want_regions,
"regions diverged (total={total})"
);
assert_eq!(
got_samples, want_samples,
"compressed buffer diverged (total={total})"
);
}
fn probs_for(total: usize, f: impl Fn(usize) -> f32) -> Vec<f32> {
(0..total.div_ceil(VAD_FRAME_SAMPLES)).map(f).collect()
}
#[test]
fn test_segmenter_matches_batch_on_shaped_sequences() {
let c = VadConfig::default();
let fs = VAD_FRAME_SAMPLES;
for period in [1usize, 2, 3, 5, 8, 16, 20, 31, 64] {
let total = 200 * fs + 137; let probs = probs_for(total, |i| if (i / period) % 2 == 0 { 0.9 } else { 0.1 });
assert_stream_matches_batch(&probs, total, &c);
}
for level in [0.1f32, 0.9] {
let total = 97 * fs;
let probs = probs_for(total, |_| level);
assert_stream_matches_batch(&probs, total, &c);
}
}
#[test]
fn test_segmenter_matches_batch_on_degenerate_configs() {
let fs = VAD_FRAME_SAMPLES;
let total = 120 * fs + 11;
let probs = probs_for(total, |i| if (i / 7) % 3 == 0 { 0.9 } else { 0.1 });
for c in [
cfg(0.5, 0, 0, 0),
cfg(0.5, 0, 0, 200),
cfg(0.5, 10, 0, 500),
cfg(0.5, 1000, 2000, 100),
cfg(0.5, 40, 40, 40),
] {
assert_stream_matches_batch(&probs, total, &c);
}
}
#[test]
fn test_segmenter_matches_batch_on_short_and_empty_inputs() {
let c = VadConfig::default();
for total in [
0usize,
1,
2,
VAD_FRAME_SAMPLES - 1,
VAD_FRAME_SAMPLES,
VAD_FRAME_SAMPLES + 1,
] {
for level in [0.1f32, 0.9] {
assert_stream_matches_batch(&probs_for(total, |_| level), total, &c);
}
}
}
#[cfg(not(miri))]
proptest::proptest! {
#![proptest_config(proptest::prelude::ProptestConfig::with_cases(256))]
#[test]
fn prop_segmenter_matches_batch(
probs in proptest::collection::vec(0.0f32..=1.0, 1..60),
tail in 1usize..=VAD_FRAME_SAMPLES,
threshold in 0.1f32..0.9,
min_silence_ms in 0u32..800,
min_speech_ms in 0u32..500,
speech_pad_ms in 0u32..400,
) {
let total = (probs.len() - 1) * VAD_FRAME_SAMPLES + tail;
let c = cfg(threshold, min_silence_ms, min_speech_ms, speech_pad_ms);
let (got_regions, got_samples, _) = stream(&probs, total, &c);
let (want_regions, want_samples) = batch(&probs, total, &c);
proptest::prop_assert_eq!(got_regions, want_regions);
proptest::prop_assert_eq!(got_samples, want_samples);
}
}
#[test]
fn test_segmenter_retains_bounded_pcm_on_unbroken_speech() {
let c = VadConfig::default();
let total = 16000 * 3600;
let probs = probs_for(total, |_| 0.9);
let (regions, out, peak) = stream(&probs, total, &c);
assert_eq!(regions, vec![(0, total)]);
assert_eq!(out.len(), total);
let bound =
VadConfig::ms_to_samples(c.min_speech_ms + c.min_silence_ms + c.speech_pad_ms)
+ VAD_FRAME_SAMPLES
+ 977; assert!(
peak <= bound,
"retained {peak} samples, expected at most {bound}"
);
}
#[test]
fn test_segmenter_retains_bounded_pcm_on_sparse_speech() {
let c = VadConfig::default();
let total = 16000 * 3600 * 3;
let probs = probs_for(total, |i| if (i / 40) % 5 == 0 { 0.9 } else { 0.1 });
let (regions, out, peak) = stream(&probs, total, &c);
assert!(!regions.is_empty());
assert_eq!(out.len(), regions.iter().map(|(s, e)| e - s).sum::<usize>());
let bound =
VadConfig::ms_to_samples(c.min_speech_ms + c.min_silence_ms + c.speech_pad_ms)
+ VAD_FRAME_SAMPLES
+ 977;
assert!(
peak <= bound,
"retained {peak} samples, expected at most {bound}"
);
}
}
#[test]
fn test_remap_no_regions_is_identity() {
assert_eq!(remap_compressed_seconds(1.5, &[], 16000.0), 1.5);
}
#[test]
fn test_remap_single_region_offsets_by_start() {
let regions = [(16000usize, 32000usize)];
assert_eq!(remap_compressed_seconds(0.0, ®ions, 16000.0), 1.0);
assert_eq!(remap_compressed_seconds(0.5, ®ions, 16000.0), 1.5);
}
#[test]
fn test_remap_second_region_skips_silence_gap() {
let regions = [(0usize, 16000usize), (48000usize, 64000usize)];
assert_eq!(remap_compressed_seconds(0.5, ®ions, 16000.0), 0.5);
assert_eq!(remap_compressed_seconds(1.5, ®ions, 16000.0), 3.5);
}
#[test]
fn test_remap_past_end_clamps_to_last_region_end() {
let regions = [(0usize, 16000usize), (48000usize, 64000usize)];
assert_eq!(remap_compressed_seconds(10.0, ®ions, 16000.0), 4.0);
}
#[test]
#[ignore = "requires the Silero VAD model at ~/.gigastt/models/vad/silero_vad.onnx"]
fn test_silero_silence_low_prob_and_runs() {
let home = std::env::var("HOME").expect("HOME");
let path = std::path::PathBuf::from(home).join(".gigastt/models/vad/silero_vad.onnx");
if !path.exists() {
eprintln!("skipping {}: Silero VAD model not present", path.display());
return;
}
let vad = SileroVad::load(&path).expect("load silero");
let silence = vec![0.0f32; 16000];
let probs = vad.frame_probs(&silence).expect("frame_probs");
assert!(!probs.is_empty(), "expected at least one frame");
for p in &probs {
assert!((0.0..=1.0).contains(p), "prob {p} out of range");
}
let max_silence = probs.iter().cloned().fold(0.0f32, f32::max);
assert!(
max_silence < 0.5,
"silence should be below threshold, got {max_silence}"
);
let tone: Vec<f32> = (0..16000)
.map(|i| 0.5 * (2.0 * std::f32::consts::PI * 200.0 * i as f32 / 16000.0).sin())
.collect();
let probs2 = vad.frame_probs(&tone).expect("frame_probs tone");
for p in &probs2 {
assert!((0.0..=1.0).contains(p), "tone prob {p} out of range");
}
assert!(
vad.speech_regions(&silence, &VadConfig::default())
.expect("regions")
.is_empty()
);
}
fn silero_model_path() -> std::path::PathBuf {
let home = std::env::var("HOME").expect("HOME");
std::path::PathBuf::from(home).join(".gigastt/models/vad/silero_vad.onnx")
}
#[test]
#[ignore = "requires the Silero VAD model at ~/.gigastt/models/vad/silero_vad.onnx"]
fn test_endpointer_buffers_subframe_chunks_across_pushes() {
let path = silero_model_path();
if !path.exists() {
eprintln!("skipping {}: Silero VAD model not present", path.display());
return;
}
let vad = SileroVad::load(&path).expect("load silero");
let c = VadConfig::default();
let mut ep = VadEndpointer::new(&c);
let part = vec![0.0f32; 200];
assert!(!ep.push(&vad, &part).expect("push part 1"));
assert!(!ep.push(&vad, &part).expect("push part 2"));
let rest = vec![0.0f32; 200];
assert!(!ep.push(&vad, &rest).expect("push part 3")); }
#[test]
#[ignore = "requires the Silero VAD model at ~/.gigastt/models/vad/silero_vad.onnx"]
fn test_endpointer_no_endpoint_on_leading_silence() {
let path = silero_model_path();
if !path.exists() {
eprintln!("skipping {}: Silero VAD model not present", path.display());
return;
}
let vad = SileroVad::load(&path).expect("load silero");
let c = VadConfig::default();
let mut ep = VadEndpointer::new(&c);
let silence = vec![0.0f32; 16000];
assert!(
!ep.push(&vad, &silence).expect("push silence"),
"leading silence must not endpoint"
);
assert!(!ep.push(&vad, &[]).expect("push empty"));
}
}