diff --git a/.github/scripts/test-rust.sh b/.github/scripts/test-rust.sh index f27637072e..88af9417b3 100755 --- a/.github/scripts/test-rust.sh +++ b/.github/scripts/test-rust.sh @@ -6,6 +6,8 @@ cd rust-api-examples ./run-version.sh +./run-silero-vad-remove-silence.sh + ./run-nemo-parakeet-en.sh ./run-zipformer-vi.sh ./run-zipformer-zh-en.sh diff --git a/rust-api-examples/Cargo.lock b/rust-api-examples/Cargo.lock index c26794bb73..d81afe9d4d 100644 --- a/rust-api-examples/Cargo.lock +++ b/rust-api-examples/Cargo.lock @@ -553,7 +553,7 @@ dependencies = [ [[package]] name = "rust-api-examples" -version = "0.1.6" +version = "0.1.7" dependencies = [ "anyhow", "clap", @@ -621,9 +621,9 @@ dependencies = [ [[package]] name = "sherpa-onnx" -version = "0.1.6" +version = "0.1.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "df1d4facedb9950eb43a3788912ebb9b2118256195f0dcc00206b1e76aa79ef7" +checksum = "79039b9e40380dd7de0fae149c38bc9c53f5c48033eb709d8a2024b95e9f957d" dependencies = [ "serde", "serde_json", @@ -632,9 +632,9 @@ dependencies = [ [[package]] name = "sherpa-onnx-sys" -version = "0.1.6" +version = "0.1.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bb4014275b85b5a4437076dd0919d74ca7641648092d2f45a48b975c5befa873" +checksum = "58a81d43817344f318a9320c6c72a8a9fb4e7cf1f9a372b831a63d9a253fe672" [[package]] name = "slab" diff --git a/rust-api-examples/Cargo.toml b/rust-api-examples/Cargo.toml index 35c87eb75e..6a1fa4c970 100644 --- a/rust-api-examples/Cargo.toml +++ b/rust-api-examples/Cargo.toml @@ -1,12 +1,12 @@ [package] name = "rust-api-examples" -version = "0.1.6" +version = "0.1.7" edition = "2021" [dependencies] anyhow = "1.0" clap = { version = "4.5", features = ["derive"] } -sherpa-onnx = "0.1.6" +sherpa-onnx = "0.1.7" # sherpa-onnx = { path = "../sherpa-onnx/rust/sherpa-onnx" } cpal = { version = "0.16", optional = true } # cross-platform audio I/O diff --git a/rust-api-examples/README.md b/rust-api-examples/README.md index 1791890e1f..21b8aca587 100644 --- a/rust-api-examples/README.md +++ b/rust-api-examples/README.md @@ -92,6 +92,19 @@ cargo run --example sense_voice -- \ --tokens ./sherpa-onnx-sense-voice-zh-en-ja-ko-yue-int8-2024-07-17/tokens.txt ``` +### Example 5: Remove silences from a file using SileroVAD + +```bash + +curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/silero_vad.onnx +curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/lei-jun-test.wav + +cargo run --example silero_vad_remove_silence -- \ + --input ./lei-jun-test.wav \ + --output ./no-silence.wav \ + --silero-vad-model ./silero_vad.onnx +``` + # Alternative rust bindings for sherpa-onnx Please see also https://github.com/thewh1teagle/sherpa-rs diff --git a/rust-api-examples/examples/silero_vad_remove_silence.rs b/rust-api-examples/examples/silero_vad_remove_silence.rs new file mode 100644 index 0000000000..ab2c3918e0 --- /dev/null +++ b/rust-api-examples/examples/silero_vad_remove_silence.rs @@ -0,0 +1,109 @@ +// Copyright (c) 2026 Xiaomi Corporation +// +// This file demonstrates how to use silero VAD with sherpa-onnx's +// Rust API to remove non-speech segments and save speech-only audio. +// +// See ../README.md for how to run it + +use clap::Parser; +use sherpa_onnx::{self, SileroVadModelConfig, VadModelConfig, VoiceActivityDetector, Wave}; + +/// Simple VAD example: remove non-speech segments from a WAV file +#[derive(Parser, Debug)] +#[command(author, version, about, long_about = None)] +struct Args { + /// Path to input WAV file + #[arg(long)] + input: String, + + /// Path to output WAV file + #[arg(long)] + output: String, + + /// Path to Silero VAD ONNX model + #[arg(long)] + silero_vad_model: String, +} + +fn main() -> anyhow::Result<()> { + let args = Args::parse(); + + // Read WAV file + let wave = Wave::read(&args.input) + .ok_or_else(|| anyhow::anyhow!("Failed to read WAV file: {}", &args.input))?; + let sample_rate = wave.sample_rate(); + let input_num_samples = wave.num_samples(); + let input_duration = input_num_samples as f32 / sample_rate as f32; + + println!( + "Input WAV: sample rate: {}, num samples: {}, duration: {:.2}s", + sample_rate, input_num_samples, input_duration + ); + + // Configure VAD + let mut silero_config = SileroVadModelConfig::default(); + silero_config.model = Some(args.silero_vad_model); + + // You can tune the values below + silero_config.threshold = 0.5; + silero_config.min_silence_duration = 0.25; + silero_config.min_speech_duration = 0.25; + silero_config.max_speech_duration = 5.0; + + let vad_config = VadModelConfig { + silero_vad: silero_config, + ten_vad: Default::default(), + sample_rate, + num_threads: 1, + provider: Some("cpu".to_string()), + debug: false, + }; + + let vad = VoiceActivityDetector::create(&vad_config, 30.0) + .expect("Failed to create VoiceActivityDetector"); + + let mut speech_samples = Vec::new(); + const WINDOW_SIZE: usize = 512; + + for chunk in wave.samples().chunks(WINDOW_SIZE) { + vad.accept_waveform(chunk); + + while let Some(seg) = vad.front() { + speech_samples.extend_from_slice(seg.samples()); + vad.pop(); + } + } + + vad.flush(); + while let Some(seg) = vad.front() { + speech_samples.extend_from_slice(seg.samples()); + vad.pop(); + } + + // Write speech-only samples to output WAV + let ok = sherpa_onnx::write(&args.output, &speech_samples, sample_rate); + if ok { + println!("Saved speech-only audio to {}", args.output); + } else { + println!("Failed to save speech-only audio to {}", args.output); + } + + // Summary + let output_num_samples = speech_samples.len(); + let output_duration = output_num_samples as f32 / sample_rate as f32; + println!("\n=== Summary ==="); + println!( + "Input: sample rate = {}, samples = {}, duration = {:.2}s", + sample_rate, input_num_samples, input_duration + ); + println!( + "Output: sample rate = {}, samples = {}, duration = {:.2}s", + sample_rate, output_num_samples, output_duration + ); + println!( + "Removed non-speech: {:.2}% of input removed", + 100.0 * (1.0 - output_duration / input_duration) + ); + + Ok(()) +} diff --git a/rust-api-examples/run-silero-vad-remove-silence.sh b/rust-api-examples/run-silero-vad-remove-silence.sh new file mode 100755 index 0000000000..eaff8cd87c --- /dev/null +++ b/rust-api-examples/run-silero-vad-remove-silence.sh @@ -0,0 +1,16 @@ +#!/usr/bin/env bash +set -ex + +# https://k2-fsa.github.io/sherpa/onnx/vad/silero-vad.html +if [ ! -f "./silero_vad.onnx" ]; then + curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/silero_vad.onnx +fi + +if [ ! -f ./lei-jun-test.wav ]; then + curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/lei-jun-test.wav +fi + +cargo run --example silero_vad_remove_silence -- \ + --input ./lei-jun-test.wav \ + --output ./no-silence.wav \ + --silero-vad-model ./silero_vad.onnx diff --git a/sherpa-onnx/c-api/c-api.cc b/sherpa-onnx/c-api/c-api.cc index b7ca93572b..315f38eeb6 100644 --- a/sherpa-onnx/c-api/c-api.cc +++ b/sherpa-onnx/c-api/c-api.cc @@ -1128,7 +1128,7 @@ struct SherpaOnnxVoiceActivityDetector { std::unique_ptr impl; }; -sherpa_onnx::VadModelConfig GetVadModelConfig( +static sherpa_onnx::VadModelConfig GetVadModelConfig( const SherpaOnnxVadModelConfig *config) { sherpa_onnx::VadModelConfig vad_config; @@ -1185,6 +1185,11 @@ sherpa_onnx::VadModelConfig GetVadModelConfig( const SherpaOnnxVoiceActivityDetector *SherpaOnnxCreateVoiceActivityDetector( const SherpaOnnxVadModelConfig *config, float buffer_size_in_seconds) { + if (!config) { + SHERPA_ONNX_LOGE("vad config is nullptr"); + return nullptr; + } + auto vad_config = GetVadModelConfig(config); if (!vad_config.Validate()) { @@ -1206,31 +1211,70 @@ void SherpaOnnxDestroyVoiceActivityDetector( void SherpaOnnxVoiceActivityDetectorAcceptWaveform( const SherpaOnnxVoiceActivityDetector *p, const float *samples, int32_t n) { + if (!p) { + SHERPA_ONNX_LOGE("vad is nullptr"); + return; + } + + if (!samples) { + SHERPA_ONNX_LOGE("samples is nullptr"); + return; + } + p->impl->AcceptWaveform(samples, n); } int32_t SherpaOnnxVoiceActivityDetectorEmpty( const SherpaOnnxVoiceActivityDetector *p) { + if (!p) { + SHERPA_ONNX_LOGE("vad is nullptr"); + return 1; // 1 means it is empty + } + return p->impl->Empty(); } int32_t SherpaOnnxVoiceActivityDetectorDetected( const SherpaOnnxVoiceActivityDetector *p) { + if (!p) { + SHERPA_ONNX_LOGE("vad is nullptr"); + return 0; + } + return p->impl->IsSpeechDetected(); } void SherpaOnnxVoiceActivityDetectorPop( const SherpaOnnxVoiceActivityDetector *p) { + if (!p) { + SHERPA_ONNX_LOGE("vad is nullptr"); + return; + } + p->impl->Pop(); } void SherpaOnnxVoiceActivityDetectorClear( const SherpaOnnxVoiceActivityDetector *p) { + if (!p) { + SHERPA_ONNX_LOGE("vad is nullptr"); + return; + } + p->impl->Clear(); } const SherpaOnnxSpeechSegment *SherpaOnnxVoiceActivityDetectorFront( const SherpaOnnxVoiceActivityDetector *p) { + if (!p) { + SHERPA_ONNX_LOGE("vad is nullptr"); + return nullptr; + } + + if (SherpaOnnxVoiceActivityDetectorEmpty(p)) { + return nullptr; + } + const sherpa_onnx::SpeechSegment &segment = p->impl->Front(); SherpaOnnxSpeechSegment *ans = new SherpaOnnxSpeechSegment; @@ -1251,11 +1295,21 @@ void SherpaOnnxDestroySpeechSegment(const SherpaOnnxSpeechSegment *p) { void SherpaOnnxVoiceActivityDetectorReset( const SherpaOnnxVoiceActivityDetector *p) { + if (!p) { + SHERPA_ONNX_LOGE("vad is nullptr"); + return; + } + p->impl->Reset(); } void SherpaOnnxVoiceActivityDetectorFlush( const SherpaOnnxVoiceActivityDetector *p) { + if (!p) { + SHERPA_ONNX_LOGE("vad is nullptr"); + return; + } + p->impl->Flush(); } diff --git a/sherpa-onnx/csrc/voice-activity-detector.cc b/sherpa-onnx/csrc/voice-activity-detector.cc index c82a77e3c2..e3c01077b4 100644 --- a/sherpa-onnx/csrc/voice-activity-detector.cc +++ b/sherpa-onnx/csrc/voice-activity-detector.cc @@ -141,7 +141,18 @@ class VoiceActivityDetector::Impl { void Clear() { std::queue().swap(segments_); } - const SpeechSegment &Front() const { return segments_.front(); } + const SpeechSegment &Front() const { + static SpeechSegment tmp; + + if (Empty()) { + SHERPA_ONNX_LOGE( + "Make sure you call this method only when Empty() returns false; " + "Return an empty segment"); + return tmp; + } + + return segments_.front(); + } void Reset() { std::queue().swap(segments_); diff --git a/sherpa-onnx/rust/sherpa-onnx-sys/Cargo.toml b/sherpa-onnx/rust/sherpa-onnx-sys/Cargo.toml index e73bf8f263..10fca8f007 100644 --- a/sherpa-onnx/rust/sherpa-onnx-sys/Cargo.toml +++ b/sherpa-onnx/rust/sherpa-onnx-sys/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "sherpa-onnx-sys" -version = "0.1.6" +version = "0.1.7" edition = "2021" description = "Raw FFI bindings to the sherpa-onnx C API" license = "Apache-2.0" diff --git a/sherpa-onnx/rust/sherpa-onnx-sys/src/lib.rs b/sherpa-onnx/rust/sherpa-onnx-sys/src/lib.rs index 596c7fad34..13efe88293 100644 --- a/sherpa-onnx/rust/sherpa-onnx-sys/src/lib.rs +++ b/sherpa-onnx/rust/sherpa-onnx-sys/src/lib.rs @@ -13,8 +13,10 @@ extern "C" { pub mod offline_asr; pub mod online_asr; +pub mod vad; pub mod wave; pub use offline_asr::*; pub use online_asr::*; +pub use vad::*; pub use wave::*; diff --git a/sherpa-onnx/rust/sherpa-onnx-sys/src/vad.rs b/sherpa-onnx/rust/sherpa-onnx-sys/src/vad.rs new file mode 100644 index 0000000000..b16b21deb9 --- /dev/null +++ b/sherpa-onnx/rust/sherpa-onnx-sys/src/vad.rs @@ -0,0 +1,85 @@ +use std::os::raw::{c_char, c_float}; + +#[repr(C)] +pub struct SileroVadModelConfig { + pub model: *const c_char, + pub threshold: c_float, + pub min_silence_duration: c_float, + pub min_speech_duration: c_float, + pub window_size: i32, + pub max_speech_duration: c_float, +} + +#[repr(C)] +pub struct TenVadModelConfig { + pub model: *const c_char, + pub threshold: c_float, + pub min_silence_duration: c_float, + pub min_speech_duration: c_float, + pub window_size: i32, + pub max_speech_duration: c_float, +} + +#[repr(C)] +pub struct VadModelConfig { + pub silero_vad: SileroVadModelConfig, + pub sample_rate: i32, + pub num_threads: i32, + pub provider: *const c_char, + pub debug: i32, + pub ten_vad: TenVadModelConfig, +} + +#[repr(C)] +pub struct CircularBuffer { + _private: [u8; 0], +} + +#[repr(C)] +pub struct SpeechSegment { + pub start: i32, + pub samples: *mut f32, + pub n: i32, +} + +#[repr(C)] +pub struct VoiceActivityDetector { + _private: [u8; 0], +} + +extern "C" { + pub fn SherpaOnnxCreateCircularBuffer(capacity: i32) -> *const CircularBuffer; + pub fn SherpaOnnxDestroyCircularBuffer(buffer: *const CircularBuffer); + pub fn SherpaOnnxCircularBufferPush(buffer: *const CircularBuffer, p: *const f32, n: i32); + pub fn SherpaOnnxCircularBufferGet( + buffer: *const CircularBuffer, + start_index: i32, + n: i32, + ) -> *const f32; + pub fn SherpaOnnxCircularBufferFree(p: *const f32); + pub fn SherpaOnnxCircularBufferPop(buffer: *const CircularBuffer, n: i32); + pub fn SherpaOnnxCircularBufferSize(buffer: *const CircularBuffer) -> i32; + pub fn SherpaOnnxCircularBufferHead(buffer: *const CircularBuffer) -> i32; + pub fn SherpaOnnxCircularBufferReset(buffer: *const CircularBuffer); + + pub fn SherpaOnnxCreateVoiceActivityDetector( + config: *const VadModelConfig, + buffer_size_in_seconds: c_float, + ) -> *const VoiceActivityDetector; + pub fn SherpaOnnxDestroyVoiceActivityDetector(p: *const VoiceActivityDetector); + pub fn SherpaOnnxVoiceActivityDetectorAcceptWaveform( + p: *const VoiceActivityDetector, + samples: *const f32, + n: i32, + ); + pub fn SherpaOnnxVoiceActivityDetectorEmpty(p: *const VoiceActivityDetector) -> i32; + pub fn SherpaOnnxVoiceActivityDetectorDetected(p: *const VoiceActivityDetector) -> i32; + pub fn SherpaOnnxVoiceActivityDetectorPop(p: *const VoiceActivityDetector); + pub fn SherpaOnnxVoiceActivityDetectorClear(p: *const VoiceActivityDetector); + pub fn SherpaOnnxVoiceActivityDetectorFront( + p: *const VoiceActivityDetector, + ) -> *const SpeechSegment; + pub fn SherpaOnnxDestroySpeechSegment(p: *const SpeechSegment); + pub fn SherpaOnnxVoiceActivityDetectorReset(p: *const VoiceActivityDetector); + pub fn SherpaOnnxVoiceActivityDetectorFlush(p: *const VoiceActivityDetector); +} diff --git a/sherpa-onnx/rust/sherpa-onnx-sys/src/wave.rs b/sherpa-onnx/rust/sherpa-onnx-sys/src/wave.rs index 702c0e65a4..c3c8f860fe 100644 --- a/sherpa-onnx/rust/sherpa-onnx-sys/src/wave.rs +++ b/sherpa-onnx/rust/sherpa-onnx-sys/src/wave.rs @@ -19,4 +19,12 @@ extern "C" { /// Free memory allocated by SherpaOnnxReadWave pub fn SherpaOnnxFreeWave(wave: *const SherpaOnnxWave); + + /// Write a WAV file. Returns 1 on success, 0 on failure. + pub fn SherpaOnnxWriteWave( + samples: *const f32, + n: i32, + sample_rate: i32, + filename: *const c_char, + ) -> i32; } diff --git a/sherpa-onnx/rust/sherpa-onnx/Cargo.toml b/sherpa-onnx/rust/sherpa-onnx/Cargo.toml index 0031078eb9..0b7228e77d 100644 --- a/sherpa-onnx/rust/sherpa-onnx/Cargo.toml +++ b/sherpa-onnx/rust/sherpa-onnx/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "sherpa-onnx" -version = "0.1.6" +version = "0.1.7" edition = "2021" description = "Safe Rust wrapper for sherpa-onnx speech recognition toolkit" license = "Apache-2.0" @@ -20,6 +20,6 @@ include = [ ] [dependencies] -sherpa-onnx-sys = { path = "../sherpa-onnx-sys", version = "0.1.6" } +sherpa-onnx-sys = { path = "../sherpa-onnx-sys", version = "0.1.7" } serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" diff --git a/sherpa-onnx/rust/sherpa-onnx/src/lib.rs b/sherpa-onnx/rust/sherpa-onnx/src/lib.rs index 07a47c75c6..ec6a5f584c 100644 --- a/sherpa-onnx/rust/sherpa-onnx/src/lib.rs +++ b/sherpa-onnx/rust/sherpa-onnx/src/lib.rs @@ -2,10 +2,12 @@ mod display; mod offline_asr; mod online_asr; mod utils; +mod vad; mod wave; pub use display::*; pub use offline_asr::*; pub use online_asr::*; pub use utils::*; +pub use vad::*; pub use wave::*; diff --git a/sherpa-onnx/rust/sherpa-onnx/src/vad.rs b/sherpa-onnx/rust/sherpa-onnx/src/vad.rs new file mode 100644 index 0000000000..33f0e35ba9 --- /dev/null +++ b/sherpa-onnx/rust/sherpa-onnx/src/vad.rs @@ -0,0 +1,235 @@ +// sherpa-onnx/src/vad.rs +use crate::utils::to_c_ptr; +use std::ffi::CString; +use std::slice; + +use sherpa_onnx_sys as sys; + +#[derive(Clone, Debug, Default)] +pub struct SileroVadModelConfig { + pub model: Option, + pub threshold: f32, + pub min_silence_duration: f32, + pub min_speech_duration: f32, + pub window_size: i32, + pub max_speech_duration: f32, +} + +impl SileroVadModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::SileroVadModelConfig { + sys::SileroVadModelConfig { + model: to_c_ptr(&self.model, cstrings), + threshold: self.threshold, + min_silence_duration: self.min_silence_duration, + min_speech_duration: self.min_speech_duration, + window_size: self.window_size, + max_speech_duration: self.max_speech_duration, + } + } +} + +#[derive(Clone, Debug, Default)] +pub struct TenVadModelConfig { + pub model: Option, + pub threshold: f32, + pub min_silence_duration: f32, + pub min_speech_duration: f32, + pub window_size: i32, + pub max_speech_duration: f32, +} + +impl TenVadModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::TenVadModelConfig { + sys::TenVadModelConfig { + model: to_c_ptr(&self.model, cstrings), + threshold: self.threshold, + min_silence_duration: self.min_silence_duration, + min_speech_duration: self.min_speech_duration, + window_size: self.window_size, + max_speech_duration: self.max_speech_duration, + } + } +} + +#[derive(Clone, Debug, Default)] +pub struct VadModelConfig { + pub silero_vad: SileroVadModelConfig, + pub ten_vad: TenVadModelConfig, + pub sample_rate: i32, + pub num_threads: i32, + pub provider: Option, + pub debug: bool, +} + +impl VadModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::VadModelConfig { + sys::VadModelConfig { + silero_vad: self + .silero_vad + .to_sys(cstrings), + ten_vad: self + .ten_vad + .to_sys(cstrings), + sample_rate: self.sample_rate, + num_threads: self.num_threads, + provider: to_c_ptr(&self.provider, cstrings), + debug: self.debug as i32, + } + } +} + +pub struct CircularBuffer { + ptr: *const sys::CircularBuffer, +} + +impl CircularBuffer { + pub fn new(capacity: i32) -> Option { + let ptr = unsafe { sys::SherpaOnnxCreateCircularBuffer(capacity) }; + if ptr.is_null() { + None + } else { + Some(Self { ptr }) + } + } + + pub fn push(&self, samples: &[f32]) { + unsafe { + sys::SherpaOnnxCircularBufferPush(self.ptr, samples.as_ptr(), samples.len() as i32) + } + } + + pub fn get(&self, start_index: i32, n: i32) -> Vec { + unsafe { + let p = sys::SherpaOnnxCircularBufferGet(self.ptr, start_index, n); + if p.is_null() { + return vec![]; + } + let slice = slice::from_raw_parts(p, n as usize); + let result = slice.to_vec(); + sys::SherpaOnnxCircularBufferFree(p); + result + } + } + + pub fn pop(&self, n: i32) { + unsafe { sys::SherpaOnnxCircularBufferPop(self.ptr, n) } + } + + pub fn size(&self) -> i32 { + unsafe { sys::SherpaOnnxCircularBufferSize(self.ptr) } + } + + pub fn head(&self) -> i32 { + unsafe { sys::SherpaOnnxCircularBufferHead(self.ptr) } + } + + pub fn reset(&self) { + unsafe { sys::SherpaOnnxCircularBufferReset(self.ptr) } + } +} + +impl Drop for CircularBuffer { + fn drop(&mut self) { + unsafe { sys::SherpaOnnxDestroyCircularBuffer(self.ptr) } + } +} + +pub struct SpeechSegment { + ptr: *const sys::SpeechSegment, +} + +impl SpeechSegment { + pub fn start(&self) -> i32 { + unsafe { (*self.ptr).start } + } + + pub fn samples(&self) -> &[f32] { + unsafe { slice::from_raw_parts((*self.ptr).samples, (*self.ptr).n as usize) } + } + + pub fn n(&self) -> i32 { + unsafe { (*self.ptr).n } + } +} + +impl Drop for SpeechSegment { + fn drop(&mut self) { + unsafe { sys::SherpaOnnxDestroySpeechSegment(self.ptr) } + } +} + +pub struct VoiceActivityDetector { + ptr: *const sys::VoiceActivityDetector, +} + +impl VoiceActivityDetector { + pub fn create(config: &VadModelConfig, buffer_size_in_seconds: f32) -> Option { + let mut cstrings = Vec::new(); + let sys_config = config.to_sys(&mut cstrings); + + let ptr = unsafe { + sys::SherpaOnnxCreateVoiceActivityDetector(&sys_config, buffer_size_in_seconds) + }; + + if ptr.is_null() { + None + } else { + Some(Self { ptr }) + } + } + + pub fn accept_waveform(&self, samples: &[f32]) { + unsafe { + sys::SherpaOnnxVoiceActivityDetectorAcceptWaveform( + self.ptr, + samples.as_ptr(), + samples.len() as i32, + ) + } + } + + pub fn is_empty(&self) -> bool { + unsafe { sys::SherpaOnnxVoiceActivityDetectorEmpty(self.ptr) != 0 } + } + + pub fn detected(&self) -> bool { + unsafe { sys::SherpaOnnxVoiceActivityDetectorDetected(self.ptr) != 0 } + } + + pub fn pop(&self) { + unsafe { sys::SherpaOnnxVoiceActivityDetectorPop(self.ptr) } + } + + pub fn clear(&self) { + unsafe { sys::SherpaOnnxVoiceActivityDetectorClear(self.ptr) } + } + + pub fn front(&self) -> Option { + if self.is_empty() { + return None; + } + + unsafe { + let ptr = sys::SherpaOnnxVoiceActivityDetectorFront(self.ptr); + if ptr.is_null() { + None + } else { + Some(SpeechSegment { ptr }) + } + } + } + + pub fn reset(&self) { + unsafe { sys::SherpaOnnxVoiceActivityDetectorReset(self.ptr) } + } + + pub fn flush(&self) { + unsafe { sys::SherpaOnnxVoiceActivityDetectorFlush(self.ptr) } + } +} + +impl Drop for VoiceActivityDetector { + fn drop(&mut self) { + unsafe { sys::SherpaOnnxDestroyVoiceActivityDetector(self.ptr) } + } +} diff --git a/sherpa-onnx/rust/sherpa-onnx/src/wave.rs b/sherpa-onnx/rust/sherpa-onnx/src/wave.rs index 66937c6132..5be88ce3a3 100644 --- a/sherpa-onnx/rust/sherpa-onnx/src/wave.rs +++ b/sherpa-onnx/rust/sherpa-onnx/src/wave.rs @@ -20,6 +20,21 @@ impl Wave { } } + /// Write the WAV to a file using SherpaOnnx C API. + /// + /// Returns true if succeeded, false otherwise. + pub fn write(&self, filename: &str) -> bool { + let c_filename = CString::new(filename).unwrap(); + unsafe { + sys::SherpaOnnxWriteWave( + (*self.inner).samples, + (*self.inner).num_samples, + (*self.inner).sample_rate, + c_filename.as_ptr(), + ) == 1 + } + } + /// Get sample rate pub fn sample_rate(&self) -> i32 { unsafe { (*self.inner).sample_rate } @@ -57,3 +72,18 @@ impl Drop for Wave { } } } + +/// Write samples directly to a WAV file without creating a Wave object. +/// +/// Returns true on success, false otherwise. +pub fn write(filename: &str, samples: &[f32], sample_rate: i32) -> bool { + let c_filename = CString::new(filename).unwrap(); + unsafe { + sys::SherpaOnnxWriteWave( + samples.as_ptr(), + samples.len() as i32, + sample_rate, + c_filename.as_ptr(), + ) == 1 + } +}