1#[derive(Debug, Clone)]
13pub struct VadConfig {
14 pub sample_rate: u32,
16 pub frame_duration_ms: u32,
18 pub start_threshold_db: f64,
20 pub stop_threshold_db: f64,
22 pub min_speech_frames: u32,
24 pub hangover_frames: u32,
26 pub speech_zcr_range: (f64, f64),
28 pub initial_noise_floor_db: f64,
30 pub pre_speech_frames: usize,
32}
33
34impl Default for VadConfig {
35 fn default() -> Self {
36 Self {
37 sample_rate: 16000,
38 frame_duration_ms: 30,
39 start_threshold_db: 15.0,
40 stop_threshold_db: 10.0,
41 min_speech_frames: 3,
42 hangover_frames: 10,
43 speech_zcr_range: (0.02, 0.5),
44 initial_noise_floor_db: -60.0,
45 pre_speech_frames: 3,
46 }
47 }
48}
49
50impl VadConfig {
51 pub fn frame_size(&self) -> usize {
53 (self.sample_rate * self.frame_duration_ms / 1000) as usize
54 }
55
56 pub fn noisy_street() -> Self {
71 Self {
72 start_threshold_db: 21.0,
73 stop_threshold_db: 16.0,
74 min_speech_frames: 1,
75 hangover_frames: 10,
76 ..Self::default()
77 }
78 }
79}
80
81#[derive(Debug, Clone, Copy, PartialEq, Eq)]
83pub enum VadState {
84 Silence,
86 PendingSpeech,
88 Speech,
90 Hangover,
92}
93
94#[derive(Debug, Clone, Copy, PartialEq, Eq)]
96pub enum VadEvent {
97 SpeechStart,
99 SpeechEnd,
101}
102
103pub struct VoiceActivityDetector {
105 config: VadConfig,
106 #[cfg(feature = "vad-wavekat")]
107 wavekat: Option<WaveKatWebRtcBackend>,
108 use_wavekat: bool,
109 state: VadState,
110 noise_floor_db: f64,
112 state_frames: u32,
114 noise_adapt_frames: u64,
116 last_probability: Option<f32>,
118 pre_speech_buf: Vec<Vec<i16>>,
120 pre_speech_idx: usize,
121}
122
123impl VoiceActivityDetector {
124 pub fn new(config: VadConfig) -> Self {
126 Self::new_with_backend(config, true)
127 }
128
129 fn new_with_backend(config: VadConfig, prefer_wavekat: bool) -> Self {
130 let frame_size = config.frame_size();
131 let pre_speech_buf: Vec<Vec<i16>> = (0..config.pre_speech_frames)
132 .map(|_| vec![0i16; frame_size])
133 .collect();
134 #[cfg(feature = "vad-wavekat")]
135 let wavekat = prefer_wavekat
136 .then(|| WaveKatWebRtcBackend::new(&config))
137 .flatten();
138 #[cfg(feature = "vad-wavekat")]
139 let use_wavekat = wavekat.is_some();
140 #[cfg(not(feature = "vad-wavekat"))]
141 let use_wavekat = {
142 let _ = prefer_wavekat;
143 false
144 };
145 Self {
146 noise_floor_db: config.initial_noise_floor_db,
147 state: VadState::Silence,
148 state_frames: 0,
149 noise_adapt_frames: 0,
150 last_probability: None,
151 pre_speech_buf,
152 pre_speech_idx: 0,
153 #[cfg(feature = "vad-wavekat")]
154 wavekat,
155 use_wavekat,
156 config,
157 }
158 }
159
160 #[cfg(test)]
161 fn new_energy(config: VadConfig) -> Self {
162 Self::new_with_backend(config, false)
163 }
164
165 pub fn state(&self) -> VadState {
167 self.state
168 }
169
170 pub fn is_speaking(&self) -> bool {
172 matches!(self.state, VadState::Speech | VadState::Hangover)
173 }
174
175 pub fn noise_floor_db(&self) -> f64 {
177 self.noise_floor_db
178 }
179
180 pub fn is_wavekat_backed(&self) -> bool {
182 self.use_wavekat
183 }
184
185 pub fn backend_name(&self) -> &'static str {
187 if self.use_wavekat {
188 "wavekat-webrtc"
189 } else {
190 "energy-zcr"
191 }
192 }
193
194 pub fn last_probability(&self) -> Option<f32> {
196 self.last_probability
197 }
198
199 pub fn drain_pre_speech(&mut self) -> Vec<Vec<i16>> {
201 let frame_size = self.config.frame_size();
202 let mut fresh: Vec<Vec<i16>> = (0..self.config.pre_speech_frames)
203 .map(|_| vec![0i16; frame_size])
204 .collect();
205 std::mem::swap(&mut self.pre_speech_buf, &mut fresh);
206 self.pre_speech_idx = 0;
207 fresh
208 }
209
210 pub fn process_frame(&mut self, samples: &[i16]) -> Option<VadEvent> {
212 let energy_db = compute_energy_db(samples);
213 let zcr = compute_zcr(samples);
214 let energy_above_noise = energy_db - self.noise_floor_db;
215
216 let wavekat_decision = self.wavekat_decision(samples);
217 let energy_speech_like = energy_above_noise > self.config.start_threshold_db
218 && zcr >= self.config.speech_zcr_range.0
219 && zcr <= self.config.speech_zcr_range.1;
220 let energy_above_stop = energy_above_noise > self.config.stop_threshold_db;
221 let is_speech_like = wavekat_decision.unwrap_or(energy_speech_like);
222 let is_above_stop = wavekat_decision.unwrap_or(energy_above_stop);
223 if wavekat_decision.is_none() {
224 self.last_probability = Some(if energy_speech_like { 1.0 } else { 0.0 });
225 }
226
227 match self.state {
228 VadState::Silence => {
229 self.update_noise_floor(energy_db);
231
232 if self.config.pre_speech_frames > 0 && !self.pre_speech_buf.is_empty() {
234 let idx = self.pre_speech_idx % self.config.pre_speech_frames;
235 let buf = &mut self.pre_speech_buf[idx];
236 buf.resize(samples.len(), 0);
237 buf.copy_from_slice(samples);
238 self.pre_speech_idx += 1;
239 }
240
241 if is_speech_like {
242 self.state_frames = 1;
243 if self.state_frames >= self.config.min_speech_frames {
245 self.state = VadState::Speech;
246 self.state_frames = 0;
247 return Some(VadEvent::SpeechStart);
248 }
249 self.state = VadState::PendingSpeech;
250 }
251 None
252 }
253
254 VadState::PendingSpeech => {
255 if is_speech_like {
256 self.state_frames += 1;
257 if self.state_frames >= self.config.min_speech_frames {
258 self.state = VadState::Speech;
259 self.state_frames = 0;
260 Some(VadEvent::SpeechStart)
261 } else {
262 None
263 }
264 } else {
265 self.state = VadState::Silence;
267 self.state_frames = 0;
268 None
269 }
270 }
271
272 VadState::Speech => {
273 if !is_above_stop {
274 self.state = VadState::Hangover;
275 self.state_frames = 1;
276 }
277 None
278 }
279
280 VadState::Hangover => {
281 if is_above_stop {
282 self.state = VadState::Speech;
284 self.state_frames = 0;
285 None
286 } else {
287 self.state_frames += 1;
288 if self.state_frames >= self.config.hangover_frames {
289 self.state = VadState::Silence;
290 self.state_frames = 0;
291 for buf in &mut self.pre_speech_buf {
292 buf.iter_mut().for_each(|s| *s = 0);
293 }
294 self.pre_speech_idx = 0;
295 Some(VadEvent::SpeechEnd)
296 } else {
297 None
298 }
299 }
300 }
301 }
302 }
303
304 fn update_noise_floor(&mut self, energy_db: f64) {
306 self.noise_adapt_frames += 1;
307 let alpha = 0.01_f64.min(1.0 / self.noise_adapt_frames as f64);
309 self.noise_floor_db = self.noise_floor_db * (1.0 - alpha) + energy_db * alpha;
310 }
311
312 #[allow(
313 clippy::unused_self,
314 reason = "self carries the backend and last probability once `vad-wavekat` is compiled in; \
315 the feature-off body is a stub behind the same call site"
316 )]
317 fn wavekat_decision(&mut self, samples: &[i16]) -> Option<bool> {
318 #[cfg(feature = "vad-wavekat")]
319 {
320 let probability = self
321 .wavekat
322 .as_mut()
323 .and_then(|backend| backend.process(samples, self.config.sample_rate));
324 self.last_probability = probability;
325 probability.map(|probability| probability >= 0.5)
326 }
327 #[cfg(not(feature = "vad-wavekat"))]
328 {
329 let _ = samples;
330 None
331 }
332 }
333
334 pub fn reset(&mut self) {
336 self.state = VadState::Silence;
337 self.state_frames = 0;
338 self.noise_adapt_frames = 0;
339 self.last_probability = None;
340 self.noise_floor_db = self.config.initial_noise_floor_db;
341 for buf in &mut self.pre_speech_buf {
342 buf.iter_mut().for_each(|s| *s = 0);
343 }
344 self.pre_speech_idx = 0;
345 }
346}
347
348#[cfg(feature = "vad-wavekat")]
349struct WaveKatWebRtcBackend {
350 detector: wavekat_vad::backends::webrtc::WebRtcVad,
351}
352
353#[cfg(feature = "vad-wavekat")]
354impl WaveKatWebRtcBackend {
355 fn new(config: &VadConfig) -> Option<Self> {
356 if !matches!(config.sample_rate, 8000 | 16000 | 32000 | 48000) {
357 return None;
358 }
359 if !matches!(config.frame_duration_ms, 10 | 20 | 30) {
360 return None;
361 }
362
363 let detector = wavekat_vad::backends::webrtc::WebRtcVad::with_frame_duration(
364 config.sample_rate,
365 wavekat_vad::backends::webrtc::WebRtcVadMode::Aggressive,
366 config.frame_duration_ms,
367 )
368 .ok()?;
369 Some(Self { detector })
370 }
371
372 fn process(&mut self, samples: &[i16], sample_rate: u32) -> Option<f32> {
373 use wavekat_vad::VoiceActivityDetector as _;
374
375 self.detector.process(samples, sample_rate).ok()
376 }
377}
378
379fn compute_energy_db(samples: &[i16]) -> f64 {
381 if samples.is_empty() {
382 return -96.0;
383 }
384
385 let sum_sq: f64 = samples.iter().map(|&s| (s as f64) * (s as f64)).sum();
386 let rms = (sum_sq / samples.len() as f64).sqrt();
387 let db = 20.0 * (rms / 32767.0).log10();
388 db.max(-96.0) }
390
391fn compute_zcr(samples: &[i16]) -> f64 {
393 if samples.len() < 2 {
394 return 0.0;
395 }
396
397 let crossings = samples
398 .windows(2)
399 .filter(|w| (w[0] >= 0) != (w[1] >= 0))
400 .count();
401
402 crossings as f64 / (samples.len() - 1) as f64
403}
404
405#[cfg(test)]
406mod tests {
407 use super::*;
408
409 fn make_vad() -> VoiceActivityDetector {
410 VoiceActivityDetector::new_energy(VadConfig {
411 sample_rate: 16000,
412 frame_duration_ms: 20,
413 start_threshold_db: 15.0,
414 stop_threshold_db: 10.0,
415 min_speech_frames: 2,
416 hangover_frames: 3,
417 speech_zcr_range: (0.01, 0.9),
418 initial_noise_floor_db: -60.0,
419 pre_speech_frames: 2,
420 })
421 }
422
423 fn silence_frame(len: usize) -> Vec<i16> {
424 vec![0i16; len]
425 }
426
427 fn speech_frame(len: usize, amplitude: i16) -> Vec<i16> {
428 (0..len)
430 .map(|i| if i % 4 < 2 { amplitude } else { -amplitude })
431 .collect()
432 }
433
434 #[test]
435 fn starts_silent() {
436 let vad = make_vad();
437 assert_eq!(vad.state(), VadState::Silence);
438 assert!(!vad.is_speaking());
439 }
440
441 #[cfg(feature = "vad-wavekat")]
442 #[test]
443 fn default_detector_uses_wavekat_for_supported_frames() {
444 let vad = VoiceActivityDetector::new(VadConfig {
445 sample_rate: 16000,
446 frame_duration_ms: 20,
447 ..VadConfig::default()
448 });
449 assert!(vad.is_wavekat_backed());
450 }
451
452 #[cfg(feature = "vad-wavekat")]
453 #[test]
454 fn unsupported_frames_fall_back_to_energy_detector() {
455 let vad = VoiceActivityDetector::new(VadConfig {
456 sample_rate: 16000,
457 frame_duration_ms: 32,
458 ..VadConfig::default()
459 });
460 assert!(!vad.is_wavekat_backed());
461 }
462
463 #[test]
464 fn silence_stays_silent() {
465 let mut vad = make_vad();
466 let frame = silence_frame(320);
467 for _ in 0..10 {
468 let event = vad.process_frame(&frame);
469 assert!(event.is_none());
470 }
471 assert_eq!(vad.state(), VadState::Silence);
472 }
473
474 #[test]
475 fn speech_detected_after_min_frames() {
476 let mut vad = make_vad();
477 let frame = speech_frame(320, 10000);
478
479 let e1 = vad.process_frame(&frame);
481 assert!(e1.is_none());
482 assert_eq!(vad.state(), VadState::PendingSpeech);
483
484 let e2 = vad.process_frame(&frame);
486 assert_eq!(e2, Some(VadEvent::SpeechStart));
487 assert_eq!(vad.state(), VadState::Speech);
488 assert!(vad.is_speaking());
489 }
490
491 #[test]
492 fn speech_end_after_hangover() {
493 let mut vad = make_vad();
494 let speech = speech_frame(320, 10000);
495 let silence = silence_frame(320);
496
497 vad.process_frame(&speech);
499 vad.process_frame(&speech);
500 assert_eq!(vad.state(), VadState::Speech);
501
502 vad.process_frame(&silence);
504 assert_eq!(vad.state(), VadState::Hangover);
505
506 vad.process_frame(&silence);
508 let e = vad.process_frame(&silence);
509 assert_eq!(e, Some(VadEvent::SpeechEnd));
510 assert_eq!(vad.state(), VadState::Silence);
511 }
512
513 #[test]
514 fn speech_resumes_during_hangover() {
515 let mut vad = make_vad();
516 let speech = speech_frame(320, 10000);
517 let silence = silence_frame(320);
518
519 vad.process_frame(&speech);
521 vad.process_frame(&speech);
522 assert_eq!(vad.state(), VadState::Speech);
523
524 vad.process_frame(&silence);
526 assert_eq!(vad.state(), VadState::Hangover);
527
528 let e = vad.process_frame(&speech);
530 assert!(e.is_none()); assert_eq!(vad.state(), VadState::Speech);
532 }
533
534 #[test]
535 fn false_alarm_returns_to_silence() {
536 let mut vad = make_vad();
537 let speech = speech_frame(320, 10000);
538 let silence = silence_frame(320);
539
540 vad.process_frame(&speech);
542 assert_eq!(vad.state(), VadState::PendingSpeech);
543
544 vad.process_frame(&silence);
546 assert_eq!(vad.state(), VadState::Silence);
547 }
548
549 #[test]
550 fn energy_db_calculation() {
551 let full_scale: Vec<i16> = (0..320).map(|_| i16::MAX).collect();
553 let db = compute_energy_db(&full_scale);
554 assert!(db > -1.0); let silence = vec![0i16; 320];
557 let db_silence = compute_energy_db(&silence);
558 assert_eq!(db_silence, -96.0);
559 }
560
561 #[test]
562 fn zcr_calculation() {
563 let alternating: Vec<i16> = (0..100)
565 .map(|i| if i % 2 == 0 { 1000 } else { -1000 })
566 .collect();
567 let zcr = compute_zcr(&alternating);
568 assert!(zcr > 0.9);
569
570 let constant = vec![1000i16; 100];
572 let zcr_const = compute_zcr(&constant);
573 assert_eq!(zcr_const, 0.0);
574 }
575
576 #[test]
577 fn noise_floor_adapts() {
578 let mut vad = make_vad();
579 let low_noise: Vec<i16> = vec![10; 320]; for _ in 0..100 {
582 vad.process_frame(&low_noise);
583 }
584 assert!(vad.noise_floor_db() > -96.0);
586 }
587
588 #[test]
589 fn reset_clears_state() {
590 let mut vad = make_vad();
591 let speech = speech_frame(320, 10000);
592 vad.process_frame(&speech);
593 vad.process_frame(&speech);
594 assert_eq!(vad.state(), VadState::Speech);
595
596 vad.reset();
597 assert_eq!(vad.state(), VadState::Silence);
598 assert_eq!(vad.noise_floor_db(), -60.0);
599 }
600
601 #[test]
602 fn noisy_street_preset_is_stricter_and_faster_than_default() {
603 let preset = VadConfig::noisy_street();
604 let default = VadConfig::default();
605 assert!(preset.start_threshold_db > default.start_threshold_db);
607 assert!(preset.stop_threshold_db < preset.start_threshold_db);
608 assert!(preset.min_speech_frames < default.min_speech_frames);
610 assert!(preset.min_speech_frames >= 1);
611 assert_eq!(preset.frame_size(), default.frame_size());
613 }
614
615 #[test]
616 fn pending_speech_emits_on_first_frame_when_min_speech_frames_one() {
617 let mut vad = VoiceActivityDetector::new_energy(VadConfig {
618 sample_rate: 16000,
619 frame_duration_ms: 20,
620 start_threshold_db: 15.0,
621 stop_threshold_db: 10.0,
622 min_speech_frames: 1,
623 hangover_frames: 3,
624 speech_zcr_range: (0.01, 0.9),
625 initial_noise_floor_db: -60.0,
626 pre_speech_frames: 2,
627 });
628
629 let frame = speech_frame(320, 10000);
630
631 let e = vad.process_frame(&frame);
633 assert_eq!(
634 e,
635 Some(VadEvent::SpeechStart),
636 "Expected SpeechStart on first speech frame when min_speech_frames=1"
637 );
638 assert_eq!(vad.state(), VadState::Speech);
639 }
640
641 #[test]
642 fn hangover_exits_after_min_frames_when_hangover_frames_one() {
643 let mut vad = VoiceActivityDetector::new_energy(VadConfig {
644 sample_rate: 16000,
645 frame_duration_ms: 20,
646 start_threshold_db: 15.0,
647 stop_threshold_db: 10.0,
648 min_speech_frames: 2,
649 hangover_frames: 1,
650 speech_zcr_range: (0.01, 0.9),
651 initial_noise_floor_db: -60.0,
652 pre_speech_frames: 2,
653 });
654
655 let speech = speech_frame(320, 10000);
656 let silence = silence_frame(320);
657
658 vad.process_frame(&speech);
660 assert_eq!(vad.state(), VadState::PendingSpeech);
661 vad.process_frame(&speech);
662 assert_eq!(vad.state(), VadState::Speech);
663
664 let e1 = vad.process_frame(&silence);
666 assert!(
667 e1.is_none(),
668 "Should not emit SpeechEnd when entering hangover"
669 );
670 assert_eq!(vad.state(), VadState::Hangover);
671
672 let e2 = vad.process_frame(&silence);
675 assert_eq!(
676 e2,
677 Some(VadEvent::SpeechEnd),
678 "Expected SpeechEnd after 1 hangover frame when hangover_frames=1"
679 );
680 assert_eq!(vad.state(), VadState::Silence);
681 }
682}