use crate::error::VadError;
use crate::{ProcessTimings, VadCapabilities, VoiceActivityDetector};
use earshot::{DefaultPredictor, Detector};
use std::time::{Duration, Instant};
pub const SAMPLE_RATE: u32 = 16_000;
pub const FRAME_SIZE: usize = 256;
pub const FRAME_DURATION_MS: u32 = 16;
pub struct EarshotVad {
detector: Box<Detector<DefaultPredictor>>,
inference_time: Duration,
frames: u64,
}
impl EarshotVad {
pub fn new() -> Self {
Self {
detector: Detector::default_boxed(),
inference_time: Duration::ZERO,
frames: 0,
}
}
}
impl Default for EarshotVad {
fn default() -> Self {
Self::new()
}
}
impl VoiceActivityDetector for EarshotVad {
fn capabilities(&self) -> VadCapabilities {
VadCapabilities {
sample_rate: SAMPLE_RATE,
frame_size: FRAME_SIZE,
frame_duration_ms: FRAME_DURATION_MS,
}
}
fn process(&mut self, samples: &[i16], sample_rate: u32) -> Result<f32, VadError> {
if sample_rate != SAMPLE_RATE {
return Err(VadError::InvalidSampleRate(sample_rate));
}
if samples.len() != FRAME_SIZE {
return Err(VadError::InvalidFrameSize {
got: samples.len(),
expected: FRAME_SIZE,
});
}
let start = Instant::now();
let score = self.detector.predict_i16(samples);
self.inference_time += start.elapsed();
self.frames += 1;
Ok(score)
}
fn reset(&mut self) {
self.detector.reset();
}
fn timings(&self) -> ProcessTimings {
ProcessTimings {
stages: vec![("inference", self.inference_time)],
frames: self.frames,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn tone(offset: usize, len: usize) -> Vec<i16> {
(offset..offset + len)
.map(|n| {
let t = n as f32 / SAMPLE_RATE as f32;
let v = (2.0 * std::f32::consts::PI * 220.0 * t).sin() * 0.4
+ (2.0 * std::f32::consts::PI * 700.0 * t).sin() * 0.25
+ (2.0 * std::f32::consts::PI * 1900.0 * t).sin() * 0.1;
(v * 12000.0) as i16
})
.collect()
}
#[test]
fn capabilities_are_16khz_256_samples() {
let vad = EarshotVad::new();
assert_eq!(
vad.capabilities(),
VadCapabilities {
sample_rate: 16_000,
frame_size: 256,
frame_duration_ms: 16,
}
);
}
#[test]
fn wrong_sample_rate_is_rejected() {
let mut vad = EarshotVad::new();
let frame = vec![0i16; FRAME_SIZE];
for rate in [8_000u32, 32_000, 44_100, 48_000, 0] {
let err = vad.process(&frame, rate).unwrap_err();
assert!(
matches!(err, VadError::InvalidSampleRate(r) if r == rate),
"rate {rate} produced {err:?}"
);
}
}
#[test]
fn wrong_frame_length_is_rejected() {
let mut vad = EarshotVad::new();
for len in [0usize, 1, 160, 255, 257, 320, 512] {
let err = vad.process(&vec![0i16; len], SAMPLE_RATE).unwrap_err();
match err {
VadError::InvalidFrameSize { got, expected } => {
assert_eq!(got, len);
assert_eq!(expected, FRAME_SIZE);
}
other => panic!("len {len} produced {other:?}"),
}
}
}
#[test]
fn sample_rate_is_checked_before_frame_size() {
let mut vad = EarshotVad::new();
let err = vad.process(&[0i16; 10], 48_000).unwrap_err();
assert!(
matches!(err, VadError::InvalidSampleRate(48_000)),
"{err:?}"
);
}
#[test]
fn scores_are_finite_and_in_unit_range() {
let mut vad = EarshotVad::new();
let audio = tone(0, FRAME_SIZE * 40);
for (i, frame) in audio.chunks_exact(FRAME_SIZE).enumerate() {
let score = vad.process(frame, SAMPLE_RATE).unwrap();
assert!(score.is_finite(), "frame {i} score {score} is not finite");
assert!(
(0.0..=1.0).contains(&score),
"frame {i} score {score} outside 0.0..=1.0"
);
}
let mut vad = EarshotVad::new();
for _ in 0..40 {
let score = vad.process(&[0i16; FRAME_SIZE], SAMPLE_RATE).unwrap();
assert!(score.is_finite() && (0.0..=1.0).contains(&score), "{score}");
}
}
#[test]
fn reset_restores_a_fresh_detectors_output_sequence() {
let audio = tone(0, FRAME_SIZE * 24);
let run = |vad: &mut EarshotVad| -> Vec<f32> {
audio
.chunks_exact(FRAME_SIZE)
.map(|f| vad.process(f, SAMPLE_RATE).unwrap())
.collect()
};
let mut fresh = EarshotVad::new();
let expected = run(&mut fresh);
let mut reused = EarshotVad::new();
let noise = tone(7_777, FRAME_SIZE * 13);
for f in noise.chunks_exact(FRAME_SIZE) {
reused.process(f, SAMPLE_RATE).unwrap();
}
let dirty = run(&mut reused);
assert_ne!(
dirty, expected,
"test is vacuous: carried state did not change the output"
);
reused.reset();
let after_reset = run(&mut reused);
assert_eq!(
after_reset, expected,
"reset() did not restore fresh-detector behaviour"
);
}
#[test]
fn instances_are_state_isolated() {
let audio_a = tone(0, FRAME_SIZE * 16);
let audio_b = tone(31_337, FRAME_SIZE * 16);
let mut solo_a = EarshotVad::new();
let expected_a: Vec<f32> = audio_a
.chunks_exact(FRAME_SIZE)
.map(|f| solo_a.process(f, SAMPLE_RATE).unwrap())
.collect();
let mut solo_b = EarshotVad::new();
let expected_b: Vec<f32> = audio_b
.chunks_exact(FRAME_SIZE)
.map(|f| solo_b.process(f, SAMPLE_RATE).unwrap())
.collect();
let mut vad_a = EarshotVad::new();
let mut vad_b = EarshotVad::new();
let mut got_a = Vec::new();
let mut got_b = Vec::new();
for (fa, fb) in audio_a
.chunks_exact(FRAME_SIZE)
.zip(audio_b.chunks_exact(FRAME_SIZE))
{
got_a.push(vad_a.process(fa, SAMPLE_RATE).unwrap());
got_b.push(vad_b.process(fb, SAMPLE_RATE).unwrap());
}
assert_eq!(got_a, expected_a, "detector A was perturbed by detector B");
assert_eq!(got_b, expected_b, "detector B was perturbed by detector A");
assert_ne!(
expected_a, expected_b,
"test is vacuous: the two streams score identically"
);
}
#[test]
fn timings_accumulate_and_survive_reset() {
let mut vad = EarshotVad::new();
assert_eq!(vad.timings().frames, 0);
for _ in 0..5 {
vad.process(&[0i16; FRAME_SIZE], SAMPLE_RATE).unwrap();
}
let t = vad.timings();
assert_eq!(t.frames, 5);
assert_eq!(t.stages.len(), 1);
assert_eq!(t.stages[0].0, "inference");
let _ = vad.process(&[0i16; 10], SAMPLE_RATE);
let _ = vad.process(&[0i16; FRAME_SIZE], 8_000);
assert_eq!(vad.timings().frames, 5);
vad.reset();
assert_eq!(vad.timings().frames, 5, "reset() must not clear timings");
}
#[test]
fn works_behind_the_frame_adapter() {
use crate::FrameAdapter;
let mut adapter = FrameAdapter::new(Box::new(EarshotVad::new()));
assert_eq!(adapter.frame_size(), FRAME_SIZE);
assert_eq!(adapter.sample_rate(), SAMPLE_RATE);
let audio = tone(0, 320 * 20);
let mut scored = 0usize;
for chunk in audio.chunks(320) {
adapter
.process_each(chunk, SAMPLE_RATE, |s| {
assert!(s.is_finite() && (0.0..=1.0).contains(&s), "{s}");
scored += 1;
})
.unwrap();
assert!(adapter.buffered_samples() < FRAME_SIZE);
}
assert_eq!(scored, (320 * 20) / FRAME_SIZE);
}
fn mean_score(audio: &[i16]) -> f32 {
let mut vad = EarshotVad::new();
let scores: Vec<f32> = audio
.chunks_exact(FRAME_SIZE)
.map(|f| vad.process(f, SAMPLE_RATE).unwrap())
.collect();
scores.iter().sum::<f32>() / scores.len() as f32
}
#[test]
fn a_signal_scores_clearly_above_digital_silence() {
let signal = mean_score(&tone(0, FRAME_SIZE * 40));
let silence = mean_score(&[0i16; FRAME_SIZE * 40]);
assert!(
signal > silence + 0.2,
"signal {signal:.3} is not separated from silence {silence:.3}"
);
assert!(
silence < 0.5,
"digital silence scored {silence:.3}, at or above the suggested threshold"
);
}
#[test]
fn silence_never_crosses_the_suggested_threshold() {
let mut vad = EarshotVad::new();
for i in 0..200 {
let score = vad.process(&[0i16; FRAME_SIZE], SAMPLE_RATE).unwrap();
assert!(score < 0.5, "silent frame {i} scored {score:.3}");
}
}
#[test]
fn rejected_frames_leave_the_stream_state_untouched() {
let audio = tone(0, FRAME_SIZE * 20);
let mut clean = EarshotVad::new();
let expected: Vec<f32> = audio
.chunks_exact(FRAME_SIZE)
.map(|f| clean.process(f, SAMPLE_RATE).unwrap())
.collect();
let mut interleaved = EarshotVad::new();
let mut got = Vec::new();
for frame in audio.chunks_exact(FRAME_SIZE) {
assert!(interleaved.process(frame, 8_000).is_err());
assert!(interleaved.process(&frame[..100], SAMPLE_RATE).is_err());
assert!(interleaved.process(&[0i16; 512], 44_100).is_err());
got.push(interleaved.process(frame, SAMPLE_RATE).unwrap());
}
assert_eq!(got, expected, "a rejected frame perturbed the stream");
}
#[test]
fn default_matches_new() {
let audio = tone(0, FRAME_SIZE * 12);
let run = |mut vad: EarshotVad| -> Vec<f32> {
audio
.chunks_exact(FRAME_SIZE)
.map(|f| vad.process(f, SAMPLE_RATE).unwrap())
.collect()
};
assert_eq!(run(EarshotVad::default()), run(EarshotVad::new()));
}
#[test]
fn two_fresh_detectors_score_identically() {
let audio = tone(1_234, FRAME_SIZE * 20);
let run = || -> Vec<f32> {
let mut vad = EarshotVad::new();
audio
.chunks_exact(FRAME_SIZE)
.map(|f| vad.process(f, SAMPLE_RATE).unwrap())
.collect()
};
assert_eq!(run(), run());
}
#[test]
fn context_carries_across_frame_boundaries() {
let audio = tone(0, FRAME_SIZE * 16);
let mut streaming = EarshotVad::new();
let continuous: Vec<f32> = audio
.chunks_exact(FRAME_SIZE)
.map(|f| streaming.process(f, SAMPLE_RATE).unwrap())
.collect();
let isolated: Vec<f32> = audio
.chunks_exact(FRAME_SIZE)
.map(|f| EarshotVad::new().process(f, SAMPLE_RATE).unwrap())
.collect();
assert_ne!(
continuous, isolated,
"per-stream context is not affecting the scores"
);
assert_eq!(continuous[0], isolated[0]);
}
}