use std::collections::HashMap;
use super::*;
use crate::{
array::Array,
audio::vad::{load::VadModel, output::SpeechSegment},
dtype::Dtype,
error::Result,
};
fn zeros(shape: &[i32]) -> Array {
Array::zeros::<f32>(&shape).expect("zeros")
}
fn zeros_dtype(shape: &[i32], dtype: Dtype) -> Array {
Array::zeros::<f32>(&shape)
.and_then(|a| a.astype(dtype))
.expect("zeros_dtype")
}
fn insert_branch_weights(
weights: &mut HashMap<String, Array>,
cfg: &BranchConfig,
branch: &str,
dtype: Dtype,
) {
let cutoff = cfg.cutoff();
let put = |w: &mut HashMap<String, Array>, suffix: &str, shape: &[i32]| {
w.insert(format!("{branch}.{suffix}"), zeros_dtype(shape, dtype));
};
put(
weights,
"stft_conv.weight",
&[2 * cutoff, cfg.filter_length(), 1],
);
put(weights, "conv1.weight", &[128, 3, cutoff]);
put(weights, "conv1.bias", &[128]);
put(weights, "conv2.weight", &[64, 3, 128]);
put(weights, "conv2.bias", &[64]);
put(weights, "conv3.weight", &[64, 3, 64]);
put(weights, "conv3.bias", &[64]);
put(weights, "conv4.weight", &[128, 3, 64]);
put(weights, "conv4.bias", &[128]);
put(weights, "lstm.Wx", &[512, 128]);
put(weights, "lstm.Wh", &[512, 128]);
put(weights, "lstm.bias", &[512]);
put(weights, "final_conv.weight", &[1, 1, 128]);
put(weights, "final_conv.bias", &[1]);
}
fn synthetic_model(config: ModelConfig) -> SileroVadModel {
let mut weights = HashMap::new();
let dtype = config.dtype();
insert_branch_weights(&mut weights, config.branch_16k(), "vad_16k", dtype);
insert_branch_weights(&mut weights, config.branch_8k(), "vad_8k", dtype);
SileroVadModel::from_weights(config, weights).expect("synthetic model")
}
#[test]
fn default_config_matches_reference() {
let cfg = ModelConfig::default();
assert_eq!(cfg.branch_16k().sample_rate(), 16_000);
assert_eq!(cfg.branch_16k().chunk_size(), 512);
assert_eq!(cfg.branch_8k().sample_rate(), 8_000);
assert_eq!(cfg.branch_8k().chunk_size(), 256);
assert_eq!(cfg.dtype(), Dtype::F32);
assert_eq!(cfg.threshold(), 0.5);
assert_eq!(cfg.min_speech_duration_ms(), 250);
assert_eq!(cfg.min_silence_duration_ms(), 100);
assert_eq!(cfg.speech_pad_ms(), 30);
}
#[test]
fn default_8k_branch_matches_reference() {
let b = BranchConfig::default_8k();
assert_eq!(b.sample_rate(), 8_000);
assert_eq!(b.filter_length(), 128);
assert_eq!(b.hop_length(), 64);
assert_eq!(b.pad(), 32);
assert_eq!(b.cutoff(), 65);
assert_eq!(b.context_size(), 32);
assert_eq!(b.chunk_size(), 256);
}
#[test]
fn config_from_json_overlays_and_resolves_dtype() {
let cfg = ModelConfig::from_json(
r#"{
"dtype": "float16",
"branch_16k": {"chunk_size": 512, "context_size": 64},
"branch_8k": {"sample_rate": 8000, "filter_length": 128}
}"#,
)
.expect("parse");
assert_eq!(cfg.dtype(), Dtype::F16);
assert_eq!(cfg.branch_16k().chunk_size(), 512);
assert_eq!(cfg.branch_16k().context_size(), 64);
assert_eq!(cfg.branch_16k().cutoff(), 129);
assert_eq!(cfg.branch_8k().filter_length(), 128);
assert_eq!(cfg.branch_8k().chunk_size(), 512);
assert_eq!(cfg.branch_8k().context_size(), 64);
}
#[test]
fn partial_branch_8k_fills_from_16k_defaults_absent_uses_8k() {
let present = ModelConfig::from_json(r#"{"branch_8k": {"hop_length": 999}}"#).expect("parse");
assert_eq!(present.branch_8k().hop_length(), 999); assert_eq!(present.branch_8k().chunk_size(), 512); assert_eq!(present.branch_8k().context_size(), 64); assert_eq!(present.branch_8k().pad(), 64); assert_eq!(present.branch_8k().cutoff(), 129);
let absent = ModelConfig::from_json("{}").expect("parse");
assert_eq!(absent.branch_8k(), &BranchConfig::default_8k());
assert_eq!(absent.branch_8k().chunk_size(), 256);
assert_eq!(absent.branch_8k().context_size(), 32);
}
#[test]
fn branch_null_is_absent_and_malformed_branch_is_rejected() {
let null8 = ModelConfig::from_json(r#"{"branch_8k": null}"#).expect("null branch_8k parses");
assert_eq!(null8.branch_8k(), &BranchConfig::default_8k());
let null16 = ModelConfig::from_json(r#"{"branch_16k": null}"#).expect("null branch_16k parses");
assert_eq!(null16.branch_16k(), &BranchConfig::default_16k());
assert!(ModelConfig::from_json(r#"{"branch_8k": []}"#).is_err());
assert!(ModelConfig::from_json(r#"{"branch_8k": "x"}"#).is_err());
assert!(ModelConfig::from_json(r#"{"branch_16k": 5}"#).is_err());
}
#[test]
fn config_from_empty_json_is_defaults() {
let cfg = ModelConfig::from_json("{}").expect("parse");
assert_eq!(cfg, ModelConfig::default());
}
#[test]
fn config_dtype_non_float16_is_f32() {
let cfg = ModelConfig::from_json(r#"{"dtype": "bfloat16"}"#).expect("parse");
assert_eq!(cfg.dtype(), Dtype::F32);
}
#[test]
fn config_rejects_non_positive_dim() {
let err = ModelConfig::from_json(r#"{"branch_16k": {"cutoff": 0}}"#);
assert!(err.is_err(), "cutoff=0 must be rejected");
}
#[test]
fn forward_shape_and_state_16k() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let x = Array::zeros::<f32>(&[2, 576])?;
let (mut out, mut state) = model.forward(&x, None, 16_000)?;
out.eval()?;
state.eval()?;
assert_eq!(out.shape(), vec![2, 1]);
assert_eq!(state.shape(), vec![2, 2, 128]);
let lo = out.min(false)?.astype(Dtype::F32)?.to_vec::<f32>()?[0];
let hi = out.max(false)?.astype(Dtype::F32)?.to_vec::<f32>()?[0];
assert!((0.0..=1.0).contains(&lo), "min {lo}");
assert!((0.0..=1.0).contains(&hi), "max {hi}");
Ok(())
}
#[test]
fn forward_shape_and_state_8k() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let x = Array::zeros::<f32>(&[1, 288])?;
let (mut out, mut state) = model.forward(&x, None, 8_000)?;
out.eval()?;
state.eval()?;
assert_eq!(out.shape(), vec![1, 1]);
assert_eq!(state.shape(), vec![2, 1, 128]);
Ok(())
}
#[test]
fn feed_updates_streaming_context() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let chunk = Array::zeros::<f32>(&[512])?;
let (mut out, state) = model.feed(&chunk, None, 16_000)?;
out.eval()?;
assert_eq!(out.shape(), vec![1, 1]);
assert_eq!(state.context().shape(), vec![1, 64]);
assert_eq!(state.sample_rate(), 16_000);
Ok(())
}
#[test]
fn predict_proba_chunk_count() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[1024])?;
let mut probs = model.predict_proba(&audio, 16_000)?;
probs.eval()?;
assert_eq!(probs.shape(), vec![2]);
Ok(())
}
#[test]
fn predict_proba_long_input_periodic_eval() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[18 * 512])?;
let mut probs = model.predict_proba(&audio, 16_000)?;
probs.eval()?;
assert_eq!(probs.shape(), vec![18]);
Ok(())
}
#[test]
fn predict_proba_empty_is_empty() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[0])?;
let mut probs = model.predict_proba(&audio, 16_000)?;
probs.eval()?;
assert_eq!(probs.shape(), vec![0]);
Ok(())
}
#[test]
fn predict_proba_rejects_bad_rank() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let scalar = Array::zeros::<f32>(&[])?;
assert!(matches!(
model.predict_proba(&scalar, 16_000),
Err(crate::error::Error::RankMismatch(_))
));
let rank3 = Array::zeros::<f32>(&[2, 2, 512])?;
assert!(matches!(
model.predict_proba(&rank3, 16_000),
Err(crate::error::Error::RankMismatch(_))
));
Ok(())
}
#[test]
fn generate_returns_output() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[512])?;
let out = model.generate(&audio, 16_000)?;
assert_eq!(out.sample_rate, 16_000);
assert_eq!(out.probabilities.shape(), vec![1]);
Ok(())
}
#[test]
fn unsupported_sample_rate_is_rejected() {
let model = synthetic_model(ModelConfig::default());
assert!(model.branch(44_100).is_err());
}
#[test]
fn probs_to_timestamps_matches_reference_vector() {
let probs = [0.1_f32, 0.8, 0.85, 0.1, 0.1];
let segs = probs_to_timestamps(
&probs,
5 * 512, 16_000, 0.5, 30, 30, 0, );
assert_eq!(segs, vec![SpeechSegment::new(512, 1536)]);
}
#[test]
fn probs_to_timestamps_coalesces_padded_overlap_like_mlx_audio() {
let probs = [0.1_f32, 0.8, 0.85, 0.1, 0.1, 0.8, 0.85, 0.1, 0.1];
let segs = probs_to_timestamps(
&probs,
9 * 512, 16_000, 0.5, 30, 30, 100, );
assert_eq!(segs, vec![SpeechSegment::new(0, 4608)]);
}
#[test]
fn probs_to_timestamps_all_silence_is_empty() {
let probs = [0.0_f32, 0.1, 0.05, 0.0];
let segs = probs_to_timestamps(&probs, 4 * 512, 16_000, 0.5, 30, 30, 0);
assert!(segs.is_empty());
}
#[test]
fn probs_to_timestamps_closes_trailing_segment() {
let probs = [0.9_f32, 0.9, 0.9, 0.9];
let segs = probs_to_timestamps(&probs, 2048, 16_000, 0.5, 30, 30, 0);
assert_eq!(segs, vec![SpeechSegment::new(0, 2048)]);
}
#[test]
fn probs_to_timestamps_applies_speech_pad() {
let probs = [0.1_f32, 0.8, 0.85, 0.1, 0.1];
let segs = probs_to_timestamps(&probs, 5 * 512, 16_000, 0.5, 30, 30, 10);
assert_eq!(segs, vec![SpeechSegment::new(352, 1696)]);
}
#[test]
fn probs_to_timestamps_8k_uses_256_chunk() {
let probs = [0.1_f32, 0.8, 0.85, 0.1, 0.1];
let segs = probs_to_timestamps(&probs, 5 * 256, 8_000, 0.5, 30, 30, 0);
assert_eq!(segs, vec![SpeechSegment::new(256, 768)]);
}
#[test]
fn sanitize_drops_val_prefixed_keys() {
let mut weights = HashMap::new();
weights.insert("vad_16k.conv1.weight".to_string(), zeros(&[1]));
weights.insert("val_loss".to_string(), zeros(&[1]));
weights.insert("val_acc.running".to_string(), zeros(&[1]));
let kept = sanitize(weights);
assert!(kept.contains_key("vad_16k.conv1.weight"));
assert!(!kept.keys().any(|k| k.starts_with("val_")));
assert_eq!(kept.len(), 1);
}
#[test]
fn feed_threads_state_across_frames() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let mut state = None;
for _ in 0..3 {
let chunk = Array::zeros::<f32>(&[512])?;
let (mut out, next) = model.feed(&chunk, state, 16_000)?;
out.eval()?;
assert_eq!(out.shape(), vec![1, 1]);
assert_eq!(next.context().shape(), vec![1, 64]);
assert!(next.state().is_some(), "feed must carry an LSTM state");
state = Some(next);
}
Ok(())
}
#[test]
fn feed_rejects_wrong_chunk_width() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let chunk = Array::zeros::<f32>(&[400])?; assert!(model.feed(&chunk, None, 16_000).is_err());
Ok(())
}
#[test]
fn forward_preserves_f16_dtype() -> Result<()> {
let cfg = ModelConfig::from_json(r#"{"dtype": "float16"}"#)?;
assert_eq!(cfg.dtype(), Dtype::F16);
let model = synthetic_model(cfg);
let x = Array::zeros::<f32>(&[1, 576])?;
let (out, state) = model.forward(&x, None, 16_000)?;
assert_eq!(out.dtype()?, Dtype::F16);
assert_eq!(state.dtype()?, Dtype::F16);
Ok(())
}
#[test]
fn predict_proba_preserves_f16_dtype() -> Result<()> {
let cfg = ModelConfig::from_json(r#"{"dtype": "float16"}"#)?;
let model = synthetic_model(cfg);
let audio = Array::zeros::<f32>(&[1536])?; let mut probs = model.predict_proba(&audio, 16_000)?;
assert_eq!(probs.dtype()?, Dtype::F16);
probs.eval()?;
assert_eq!(probs.shape(), vec![3]);
Ok(())
}
#[test]
fn predict_proba_batched_returns_per_row_frames() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[2, 1024])?; let mut probs = model.predict_proba(&audio, 16_000)?;
probs.eval()?;
assert_eq!(probs.shape(), vec![2, 2]); Ok(())
}
#[test]
fn predict_proba_8k_chunk_count() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[1024])?;
let mut probs = model.predict_proba(&audio, 8_000)?;
probs.eval()?;
assert_eq!(probs.shape(), vec![4]); Ok(())
}
#[test]
fn predict_proba_exact_eval_every_multiple() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[16 * 512])?;
let mut probs = model.predict_proba(&audio, 16_000)?;
probs.eval()?;
assert_eq!(probs.shape(), vec![16]);
Ok(())
}
#[test]
fn from_weights_rejects_quantized_checkpoint() {
let config = ModelConfig::default();
let mut weights = HashMap::new();
insert_branch_weights(&mut weights, config.branch_16k(), "vad_16k", config.dtype());
insert_branch_weights(&mut weights, config.branch_8k(), "vad_8k", config.dtype());
weights.insert("vad_16k.conv1.scales".to_string(), zeros(&[128, 1]));
assert!(matches!(
SileroVadModel::from_weights(config, weights),
Err(crate::error::Error::OutOfRange(_))
));
}
#[test]
fn from_weights_rejects_malformed_lstm_shape() {
let config = ModelConfig::default();
let mut weights = HashMap::new();
insert_branch_weights(&mut weights, config.branch_16k(), "vad_16k", config.dtype());
insert_branch_weights(&mut weights, config.branch_8k(), "vad_8k", config.dtype());
weights.insert(
"vad_16k.lstm.Wx".to_string(),
zeros_dtype(&[513, 128], config.dtype()),
);
assert!(SileroVadModel::from_weights(config, weights).is_err());
}
#[test]
fn prepare_audio_downmixes_stereo_and_keeps_batch() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let stereo = Array::zeros::<f32>(&[1000, 2])?;
let (prepared, sr) = model.prepare_audio(&stereo, 16_000)?;
assert_eq!(prepared.shape(), vec![1000]);
assert_eq!(sr, 16_000);
let batched = Array::zeros::<f32>(&[2, 1000])?;
let (prepared_b, _) = model.prepare_audio(&batched, 16_000)?;
assert_eq!(prepared_b.shape(), vec![2, 1000]);
Ok(())
}
#[test]
fn prepare_audio_resamples_unsupported_rate_to_16k() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[1024])?;
let (prepared, sr) = model.prepare_audio(&audio, 32_000)?;
assert_eq!(sr, 16_000);
let n = prepared.shape()[0];
assert!(
(400..=600).contains(&n),
"32k→16k resample of 1024 → ~512, got {n}"
);
Ok(())
}
#[test]
fn speech_segment_seconds_accessors() {
let seg = SpeechSegment::new(16_000, 32_000);
assert_eq!(seg.start_seconds(16_000), 1.0);
assert_eq!(seg.end_seconds(16_000), 2.0);
assert_eq!(seg.start_seconds(8_000), 2.0);
}
#[test]
fn get_speech_timestamps_default_and_override() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[16 * 512])?;
let segs = model.get_speech_timestamps(&audio, 16_000, SpeechTimestampOptions::default())?;
assert!(!segs.is_empty(), "all-0.5 probs at threshold 0.5 → speech");
let opts = SpeechTimestampOptions {
threshold: Some(0.6),
..Default::default()
};
let none = model.get_speech_timestamps(&audio, 16_000, opts)?;
assert!(none.is_empty(), "all-0.5 probs at threshold 0.6 → silence");
Ok(())
}
#[test]
fn predict_and_reset_state_smoke() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[1024])?;
let mut probs = model.predict(&audio, 16_000)?;
probs.eval()?;
assert_eq!(probs.shape(), vec![2]);
let state = model.reset_state(1, 16_000)?;
assert_eq!(state.context().shape(), vec![1, 64]);
assert!(state.state().is_none());
Ok(())
}
#[test]
fn get_speech_timestamps_rejects_negative_override() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[4 * 512])?;
let opts = SpeechTimestampOptions {
speech_pad_ms: Some(-100),
..Default::default()
};
assert!(matches!(
model.get_speech_timestamps(&audio, 16_000, opts),
Err(crate::error::Error::OutOfRange(_))
));
Ok(())
}
#[test]
fn from_weights_rejects_quantized_lstm_scales() {
let config = ModelConfig::default();
let mut weights = HashMap::new();
insert_branch_weights(&mut weights, config.branch_16k(), "vad_16k", config.dtype());
insert_branch_weights(&mut weights, config.branch_8k(), "vad_8k", config.dtype());
weights.insert("vad_16k.lstm.Wx.scales".to_string(), zeros(&[512, 1]));
assert!(matches!(
SileroVadModel::from_weights(config, weights),
Err(crate::error::Error::OutOfRange(_))
));
}
#[test]
fn prepare_audio_zero_row_batch_at_unsupported_rate() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let empty_batch = Array::zeros::<f32>(&[0, 1000])?;
let (prepared, sr) = model.prepare_audio(&empty_batch, 32_000)?;
assert_eq!(sr, 16_000);
assert_eq!(prepared.shape()[0], 0);
Ok(())
}
#[test]
fn empty_batches_yield_empty_timestamps_end_to_end() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
for (r, c) in [(0usize, 1000usize), (2, 0)] {
for rate in [16_000u32, 32_000] {
let audio = Array::zeros::<f32>(&[r as i32, c as i32])?;
let out = model.generate(&audio, rate)?;
assert!(
out.timestamps.is_empty(),
"generate on ({r}, {c}) at {rate} Hz must yield no timestamps"
);
let segs = model.get_speech_timestamps(&audio, rate, SpeechTimestampOptions::default())?;
assert!(
segs.is_empty(),
"get_speech_timestamps on ({r}, {c}) at {rate} Hz must yield no segments"
);
}
}
Ok(())
}
#[test]
fn zero_row_unsupported_rate_resolves_resampled_width() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let audio = Array::zeros::<f32>(&[0, 1000])?;
let (prepared, sr) = model.prepare_audio(&audio, 32_000)?;
assert_eq!(sr, 16_000);
let w = prepared.shape()[1];
assert!(
(400..=600).contains(&w),
"the empty batch's width must be the RESAMPLED length (~500), got {w}"
);
let out = model.generate(&audio, 32_000)?;
let frames = *out.probabilities.shape().last().unwrap_or(&usize::MAX);
assert_eq!(
frames,
w.div_ceil(512),
"probability frames must chunk the resampled width"
);
assert!(out.timestamps.is_empty());
Ok(())
}
#[test]
fn zero_row_huge_width_is_shape_arithmetic_only() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let huge = 1_000_000_000i32; let audio = Array::zeros::<f32>(&[0, huge])?;
let (prepared, sr) = model.prepare_audio(&audio, 32_000)?;
assert_eq!(sr, 16_000);
assert_eq!(
prepared.shape(),
vec![0, (huge as usize) / 2],
"width must be the arithmetic in*to/from resample length"
);
Ok(())
}
#[test]
fn zero_row_predict_proba_short_circuits_frame_shape() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let huge = 1_000_000_000i32;
let audio = Array::zeros::<f32>(&[0, huge])?;
let probs = model.predict_proba(&audio, 16_000)?;
assert_eq!(
probs.shape(),
vec![0, (huge as usize).div_ceil(512)],
"zero-row probabilities must carry the reference frame count"
);
Ok(())
}
#[test]
fn zero_sample_rate_is_typed_error() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
for shape in [[0i32, 1000], [2, 1000]] {
let audio = Array::zeros::<f32>(&shape)?;
assert!(
matches!(
model.prepare_audio(&audio, 0),
Err(crate::error::Error::OutOfRange(_))
),
"sample_rate 0 on {shape:?} must be a typed error"
);
}
Ok(())
}
#[test]
fn prepare_audio_small_batch_is_not_downmixed() -> Result<()> {
let model = synthetic_model(ModelConfig::default());
let small = Array::zeros::<f32>(&[3, 2])?;
let (prepared, _) = model.prepare_audio(&small, 16_000)?;
assert_eq!(
prepared.shape(),
vec![3, 2],
"rows <= 8 must not trigger the stereo downmix (reference: 8 < rows)"
);
Ok(())
}