mod common;
use coremlit::{
ComputeUnits,
audio::vad::{CHUNK_SAMPLES, CONTEXT_SAMPLES, VadModel, VadModelOptions, VadState},
};
fn load() -> VadModel {
VadModel::load_with(
common::model_path(),
VadModelOptions::new().with_compute(ComputeUnits::CpuOnly),
)
.expect("load vad model")
}
fn two_chunks() -> (Vec<f32>, Vec<f32>) {
let samples = common::load_wav_16k_mono(&common::fixture_wav_path("02_pyannote_sample"));
assert!(
samples.len() >= 2 * CHUNK_SAMPLES,
"fixture must have at least two full chunks"
);
(
samples[..CHUNK_SAMPLES].to_vec(),
samples[CHUNK_SAMPLES..2 * CHUNK_SAMPLES].to_vec(),
)
}
#[test]
#[ignore = "requires local vadkit models (VADKIT_TEST_MODELS)"]
fn reset_returns_to_initial_state() {
let (chunk0, chunk1) = two_chunks();
let mut model = load();
assert_eq!(model.state(), &VadState::initial(), "starts at initial");
let p_first = model.predict_chunk(&chunk0).expect("chunk 0");
assert_ne!(
model.state(),
&VadState::initial(),
"one chunk must advance the state"
);
model.predict_chunk(&chunk1).expect("chunk 1");
model.reset();
assert_eq!(
model.state(),
&VadState::initial(),
"reset must restore the initial state exactly"
);
let p_after_reset = model.predict_chunk(&chunk0).expect("chunk 0 again");
assert_eq!(
p_after_reset, p_first,
"after reset, chunk 0 must reproduce its first probability bit-for-bit"
);
}
#[test]
#[ignore = "requires local vadkit models (VADKIT_TEST_MODELS)"]
fn state_round_trips_across_chunks() {
let (chunk0, chunk1) = two_chunks();
let model = load();
let (p0, s1) = model
.predict_chunk_with_state(&chunk0, &VadState::initial())
.expect("chunk 0");
let (p1, s2) = model
.predict_chunk_with_state(&chunk1, &s1)
.expect("chunk 1 from s1");
let (p1_again, s2_again) = model
.predict_chunk_with_state(&chunk1, &s1)
.expect("chunk 1 from saved s1");
assert_eq!(p1, p1_again, "same input state → same probability");
assert_eq!(s2, s2_again, "same input state → same output state");
let mut streaming = load();
assert_eq!(streaming.predict_chunk(&chunk0).expect("stream 0"), p0);
assert_eq!(streaming.predict_chunk(&chunk1).expect("stream 1"), p1);
assert_eq!(
streaming.state(),
&s2,
"streaming state tracks the explicit one"
);
}
#[test]
#[ignore = "requires local vadkit models (VADKIT_TEST_MODELS)"]
fn misaligned_context_changes_the_probability() {
let samples = common::load_wav_16k_mono(&common::fixture_wav_path("02_pyannote_sample"));
let (chunks, _tail) = samples.as_chunks::<CHUNK_SAMPLES>();
assert!(chunks.len() >= 40, "need enough chunks to scan");
let model = load();
let (_p0, s1) = model
.predict_chunk_with_state(&chunks[0], &VadState::initial())
.expect("chunk 0");
assert_eq!(
&s1.context()[..],
&chunks[0][CHUNK_SAMPLES - CONTEXT_SAMPLES..],
"carried context must be the previous chunk's last 64 samples, no skew"
);
let mut state = VadState::initial();
let (mut zeroed_diffs, mut skew_diffs) = (0usize, 0usize);
let (mut zeroed_max, mut skew_max) = (0.0f64, 0.0f64);
for chunk in chunks {
let (correct, next) = model
.predict_chunk_with_state(chunk, &state)
.expect("correct context");
let zeroed = VadState::from_parts(*state.hidden(), *state.cell(), [0.0f32; CONTEXT_SAMPLES]);
let (p_zeroed, _) = model
.predict_chunk_with_state(chunk, &zeroed)
.expect("zeroed");
let mut skewed_ctx = *state.context();
skewed_ctx.rotate_left(1);
let skewed = VadState::from_parts(*state.hidden(), *state.cell(), skewed_ctx);
let (p_skewed, _) = model
.predict_chunk_with_state(chunk, &skewed)
.expect("skewed");
if p_zeroed != correct {
zeroed_diffs += 1;
}
if p_skewed != correct {
skew_diffs += 1;
}
zeroed_max = zeroed_max.max((f64::from(p_zeroed) - f64::from(correct)).abs());
skew_max = skew_max.max((f64::from(p_skewed) - f64::from(correct)).abs());
state = next;
}
println!(
"[misaligned] {} chunks | zeroed-context differs on {zeroed_diffs} (max |Δ| {zeroed_max:.3e}) \
| 1-sample skew differs on {skew_diffs} (max |Δ| {skew_max:.3e})",
chunks.len()
);
assert!(
zeroed_diffs > 0,
"a zeroed context must change the probability on at least one chunk — the graph consumes it"
);
assert!(
skew_diffs > 0,
"a one-sample context skew must be observable on at least one chunk (corroborating the trace gate)"
);
}