Files
rustytorch/crates/models/rtx-csm/src/silero_vad.rs
T
osobhandClaude Opus 4.7 a5cedfb46a rtx-csm: emotional_speech_guide — CREMA-D vs RAVDESS firdhokk verdict
8-gen bench (4 emotions × 2 corpora) at seed=42 against firdhokk
Whisper-LV3:

  target    RAVDESS              CREMA-D
  happy     happy (0.999) ✓      happy (0.999) ✓
  angry     neutral (0.92)       sad (0.99)
  fearful   happy (0.998)        fearful (0.984) ✓
  sad       angry (0.99)         fearful (0.99)

CREMA-D 2/4 vs RAVDESS 1/4. Larger / more naturalistic corpus
produces more class-pure fearful direction. Neither corpus solves
angry or sad — recipe shifts into 'vague expressivity' rather than
class-specific corners.

Practical: prefer CREMA-D when available; A/B both per emotion if
class precision matters.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
2026-04-30 00:01:02 -07:00

359 lines
12 KiB
Rust

//! Silero V5 VAD candle port — pure Rust + candle 0.9, NO `ort` /
//! ONNX runtime, NO sentencepiece-protobuf conflict (the issue that
//! deferred this in Phase 8.1.3).
//!
//! Direct port of `Snakers4/silero-vad/src/silero_vad/tinygrad_model.py`
//! (71 LOC, canonical small-framework reference). Architecture facts
//! captured in `docs/silero_vad_port_notes.md`.
//!
//! ## Usage
//!
//! ```ignore
//! let vad = SileroVad::load_default(&device)?;
//! let mut state = SileroState::zero(&device)?;
//! // 512 samples per chunk @ 16 kHz = 32 ms; +64 context prefix.
//! for chunk in audio.chunks(512) {
//! let p = vad.forward_step(chunk, &mut state)?;
//! println!("speech prob = {p:.3}");
//! }
//! ```
use candle_core::{D, DType, Device, Module, Result, Tensor};
use candle_nn::{Conv1d, Conv1dConfig, VarBuilder};
const N_FFT: usize = 256;
const STRIDE: usize = 128;
const PAD_RIGHT: usize = 64;
const CHUNK_SAMPLES: usize = 512;
const CONTEXT_SAMPLES: usize = 64;
const CUTOFF: usize = N_FFT / 2 + 1; // 129
/// LSTM hidden + cell state for stateful per-stream detection. Pass
/// the same `SileroState` across consecutive `forward_step` calls on
/// the same audio stream.
pub struct SileroState {
h: Tensor, // (1, 128)
c: Tensor, // (1, 128)
}
impl SileroState {
pub fn zero(device: &Device) -> Result<Self> {
Ok(Self {
h: Tensor::zeros((1, 128), DType::F32, device)?,
c: Tensor::zeros((1, 128), DType::F32, device)?,
})
}
}
/// Single LSTM cell port. Uses fused weight_ih (input -> 4*hidden)
/// and weight_hh (hidden -> 4*hidden) tensors as PyTorch does.
struct LstmCell {
weight_ih: Tensor, // (4*hidden, input)
weight_hh: Tensor, // (4*hidden, hidden)
bias_ih: Tensor, // (4*hidden,)
bias_hh: Tensor, // (4*hidden,)
hidden: usize,
}
impl LstmCell {
fn new(input_size: usize, hidden: usize, vb: VarBuilder) -> Result<Self> {
let weight_ih = vb.get((4 * hidden, input_size), "weight_ih")?;
let weight_hh = vb.get((4 * hidden, hidden), "weight_hh")?;
let bias_ih = vb.get(4 * hidden, "bias_ih")?;
let bias_hh = vb.get(4 * hidden, "bias_hh")?;
Ok(Self {
weight_ih,
weight_hh,
bias_ih,
bias_hh,
hidden,
})
}
/// `x`: `(B, input)`. `state.h`/`state.c`: `(B, hidden)`.
/// Returns updated `(h, c)`.
fn forward(&self, x: &Tensor, state: &SileroState) -> Result<(Tensor, Tensor)> {
// gates = x @ Wih.T + bih + h @ Whh.T + bhh
let gx = x
.matmul(&self.weight_ih.t()?)?
.broadcast_add(&self.bias_ih)?;
let gh = state
.h
.matmul(&self.weight_hh.t()?)?
.broadcast_add(&self.bias_hh)?;
let gates = (gx + gh)?;
// Split into [i, f, g, o] along the last dim.
let h = self.hidden;
let i = candle_nn::ops::sigmoid(&gates.narrow(D::Minus1, 0, h)?)?;
let f = candle_nn::ops::sigmoid(&gates.narrow(D::Minus1, h, h)?)?;
let g = gates.narrow(D::Minus1, 2 * h, h)?.tanh()?;
let o = candle_nn::ops::sigmoid(&gates.narrow(D::Minus1, 3 * h, h)?)?;
// c' = f * c + i * g
let c_new = ((f * &state.c)? + (i * g)?)?;
// h' = o * tanh(c')
let h_new = (o * c_new.tanh()?)?;
Ok((h_new, c_new))
}
}
/// Silero V5 VAD model.
pub struct SileroVad {
stft_conv: Conv1d,
conv1: Conv1d,
conv2: Conv1d,
conv3: Conv1d,
conv4: Conv1d,
lstm: LstmCell,
final_conv: Conv1d,
}
impl SileroVad {
pub fn new(vb: VarBuilder) -> Result<Self> {
let stft_conv = candle_nn::conv1d_no_bias(
1,
258,
N_FFT,
Conv1dConfig {
padding: 0,
stride: STRIDE,
dilation: 1,
groups: 1,
cudnn_fwd_algo: None,
},
vb.pp("stft_conv"),
)?;
let conv1 = candle_nn::conv1d(
CUTOFF, // 129
128,
3,
Conv1dConfig {
padding: 1,
stride: 1,
dilation: 1,
groups: 1,
cudnn_fwd_algo: None,
},
vb.pp("conv1"),
)?;
let conv2 = candle_nn::conv1d(
128,
64,
3,
Conv1dConfig {
padding: 1,
stride: 2,
dilation: 1,
groups: 1,
cudnn_fwd_algo: None,
},
vb.pp("conv2"),
)?;
let conv3 = candle_nn::conv1d(
64,
64,
3,
Conv1dConfig {
padding: 1,
stride: 2,
dilation: 1,
groups: 1,
cudnn_fwd_algo: None,
},
vb.pp("conv3"),
)?;
let conv4 = candle_nn::conv1d(
64,
128,
3,
Conv1dConfig {
padding: 1,
stride: 1,
dilation: 1,
groups: 1,
cudnn_fwd_algo: None,
},
vb.pp("conv4"),
)?;
let lstm = LstmCell::new(128, 128, vb.pp("lstm_cell"))?;
let final_conv = candle_nn::conv1d(
128,
1,
1,
Conv1dConfig {
padding: 0,
stride: 1,
dilation: 1,
groups: 1,
cudnn_fwd_algo: None,
},
vb.pp("final_conv"),
)?;
Ok(Self {
stft_conv,
conv1,
conv2,
conv3,
conv4,
lstm,
final_conv,
})
}
/// Download (cached) the upstream `Snakers4/silero-vad` 16 kHz
/// safetensors and load. ~1.24 MB; cached under
/// `~/.cache/silero-vad/silero_vad_16k.safetensors` after first use.
pub fn load_default(device: &Device) -> Result<Self> {
let path = ensure_default_weights()?;
Self::load_from_file(&path, device)
}
pub fn load_from_file(path: &std::path::Path, device: &Device) -> Result<Self> {
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[path], DType::F32, device) }?;
Self::new(vb)
}
/// Forward one chunk. `samples` must contain `CONTEXT_SAMPLES + CHUNK_SAMPLES`
/// samples (576 total). Returns the speech probability in `[0, 1]`
/// and updates the LSTM state in place.
pub fn forward_step(&self, samples: &[f32], state: &mut SileroState) -> Result<f32> {
if samples.len() != CHUNK_SAMPLES + CONTEXT_SAMPLES {
return Err(candle_core::Error::Msg(format!(
"silero_vad: expected {} samples, got {}",
CHUNK_SAMPLES + CONTEXT_SAMPLES,
samples.len()
)));
}
let device = state.h.device();
// Build (B=1, T=576) tensor and reflect-pad on the right by 64.
let mut padded = Vec::with_capacity(samples.len() + PAD_RIGHT);
padded.extend_from_slice(samples);
// Reflect padding: last sample, second-to-last, ... up to PAD_RIGHT.
let n = samples.len();
for i in 0..PAD_RIGHT {
// PyTorch reflect skips the boundary sample (idx n-1) for the
// first reflected sample — i.e. mirror around idx n-1.
// padded[n + i] = samples[n - 2 - i] for i in 0..pad_right.
let src = (n as isize - 2 - i as isize).max(0) as usize;
padded.push(samples[src.min(n - 1)]);
}
let x = Tensor::from_vec(padded, (1, 1, samples.len() + PAD_RIGHT), device)?;
// STFT conv -> (1, 258, n_frames)
let x = self.stft_conv.forward(&x)?;
// Magnitude: sqrt(real² + imag²) for the first 129 bins; output
// has shape (1, 129, n_frames).
let real = x.narrow(1, 0, CUTOFF)?;
let imag = x.narrow(1, CUTOFF, CUTOFF)?;
let mag = (real.sqr()? + imag.sqr()?)?.sqrt()?;
// 4-layer Conv1d feature extractor.
let h = self.conv1.forward(&mag)?.relu()?;
let h = self.conv2.forward(&h)?.relu()?;
let h = self.conv3.forward(&h)?.relu()?;
let h = self.conv4.forward(&h)?.relu()?;
// Collapse the time axis (it's now 1 frame after the strides).
// Shape: (1, 128, 1) -> (1, 128).
let h = h.squeeze(D::Minus1)?;
// LSTM step.
let (h_new, c_new) = self.lstm.forward(&h, state)?;
// unsqueeze -> (1, 128, 1) -> ReLU -> final 1x1 conv -> sigmoid
let y = h_new.unsqueeze(D::Minus1)?;
let y = y.relu()?;
let y = self.final_conv.forward(&y)?;
let y = candle_nn::ops::sigmoid(&y)?;
// y shape: (1, 1, 1). Squeeze + mean -> scalar.
let prob: f32 = y.flatten_all()?.to_vec1::<f32>()?[0];
// Update state.
state.h = h_new;
state.c = c_new;
Ok(prob)
}
/// Sweep the model over a multi-chunk audio buffer. Returns one
/// probability per `CHUNK_SAMPLES` chunk. Pads the input with
/// `CONTEXT_SAMPLES` zeros at the front and zero-pads the tail to
/// a whole-chunk multiple. State is owned internally.
pub fn forward_audio(&self, samples: &[f32], device: &Device) -> Result<Vec<f32>> {
let mut state = SileroState::zero(device)?;
// Front-pad with CONTEXT_SAMPLES zeros (the "context" the model
// expects on the first chunk).
let mut padded: Vec<f32> = vec![0.0; CONTEXT_SAMPLES];
padded.extend_from_slice(samples);
// Tail-pad to whole-chunk multiple.
let extra = (CHUNK_SAMPLES - (samples.len() % CHUNK_SAMPLES)) % CHUNK_SAMPLES;
padded.extend(std::iter::repeat_n(0.0, extra));
let mut probs = Vec::new();
let mut i = 0;
while i + CHUNK_SAMPLES + CONTEXT_SAMPLES <= padded.len() + CONTEXT_SAMPLES {
// Window: padded[i .. i + CHUNK_SAMPLES + CONTEXT_SAMPLES]
let end = i + CHUNK_SAMPLES + CONTEXT_SAMPLES;
if end > padded.len() {
break;
}
let window = &padded[i..end];
let p = self.forward_step(window, &mut state)?;
probs.push(p);
i += CHUNK_SAMPLES;
}
Ok(probs)
}
}
/// Ensure the upstream Silero V5 safetensors is cached locally and
/// return its path. Downloads from the Snakers4/silero-vad GitHub
/// raw URL on first use; ~1.24 MB. Cache lives at
/// `~/.cache/rtx-csm/silero_vad_16k.safetensors`.
pub fn ensure_default_weights() -> Result<std::path::PathBuf> {
const URL: &str = "https://raw.githubusercontent.com/snakers4/silero-vad/master/src/silero_vad/data/silero_vad_16k.safetensors";
let cache_dir = dirs_cache_dir().join("rtx-csm");
let path = cache_dir.join("silero_vad_16k.safetensors");
if path.exists() {
return Ok(path);
}
std::fs::create_dir_all(&cache_dir)
.map_err(|e| candle_core::Error::Msg(format!("create cache dir: {e}")))?;
let bytes = ureq::get(URL)
.call()
.map_err(|e| candle_core::Error::Msg(format!("fetch silero weights: {e}")))?
.into_body()
.read_to_vec()
.map_err(|e| candle_core::Error::Msg(format!("read silero body: {e}")))?;
std::fs::write(&path, &bytes)
.map_err(|e| candle_core::Error::Msg(format!("write silero cache: {e}")))?;
Ok(path)
}
fn dirs_cache_dir() -> std::path::PathBuf {
if let Some(home) = std::env::var_os("HOME") {
std::path::PathBuf::from(home).join(".cache")
} else {
std::env::temp_dir()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn loads_default_weights_and_runs_one_step() {
let device = Device::Cpu;
let model = SileroVad::load_default(&device).expect("load weights");
let mut state = SileroState::zero(&device).expect("state");
// Random-ish input chunk (576 samples).
let samples: Vec<f32> = (0..576).map(|i| ((i as f32 * 0.01).sin() * 0.1)).collect();
let prob = model
.forward_step(&samples, &mut state)
.expect("forward_step");
assert!(
(0.0..=1.0).contains(&prob),
"speech prob out of range: {prob}"
);
}
}