use super::*;
use crate::ComputeUnits;
#[test]
fn chunk_segmentation_range_hand_values() {
assert_eq!(chunk_segmentation_range(0, 4), 0..12);
assert_eq!(chunk_segmentation_range(2, 4), 24..36);
}
#[test]
fn embedding_range_hand_values() {
assert_eq!(embedding_range(0, 0), 0..256);
assert_eq!(embedding_range(1, 2), 1280..1536);
}
#[test]
fn chunk_and_slot_embedding_ranges_compose_to_embedding_range() {
for c in 0..4 {
let block = chunk_embedding_range(c);
assert_eq!(block.len(), SEG_NUM_SLOTS * EMBEDDING_DIM);
for s in 0..SEG_NUM_SLOTS {
let slot = slot_embedding_range(s);
assert_eq!(slot.len(), EMBEDDING_DIM);
assert_eq!(
block.start + slot.start..block.start + slot.end,
embedding_range(c, s),
"chunk {c} slot {s}"
);
}
assert_eq!(block.end, chunk_embedding_range(c + 1).start);
}
}
#[test]
fn fill_padded_chunk_middle_chunk_full_copy() {
let samples: Vec<f32> = (0..SEG_CHUNK_SAMPLES + 10)
.map(|i| (i + 1) as f32)
.collect();
let mut padded = vec![0.0f32; SEG_CHUNK_SAMPLES];
fill_padded_chunk(&mut padded, &samples, 5);
assert_eq!(padded.len(), SEG_CHUNK_SAMPLES);
assert_eq!(padded[0], 6.0); assert_eq!(
padded[SEG_CHUNK_SAMPLES - 1],
(SEG_CHUNK_SAMPLES + 5) as f32
); }
#[test]
fn fill_padded_chunk_final_chunk_partial_with_zero_tail() {
let samples: Vec<f32> = (0..SEG_CHUNK_SAMPLES + 5).map(|i| (i + 1) as f32).collect();
let mut padded = vec![0.0f32; SEG_CHUNK_SAMPLES];
fill_padded_chunk(&mut padded, &samples, 10);
assert_eq!(padded[0], 11.0); assert_eq!(padded[159_994], (SEG_CHUNK_SAMPLES + 5) as f32); assert!(
padded[159_995..].iter().all(|v| *v == 0.0),
"out-of-range tail must be zero"
);
}
#[test]
fn fill_padded_chunk_start_beyond_samples_is_all_zero() {
let samples = vec![1.0f32, 2.0, 3.0];
let mut padded = vec![0.0f32; SEG_CHUNK_SAMPLES];
fill_padded_chunk(&mut padded, &samples, 2_000);
assert!(padded.iter().all(|v| *v == 0.0));
}
#[test]
fn fill_padded_chunk_samples_shorter_than_window() {
let samples: Vec<f32> = (0..500).map(|i| (i + 1) as f32).collect();
let mut padded = vec![0.0f32; SEG_CHUNK_SAMPLES];
fill_padded_chunk(&mut padded, &samples, 0);
assert_eq!(padded[0], 1.0);
assert_eq!(padded[499], 500.0);
assert!(padded[500..].iter().all(|v| *v == 0.0));
}
#[test]
fn zero_slot_column_zeroes_only_the_named_column() {
let mut slab = vec![1.0f64, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
zero_slot_column(&mut slab, 3, 1);
assert_eq!(slab, vec![1.0, 0.0, 3.0, 4.0, 0.0, 6.0, 7.0, 0.0, 9.0]);
}
fn logits_for_classes(classes: &[usize]) -> Vec<f32> {
let mut out =
Vec::with_capacity(classes.len() * crate::audio::speaker::segment::POWERSET_CLASSES);
for &c in classes {
let mut row = [0.0f32; crate::audio::speaker::segment::POWERSET_CLASSES];
row[c] = 5.0;
out.extend_from_slice(&row);
}
out
}
fn classes_to_slab(classes: &[usize]) -> Vec<f64> {
crate::audio::speaker::segment::multilabel(&logits_for_classes(classes), classes.len())
}
fn embed6(mask: [bool; 6]) -> SlotPlan {
SlotPlan::Embed(mask.to_vec())
}
#[test]
fn derive_slot_plans_s1_no_overlap() {
let slab = classes_to_slab(&[1, 1, 1, 2, 2, 0]);
assert_eq!(
derive_slot_plans(&slab, 6, 0.5),
[
embed6([true, true, true, false, false, false]),
embed6([false, false, false, true, true, false]),
SlotPlan::Skip,
]
);
}
#[test]
fn derive_slot_plans_s2_full_overlap_falls_back() {
let slab = classes_to_slab(&[4, 4, 4, 4, 4, 4]);
assert_eq!(
derive_slot_plans(&slab, 6, 0.5),
[
embed6([true, true, true, true, true, true]),
embed6([true, true, true, true, true, true]),
SlotPlan::Skip,
]
);
}
#[test]
fn derive_slot_plans_s3_exactly_two_clean_frames_falls_back() {
let slab = classes_to_slab(&[1, 1, 4, 4, 0, 0]);
assert_eq!(
derive_slot_plans(&slab, 6, 0.5),
[
embed6([true, true, true, true, false, false]),
embed6([false, false, true, true, false, false]),
SlotPlan::Skip,
]
);
}
#[test]
fn derive_slot_plans_s4_three_clean_frames_uses_clean_mask() {
let slab = classes_to_slab(&[1, 1, 4, 4, 1, 0]);
assert_eq!(
derive_slot_plans(&slab, 6, 0.5),
[
embed6([true, true, false, false, true, false]),
embed6([false, false, true, true, false, false]),
SlotPlan::Skip,
]
);
}
#[test]
fn derive_slot_plans_s5_fallback_is_per_slot_not_whole_chunk() {
let slab = classes_to_slab(&[1, 1, 1, 4, 4, 0]);
assert_eq!(
derive_slot_plans(&slab, 6, 0.5),
[
embed6([true, true, true, false, false, false]),
embed6([false, false, false, true, true, false]),
SlotPlan::Skip,
]
);
}
#[test]
fn derive_slot_plans_s6_single_speaker() {
let slab = classes_to_slab(&[1, 1, 0, 0, 0, 0]);
assert_eq!(
derive_slot_plans(&slab, 6, 0.5),
[
embed6([true, true, false, false, false, false]),
SlotPlan::Skip,
SlotPlan::Skip,
]
);
}
#[test]
fn derive_slot_plans_s7_empty_chunk_all_skip() {
let slab = classes_to_slab(&[0, 0, 0, 0, 0, 0]);
assert_eq!(
derive_slot_plans(&slab, 6, 0.5),
[SlotPlan::Skip, SlotPlan::Skip, SlotPlan::Skip]
);
}
#[test]
#[should_panic(expected = "chunk_segs.len() must equal num_frames * SEG_NUM_SLOTS")]
fn derive_slot_plans_panics_on_length_mismatch() {
let _ = derive_slot_plans(&[0.0f64; 5], 2, 0.5);
}
#[test]
fn geometry_pipeline_three_chunks_hand_derived_count() {
let num_chunks = 3;
let num_frames = 4;
let mut segmentations = vec![0.0f64; num_chunks * num_frames * SEG_NUM_SLOTS];
let chunk_classes = [[1, 4, 0, 2], [1, 1, 6, 0], [0, 5, 3, 1]];
for (c, classes) in chunk_classes.iter().enumerate() {
let slab = classes_to_slab(classes);
segmentations[chunk_segmentation_range(c, num_frames)].copy_from_slice(&slab);
}
let count = crate::audio::speaker::window::count_from_segmentations(
&segmentations,
num_chunks,
num_frames,
SEG_NUM_SLOTS,
0.5,
SlidingWindow::new(0.0, 4.0, 2.0),
SlidingWindow::new(0.0, 1.0, 1.0),
);
assert_eq!(count, vec![1, 2, 0, 1, 1, 1, 1, 1, 0]);
assert_eq!(count.len(), 9); }
#[test]
fn compute_options_new_matches_default() {
assert_eq!(ComputeOptions::new(), ComputeOptions::default());
}
#[test]
fn compute_options_defaults_match_crate_consts() {
let o = ComputeOptions::new();
assert_eq!(
o.segmenter(),
crate::audio::speaker::segment::DEFAULT_SEGMENT_COMPUTE
);
assert_eq!(
o.embedder(),
crate::audio::speaker::embed::DEFAULT_EMBED_COMPUTE
);
assert_eq!(o.segmenter(), ComputeUnits::All);
assert_eq!(o.embedder(), ComputeUnits::All);
}
#[test]
fn compute_options_builders_and_setters() {
let o = ComputeOptions::new()
.with_segmenter(ComputeUnits::CpuOnly)
.with_embedder(ComputeUnits::CpuAndNeuralEngine);
assert_eq!(o.segmenter(), ComputeUnits::CpuOnly);
assert_eq!(o.embedder(), ComputeUnits::CpuAndNeuralEngine);
let mut m = ComputeOptions::new();
m.set_segmenter(ComputeUnits::CpuAndGpu);
m.set_embedder(ComputeUnits::CpuOnly);
assert_eq!(m.segmenter(), ComputeUnits::CpuAndGpu);
assert_eq!(m.embedder(), ComputeUnits::CpuOnly);
}
#[test]
fn compute_options_display_pins_the_spelling() {
assert_eq!(
ComputeOptions::new().to_string(),
"segmenter=all,embedder=all"
);
assert_eq!(
ComputeOptions::new()
.with_segmenter(ComputeUnits::CpuOnly)
.with_embedder(ComputeUnits::CpuAndGpu)
.to_string(),
"segmenter=cpu_only,embedder=cpu_and_gpu"
);
}
#[test]
fn options_new_matches_default() {
assert_eq!(Options::new(), Options::default());
}
#[test]
fn options_defaults_delegate_to_components() {
let o = Options::new();
assert_eq!(o.window(), WindowOptions::new());
assert_eq!(o.compute(), ComputeOptions::new());
assert_eq!(o.source(), Source::default());
assert_eq!(o.source(), Source::FluidAudio);
}
#[test]
fn options_builders_and_setters() {
let window = WindowOptions::new().with_onset(0.25);
let compute = ComputeOptions::new().with_segmenter(ComputeUnits::CpuOnly);
let source = Source::Argmax;
let o = Options::new()
.with_window(window)
.with_compute(compute)
.with_source(source);
assert_eq!(o.window(), window);
assert_eq!(o.compute(), compute);
assert_eq!(o.source(), source);
let mut m = Options::new();
m.set_window(window);
m.set_compute(compute);
m.set_source(source);
assert_eq!(m.window(), window);
assert_eq!(m.compute(), compute);
assert_eq!(m.source(), source);
}
#[test]
fn options_display_pins_the_composed_spelling() {
assert_eq!(
Options::new().to_string(),
"window=(step_samples=16000,onset=0.5),compute=(segmenter=all,embedder=all),source=fluid_audio"
);
let built = Options::new()
.with_window(
WindowOptions::new()
.with_step_samples(8_000)
.with_onset(0.75),
)
.with_compute(
ComputeOptions::new()
.with_segmenter(ComputeUnits::CpuOnly)
.with_embedder(ComputeUnits::CpuAndGpu),
)
.with_source(Source::Argmax);
assert_eq!(
built.to_string(),
"window=(step_samples=8000,onset=0.75),compute=(segmenter=cpu_only,embedder=cpu_and_gpu),source=argmax"
);
}
#[test]
fn extractor_new_matches_default_and_holds_default_options() {
assert_eq!(Extractor::new(), Extractor::default());
assert_eq!(*Extractor::new().options_ref(), Options::new());
}
#[test]
fn extractor_with_options_round_trips() {
let options = Options::new().with_window(WindowOptions::new().with_step_samples(40_000));
let extractor = Extractor::with_options(options);
assert_eq!(*extractor.options_ref(), options);
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_empty_object_is_full_defaults() {
let o: Options = serde_json::from_str("{}").unwrap();
assert_eq!(o, Options::new());
assert_eq!(o.source(), Source::FluidAudio);
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_partial_window_keeps_step_default() {
let o: Options = serde_json::from_str(r#"{"window":{"onset":0.25}}"#).unwrap();
assert_eq!(o.window().onset(), 0.25);
assert_eq!(o.window().step_samples(), 16_000);
assert_eq!(o.compute(), ComputeOptions::new());
assert_eq!(o.source(), Source::default());
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_partial_compute_defaults_other_unit() {
let o: Options = serde_json::from_str(r#"{"compute":{"segmenter":"cpu_only"}}"#).unwrap();
assert_eq!(o.compute().segmenter(), ComputeUnits::CpuOnly);
assert_eq!(o.compute().embedder(), ComputeUnits::All);
assert_eq!(o.window(), WindowOptions::new());
assert_eq!(o.source(), Source::default());
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_partial_source_defaults_others() {
let o: Options = serde_json::from_str(r#"{"source":"argmax"}"#).unwrap();
assert_eq!(o.source(), Source::Argmax);
assert_eq!(o.window(), WindowOptions::new());
assert_eq!(o.compute(), ComputeOptions::new());
}
#[cfg(feature = "serde")]
#[test]
fn options_serde_round_trips() {
let o = Options::new()
.with_window(
WindowOptions::new()
.with_step_samples(40_000)
.with_onset(0.7),
)
.with_compute(ComputeOptions::new().with_segmenter(ComputeUnits::CpuOnly))
.with_source(Source::Argmax);
let json = serde_json::to_string(&o).unwrap();
let back: Options = serde_json::from_str(&json).unwrap();
assert_eq!(back, o);
}
#[test]
fn into_offline_input_round_trips_against_real_dia() {
let e = tiny_extraction();
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let input = e.into_offline_input(&plda);
assert_eq!(input.raw_embeddings(), e.raw_embeddings());
assert_eq!(input.num_chunks(), e.num_chunks());
assert_eq!(input.num_speakers(), 3);
assert_eq!(input.num_speakers(), e.num_speakers());
assert_eq!(input.segmentations(), e.segmentations());
assert_eq!(input.num_frames_per_chunk(), e.num_frames_per_chunk());
assert_eq!(input.count(), e.count());
assert_eq!(input.num_output_frames(), e.num_output_frames());
let cs = input.chunks_sw();
assert_eq!(cs.start(), e.chunks_sw().start());
assert_eq!(cs.duration(), e.chunks_sw().duration());
assert_eq!(cs.step(), e.chunks_sw().step());
let fs = input.frames_sw();
assert_eq!(fs.start(), e.frames_sw().start());
assert_eq!(fs.duration(), e.frames_sw().duration());
assert_eq!(fs.step(), e.frames_sw().step());
assert!(std::ptr::eq(input.plda(), &plda));
}
#[test]
fn diarize_matches_manual_into_offline_input_pipeline() {
let e = tiny_extraction();
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let via_public = e.diarize(&plda);
let via_manual = diaric::offline::diarize_offline(&e.into_offline_input(&plda));
match (via_public, via_manual) {
(Ok(pub_out), Ok(man_out)) => {
let spans = |o: &diaric::offline::OfflineOutput| -> Vec<(f64, f64, usize)> {
o.spans_slice()
.iter()
.map(|s| (s.start(), s.end(), s.cluster()))
.collect()
};
assert_eq!(
spans(&pub_out),
spans(&man_out),
"diarize() spans diverged from into_offline_input → diarize_offline"
);
}
(Err(pub_err), Err(man_err)) => {
assert_eq!(
format!("{pub_err:?}"),
format!("{man_err:?}"),
"diarize() and the manual plumbing refused differently"
);
}
(pub_res, man_res) => panic!(
"diarize() ({}) diverged from manual into_offline_input → diarize_offline ({})",
if pub_res.is_ok() { "Ok" } else { "Err" },
if man_res.is_ok() { "Ok" } else { "Err" },
),
}
}
fn tiny_extraction() -> Extraction {
Extraction {
raw_embeddings: (0..(SEG_NUM_SLOTS * EMBEDDING_DIM))
.map(|i| i as f32 * 0.25 - 3.0)
.collect(),
segmentations: vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0],
count: vec![1, 1, 0, 0],
num_chunks: 1,
num_frames_per_chunk: 2,
num_output_frames: 4,
chunks_sw: crate::audio::speaker::window::chunk_sliding_window(&WindowOptions::new())
.with_duration(3.0 * crate::audio::speaker::window::FRAME_STEP_S),
frames_sw: crate::audio::speaker::window::frame_sliding_window(),
}
}
#[test]
fn diarize_with_offline_routes_the_backend_options() {
let e = tiny_extraction();
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let opts = crate::audio::speaker::cluster::OfflineOptions::new()
.with_threshold(0.55)
.with_fa(0.09)
.with_fb(0.71)
.with_max_iters(33)
.with_min_duration_off(1.25);
let via_public = e.diarize_with(&plda, ClusterBackend::Offline(opts));
let via_manual = diaric::offline::diarize_offline(&opts.apply_to(e.into_offline_input(&plda)));
match (via_public, via_manual) {
(Ok(pub_out), Ok(man_out)) => {
let spans = |o: &diaric::offline::OfflineOutput| -> Vec<(f64, f64, usize)> {
o.spans_slice()
.iter()
.map(|s| (s.start(), s.end(), s.cluster()))
.collect()
};
assert_eq!(
spans(&pub_out),
spans(&man_out),
"diarize_with routed a different OfflineInput than apply_to"
);
}
(Err(pub_err), Err(man_err)) => {
assert_eq!(
format!("{pub_err:?}"),
format!("{man_err:?}"),
"diarize_with and the apply_to path refused differently"
);
}
(p, m) => panic!(
"diarize_with ({}) diverged from the apply_to path ({})",
if p.is_ok() { "Ok" } else { "Err" },
if m.is_ok() { "Ok" } else { "Err" },
),
}
}
fn online_extraction() -> Extraction {
const F: usize = 4;
let seg_idx = |c: usize, f: usize, s: usize| (c * F + f) * SEG_NUM_SLOTS + s;
let mut segmentations = vec![0.0f64; 2 * F * SEG_NUM_SLOTS];
for f in 0..2 {
segmentations[seg_idx(0, f, 0)] = 1.0; }
for f in 2..4 {
segmentations[seg_idx(0, f, 1)] = 1.0; }
for f in 0..4 {
segmentations[seg_idx(1, f, 0)] = 1.0; }
for f in 0..2 {
segmentations[seg_idx(1, f, 1)] = 1.0; }
for f in 2..4 {
segmentations[seg_idx(1, f, 2)] = 1.0; }
let mut raw_embeddings = vec![0.0f32; 2 * SEG_NUM_SLOTS * EMBEDDING_DIM];
let mut set_block = |c: usize, s: usize, block: usize| {
let base = (c * SEG_NUM_SLOTS + s) * EMBEDDING_DIM;
raw_embeddings[(base + block * 64)..(base + (block + 1) * 64)].fill(1.0);
};
set_block(0, 0, 0); set_block(0, 1, 1); set_block(1, 0, 0); set_block(1, 1, 0); set_block(1, 2, 2);
let mut count = vec![0u8; 63];
count[0..4].fill(1);
count[59..63].fill(2);
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&WindowOptions::new())
.with_duration((F as f64 - 1.0) * crate::audio::speaker::window::FRAME_STEP_S);
Extraction {
raw_embeddings,
segmentations,
count,
num_chunks: 2,
num_frames_per_chunk: F,
num_output_frames: 63,
chunks_sw,
frames_sw: crate::audio::speaker::window::frame_sliding_window(),
}
}
#[test]
fn diarize_online_labels_slots_and_reconstructs_spans() {
let e = online_extraction();
let opts = OnlineOptions::new().with_min_speech_duration(0.0);
let out = e
.diarize_online(opts)
.expect("online reconstruction succeeds on a valid extraction");
assert_eq!(
out.hard_clusters_slice(),
&[[0, 1, -2], [0, 0, 2]],
"online per-slot labels (chunk order, slot order) diverged"
);
assert_eq!(out.num_clusters(), 3);
let spans = out.spans_slice();
assert!(!spans.is_empty(), "reconstruction produced no spans");
assert!(
spans.iter().all(|s| s.cluster() < 3),
"a span named a cluster outside the online roster: {:?}",
spans.iter().map(|s| s.cluster()).collect::<Vec<_>>()
);
}
#[test]
fn diarize_with_online_routes_to_diarize_online_ignoring_plda() {
let e = online_extraction();
let opts = OnlineOptions::new().with_min_speech_duration(0.0);
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let via_online = e.diarize_online(opts).expect("diarize_online ok");
let via_with = e
.diarize_with(&plda, ClusterBackend::Online(opts))
.expect("diarize_with(Online) ok");
assert_eq!(
via_online.hard_clusters_slice(),
via_with.hard_clusters_slice(),
"diarize_with(Online) routed to a different labelling than diarize_online"
);
assert_eq!(via_online.num_clusters(), via_with.num_clusters());
let spans = |o: &diaric::offline::OfflineOutput| -> Vec<(f64, f64, usize)> {
o.spans_slice()
.iter()
.map(|s| (s.start(), s.end(), s.cluster()))
.collect()
};
assert_eq!(spans(&via_online), spans(&via_with));
}
#[test]
fn diarize_online_default_options_drops_subsecond_slots() {
let e = online_extraction();
let out = e
.diarize_online(OnlineOptions::default())
.expect("online reconstruction succeeds even with all slots dropped");
assert_eq!(
out.hard_clusters_slice(),
&[[-2, -2, -2], [-2, -2, -2]],
"default min_speech_duration should drop every sub-second slot"
);
assert!(
out.spans_slice().is_empty(),
"all-dropped extraction must produce no spans"
);
}
fn online_extraction_default_gate() -> Extraction {
const F: usize = 64;
const ABOVE: usize = 64; const BELOW: usize = 20; let seg_idx = |c: usize, f: usize, s: usize| (c * F + f) * SEG_NUM_SLOTS + s;
let mut segmentations = vec![0.0f64; F * SEG_NUM_SLOTS];
for f in 0..ABOVE {
segmentations[seg_idx(0, f, 0)] = 1.0; }
for f in 0..BELOW {
segmentations[seg_idx(0, f, 1)] = 1.0; }
for f in 0..ABOVE {
segmentations[seg_idx(0, f, 2)] = 1.0; }
let mut raw_embeddings = vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM];
let mut set_block = |s: usize, block: usize| {
let base = s * EMBEDDING_DIM;
raw_embeddings[(base + block * 64)..(base + (block + 1) * 64)].fill(1.0);
};
set_block(0, 0); set_block(1, 1); set_block(2, 2);
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&WindowOptions::new())
.with_duration((F as f64 - 1.0) * crate::audio::speaker::window::FRAME_STEP_S);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let count = crate::audio::speaker::window::try_count_from_segmentations(
&segmentations,
1,
F,
SEG_NUM_SLOTS,
0.5,
chunks_sw,
frames_sw,
)
.expect("fixture chunk/frame geometry yields exactly F output frames");
Extraction {
raw_embeddings,
segmentations,
count,
num_chunks: 1,
num_frames_per_chunk: F,
num_output_frames: F,
chunks_sw,
frames_sw,
}
}
#[test]
fn diarize_online_default_gate_keeps_above_threshold_drops_below() {
let e = online_extraction_default_gate();
let out = e
.diarize_online(OnlineOptions::default())
.expect("online reconstruction succeeds on the default-gate fixture");
assert_eq!(
out.hard_clusters_slice(),
&[[0, -2, 1]],
"default-gate labels: above-threshold slots create speakers, the sub-second slot drops"
);
assert_eq!(out.num_clusters(), 2, "two above-threshold speakers");
let fs = e.frames_sw();
let center_offset = fs.duration() / 2.0;
let n = e.num_output_frames() as f64;
let span_start = fs.start() + center_offset; let span_end = fs.start() + (n - 1.0) * fs.step() + center_offset;
let span_dur = span_end - span_start;
let got: Vec<(usize, f64, f64)> = out
.spans_slice()
.iter()
.map(|s| (s.cluster(), s.start(), s.duration()))
.collect();
assert_eq!(
got,
vec![(0, span_start, span_dur), (1, span_start, span_dur)],
"default-gate spans: exactly clusters 0 and 1, each spanning the full output grid"
);
}
fn many_cluster_online_extraction(num_clusters: usize, num_frames_per_chunk: usize) -> Extraction {
assert!(
num_clusters <= 2 * EMBEDDING_DIM,
"`{{±e_i}}` yields at most 2*EMBEDDING_DIM ({}) distinct far vectors",
2 * EMBEDDING_DIM
);
let num_chunks = num_clusters.div_ceil(SEG_NUM_SLOTS);
let f = num_frames_per_chunk;
let mut segmentations = vec![0.0f64; num_chunks * f * SEG_NUM_SLOTS];
let mut raw_embeddings = vec![0.0f32; num_chunks * SEG_NUM_SLOTS * EMBEDDING_DIM];
for g in 0..num_clusters {
let c = g / SEG_NUM_SLOTS;
let s = g % SEG_NUM_SLOTS;
let pos = g % EMBEDDING_DIM;
let sign = if g < EMBEDDING_DIM { 1.0f32 } else { -1.0f32 };
raw_embeddings[(c * SEG_NUM_SLOTS + s) * EMBEDDING_DIM + pos] = sign;
for ff in 0..f {
segmentations[(c * f + ff) * SEG_NUM_SLOTS + s] = 1.0;
}
}
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&WindowOptions::new())
.with_duration((f as f64 - 1.0) * crate::audio::speaker::window::FRAME_STEP_S);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let count = crate::audio::speaker::window::count_from_segmentations(
&segmentations,
num_chunks,
f,
SEG_NUM_SLOTS,
0.5,
chunks_sw,
frames_sw,
);
Extraction::from_parts(
raw_embeddings,
segmentations,
count,
num_chunks,
f,
chunks_sw,
frames_sw,
)
}
#[test]
fn diarize_online_many_clusters_use_no_cluster_axis_allocation() {
const NUM_CLUSTERS: usize = 380;
const F: usize = 4;
let e = many_cluster_online_extraction(NUM_CLUSTERS, F);
assert_eq!(e.num_chunks(), 127, "ceil(380/3) chunks");
let out = e
.diarize_online(OnlineOptions::new().with_min_speech_duration(0.0))
.expect("high-churn online reconstruction succeeds with no cluster-axis allocation");
assert_eq!(
out.num_clusters(),
NUM_CLUSTERS,
"every distinct far embedding seeds its own global speaker"
);
let hc = out.hard_clusters_slice();
assert_eq!(hc.len(), 127);
assert_eq!(hc[0], [0, 1, 2], "first chunk seeds labels 0,1,2");
assert_eq!(
hc[126],
[378, 379, -2],
"last chunk: two labels + the dropped tail"
);
let spans = out.spans_slice();
assert!(
!spans.is_empty(),
"reconstruction produced spans for the many clusters"
);
assert!(
spans.iter().all(|s| s.cluster() < NUM_CLUSTERS),
"every span names a cluster inside the 380-speaker roster"
);
}
#[test]
fn diarize_online_over_cap_grid_is_a_typed_reconstruct_error_not_an_oom() {
const NUM_CLUSTERS: usize = 380;
const F: usize = 8300; let e = many_cluster_online_extraction(NUM_CLUSTERS, F);
let err = e
.diarize_online(OnlineOptions::new().with_min_speech_duration(0.0))
.expect_err("an over-cap clustered grid must be a typed reconstruct error, not an OOM/panic");
assert!(
matches!(
err,
diaric::offline::Error::Reconstruct(diaric::reconstruct::Error::Shape(
diaric::reconstruct::ShapeError::OutputGridTooLarge { .. }
))
),
"expected Reconstruct(Shape(OutputGridTooLarge)), got {err:?}"
);
}
fn all_new_online_extraction(
num_speakers: usize,
nan_cell: Option<(usize, usize, usize)>,
) -> Extraction {
const F: usize = 4;
let num_chunks = num_speakers.div_ceil(SEG_NUM_SLOTS);
let mut segmentations = vec![0.0f64; num_chunks * F * SEG_NUM_SLOTS];
let mut raw_embeddings = vec![0.0f32; num_chunks * SEG_NUM_SLOTS * EMBEDDING_DIM];
for g in 0..num_speakers {
let c = g / SEG_NUM_SLOTS;
let s = g % SEG_NUM_SLOTS;
let base = (c * SEG_NUM_SLOTS + s) * EMBEDDING_DIM;
raw_embeddings[base..base + EMBEDDING_DIM].fill(1.0);
for ff in 0..F {
segmentations[(c * F + ff) * SEG_NUM_SLOTS + s] = 1.0;
}
}
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&WindowOptions::new())
.with_duration((F as f64 - 1.0) * crate::audio::speaker::window::FRAME_STEP_S);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let count = crate::audio::speaker::window::count_from_segmentations(
&segmentations,
num_chunks,
F,
SEG_NUM_SLOTS,
0.5,
chunks_sw,
frames_sw,
);
if let Some((c, ff, s)) = nan_cell {
segmentations[(c * F + ff) * SEG_NUM_SLOTS + s] = f64::NAN;
}
Extraction::from_parts(
raw_embeddings,
segmentations,
count,
num_chunks,
F,
chunks_sw,
frames_sw,
)
}
#[test]
fn diarize_online_early_cap_not_late_reconstruction_rejection() {
const NUM_SPEAKERS: usize = 1200; let ceiling = diaric::reconstruct::MAX_CLUSTER_ID as usize + 1;
assert!(
NUM_SPEAKERS > ceiling,
"fixture ({NUM_SPEAKERS} speakers) must exceed the {ceiling}-speaker ceiling to reach the cap"
);
let e = all_new_online_extraction(NUM_SPEAKERS, Some((350, 0, 0)));
let opts = OnlineOptions::default()
.with_speaker_threshold(0.0)
.with_min_speech_duration(0.0);
let err = e
.diarize_online(opts)
.expect_err("past MAX_CLUSTER_ID the online loop must return the typed cap error early");
assert!(
matches!(
err,
diaric::offline::Error::Reconstruct(diaric::reconstruct::Error::Shape(
diaric::reconstruct::ShapeError::HardClustersIdAboveMax
))
),
"expected an EARLY Reconstruct(Shape(HardClustersIdAboveMax)) from the assign-loop cap \
(removing the guard surfaces NonFinite(Segmentations) from the planted NaN instead), got {err:?}"
);
}
#[test]
fn diarize_online_accepts_exactly_max_cluster_id_plus_one_speakers() {
let ceiling = diaric::reconstruct::MAX_CLUSTER_ID as usize + 1;
assert_eq!(
ceiling, 1024,
"diaric's reconstruction ceiling is MAX_CLUSTER_ID + 1"
);
let e = all_new_online_extraction(ceiling, None);
let out = e
.diarize_online(
OnlineOptions::default()
.with_speaker_threshold(0.0)
.with_min_speech_duration(0.0),
)
.expect("exactly MAX_CLUSTER_ID + 1 speakers sit ON the ceiling and must reconstruct");
assert_eq!(
out.num_clusters(),
ceiling,
"every one of the 1024 all-New slots keeps its own cluster (labels 0..=1023)"
);
}
fn models_dir() -> std::path::PathBuf {
std::env::var_os("SPEAKERKIT_TEST_MODELS").map_or_else(
|| crate::tests::models_root().join("speakerkit"),
std::path::PathBuf::from,
)
}
fn load_seg_model() -> SegmentModel {
SegmentModel::from_file_with(
models_dir().join("pyannote_segmentation.mlmodelc"),
crate::audio::speaker::segment::SegmentModelOptions::new().with_compute(ComputeUnits::CpuOnly),
)
.expect("load pyannote_segmentation.mlmodelc")
}
fn load_embed_model() -> EmbedModel {
EmbedModel::from_file_with(
models_dir().join("wespeaker_v2.mlmodelc"),
crate::audio::speaker::embed::EmbedModelOptions::new().with_compute(ComputeUnits::CpuOnly),
)
.expect("load wespeaker_v2.mlmodelc")
}
fn load_ted_60() -> Vec<f32> {
let path = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests/whisper/fixtures/audio/ted_60.wav");
let mut reader = hound::WavReader::open(&path).expect("ted_60.wav opens");
let spec = reader.spec();
assert_eq!(spec.channels, 1, "fixture must be mono");
assert_eq!(spec.sample_rate, 16_000, "fixture must be 16 kHz");
assert_eq!(spec.sample_format, hound::SampleFormat::Int);
reader
.samples::<i16>()
.map(|s| f32::from(s.expect("valid sample")) / 32_768.0)
.collect()
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn extract_ted30_invariants() {
let seg = load_seg_model();
let embed = load_embed_model();
let all = load_ted_60();
assert_eq!(all.len(), 960_000, "ted_60.wav is 60 s at 16 kHz");
let samples = &all[..480_000];
let extraction = Extractor::new()
.extract(&seg, &embed, samples)
.expect("extract on 30 s of ted_60");
let f = seg.num_frames();
assert_eq!(extraction.num_chunks(), 21);
assert_eq!(extraction.num_frames_per_chunk(), f);
assert_eq!(extraction.num_speakers(), 3);
assert_eq!(extraction.raw_embeddings().len(), 21 * 3 * EMBEDDING_DIM);
assert_eq!(extraction.segmentations().len(), 21 * f * 3);
assert_eq!(extraction.count().len(), extraction.num_output_frames());
assert_eq!(extraction.num_output_frames(), 1779);
assert!(
extraction.count().iter().all(|c| *c <= 3),
"count never exceeds SEG_NUM_SLOTS = 3"
);
assert!(
extraction.raw_embeddings().iter().all(|v| v.is_finite()),
"every raw embedding value is finite"
);
assert!(
extraction
.segmentations()
.iter()
.all(|v| *v == 0.0 || *v == 1.0),
"hard multilabel: every segmentation value is exactly 0.0 or 1.0"
);
assert!(
(0..extraction.num_chunks() * 3).any(|i| extraction.raw_embeddings()
[i * EMBEDDING_DIM..(i + 1) * EMBEDDING_DIM]
.iter()
.any(|v| *v != 0.0)),
"at least one embedding row is non-zero (real speech survives the drop paths)"
);
for c in 0..extraction.num_chunks() {
for s in 0..3 {
let row = &extraction.raw_embeddings()[embedding_range(c, s)];
let row_zero = row.iter().all(|v| *v == 0.0);
let col_zero =
(0..f).all(|frame| extraction.segmentations()[(c * f + frame) * SEG_NUM_SLOTS + s] == 0.0);
assert_eq!(
row_zero, col_zero,
"chunk {c} slot {s}: embedding-row-zero must match segmentation-column-zero"
);
}
}
}
fn caller_padded_chunk(samples: &[f32], start: usize) -> Vec<f32> {
let mut padded = vec![0.0f32; SEG_CHUNK_SAMPLES];
let end = (start + SEG_CHUNK_SAMPLES).min(samples.len());
let lo = start.min(samples.len());
padded[..end - lo].copy_from_slice(&samples[lo..end]);
padded
}
fn first_bit_divergence_f32(a: &[f32], b: &[f32]) -> Option<usize> {
a.iter()
.zip(b.iter())
.position(|(x, y)| x.to_bits() != y.to_bits())
}
fn first_bit_divergence_f64(a: &[f64], b: &[f64]) -> Option<usize> {
a.iter()
.zip(b.iter())
.position(|(x, y)| x.to_bits() != y.to_bits())
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn extract_chunk_embeddings_reassembles_extract_byte_for_byte() {
use crate::audio::speaker::window::{
chunk_sliding_window, chunk_starts, count_from_segmentations, frame_sliding_window,
};
let seg = load_seg_model();
let embed = load_embed_model();
let all = load_ted_60();
let samples = &all[..480_000];
let extractor = Extractor::new();
let fused = extractor
.extract(&seg, &embed, samples)
.expect("the fused path extracts 30 s of ted_60");
let w = extractor.options_ref().window();
let num_frames = embed.num_mask_frames();
let starts = chunk_starts(samples.len(), &w);
let mut raw_embeddings: Vec<f32> = Vec::new();
let mut segmentations: Vec<f64> = Vec::new();
for &start in &starts {
let padded = caller_padded_chunk(samples, start);
let logits = seg.infer(&padded).expect("segment a chunk");
let mut slab = crate::audio::speaker::segment::multilabel(&logits, num_frames);
let rows = extractor
.extract_chunk_embeddings(&embed, &padded, &mut slab)
.expect("the embed door runs a chunk");
assert_eq!(rows.len(), SEG_NUM_SLOTS * EMBEDDING_DIM);
assert_eq!(slab.len(), num_frames * SEG_NUM_SLOTS);
raw_embeddings.extend_from_slice(&rows);
segmentations.extend_from_slice(&slab);
}
let chunks_sw = chunk_sliding_window(&w);
let frames_sw = frame_sliding_window();
let count = count_from_segmentations(
&segmentations,
starts.len(),
num_frames,
SEG_NUM_SLOTS,
w.onset(),
chunks_sw,
frames_sw,
);
let split = Extraction::try_from_parts(ExtractionParts {
raw_embeddings,
segmentations,
count,
num_chunks: starts.len(),
num_frames_per_chunk: num_frames,
chunks_sw,
frames_sw,
})
.expect("the split road's parts pass every try_from_parts check");
assert_eq!(split.num_chunks(), fused.num_chunks());
assert_eq!(split.num_frames_per_chunk(), fused.num_frames_per_chunk());
assert_eq!(split.num_speakers(), fused.num_speakers());
assert_eq!(split.num_output_frames(), fused.num_output_frames());
assert_eq!(split.chunks_sw(), fused.chunks_sw());
assert_eq!(split.frames_sw(), fused.frames_sw());
assert_eq!(split.count(), fused.count(), "count tensors differ");
assert_eq!(
split.raw_embeddings().len(),
fused.raw_embeddings().len(),
"raw_embeddings lengths differ"
);
assert_eq!(
first_bit_divergence_f32(split.raw_embeddings(), fused.raw_embeddings()),
None,
"raw_embeddings diverge at the reported index (split vs fused)"
);
assert_eq!(
split.segmentations().len(),
fused.segmentations().len(),
"segmentations lengths differ"
);
assert_eq!(
first_bit_divergence_f64(split.segmentations(), fused.segmentations()),
None,
"segmentations diverge at the reported index (split vs fused)"
);
assert!(
fused.raw_embeddings().iter().any(|v| *v != 0.0),
"the fused reference must carry real embeddings for this to prove anything"
);
assert!(
fused.segmentations().iter().any(|v| *v != 0.0),
"the fused reference must carry real activity for this to prove anything"
);
assert_eq!(split, fused);
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn extract_chunk_embeddings_refuses_a_chunk_that_is_not_the_window_length() {
use crate::audio::speaker::error::{InferError, InputLength};
let embed = load_embed_model();
let num_frames = embed.num_mask_frames();
let mut slab = vec![0.0f64; num_frames * SEG_NUM_SLOTS];
for got in [0usize, SEG_CHUNK_SAMPLES - 1, SEG_CHUNK_SAMPLES + 1] {
assert_eq!(
Extractor::new().extract_chunk_embeddings(&embed, &vec![0.0f32; got], &mut slab),
Err(ExtractError::Infer(InferError::InputLength(
InputLength::new(got, SEG_CHUNK_SAMPLES)
))),
"a {got}-sample chunk must be refused"
);
}
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn extract_chunk_embeddings_refuses_a_slab_the_embedder_disagrees_with() {
use crate::audio::speaker::error::{ExtractionLenMismatch, ExtractionPart};
let embed = load_embed_model();
let expected = embed.num_mask_frames() * SEG_NUM_SLOTS;
let padded = vec![0.0f32; SEG_CHUNK_SAMPLES];
for got in [0usize, expected - 1, expected + 1] {
let mut slab = vec![0.0f64; got];
assert_eq!(
Extractor::new().extract_chunk_embeddings(&embed, &padded, &mut slab),
Err(ExtractError::ExtractionLenMismatch(
ExtractionLenMismatch::new(ExtractionPart::Segmentations, got, expected)
)),
"a {got}-value slab must be refused against the embedder's {expected}"
);
}
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn extract_empty_samples_errors() {
let seg = load_seg_model();
let embed = load_embed_model();
assert_eq!(
Extractor::new().extract(&seg, &embed, &[]),
Err(ExtractError::EmptySamples)
);
}
use crate::audio::speaker::error::ExtractionPart;
fn valid_parts() -> ExtractionParts {
let e = tiny_extraction();
ExtractionParts {
raw_embeddings: e.raw_embeddings().to_vec(),
segmentations: e.segmentations().to_vec(),
count: e.count().to_vec(),
num_chunks: e.num_chunks(),
num_frames_per_chunk: e.num_frames_per_chunk(),
chunks_sw: e.chunks_sw(),
frames_sw: e.frames_sw(),
}
}
fn rebuild_through_public_api(e: &Extraction) -> Result<Extraction, ExtractError> {
Extraction::try_from_parts(ExtractionParts {
raw_embeddings: e.raw_embeddings().to_vec(),
segmentations: e.segmentations().to_vec(),
count: e.count().to_vec(),
num_chunks: e.num_chunks(),
num_frames_per_chunk: e.num_frames_per_chunk(),
chunks_sw: e.chunks_sw(),
frames_sw: e.frames_sw(),
})
}
type OutputFingerprint = (
Vec<(f64, f64, usize)>,
Vec<diaric::pipeline::ChunkAssignment>,
usize,
Vec<f32>,
);
fn output_fingerprint(o: &diaric::offline::OfflineOutput) -> OutputFingerprint {
(
o.spans_slice()
.iter()
.map(|s| (s.start(), s.end(), s.cluster()))
.collect(),
o.hard_clusters_slice().to_vec(),
o.num_clusters(),
o.discrete_diarization_slice().to_vec(),
)
}
#[test]
fn try_from_parts_round_trips_an_extraction_through_the_public_api() {
let original = tiny_extraction();
let rebuilt = rebuild_through_public_api(&original).expect("a real Extraction's own parts");
assert_eq!(
rebuilt, original,
"rebuilt Extraction diverged from the original"
);
assert_eq!(rebuilt.num_output_frames(), original.count().len());
assert_eq!(rebuilt.num_speakers(), SEG_NUM_SLOTS);
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let backend = ClusterBackend::default();
let from_original = original.diarize_with(&plda, backend);
let from_rebuilt = rebuilt.diarize_with(&plda, backend);
match (from_original, from_rebuilt) {
(Ok(a), Ok(b)) => assert_eq!(
output_fingerprint(&a),
output_fingerprint(&b),
"rebuilt Extraction produced a different OfflineOutput"
),
(Err(a), Err(b)) => assert_eq!(
format!("{a:?}"),
format!("{b:?}"),
"rebuilt Extraction refused differently"
),
(a, b) => panic!(
"rebuilt Extraction diverged: original {} vs rebuilt {}",
if a.is_ok() { "Ok" } else { "Err" },
if b.is_ok() { "Ok" } else { "Err" },
),
}
}
#[test]
fn try_from_parts_round_trips_the_online_backend_too() {
let original = online_extraction();
let rebuilt = rebuild_through_public_api(&original).expect("a real Extraction's own parts");
assert_eq!(rebuilt, original);
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let backend = ClusterBackend::Online(OnlineOptions::new().with_min_speech_duration(0.0));
let a = original
.diarize_with(&plda, backend)
.expect("online reconstruction succeeds on a valid extraction");
let b = rebuilt
.diarize_with(&plda, backend)
.expect("online reconstruction succeeds on the rebuilt extraction");
assert_eq!(
output_fingerprint(&a),
output_fingerprint(&b),
"rebuilt Extraction produced a different online OfflineOutput"
);
}
#[test]
fn try_from_parts_accepts_self_consistent_parts_and_derives_num_output_frames() {
let mut parts = valid_parts();
parts.chunks_sw = parts
.chunks_sw
.with_duration(6.0 * crate::audio::speaker::window::FRAME_STEP_S);
parts.count = vec![1, 1, 0, 0, 0, 0, 0];
let e = Extraction::try_from_parts(parts).expect("self-consistent parts");
assert_eq!(e.num_output_frames(), 7);
assert_eq!(e.count().len(), 7);
assert_eq!(e.num_speakers(), SEG_NUM_SLOTS);
assert_eq!(e.num_chunks(), 1);
assert_eq!(e.num_frames_per_chunk(), 2);
}
#[test]
fn try_from_parts_rejects_zero_num_chunks() {
let parts = ExtractionParts {
raw_embeddings: Vec::new(),
segmentations: Vec::new(),
num_chunks: 0,
..valid_parts()
};
assert_eq!(
Extraction::try_from_parts(parts).unwrap_err(),
ExtractError::ZeroExtractionDimension(ExtractionPart::NumChunks)
);
}
#[test]
fn try_from_parts_rejects_zero_num_frames_per_chunk() {
let parts = ExtractionParts {
segmentations: Vec::new(),
num_frames_per_chunk: 0,
..valid_parts()
};
assert_eq!(
Extraction::try_from_parts(parts).unwrap_err(),
ExtractError::ZeroExtractionDimension(ExtractionPart::NumFramesPerChunk)
);
}
#[test]
fn try_from_parts_rejects_empty_count() {
let parts = ExtractionParts {
count: Vec::new(),
..valid_parts()
};
assert_eq!(
Extraction::try_from_parts(parts).unwrap_err(),
ExtractError::ZeroExtractionDimension(ExtractionPart::Count)
);
}
#[test]
fn try_from_parts_rejects_non_positive_chunks_sw_step() {
let base = valid_parts();
let parts = ExtractionParts {
chunks_sw: base.chunks_sw.with_step(0.0),
..base
};
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::InvalidSlidingWindow(w) = err else {
panic!("expected InvalidSlidingWindow, got {err:?}")
};
assert_eq!(w.part(), ExtractionPart::ChunksSw);
assert_eq!(w.window().step(), 0.0);
}
#[test]
fn try_from_parts_rejects_non_finite_frames_sw_duration() {
let base = valid_parts();
let parts = ExtractionParts {
frames_sw: base.frames_sw.with_duration(f64::NAN),
..base
};
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::InvalidSlidingWindow(w) = err else {
panic!("expected InvalidSlidingWindow, got {err:?}")
};
assert_eq!(w.part(), ExtractionPart::FramesSw);
assert!(w.window().duration().is_nan());
}
#[test]
fn try_from_parts_rejects_infinite_frames_sw_duration() {
let base = valid_parts();
let parts = ExtractionParts {
frames_sw: base.frames_sw.with_duration(f64::INFINITY),
..base
};
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::InvalidSlidingWindow(w) = err else {
panic!("expected InvalidSlidingWindow, got {err:?}")
};
assert_eq!(w.part(), ExtractionPart::FramesSw);
assert_eq!(w.window().duration(), f64::INFINITY);
}
#[test]
fn try_from_parts_rejects_zero_frames_sw_duration() {
let base = valid_parts();
let parts = ExtractionParts {
frames_sw: base.frames_sw.with_duration(0.0),
..base
};
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::InvalidSlidingWindow(w) = err else {
panic!("expected InvalidSlidingWindow, got {err:?}")
};
assert_eq!(w.part(), ExtractionPart::FramesSw);
assert_eq!(w.window().duration(), 0.0);
}
#[test]
fn try_from_parts_rejects_infinite_frames_sw_step() {
let base = valid_parts();
let parts = ExtractionParts {
frames_sw: base.frames_sw.with_step(f64::INFINITY),
..base
};
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::InvalidSlidingWindow(w) = err else {
panic!("expected InvalidSlidingWindow, got {err:?}")
};
assert_eq!(w.part(), ExtractionPart::FramesSw);
assert_eq!(w.window().step(), f64::INFINITY);
}
#[test]
fn try_from_parts_rejects_non_finite_sliding_window_start() {
let base = valid_parts();
let parts = ExtractionParts {
chunks_sw: base.chunks_sw.with_start(f64::INFINITY),
..base
};
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::InvalidSlidingWindow(w) = err else {
panic!("expected InvalidSlidingWindow, got {err:?}")
};
assert_eq!(w.part(), ExtractionPart::ChunksSw);
assert!(w.window().start().is_infinite());
}
#[test]
fn try_from_parts_rejects_raw_embeddings_geometry_overflow() {
let parts = ExtractionParts {
raw_embeddings: Vec::new(),
segmentations: Vec::new(),
num_chunks: 1usize << 60,
num_frames_per_chunk: 1,
..valid_parts()
};
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::ExtractionGeometryOverflow(g) = err else {
panic!("expected ExtractionGeometryOverflow, got {err:?}")
};
assert_eq!(g.part(), ExtractionPart::RawEmbeddings);
assert_eq!(g.num_chunks(), 1usize << 60);
}
#[test]
fn try_from_parts_rejects_segmentations_geometry_overflow() {
let parts = ExtractionParts {
raw_embeddings: Vec::new(),
segmentations: Vec::new(),
num_chunks: 1usize << 32,
num_frames_per_chunk: 1usize << 32,
..valid_parts()
};
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::ExtractionGeometryOverflow(g) = err else {
panic!("expected ExtractionGeometryOverflow, got {err:?}")
};
assert_eq!(g.part(), ExtractionPart::Segmentations);
assert_eq!(g.num_chunks(), 1usize << 32);
assert_eq!(g.num_frames_per_chunk(), 1usize << 32);
}
#[test]
fn try_from_parts_rejects_raw_embeddings_len_mismatch() {
let mut parts = valid_parts();
parts.raw_embeddings.pop();
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::ExtractionLenMismatch(m) = err else {
panic!("expected ExtractionLenMismatch, got {err:?}")
};
assert_eq!(m.part(), ExtractionPart::RawEmbeddings);
assert_eq!(
(m.got(), m.expected()),
(767, SEG_NUM_SLOTS * EMBEDDING_DIM)
);
let rendered = err.to_string();
assert!(rendered.contains("raw_embeddings"), "{rendered}");
assert!(rendered.contains("767"), "{rendered}");
assert!(rendered.contains("768"), "{rendered}");
}
#[test]
fn try_from_parts_rejects_segmentations_len_mismatch() {
let mut parts = valid_parts();
parts.segmentations.pop();
let err = Extraction::try_from_parts(parts).unwrap_err();
let ExtractError::ExtractionLenMismatch(m) = err else {
panic!("expected ExtractionLenMismatch, got {err:?}")
};
assert_eq!(m.part(), ExtractionPart::Segmentations);
assert_eq!((m.got(), m.expected()), (5, 6));
let rendered = err.to_string();
assert!(rendered.contains("segmentations"), "{rendered}");
assert!(rendered.contains('5'), "{rendered}");
assert!(rendered.contains('6'), "{rendered}");
}
#[test]
fn try_from_parts_rejects_geometry_whose_output_frame_count_overflows() {
let parts = ExtractionParts {
chunks_sw: SlidingWindow::new(0.0, 1e300, 1.0),
frames_sw: SlidingWindow::new(0.0, 0.1, 1e-300),
..valid_parts()
};
assert_eq!(
Extraction::try_from_parts(parts).unwrap_err(),
ExtractError::OutputFrameCountOverflow
);
}
#[test]
fn try_from_parts_guarantee_makes_diarize_online_panic_free() {
let e = Extraction::try_from_parts(valid_parts()).expect("self-consistent parts");
let _ = e.diarize_online(OnlineOptions::new());
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let _ = e.diarize(&plda);
}
#[test]
fn diarize_online_refuses_an_oversized_derived_grid_before_allocating_it() {
let chunks_sw = SlidingWindow::new(0.0, 1e13, 1.0);
let frames_sw = SlidingWindow::new(0.0, 0.06, 0.01);
let err = refused(ExtractionParts {
chunks_sw,
frames_sw,
..valid_parts()
});
let ExtractError::ExtractionLenMismatch(m) = err else {
panic!("expected an ExtractionLenMismatch on count, got {err:?}")
};
assert_eq!(m.part(), ExtractionPart::Count);
assert_eq!(m.got(), 4);
let base = tiny_extraction();
let e = Extraction {
chunks_sw,
frames_sw,
..base
};
assert_eq!(e.num_output_frames(), 4);
let err = e
.diarize_online(OnlineOptions::new())
.expect_err("a derived grid that does not match num_output_frames must be refused");
assert!(
matches!(
err,
diaric::offline::Error::Reconstruct(diaric::reconstruct::Error::Shape(
diaric::reconstruct::ShapeError::CountLenMismatch
))
),
"expected Reconstruct(Shape(CountLenMismatch)), got {err:?}"
);
}
#[test]
fn try_from_parts_cannot_detect_mutually_inconsistent_parts() {
let a = online_extraction();
let mut other_track_embeddings = vec![0.0f32; 2 * SEG_NUM_SLOTS * EMBEDDING_DIM];
for c in 0..2 {
for s in 0..SEG_NUM_SLOTS {
let base = (c * SEG_NUM_SLOTS + s) * EMBEDDING_DIM;
other_track_embeddings[base..base + 64].fill(1.0);
}
}
assert_eq!(
other_track_embeddings.len(),
a.raw_embeddings().len(),
"the two tracks must be shape-identical for this to be about provenance"
);
let crossed = Extraction::try_from_parts(ExtractionParts {
raw_embeddings: other_track_embeddings,
segmentations: a.segmentations().to_vec(),
count: a.count().to_vec(),
num_chunks: a.num_chunks(),
num_frames_per_chunk: a.num_frames_per_chunk(),
chunks_sw: a.chunks_sw(),
frames_sw: a.frames_sw(),
})
.expect("shape-valid parts from two tracks are ACCEPTED — the gap this test pins");
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let online = ClusterBackend::Online(OnlineOptions::new().with_min_speech_duration(0.0));
let real = output_fingerprint(&a.diarize_with(&plda, online).expect("real track diarizes"));
let mixed = output_fingerprint(&crossed.diarize_with(&plda, online).expect("mixed diarizes"));
assert_eq!(real.2, 3, "the real track has three online clusters");
assert_eq!(mixed.2, 1, "the mixed one collapses to a single speaker");
assert_ne!(
real, mixed,
"a silently-accepted cross-track mix-up must at least be observable here"
);
}
fn extract_shaped_extraction(num_chunks: usize, num_frames_per_chunk: usize) -> Extraction {
let w = WindowOptions::new();
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let mut segmentations = vec![0.0f64; num_chunks * num_frames_per_chunk * SEG_NUM_SLOTS];
for c in 0..num_chunks {
for f in 0..num_frames_per_chunk {
segmentations[(c * num_frames_per_chunk + f) * SEG_NUM_SLOTS + (c % SEG_NUM_SLOTS)] = 1.0;
}
}
let count = crate::audio::speaker::window::try_count_from_segmentations(
&segmentations,
num_chunks,
num_frames_per_chunk,
SEG_NUM_SLOTS,
w.onset(),
chunks_sw,
frames_sw,
)
.expect("this geometry's output-frame count fits usize");
let raw_embeddings: Vec<f32> = (0..(num_chunks * SEG_NUM_SLOTS * EMBEDDING_DIM))
.map(|i| ((i % 64) as f32).mul_add(0.01, 0.5))
.collect();
Extraction::try_from_parts(ExtractionParts {
raw_embeddings,
segmentations,
count,
num_chunks,
num_frames_per_chunk,
chunks_sw,
frames_sw,
})
.unwrap_or_else(|err| panic!("extract()-shaped parts must be accepted: {err}"))
}
#[test]
fn try_from_parts_round_trips_an_extract_shaped_extraction() {
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
for (num_chunks, num_frames_per_chunk) in [(1usize, 4usize), (2, 8), (5, 3)] {
let original = extract_shaped_extraction(num_chunks, num_frames_per_chunk);
let rebuilt =
rebuild_through_public_api(&original).expect("an extract()-shaped Extraction's own parts");
assert_eq!(rebuilt, original, "({num_chunks}, {num_frames_per_chunk})");
for backend in [
ClusterBackend::default(),
ClusterBackend::Online(OnlineOptions::new().with_min_speech_duration(0.0)),
] {
match (
original.diarize_with(&plda, backend),
rebuilt.diarize_with(&plda, backend),
) {
(Ok(a), Ok(b)) => assert_eq!(
output_fingerprint(&a),
output_fingerprint(&b),
"({num_chunks}, {num_frames_per_chunk}) diverged on {backend:?}"
),
(Err(a), Err(b)) => assert_eq!(format!("{a:?}"), format!("{b:?}")),
(a, b) => panic!(
"({num_chunks}, {num_frames_per_chunk}) diverged on {backend:?}: original {} vs \
rebuilt {}",
if a.is_ok() { "Ok" } else { "Err" },
if b.is_ok() { "Ok" } else { "Err" },
),
}
}
}
}
#[test]
fn diarize_online_never_refuses_a_count_derived_the_way_extract_derives_it() {
for (num_chunks, num_frames_per_chunk) in [(1usize, 4usize), (2, 8), (5, 3)] {
let e = extract_shaped_extraction(num_chunks, num_frames_per_chunk);
let outcome = e.diarize_online(OnlineOptions::new());
assert!(
!matches!(
outcome,
Err(diaric::offline::Error::Reconstruct(
diaric::reconstruct::Error::Shape(diaric::reconstruct::ShapeError::CountLenMismatch)
))
),
"the derived-grid guard refused an extract()-derived count at \
({num_chunks} chunks, {num_frames_per_chunk} frames/chunk)"
);
}
}
fn one_usable_slot_row(slot: usize) -> Vec<f32> {
let mut raw = vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM];
let base = slot * EMBEDDING_DIM;
raw[base..base + 64].fill(1.0);
raw
}
fn unit_sw() -> SlidingWindow {
SlidingWindow::new(0.0, 1.0, 1.0)
}
#[track_caller]
fn refused(parts: ExtractionParts) -> ExtractError {
match Extraction::try_from_parts(parts) {
Err(e) => e,
Ok(e) => panic!(
"try_from_parts ACCEPTED these parts: num_chunks={}, num_frames_per_chunk={}, \
num_output_frames={}, chunks_sw={:?}, frames_sw={:?}",
e.num_chunks(),
e.num_frames_per_chunk(),
e.num_output_frames(),
e.chunks_sw(),
e.frames_sw()
),
}
}
#[test]
fn try_from_parts_rejects_an_active_slot_whose_embedding_row_is_unusable() {
let mut raw = vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM];
raw[0] = f32::NAN;
let parts = ExtractionParts {
raw_embeddings: raw,
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 1],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
};
let err = refused(parts);
let ExtractError::ActiveSlotWithoutEmbedding(a) = err else {
panic!("expected ActiveSlotWithoutEmbedding, got {err:?}")
};
assert_eq!((a.chunk(), a.slot()), (0, 0));
let parts = ExtractionParts {
raw_embeddings: vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM],
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 1],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
};
let err = refused(parts);
assert!(
matches!(err, ExtractError::ActiveSlotWithoutEmbedding(a) if (a.chunk(), a.slot()) == (0, 0)),
"expected ActiveSlotWithoutEmbedding(0, 0), got {err:?}"
);
}
#[test]
fn try_from_parts_rejects_a_count_the_segmentations_cannot_support() {
for (claimed, supported) in [(4u8, 1u8), (3, 1)] {
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![claimed, claimed],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
};
let err = refused(parts);
let ExtractError::CountNotSegmentationDerived(c) = err else {
panic!("expected CountNotSegmentationDerived for count {claimed}, got {err:?}")
};
assert_eq!((c.frame(), c.got(), c.expected()), (0, claimed, supported));
}
}
#[test]
fn try_from_parts_rejects_a_count_length_the_geometry_does_not_derive() {
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0],
count: vec![1],
num_chunks: 1,
num_frames_per_chunk: 1,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
};
let err = refused(parts);
let ExtractError::ExtractionLenMismatch(m) = err else {
panic!("expected ExtractionLenMismatch, got {err:?}")
};
assert_eq!(m.part(), ExtractionPart::Count);
assert_eq!((m.got(), m.expected()), (1, 2));
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0],
count: vec![1; 10],
num_chunks: 1,
num_frames_per_chunk: 1,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
};
let err = refused(parts);
assert!(
matches!(err, ExtractError::ExtractionLenMismatch(m)
if m.part() == ExtractionPart::Count && (m.got(), m.expected()) == (10, 2)),
"expected an ExtractionLenMismatch(Count, 10, 2), got {err:?}"
);
}
#[test]
fn try_from_parts_rejects_an_uncancelled_sliding_window_origin() {
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 1],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: SlidingWindow::new(-1.0, 1.0, 1.0),
frames_sw: unit_sw(),
};
let err = refused(parts);
let ExtractError::MisalignedChunkPlacement(m) = err else {
panic!("expected MisalignedChunkPlacement, got {err:?}")
};
assert_eq!((m.chunk(), m.aggregated(), m.reconstructed()), (0, 0, -1));
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 1],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: unit_sw(),
frames_sw: SlidingWindow::new(1.0, 1.0, 1.0),
};
let err = refused(parts);
assert!(
matches!(err, ExtractError::MisalignedChunkPlacement(m)
if (m.chunk(), m.aggregated(), m.reconstructed()) == (0, 0, -1)),
"expected MisalignedChunkPlacement(0, 0, -1), got {err:?}"
);
}
#[test]
fn try_from_parts_rejects_a_frame_step_that_does_not_survive_f32_narrowing() {
let step = 7.0e-46_f64;
assert_eq!(step as f32, 0.0, "the premise: this step narrows to zero");
assert_eq!(1e-300_f64 as f32, 0.0);
for step in [step, 1e-300_f64] {
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 1, 1],
num_chunks: 1,
num_frames_per_chunk: 3,
chunks_sw: SlidingWindow::new(0.0, 2.0 * step, step),
frames_sw: SlidingWindow::new(0.0, 2.0 * step, step),
};
let err = refused(parts);
let ExtractError::FrameStepNotRepresentableInF32(w) = err else {
panic!("expected FrameStepNotRepresentableInF32 for step {step:e}, got {err:?}")
};
assert_eq!(w.part(), ExtractionPart::FramesSw);
assert_eq!(w.window().step(), step);
}
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 1],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: SlidingWindow::new(0.0, 1e300, 1e300),
frames_sw: SlidingWindow::new(0.0, 1e300, 1e300),
};
let err = refused(parts);
assert!(
matches!(err, ExtractError::FrameStepNotRepresentableInF32(w) if w.window().step() == 1e300),
"expected FrameStepNotRepresentableInF32 for a step above f32::MAX, got {err:?}"
);
}
#[test]
fn try_from_parts_rejects_an_output_frame_grid_above_the_allocation_cap() {
let n = MAX_OUTPUT_FRAMES + 1;
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0],
count: vec![0; n],
num_chunks: 1,
num_frames_per_chunk: 1,
chunks_sw: SlidingWindow::new(0.0, (n - 1) as f64, 1.0),
frames_sw: SlidingWindow::new(0.0, 1.0, 1.0),
};
let err = refused(parts);
assert_eq!(err, ExtractError::OutputFrameCountTooLarge(n));
let n = MAX_OUTPUT_FRAMES;
let mut count = vec![0u8; n];
count[0] = 1;
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0],
count,
num_chunks: 1,
num_frames_per_chunk: 1,
chunks_sw: SlidingWindow::new(0.0, (n - 1) as f64, 1.0),
frames_sw: SlidingWindow::new(0.0, 1.0, 1.0),
};
Extraction::try_from_parts(parts).expect("exactly at the cap is accepted");
}
#[test]
fn try_from_parts_refuses_the_soft_active_slot_rounds_1_and_2_argued_over() {
let mut raw = one_usable_slot_row(0);
raw[EMBEDDING_DIM..EMBEDDING_DIM + 64].fill(1.0); let soft = ExtractionParts {
raw_embeddings: raw,
segmentations: vec![1.0, 0.0, 0.0, 0.0, 0.3, 0.0],
count: vec![1, 1, 0, 0],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: crate::audio::speaker::window::chunk_sliding_window(&WindowOptions::new())
.with_duration(3.0 * crate::audio::speaker::window::FRAME_STEP_S),
frames_sw: crate::audio::speaker::window::frame_sliding_window(),
};
for count in [vec![1u8, 1, 0, 0], vec![1, 0, 0, 0]] {
let err = refused(ExtractionParts {
count,
..soft.clone()
});
let ExtractError::NonBinarySegmentation(n) = err else {
panic!("expected NonBinarySegmentation, got {err:?}")
};
assert_eq!(
(n.index(), n.value(), n.slot()),
(SEG_NUM_SLOTS + 1, 0.3, 1)
);
}
let hard = ExtractionParts {
segmentations: vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0],
..soft.clone()
};
let e = Extraction::try_from_parts(hard.clone()).expect("the hard-binary twin is accepted");
assert_eq!(e.num_output_frames(), 4);
let err = refused(ExtractionParts {
count: vec![1, 0, 0, 0],
..hard.clone()
});
assert!(
matches!(err, ExtractError::CountNotSegmentationDerived(c)
if (c.frame(), c.got(), c.expected()) == (1, 0, 1)),
"expected CountNotSegmentationDerived(1, 0, 1), got {err:?}"
);
let derived_at = |onset: f32| {
crate::audio::speaker::window::try_count_from_segmentations(
&hard.segmentations,
hard.num_chunks,
hard.num_frames_per_chunk,
SEG_NUM_SLOTS,
onset,
hard.chunks_sw,
hard.frames_sw,
)
.expect("count")
};
for onset in [f32::MIN_POSITIVE, 0.01, 0.3, 0.5, 0.9, 1.0] {
assert!(
crate::audio::speaker::window::check_onset(onset),
"onset={onset} must be VALID for this to prove the question is moot"
);
assert_eq!(
derived_at(onset),
vec![1u8, 1, 0, 0],
"onset={onset} must derive the same count as `seg > 0.0` on a hard buffer"
);
}
let broken = ExtractionParts {
raw_embeddings: one_usable_slot_row(0), ..hard
};
let err = refused(broken);
assert!(
matches!(err, ExtractError::ActiveSlotWithoutEmbedding(a) if (a.chunk(), a.slot()) == (0, 1)),
"expected ActiveSlotWithoutEmbedding(0, 1), got {err:?}"
);
}
#[test]
fn try_from_parts_requires_the_count_the_segmentations_derive_not_merely_one_they_support() {
const F: usize = 589;
let mut segmentations = vec![0.0f64; F * SEG_NUM_SLOTS];
for f in 0..F {
segmentations[f * SEG_NUM_SLOTS] = 1.0; }
let parts = ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations,
count: vec![0u8; 594],
num_chunks: 1,
num_frames_per_chunk: F,
chunks_sw: crate::audio::speaker::window::chunk_sliding_window(&WindowOptions::new()),
frames_sw: crate::audio::speaker::window::frame_sliding_window(),
};
let err = refused(parts.clone());
let ExtractError::CountNotSegmentationDerived(c) = err else {
panic!("expected CountNotSegmentationDerived for the all-zero count, got {err:?}")
};
assert_eq!((c.frame(), c.got(), c.expected()), (0, 0, 1));
let mut derived = vec![0u8; 594];
derived[..F].fill(1);
let e = Extraction::try_from_parts(ExtractionParts {
count: derived,
..parts
})
.expect("the derived count is the one accepted");
assert_eq!(e.num_output_frames(), 594);
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let offline = e
.diarize_with(
&plda,
ClusterBackend::Offline(crate::audio::speaker::cluster::OfflineOptions::new()),
)
.expect("offline accepts it");
let online = e
.diarize_online(OnlineOptions::new().with_min_speech_duration(0.0))
.expect("online accepts it");
assert!(
!offline.spans_slice().is_empty(),
"offline must now emit the speaker its segmentations declare"
);
assert!(
!online.spans_slice().is_empty(),
"online emits the speaker too"
);
}
#[test]
fn try_from_parts_rejects_misaligned_chunk_placement_even_at_a_zero_origin() {
let d = 0.04218750000000001_f64;
let mut raw = vec![0.0f32; 2 * SEG_NUM_SLOTS * EMBEDDING_DIM];
raw[0..64].fill(1.0); raw[(SEG_NUM_SLOTS * EMBEDDING_DIM)..(SEG_NUM_SLOTS * EMBEDDING_DIM + 64)].fill(1.0); let parts = ExtractionParts {
raw_embeddings: raw,
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 0, 0, 1, 0, 0],
num_chunks: 2,
num_frames_per_chunk: 1,
chunks_sw: SlidingWindow::new(0.0, d, d),
frames_sw: SlidingWindow::new(0.0, 0.0619375, 0.016875),
};
let err = refused(parts);
let ExtractError::MisalignedChunkPlacement(m) = err else {
panic!("expected MisalignedChunkPlacement at a zero origin, got {err:?}")
};
assert_eq!(
(m.chunk(), m.aggregated(), m.reconstructed()),
(1, 3, 2),
"the two mappings disagree about chunk 1"
);
}
#[test]
fn try_from_parts_accepts_equal_non_zero_origins_that_place_every_chunk_identically() {
let sw = SlidingWindow::new(1.0, 1.0, 1.0);
let e = Extraction::try_from_parts(ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 1],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: sw,
frames_sw: sw,
})
.expect("equal, cancelling origins place every chunk identically");
assert_eq!(e.num_output_frames(), 2);
assert_eq!(e.chunks_sw().start(), 1.0);
}
fn misaligned_extract_geometry() -> (usize, SlidingWindow, SlidingWindow) {
let w = WindowOptions::new().with_step_samples(31_995);
let num_chunks = crate::audio::speaker::window::chunk_starts(160_001, &w).len();
(
num_chunks,
crate::audio::speaker::window::chunk_sliding_window(&w),
crate::audio::speaker::window::frame_sliding_window(),
)
}
#[test]
fn extract_derives_a_geometry_whose_two_frame_mappings_disagree() {
let (num_chunks, chunks_sw, frames_sw) = misaligned_extract_geometry();
assert_eq!(num_chunks, 2, "160_001 samples over a 31_995 step");
assert_eq!(chunks_sw.step(), 1.9996875);
assert_eq!(
crate::audio::speaker::window::aggregate_chunk_start_frame(
1,
chunks_sw.step(),
frames_sw.step()
),
118,
"the count aggregation places chunk 1 at frame 118"
);
assert_eq!(
crate::audio::speaker::window::reconstruct_chunk_start_frame(1, chunks_sw, frames_sw),
119,
"diaric's reconstruction places the same chunk at frame 119"
);
let m = crate::audio::speaker::window::first_misaligned_chunk(num_chunks, chunks_sw, frames_sw)
.expect("the shared guard must see the same disagreement");
assert_eq!(
(m.chunk(), m.aggregated(), m.reconstructed()),
(1, 118, 119)
);
}
#[test]
fn a_misaligned_geometry_shifts_the_emitted_span_by_a_whole_frame() {
let (num_chunks, chunks_sw, frames_sw) = misaligned_extract_geometry();
let nf = 589; let mut segmentations = vec![0.0f64; num_chunks * nf * SEG_NUM_SLOTS];
for f in 0..100 {
segmentations[f * SEG_NUM_SLOTS] = 1.0;
}
segmentations[nf * SEG_NUM_SLOTS + 1] = 1.0;
segmentations[nf * SEG_NUM_SLOTS + 2] = 1.0;
let count = crate::audio::speaker::window::try_count_from_segmentations(
&segmentations,
num_chunks,
nf,
SEG_NUM_SLOTS,
WindowOptions::new().onset(),
chunks_sw,
frames_sw,
)
.expect("this geometry's output-frame count fits usize");
assert_eq!(
(count[118], count[119]),
(1, 0),
"the count marks frame 118 and leaves 119 empty"
);
let aligned = {
let mut agg = vec![0.0f64; count.len()];
let mut cov = vec![0.0f64; count.len()];
for c in 0..num_chunks {
let start =
crate::audio::speaker::window::reconstruct_chunk_start_frame(c, chunks_sw, frames_sw);
for f in 0..nf {
let Ok(t) = usize::try_from(start + f as i64) else {
continue;
};
if t >= count.len() {
continue;
}
agg[t] += segmentations[((c * nf + f) * SEG_NUM_SLOTS)..][..SEG_NUM_SLOTS]
.iter()
.filter(|v| **v > 0.0)
.count() as f64;
cov[t] += 1.0;
}
}
(0..count.len())
.map(|t| {
if cov[t] > 0.0 {
(agg[t] / cov[t]).round_ties_even() as u8
} else {
0
}
})
.collect::<Vec<u8>>()
};
assert_eq!(
(aligned[118], aligned[119]),
(0, 1),
"placed the way the activations are, the same frame's count belongs at 119"
);
let mut raw_embeddings = vec![0.0f32; num_chunks * SEG_NUM_SLOTS * EMBEDDING_DIM];
raw_embeddings[..64].fill(1.0);
let c1 = SEG_NUM_SLOTS * EMBEDDING_DIM;
raw_embeddings[c1 + EMBEDDING_DIM..c1 + EMBEDDING_DIM + 64].fill(-1.0);
for k in 0..64 {
raw_embeddings[c1 + 2 * EMBEDDING_DIM + 2 * k] = 1.0;
}
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let span_start = |count: Vec<u8>| -> f64 {
Extraction::from_parts(
raw_embeddings.clone(),
segmentations.clone(),
count,
num_chunks,
nf,
chunks_sw,
frames_sw,
)
.diarize_with(&plda, ClusterBackend::default())
.expect("both counts diarize")
.spans_slice()
.last()
.expect("a trailing span for chunk 1's lone active frame")
.start()
};
let shifted = span_start(count.clone());
let honest = span_start(aligned);
assert!(
(honest - shifted - frames_sw.step()).abs() < 1e-12,
"the count's placement moves the emitted span a whole frame step earlier: \
{shifted} vs {honest}"
);
let err = refused(ExtractionParts {
raw_embeddings,
segmentations,
count,
num_chunks,
num_frames_per_chunk: nf,
chunks_sw,
frames_sw,
});
assert!(
matches!(err, ExtractError::MisalignedChunkPlacement(m)
if (m.chunk(), m.aggregated(), m.reconstructed()) == (1, 118, 119)),
"expected MisalignedChunkPlacement(1, 118, 119), got {err:?}"
);
}
#[test]
fn every_shipping_extract_geometry_places_its_chunks_identically() {
let w = WindowOptions::new();
assert_eq!(w.step_samples() % 2, 0, "the default step is even");
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
assert_eq!(
crate::audio::speaker::window::first_misaligned_chunk(100_000, chunks_sw, frames_sw),
None,
"the default geometry must survive 100 000 chunks (~27.7 h of audio)"
);
for step in (2..=SEG_CHUNK_SAMPLES as u32).step_by(1_998) {
assert_eq!(step % 2, 0, "the sweep must stay on even steps");
let w = WindowOptions::new().with_step_samples(step);
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
assert_eq!(
crate::audio::speaker::window::first_misaligned_chunk(4_096, chunks_sw, frames_sw),
None,
"even step_samples={step} must never tie"
);
}
let odd = crate::audio::speaker::window::chunk_sliding_window(
&WindowOptions::new().with_step_samples(31_995),
);
assert!(
crate::audio::speaker::window::first_misaligned_chunk(2, odd, frames_sw).is_some(),
"step_samples=31_995 must still be caught"
);
}
#[test]
#[ignore = "requires local speakerkit models (SPEAKERKIT_TEST_MODELS)"]
fn extract_refuses_a_geometry_whose_two_frame_mappings_disagree() {
let seg = load_seg_model();
let embed = load_embed_model();
let options = Options::new().with_window(WindowOptions::new().with_step_samples(31_995));
match Extractor::with_options(options).extract(&seg, &embed, &vec![0.0f32; 160_001]) {
Err(ExtractError::MisalignedChunkPlacement(m)) => {
assert_eq!(
(m.chunk(), m.aggregated(), m.reconstructed()),
(1, 118, 119)
);
}
Err(other) => panic!("expected MisalignedChunkPlacement(1, 118, 119), got {other:?}"),
Ok(e) => panic!(
"extract ACCEPTED a misaligned geometry: {} chunks on chunks_sw={:?}, which the count \
aggregation places at frame {} and diaric's reconstruction at frame {}",
e.num_chunks(),
e.chunks_sw(),
crate::audio::speaker::window::aggregate_chunk_start_frame(
1,
e.chunks_sw().step(),
e.frames_sw().step()
),
crate::audio::speaker::window::reconstruct_chunk_start_frame(1, e.chunks_sw(), e.frames_sw()),
),
}
Extractor::new()
.extract(&seg, &embed, &vec![0.0f32; 160_001])
.expect("the default geometry places every chunk identically");
}
#[test]
fn try_from_parts_rejects_an_active_row_below_plda_norm_that_only_online_tolerates() {
let w = WindowOptions::new();
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let nf = 589;
let mut segmentations = vec![0.0f64; nf * SEG_NUM_SLOTS];
for f in 0..nf {
segmentations[f * SEG_NUM_SLOTS] = 1.0;
}
let count = crate::audio::speaker::window::try_count_from_segmentations(
&segmentations,
1,
nf,
SEG_NUM_SLOTS,
w.onset(),
chunks_sw,
frames_sw,
)
.expect("this geometry's output-frame count fits usize");
let mut raw_embeddings = vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM];
raw_embeddings[0] = 0.005;
let mut row = [0.0f32; EMBEDDING_DIM];
row.copy_from_slice(&raw_embeddings[..EMBEDDING_DIM]);
assert!(
diaric::embed::Embedding::normalize_from(row).is_some(),
"the online engine's own test ACCEPTS this row — matching it is the defect"
);
assert!(
matches!(
diaric::plda::RawEmbedding::from_wespeaker(row),
Err(diaric::plda::Error::DegenerateInput)
),
"PLDA's raw boundary refuses it at the 0.01 floor"
);
let parts = ExtractionParts {
raw_embeddings,
segmentations,
count,
num_chunks: 1,
num_frames_per_chunk: nf,
chunks_sw,
frames_sw,
};
let unchecked = Extraction::from_parts(
parts.raw_embeddings.clone(),
parts.segmentations.clone(),
parts.count.clone(),
parts.num_chunks,
parts.num_frames_per_chunk,
parts.chunks_sw,
parts.frames_sw,
);
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
assert!(
matches!(
unchecked.diarize_with(&plda, ClusterBackend::default()),
Err(diaric::offline::Error::Plda(
diaric::plda::Error::DegenerateInput
))
),
"offline must fail on this row"
);
let online = unchecked
.diarize_online(OnlineOptions::new())
.expect("online accepts it");
assert_eq!(
online.spans_slice().len(),
1,
"online manufactures a speaker from the same row"
);
let err = refused(parts);
assert!(
matches!(err, ExtractError::ActiveSlotWithoutEmbedding(a) if (a.chunk(), a.slot()) == (0, 0)),
"expected ActiveSlotWithoutEmbedding(0, 0), got {err:?}"
);
}
#[test]
fn the_active_row_floor_is_the_one_both_in_crate_producers_drop_at() {
assert_eq!(PLDA_MIN_NORM, 0.01);
let plda = shared_plda_transform().expect("diaric's PLDA weights are embedded");
let mut disagreements = 0;
for micro in 9_000..11_000u32 {
let v = f64::from(micro) / 1_000_000.0;
let mut row = [0.0f32; EMBEDDING_DIM];
row[0] = v as f32;
let mine = raw_embedding_reaches_plda(plda, &row);
let admitted = diaric::plda::RawEmbedding::from_wespeaker(row).is_ok();
if mine != admitted {
disagreements += 1;
}
}
assert_eq!(
disagreements, 0,
"the crate's row predicate and PLDA's own boundary must admit the same rows"
);
let mut between = [0.0f32; EMBEDDING_DIM];
between[0] = 0.005;
assert!(!raw_embedding_reaches_plda(plda, &between));
assert!(diaric::embed::Embedding::normalize_from(between).is_some());
let mut above = [0.0f32; EMBEDDING_DIM];
above[0] = 0.02;
assert!(raw_embedding_reaches_plda(plda, &above));
let mut infinite = above;
infinite[1] = f32::INFINITY;
assert!(!raw_embedding_reaches_plda(plda, &infinite));
assert!(diaric::plda::RawEmbedding::from_wespeaker(infinite).is_err());
let mut nan = above;
nan[1] = f32::NAN;
assert!(!raw_embedding_reaches_plda(plda, &nan));
assert!(diaric::plda::RawEmbedding::from_wespeaker(nan).is_err());
}
#[test]
fn try_from_parts_rejects_an_active_row_whose_norm_overflows_f32_for_the_online_engine() {
let w = WindowOptions::new();
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let nf = 60;
let mut segmentations = vec![0.0f64; nf * SEG_NUM_SLOTS];
for f in 0..nf {
segmentations[f * SEG_NUM_SLOTS] = 1.0;
}
let count = crate::audio::speaker::window::try_count_from_segmentations(
&segmentations,
1,
nf,
SEG_NUM_SLOTS,
w.onset(),
chunks_sw,
frames_sw,
)
.expect("this geometry's output-frame count fits usize");
let mut raw_embeddings = vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM];
raw_embeddings[0] = f32::MAX;
raw_embeddings[1] = f32::MAX;
let mut row = [0.0f32; EMBEDDING_DIM];
row.copy_from_slice(&raw_embeddings[..EMBEDDING_DIM]);
let f64_norm: f64 = row
.iter()
.map(|v| f64::from(*v) * f64::from(*v))
.sum::<f64>()
.sqrt();
assert!(
f64_norm > PLDA_MIN_NORM && f64_norm.is_finite(),
"the f64 norm ({f64_norm:e}) is finite and far above the floor — a f64 \
comparison accepts this row"
);
assert!(
!(f64_norm as f32).is_finite(),
"and the SAME norm is +inf once narrowed to f32, which is what \
`normalize_from` compares"
);
assert!(
diaric::plda::RawEmbedding::from_wespeaker(row).is_ok(),
"PLDA's raw boundary ACCEPTS this row — matching it alone is the defect"
);
assert!(
diaric::embed::Embedding::normalize_from(row).is_none(),
"the online engine's own test REFUSES it, and `None` is its dropped-slot \
sentinel"
);
let parts = ExtractionParts {
raw_embeddings,
segmentations,
count,
num_chunks: 1,
num_frames_per_chunk: nf,
chunks_sw,
frames_sw,
};
let unchecked = Extraction::from_parts(
parts.raw_embeddings.clone(),
parts.segmentations.clone(),
parts.count.clone(),
parts.num_chunks,
parts.num_frames_per_chunk,
parts.chunks_sw,
parts.frames_sw,
);
let online = unchecked
.diarize_online(OnlineOptions::new())
.expect("online returns Ok");
assert_eq!(
online.spans_slice().len(),
0,
"online silently drops the active slot — the failure this constructor exists \
to make impossible"
);
let err = refused(parts);
assert!(
matches!(err, ExtractError::ActiveSlotWithoutEmbedding(a) if (a.chunk(), a.slot()) == (0, 0)),
"expected ActiveSlotWithoutEmbedding(0, 0), got {err:?}"
);
}
#[test]
fn the_row_predicate_is_the_two_backend_functions_not_a_description_of_them() {
let mut probes: Vec<[f32; EMBEDDING_DIM]> = Vec::new();
let mut push = |first: f32, second: f32| {
let mut row = [0.0f32; EMBEDDING_DIM];
row[0] = first;
row[1] = second;
probes.push(row);
};
push(0.0, 0.0);
push(f32::from_bits(1), 0.0);
push(f32::MIN_POSITIVE, 0.0);
push(1e-13, 0.0);
push(1e-12, 0.0);
push(1e-11, 0.0);
push(0.005, 0.0);
push(0.009_999, 0.0);
push(0.01, 0.0);
push(0.010_001, 0.0);
push(2.07, 0.0);
push(1e19, 1e19);
push(f32::MAX, 0.0);
push(f32::MAX, f32::MAX);
push(f32::MAX, f32::MAX / 2.0);
push(f32::INFINITY, 0.0);
push(f32::NEG_INFINITY, 0.0);
push(f32::NAN, 0.0);
push(2.07, f32::NAN);
push(-2.07, 0.0);
probes.push(diaric_mean1_as_f32());
let plda = shared_plda_transform().expect("diaric's PLDA weights are embedded");
let projects = |row: &[f32; EMBEDDING_DIM]| {
diaric::plda::RawEmbedding::from_wespeaker(*row).is_ok_and(|raw| plda.project(&raw).is_ok())
};
for row in &probes {
let online = diaric::embed::Embedding::normalize_from(*row).is_some();
let offline = diaric::plda::RawEmbedding::from_wespeaker(*row).is_ok();
let projected = projects(row);
assert_eq!(
raw_embedding_reaches_plda(plda, row),
online && offline && projected,
"row [{}, {}, 0, …] — online accepts {online}, offline accepts {offline}, \
projection accepts {projected}",
row[0],
row[1]
);
}
assert!(
probes.iter().any(|r| {
diaric::embed::Embedding::normalize_from(*r).is_some()
&& diaric::plda::RawEmbedding::from_wespeaker(*r).is_err()
}),
"the probe set must contain a row ONLY the online engine accepts"
);
assert!(
probes.iter().any(|r| {
diaric::embed::Embedding::normalize_from(*r).is_none()
&& diaric::plda::RawEmbedding::from_wespeaker(*r).is_ok()
}),
"the probe set must contain a row ONLY the offline boundary accepts"
);
assert!(
probes.iter().any(|r| {
diaric::embed::Embedding::normalize_from(*r).is_some()
&& diaric::plda::RawEmbedding::from_wespeaker(*r).is_ok()
&& !projects(r)
}),
"the probe set must contain a row BOTH admission functions accept and the \
PROJECTION refuses — otherwise the third clause is vacuous"
);
assert!(!raw_embedding_reaches_plda(
plda,
&[1.0f32; EMBEDDING_DIM - 1]
));
assert!(!raw_embedding_reaches_plda(
plda,
&[1.0f32; EMBEDDING_DIM + 1]
));
}
#[test]
fn plda_min_norm_is_diarics_own_floor_measured_not_copied() {
let admits = |v: f32| {
let mut row = [0.0f32; EMBEDDING_DIM];
row[0] = v;
diaric::plda::RawEmbedding::from_wespeaker(row).is_ok()
};
let (mut lo, mut hi) = (0.0f32, 1.0f32);
assert!(!admits(lo) && admits(hi), "the floor lies inside (0, 1]");
for _ in 0..200 {
let mid = f32::from_bits(lo.to_bits().midpoint(hi.to_bits()));
if mid == lo || mid == hi {
break;
}
if admits(mid) { hi = mid } else { lo = mid }
}
assert!(
f64::from(lo) < PLDA_MIN_NORM && f64::from(hi) >= PLDA_MIN_NORM,
"diaric's measured raw-embedding floor is in ({lo:e}, {hi:e}] but \
PLDA_MIN_NORM says {PLDA_MIN_NORM:e} — the published constant no longer \
names the number diaric enforces"
);
assert_eq!(hi.to_bits() - lo.to_bits(), 1, "search did not converge");
}
#[test]
fn plda_transform_is_available() {
let mine = shared_plda_transform().expect("diaric's PLDA weights are compile-time embedded");
assert!(
std::ptr::eq(
mine,
shared_plda_transform().expect("the cached transform is still there")
),
"shared_plda_transform must hand out one process-wide transform"
);
let theirs = diaric::plda::PldaTransform::new().expect("a caller's own transform");
assert_eq!(
mine.phi(),
theirs.phi(),
"the eigenvalue diagonal must match"
);
let mut row = [0.0f32; EMBEDDING_DIM];
row[0] = 2.07; row[1] = -0.5;
let raw = diaric::plda::RawEmbedding::from_wespeaker(row).expect("an in-distribution row");
assert_eq!(
mine
.project(&raw)
.expect("the cached transform projects it"),
theirs
.project(&raw)
.expect("the caller's transform projects it"),
"the cached transform must project exactly as a caller's own does"
);
}
fn three_chunk_slot0_parts(slot0_rows: [[f32; EMBEDDING_DIM]; 3]) -> ExtractionParts {
let mut raw_embeddings = vec![0.0f32; 3 * SEG_NUM_SLOTS * EMBEDDING_DIM];
for (c, row) in slot0_rows.iter().enumerate() {
raw_embeddings[embedding_range(c, 0)].copy_from_slice(row);
}
ExtractionParts {
raw_embeddings,
segmentations: vec![1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0],
count: vec![1, 0, 1, 0],
num_chunks: 3,
num_frames_per_chunk: 1,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
}
}
fn planar_row(deg: f64) -> [f32; EMBEDDING_DIM] {
let mut row = [0.0f32; EMBEDDING_DIM];
row[0] = deg.to_radians().cos() as f32;
row[1] = deg.to_radians().sin() as f32;
row
}
#[test]
fn an_inactive_slots_row_cannot_change_the_online_result() {
let a = planar_row(0.0);
let b = planar_row(55.0);
let c = planar_row(-50.0);
let zero = [0.0f32; EMBEDDING_DIM];
let with_row = Extraction::try_from_parts(three_chunk_slot0_parts([a, b, c]))
.expect("an inactive slot carrying a usable row is admitted by construction");
let without_row = Extraction::try_from_parts(three_chunk_slot0_parts([a, zero, c]))
.expect("the same parts with the dropped slot's row zeroed");
let clusters = |e: &Extraction| -> Vec<(u64, u64, usize)> {
e.diarize_online(OnlineOptions::new())
.expect("this geometry clusters")
.spans_slice()
.iter()
.map(|s| (s.start().to_bits(), s.end().to_bits(), s.cluster()))
.collect()
};
assert_eq!(
clusters(&with_row),
clusters(&without_row),
"an INACTIVE slot's raw-embedding row changed the online clustering"
);
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let spans = |e: &Extraction| -> Vec<(u64, u64, usize)> {
e.diarize_with(&plda, ClusterBackend::default())
.expect("this geometry clusters offline")
.spans_slice()
.iter()
.map(|s| (s.start().to_bits(), s.end().to_bits(), s.cluster()))
.collect()
};
assert_eq!(
spans(&with_row),
spans(&without_row),
"the OFFLINE route was supposed to be blind to an inactive slot's row"
);
}
fn diaric_mean1_as_f32() -> [f32; EMBEDDING_DIM] {
let cargo = option_env!("CARGO").unwrap_or("cargo");
let out = std::process::Command::new(cargo)
.args([
"metadata",
"--format-version",
"1",
"--all-features",
"--manifest-path",
concat!(env!("CARGO_MANIFEST_DIR"), "/Cargo.toml"),
])
.output()
.expect("`cargo metadata` must run: this test reads diaric's shipped PLDA weights");
assert!(
out.status.success(),
"cargo metadata failed: {}",
String::from_utf8_lossy(&out.stderr)
);
let meta: serde_json::Value =
serde_json::from_slice(&out.stdout).expect("cargo metadata emits JSON");
let manifest = meta["packages"]
.as_array()
.expect("metadata.packages is an array")
.iter()
.find(|p| p["name"] == "diaric")
.and_then(|p| p["manifest_path"].as_str())
.expect("diaric is an --all-features dependency of this crate");
let blob = std::path::Path::new(manifest)
.parent()
.expect("a manifest path has a parent")
.join("models/plda/mean1.bin");
let bytes =
std::fs::read(&blob).unwrap_or_else(|e| panic!("diaric ships {}: {e}", blob.display()));
assert_eq!(
bytes.len(),
EMBEDDING_DIM * 8,
"mean1.bin is {EMBEDDING_DIM} little-endian f64 values"
);
let mut row = [0.0f32; EMBEDDING_DIM];
for (i, slot) in row.iter_mut().enumerate() {
let mut le = [0u8; 8];
le.copy_from_slice(&bytes[i * 8..i * 8 + 8]);
*slot = f64::from_le_bytes(le) as f32;
}
row
}
#[test]
fn try_from_parts_rejects_an_active_row_plda_projection_refuses() {
let row = diaric_mean1_as_f32();
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
assert!(
diaric::embed::Embedding::normalize_from(row).is_some(),
"the ONLINE engine accepts this row"
);
let raw = diaric::plda::RawEmbedding::from_wespeaker(row)
.expect("PLDA's RAW boundary accepts this row (norm ~1.42)");
assert!(
matches!(
plda.project(&raw),
Err(diaric::plda::Error::DegenerateInput)
),
"the PROJECTION that follows must refuse it — otherwise this is not a witness"
);
let mut raw_embeddings = vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM];
raw_embeddings[..EMBEDDING_DIM].copy_from_slice(&row);
let parts = ExtractionParts {
raw_embeddings,
segmentations: vec![1.0, 0.0, 0.0],
count: vec![1, 0],
num_chunks: 1,
num_frames_per_chunk: 1,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
};
let unchecked = Extraction::from_parts(
parts.raw_embeddings.clone(),
parts.segmentations.clone(),
parts.count.clone(),
parts.num_chunks,
parts.num_frames_per_chunk,
parts.chunks_sw,
parts.frames_sw,
);
assert!(
matches!(
unchecked.diarize_with(&plda, ClusterBackend::default()),
Err(diaric::offline::Error::Plda(
diaric::plda::Error::DegenerateInput
))
),
"offline must fail on this row"
);
assert_eq!(
unchecked
.diarize_online(OnlineOptions::new())
.expect("online accepts it")
.spans_slice()
.len(),
1,
"online manufactures a speaker from the same row"
);
let err = refused(parts);
assert!(
matches!(err, ExtractError::ActiveSlotWithoutEmbedding(a) if (a.chunk(), a.slot()) == (0, 0)),
"expected ActiveSlotWithoutEmbedding(0, 0), got {err:?}"
);
}
fn one_chunk_with_poisoned_slot(slot: usize, dimension: usize, value: f32) -> ExtractionParts {
let mut raw_embeddings = one_usable_slot_row(0);
raw_embeddings[embedding_range(0, slot)][dimension] = value;
ExtractionParts {
raw_embeddings,
segmentations: vec![1.0, 0.0, 0.0],
count: vec![1, 0],
num_chunks: 1,
num_frames_per_chunk: 1,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
}
}
#[test]
fn try_from_parts_rejects_a_non_finite_row_under_an_inactive_column() {
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
for bad in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
let parts = one_chunk_with_poisoned_slot(2, 0, bad);
let unchecked = Extraction::from_parts(
parts.raw_embeddings.clone(),
parts.segmentations.clone(),
parts.count.clone(),
parts.num_chunks,
parts.num_frames_per_chunk,
parts.chunks_sw,
parts.frames_sw,
);
assert!(
matches!(
unchecked.diarize_with(&plda, ClusterBackend::default()),
Err(diaric::offline::Error::Pipeline(
diaric::pipeline::Error::NonFinite(diaric::pipeline::error::NonFiniteField::Embeddings)
))
),
"offline must fail the whole extraction on {bad:?} in an inactive slot's row, got {:?}",
unchecked.diarize_with(&plda, ClusterBackend::default())
);
assert_eq!(
unchecked
.diarize_online(OnlineOptions::new())
.expect("online never reads an inactive slot's row, so it returns Ok")
.spans_slice()
.len(),
1,
"online must be blind to the same value the offline route dies on"
);
let err = refused(parts);
let expected = embedding_range(0, 2).start;
assert!(
matches!(err, ExtractError::NonFiniteRawEmbedding(i) if i == expected),
"expected NonFiniteRawEmbedding({expected}) for {bad:?}, got {err:?}"
);
}
for ok in [0.0f32, -0.0, 7.5, f32::MAX, f32::MIN_POSITIVE] {
Extraction::try_from_parts(one_chunk_with_poisoned_slot(2, 0, ok)).unwrap_or_else(|e| {
panic!("a finite {ok:?} in an inactive slot's row must be accepted: {e:?}")
});
}
}
#[test]
fn the_finiteness_check_covers_the_whole_buffer_not_only_active_rows() {
for slot in 0..SEG_NUM_SLOTS {
for dimension in [0, 1, EMBEDDING_DIM - 1] {
for bad in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
let err = refused(one_chunk_with_poisoned_slot(slot, dimension, bad));
let flat = embedding_range(0, slot).start + dimension;
if slot == 0 {
assert!(
matches!(err, ExtractError::ActiveSlotWithoutEmbedding(a)
if (a.chunk(), a.slot()) == (0, 0)),
"an ACTIVE slot's non-finite row must keep check 9's diagnosis \
(slot {slot}, dimension {dimension}, {bad:?}), got {err:?}"
);
} else {
assert!(
matches!(err, ExtractError::NonFiniteRawEmbedding(i) if i == flat),
"expected NonFiniteRawEmbedding({flat}) for slot {slot}, dimension \
{dimension}, {bad:?}, got {err:?}"
);
}
}
}
}
}
#[test]
fn no_producer_can_emit_a_buffer_the_finiteness_check_refuses() {
assert!(0.0f32.is_finite(), "an unwritten row is all-zero");
let plda = shared_plda_transform().expect("hermetic PLDA weights load");
let usable = planar_row(30.0);
assert!(
raw_embedding_reaches_plda(plda, &usable),
"non-vacuity: the unpoisoned row must be one a producer WOULD write"
);
for dimension in [0, 1, EMBEDDING_DIM - 1] {
for bad in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
let mut row = usable;
row[dimension] = bad;
assert!(
!raw_embedding_reaches_plda(plda, &row),
"a producer must never write a row holding {bad:?} at dimension {dimension}"
);
}
}
Extraction::try_from_parts(valid_parts()).expect("the reference parts stay accepted");
Extraction::try_from_parts(ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0],
count: vec![1, 0],
num_chunks: 1,
num_frames_per_chunk: 1,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
})
.expect("two all-zero rows under two all-zero columns is the producers' own output shape");
}
fn axis_row(i: usize, v: f32) -> [f32; EMBEDDING_DIM] {
let mut r = [0.0f32; EMBEDDING_DIM];
r[i] = v;
r
}
fn six_chunk_parts(probe: [f32; EMBEDDING_DIM], chunk: usize, slot: usize) -> ExtractionParts {
let slot0: [[f32; EMBEDDING_DIM]; 6] = [
planar_row(0.0),
planar_row(4.0),
planar_row(88.0),
planar_row(92.0),
axis_row(2, 1.0),
{
let mut r = axis_row(2, 1.0);
r[3] = 0.07;
r
},
];
let n = slot0.len();
let mut raw_embeddings = vec![0.0f32; n * SEG_NUM_SLOTS * EMBEDDING_DIM];
for (c, row) in slot0.iter().enumerate() {
raw_embeddings[embedding_range(c, 0)].copy_from_slice(row);
}
raw_embeddings[embedding_range(chunk, slot)].copy_from_slice(&probe);
let other = (chunk + 2) % n;
raw_embeddings[embedding_range(other, 1)].copy_from_slice(&probe);
let mut anti = probe;
for v in anti.iter_mut() {
*v = -*v;
}
raw_embeddings[embedding_range(other, 2)].copy_from_slice(&anti);
let mut segmentations = vec![0.0f64; n * SEG_NUM_SLOTS];
for c in 0..n {
segmentations[c * SEG_NUM_SLOTS] = 1.0;
}
let mut count = vec![1u8; n];
count.push(0);
ExtractionParts {
raw_embeddings,
segmentations,
count,
num_chunks: n,
num_frames_per_chunk: 1,
chunks_sw: unit_sw(),
frames_sw: unit_sw(),
}
}
#[test]
fn an_inactive_slots_row_cannot_change_the_offline_result() {
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let fingerprint = |parts: ExtractionParts| -> OutputFingerprint {
output_fingerprint(
&Extraction::try_from_parts(parts)
.expect("the probe geometry is self-consistent")
.diarize_with(&plda, ClusterBackend::default())
.expect("this geometry clusters offline"),
)
};
let zero = [0.0f32; EMBEDDING_DIM];
let base = fingerprint(six_chunk_parts(zero, 0, 1));
assert_eq!(base.2, 3, "the probe geometry must reach three clusters");
assert!(
base.1.iter().all(|row| row[1] >= 0 && row[2] >= 0),
"the inactive slots must draw real cluster labels, got {:?}",
base.1
);
let probes = [
planar_row(180.0),
planar_row(270.0),
planar_row(45.0),
axis_row(2, -1.0),
axis_row(7, 1.0),
axis_row(7, -3.5),
];
for slot in [1usize, 2] {
for chunk in [0usize, 3, 5] {
let at = fingerprint(six_chunk_parts(zero, chunk, slot));
assert_eq!(
at, base,
"moving the all-zero probe to (chunk {chunk}, slot {slot}) changed the offline output"
);
for (i, probe) in probes.iter().enumerate() {
assert_eq!(
fingerprint(six_chunk_parts(*probe, chunk, slot)),
base,
"probe {i} in INACTIVE (chunk {chunk}, slot {slot}) changed the offline output"
);
}
}
}
}
fn split_by_magnitude_parts(v: f64) -> ExtractionParts {
let mut raw_embeddings = one_usable_slot_row(0);
raw_embeddings[embedding_range(0, 1)][64..128].fill(1.0);
let mut segmentations = vec![0.0f64; 4 * SEG_NUM_SLOTS];
for f in 0..2 {
segmentations[f * SEG_NUM_SLOTS] = v; }
for f in 2..4 {
segmentations[f * SEG_NUM_SLOTS + 1] = v; }
ExtractionParts {
raw_embeddings,
segmentations,
count: vec![1, 1, 1, 1],
num_chunks: 1,
num_frames_per_chunk: 4,
chunks_sw: SlidingWindow::new(0.0, 3.0, 1.0),
frames_sw: SlidingWindow::new(0.0, 1.0, 1.0),
}
}
type BackendAnswer = (usize, Vec<(f64, f64, usize)>);
fn both_backends(e: &Extraction) -> (BackendAnswer, BackendAnswer) {
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let read = |o: &diaric::offline::OfflineOutput| {
(
o.num_clusters(),
o.spans_slice()
.iter()
.map(|s| (s.start(), s.end(), s.cluster()))
.collect::<Vec<_>>(),
)
};
(
read(
&e.diarize_with(&plda, ClusterBackend::default())
.expect("this geometry clusters offline"),
),
read(
&e.diarize_online(OnlineOptions::new())
.expect("this geometry clusters online"),
),
)
}
#[test]
fn a_fractional_segmentation_splits_the_two_backends() {
let parts = split_by_magnitude_parts(0.1);
let unchecked = Extraction::from_parts(
parts.raw_embeddings.clone(),
parts.segmentations.clone(),
parts.count.clone(),
parts.num_chunks,
parts.num_frames_per_chunk,
parts.chunks_sw,
parts.frames_sw,
);
let (offline, online) = both_backends(&unchecked);
assert_eq!(
offline,
(1, vec![(0.5, 3.5, 0)]),
"offline must read the 0.1 cells as magnitudes and merge the two slots"
);
assert_eq!(
online,
(2, vec![(0.5, 2.5, 0), (2.5, 3.5, 1)]),
"online must read the same cells as booleans and split the two slots"
);
assert_ne!(offline, online, "the two backends must actually disagree");
let hard = Extraction::try_from_parts(split_by_magnitude_parts(1.0))
.expect("the hard-binary twin of the same geometry is accepted");
let (offline_hard, online_hard) = both_backends(&hard);
assert_eq!(
offline_hard, online_hard,
"on the hard-binary domain the two backends must agree"
);
assert_eq!(offline_hard.0, 2, "and both must find the two speakers");
let err = refused(parts);
let ExtractError::NonBinarySegmentation(n) = err else {
panic!("expected NonBinarySegmentation, got {err:?}")
};
assert_eq!((n.index(), n.value(), n.slot()), (0, 0.1, 0));
}
#[test]
fn the_segmentation_domain_check_is_an_equality_over_the_whole_buffer() {
let base = split_by_magnitude_parts(1.0);
Extraction::try_from_parts(base.clone()).expect("the all-hard buffer is accepted");
for (cell, bad) in [
(0usize, 0.5f64),
(1, 0.5), (2, -1.0), (5, 1.5), (7, f64::NAN),
(10, f64::INFINITY),
(11, f64::NEG_INFINITY),
(11, f64::MIN_POSITIVE),
] {
let mut parts = base.clone();
parts.segmentations[cell] = bad;
let err = refused(parts);
let ExtractError::NonBinarySegmentation(n) = err else {
panic!("expected NonBinarySegmentation for {bad} at cell {cell}, got {err:?}")
};
assert_eq!(n.index(), cell, "the FIRST offending cell, for {bad}");
assert_eq!(n.slot(), cell % SEG_NUM_SLOTS);
assert_eq!(n.value().to_bits(), bad.to_bits(), "the value, verbatim");
}
let mut negative_zero = base;
negative_zero.segmentations[1] = -0.0;
Extraction::try_from_parts(negative_zero).expect("-0.0 == 0.0 and both backends read it so");
}
#[test]
fn no_producer_can_emit_a_segmentation_cell_the_domain_check_refuses() {
for class in 0..crate::audio::speaker::segment::POWERSET_CLASSES {
let mut logits = vec![0.0f32; crate::audio::speaker::segment::POWERSET_CLASSES];
logits[class] = 1.0;
for v in crate::audio::speaker::segment::multilabel(&logits, 1) {
assert!(
v == 0.0 || v == 1.0,
"multilabel class {class} emitted {v}, which check 9 would refuse"
);
}
}
for v in [0.0f32, 1.0] {
let widened = f64::from(crate::f16::from_f32(v).to_f32());
assert!(
widened == 0.0 || widened == 1.0,
"the f16 speaker_ids value {v} must widen to exactly {v}"
);
}
}
fn collapsing_frame_grid_parts() -> ExtractionParts {
ExtractionParts {
raw_embeddings: one_usable_slot_row(0),
segmentations: vec![1.0, 0.0, 0.0, 0.0, 0.0, 0.0],
count: vec![1, 0, 0, 0],
num_chunks: 1,
num_frames_per_chunk: 2,
chunks_sw: SlidingWindow::new(1e9, 3e-8, 1e-8),
frames_sw: SlidingWindow::new(1e9, 1e-8, 1e-8),
}
}
#[test]
fn try_from_parts_rejects_a_frames_sw_that_collapses_adjacent_frame_centers() {
assert_eq!(
f64::from_bits(1e9f64.to_bits() + 1) - 1e9,
1.1920928955078125e-7
);
assert_eq!(1e9f64 + 1e-8, 1e9);
let parts = collapsing_frame_grid_parts();
let separated = ExtractionParts {
chunks_sw: SlidingWindow::new(0.0, 3e-8, 1e-8),
frames_sw: SlidingWindow::new(0.0, 1e-8, 1e-8),
..parts.clone()
};
let ok = Extraction::try_from_parts(separated).expect("the same grid at origin 0.0 is accepted");
assert_eq!(ok.num_output_frames(), 4);
let unchecked = Extraction::from_parts(
parts.raw_embeddings.clone(),
parts.segmentations.clone(),
parts.count.clone(),
parts.num_chunks,
parts.num_frames_per_chunk,
parts.chunks_sw,
parts.frames_sw,
);
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let offline = unchecked
.diarize_with(&plda, ClusterBackend::default())
.expect("offline returns Ok");
assert_eq!(
offline
.spans_slice()
.iter()
.map(|s| (s.start(), s.duration()))
.collect::<Vec<_>>(),
vec![(1e9, 0.0)],
"the active run closes at identical endpoints: a zero-duration span"
);
assert_eq!(
unchecked
.diarize_online(OnlineOptions::new())
.expect("online returns Ok")
.spans_slice()
.len(),
0
);
assert_eq!(
ok.diarize_with(&plda, ClusterBackend::default())
.expect("offline returns Ok")
.spans_slice()
.iter()
.map(|s| (s.start(), s.duration()))
.collect::<Vec<_>>(),
vec![(5e-9, 1.0000000000000002e-8)],
"a real duration — the 2-ULP tail is the same rounding, now visible instead \
of annihilating"
);
let err = refused(parts);
let ExtractError::CollapsedFrameCenter(c) = err else {
panic!("expected CollapsedFrameCenter, got {err:?}")
};
assert_eq!(
(c.frame(), c.center(), c.previous()),
(1, 1e9, 1e9),
"frame 1 lands on frame 0's center"
);
}
#[test]
fn no_producer_can_emit_a_frame_grid_the_center_check_refuses() {
assert_eq!(
crate::audio::speaker::window::first_collapsed_frame_center(
MAX_OUTPUT_FRAMES,
crate::audio::speaker::window::frame_sliding_window()
),
None
);
let last = crate::audio::speaker::window::frame_center(
MAX_OUTPUT_FRAMES - 1,
crate::audio::speaker::window::frame_sliding_window(),
);
let ulp = f64::from_bits(last.to_bits() + 1) - last;
assert!(
crate::audio::speaker::window::FRAME_STEP_S / ulp > 1e9,
"step {} against ULP {ulp:e} at the last center {last}",
crate::audio::speaker::window::FRAME_STEP_S
);
}
#[test]
fn assemble_checked_reaches_the_same_verdict_as_try_from_parts() {
let onset = WindowOptions::new().onset();
let good_row = one_usable_slot_row(0);
let unit = unit_sw();
type DoorCase = (
&'static str,
Vec<f32>,
Vec<f64>,
usize,
usize,
SlidingWindow,
SlidingWindow,
);
let cases: [DoorCase; 8] = [
(
"accepted",
good_row.clone(),
vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
1,
2,
unit,
unit,
),
(
"check 9: a fractional cell",
good_row.clone(),
vec![0.5, 0.0, 0.0, 1.0, 0.0, 0.0],
1,
2,
unit,
unit,
),
(
"check 10: an active slot with an all-zero row",
vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM],
vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
1,
2,
unit,
unit,
),
(
"check 12: a non-finite row under an INACTIVE column",
{
let mut r = good_row.clone();
r[2 * EMBEDDING_DIM] = f32::NAN;
r
},
vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
1,
2,
unit,
unit,
),
(
"check 13: a frames_sw that collapses adjacent centers",
good_row.clone(),
vec![1.0, 0.0, 0.0, 0.0, 0.0, 0.0],
1,
2,
SlidingWindow::new(1e9, 3e-8, 1e-8),
SlidingWindow::new(1e9, 1e-8, 1e-8),
),
(
"check 8: the two grids place chunk 1 differently",
{
let mut r = vec![0.0f32; 2 * SEG_NUM_SLOTS * EMBEDDING_DIM];
r[..64].fill(1.0);
r
},
vec![1.0, 0.0, 0.0, 0.0, 0.0, 0.0],
2,
1,
SlidingWindow::new(0.0, 0.04218750000000001, 0.04218750000000001),
crate::audio::speaker::window::frame_sliding_window(),
),
(
"check 7: a frames_sw step that vanishes in f32",
good_row.clone(),
vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
1,
2,
SlidingWindow::new(0.0, 1e-300, 1e-300),
SlidingWindow::new(0.0, 1.0, 1e-300),
),
(
"check 14: a chunk longer than the grid it derives",
good_row,
vec![1.0, 0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 0.0],
1,
3,
unit,
unit,
),
];
let expected_variant = |name: &str| -> Option<&'static str> {
match name
.split(':')
.next()
.expect("split yields at least one part")
{
"accepted" => None,
"check 7" => Some("FrameStepNotRepresentableInF32"),
"check 8" => Some("MisalignedChunkPlacement"),
"check 9" => Some("NonBinarySegmentation"),
"check 10" => Some("ActiveSlotWithoutEmbedding"),
"check 12" => Some("NonFiniteRawEmbedding"),
"check 13" => Some("CollapsedFrameCenter"),
"check 14" => Some("UncoveredLastChunk"),
other => panic!("unmapped case label `{other}`"),
}
};
for (
name,
raw_embeddings,
segmentations,
num_chunks,
num_frames_per_chunk,
chunks_sw,
frames_sw,
) in cases
{
let count = crate::audio::speaker::window::try_count_from_segmentations(
&segmentations,
num_chunks,
num_frames_per_chunk,
SEG_NUM_SLOTS,
onset,
chunks_sw,
frames_sw,
)
.expect("every case's geometry derives a count that fits usize");
let public = Extraction::try_from_parts(ExtractionParts {
raw_embeddings: raw_embeddings.clone(),
segmentations: segmentations.clone(),
count,
num_chunks,
num_frames_per_chunk,
chunks_sw,
frames_sw,
});
let producer = Extraction::assemble_checked(
raw_embeddings,
segmentations,
num_chunks,
num_frames_per_chunk,
onset,
chunks_sw,
frames_sw,
);
assert_eq!(
public.as_ref().err(),
producer.as_ref().err(),
"the two doors disagreed about `{name}`"
);
assert_eq!(
public.ok(),
producer.clone().ok(),
"the two doors assembled different Extractions for `{name}`"
);
match (expected_variant(name), producer) {
(None, Ok(_)) => {}
(Some(variant), Err(e)) => assert!(
format!("{e:?}").starts_with(variant),
"`{name}` must be refused by {variant}, got {e:?}"
),
(expected, got) => panic!("`{name}`: expected {expected:?}, got {got:?}"),
}
}
}
#[test]
fn the_frame_cap_refuses_before_the_allocation_it_bounds() {
use crate::{audio::speaker::window, tests::alloc_probe};
const SMALLEST_OVER_CAP_SAMPLES: usize = 1_132_448_001;
const OVER_CAP_FRAMES: usize = 4_194_312;
const FRAMES_PER_CHUNK: usize = 589;
assert_eq!(
FRAMES_PER_CHUNK,
crate::audio::speaker::source::argmax::ARGMAX_FRAMES_PER_WINDOW
);
const { assert!(OVER_CAP_FRAMES > MAX_OUTPUT_FRAMES) };
let extractor = Extractor::new();
assert_eq!(
extractor.checked_geometry(SMALLEST_OVER_CAP_SAMPLES, FRAMES_PER_CHUNK),
Err(ExtractError::OutputFrameCountTooLarge(OVER_CAP_FRAMES)),
"the cap is reachable from `samples.len()` and the options alone — no \
tensor, no model"
);
let w = WindowOptions::new();
let (num_chunks, chunks_sw, frames_sw) = extractor
.checked_geometry(SMALLEST_OVER_CAP_SAMPLES - 1, FRAMES_PER_CHUNK)
.expect("one sample fewer derives a grid inside the cap");
assert_eq!(
(
num_chunks,
checked_output_frame_count(num_chunks, chunks_sw, frames_sw)
),
(70_769, Ok(4_194_253))
);
let num_chunks = window::num_chunks(SMALLEST_OVER_CAP_SAMPLES, &w);
assert_eq!(num_chunks, 70_770);
let chunks_sw = window::chunk_sliding_window(&w);
let frames_sw = window::frame_sliding_window();
let (tensors, built) = alloc_probe::measure(|| {
(
vec![0.0f32; num_chunks * SEG_NUM_SLOTS * EMBEDDING_DIM],
vec![0.0f64; num_chunks * FRAMES_PER_CHUNK * SEG_NUM_SLOTS],
)
});
assert_eq!(
(built.total, built.peak),
(1_217_810_160, 1_217_810_160),
"the two extraction tensors a producer builds before assembly"
);
let (err, scratch) = alloc_probe::measure(|| {
Extraction::assemble_checked(
tensors.0,
tensors.1,
num_chunks,
FRAMES_PER_CHUNK,
w.onset(),
chunks_sw,
frames_sw,
)
.expect_err("a 19.66 h clip derives a grid past MAX_OUTPUT_FRAMES")
});
assert_eq!(err, ExtractError::OutputFrameCountTooLarge(OVER_CAP_FRAMES));
assert_eq!(
(scratch.total, scratch.peak),
(0, 0),
"the door allocated before refusing a geometry it could refuse from \
`num_chunks` and the two windows alone (404 771 544 bytes before round 9)"
);
}
#[test]
fn the_geometry_preflight_and_the_assembled_check_never_disagree() {
let onset = WindowOptions::new().onset();
let unit = unit_sw();
let second = SlidingWindow::new(0.0, 1.0, 1.0);
let cases: [(usize, SlidingWindow, SlidingWindow); 6] = [
(1, unit, unit),
(2, unit, unit),
(
1,
SlidingWindow::new(0.0, (MAX_OUTPUT_FRAMES - 2) as f64, 1.0),
second,
),
(
1,
SlidingWindow::new(0.0, (MAX_OUTPUT_FRAMES - 1) as f64, 1.0),
second,
),
(
1,
SlidingWindow::new(0.0, MAX_OUTPUT_FRAMES as f64, 1.0),
second,
),
(
1,
SlidingWindow::new(0.0, 1e300, 1e300),
SlidingWindow::new(0.0, 1e-300, 1e-300),
),
];
for (num_chunks, chunks_sw, frames_sw) in cases {
let preflight = checked_output_frame_count(num_chunks, chunks_sw, frames_sw);
let assembled = Extraction::assemble_checked(
vec![0.0f32; num_chunks * SEG_NUM_SLOTS * EMBEDDING_DIM],
vec![0.0f64; num_chunks * SEG_NUM_SLOTS],
num_chunks,
1,
onset,
chunks_sw,
frames_sw,
);
match (preflight, assembled) {
(Err(pre), Err(post)) => assert_eq!(
pre, post,
"num_chunks={num_chunks} chunks_sw={chunks_sw:?}: the preflight and the \
sequence named different errors"
),
(Ok(n), Ok(e)) => assert_eq!(
n,
e.num_output_frames(),
"num_chunks={num_chunks} chunks_sw={chunks_sw:?}: the preflight derived \
a different grid than the one assembled"
),
(pre, post) => panic!(
"num_chunks={num_chunks} chunks_sw={chunks_sw:?}: preflight {pre:?} \
disagrees with the sequence {post:?}"
),
}
}
}
#[test]
fn the_chunk_axis_cap_refuses_before_the_allocation_it_bounds() {
use crate::{audio::speaker::window, tests::alloc_probe};
const STEP: u32 = 2;
const SAMPLES: usize = 9_600_000;
const FRAMES_PER_CHUNK: usize = 589;
const NUM_CHUNKS: usize = 4_720_001;
const TENSOR_BYTES: usize = 81_221_777_208;
let w = WindowOptions::new().with_step_samples(STEP);
let extractor = Extractor::with_options(Options::new().with_window(w));
let chunks_sw = window::chunk_sliding_window(&w);
let frames_sw = window::frame_sliding_window();
let num_chunks = window::num_chunks(SAMPLES, &w);
assert_eq!(num_chunks, NUM_CHUNKS);
assert_eq!(
checked_output_frame_count(num_chunks, chunks_sw, frames_sw),
Ok(35_557),
"the output-frame cap sees a ten-minute clip and passes it"
);
const { assert!(35_557 * 100 / MAX_OUTPUT_FRAMES == 0) };
assert_eq!(
window::first_misaligned_chunk(num_chunks, chunks_sw, frames_sw),
None,
"and the even stride clears the placement scan"
);
assert_eq!(
extractor.checked_geometry(SAMPLES, FRAMES_PER_CHUNK),
Err(ExtractError::ExtractionChunkCountTooLarge(NUM_CHUNKS)),
"the chunk-axis bounds are reachable from `samples.len()`, `step_samples` \
and the segmenter's frame count alone — no tensor, no model"
);
assert_eq!(
checked_extraction_tensor_bytes(num_chunks, FRAMES_PER_CHUNK),
Err(ExtractError::ExtractionTensorBytesTooLarge(TENSOR_BYTES)),
"and the memory bound refuses it independently"
);
let (tensors, built) = alloc_probe::measure(|| {
(
vec![0.0f32; num_chunks * SEG_NUM_SLOTS * EMBEDDING_DIM],
vec![0.0f64; num_chunks * FRAMES_PER_CHUNK * SEG_NUM_SLOTS],
)
});
assert_eq!(
(built.total, built.peak),
(TENSOR_BYTES, TENSOR_BYTES),
"the two extraction tensors this geometry used to reach the chunk loop with"
);
assert_eq!(
(
tensors.0.len() * size_of::<f32>(),
tensors.1.len() * size_of::<f64>()
),
(14_499_843_072, 66_721_934_136),
"raw_embeddings and segmentations, the two terms the bound adds"
);
drop(tensors);
let (err, spent) = alloc_probe::measure(|| {
extractor
.checked_geometry(SAMPLES, FRAMES_PER_CHUNK)
.expect_err("a 4 720 001-chunk grid is past both chunk-axis bounds")
});
assert_eq!(err, ExtractError::ExtractionChunkCountTooLarge(NUM_CHUNKS));
assert_eq!(
(spent.total, spent.peak),
(0, 0),
"the preflight reaches the verdict from three `usize`s"
);
}
#[test]
fn the_chunk_axis_cap_is_a_boundary_and_precedes_the_placement_scan() {
use crate::audio::speaker::window;
const FRAMES_PER_CHUNK: usize = 589;
const SMALLEST_REFUSED_SAMPLES: usize = 230_770;
let w = WindowOptions::new().with_step_samples(1);
let extractor = Extractor::with_options(Options::new().with_window(w));
let chunks_sw = window::chunk_sliding_window(&w);
let frames_sw = window::frame_sliding_window();
let num_chunks = window::num_chunks(SMALLEST_REFUSED_SAMPLES, &w);
assert_eq!(num_chunks, 70_771);
assert!(
checked_output_frame_count(num_chunks, chunks_sw, frames_sw).is_ok(),
"a 14.42 s clip is nowhere near MAX_OUTPUT_FRAMES"
);
assert_eq!(
window::first_misaligned_chunk(num_chunks, chunks_sw, frames_sw)
.map(|m| m.chunk())
.expect("an odd stride ties, so the placement scan has something to say"),
135
);
assert_eq!(
extractor.checked_geometry(SMALLEST_REFUSED_SAMPLES, FRAMES_PER_CHUNK),
Err(ExtractError::ExtractionChunkCountTooLarge(70_771)),
"the chunk bound runs BEFORE the placement scan, which is O(num_chunks) \
over the very axis it limits"
);
assert_eq!(window::num_chunks(SMALLEST_REFUSED_SAMPLES - 1, &w), 70_770);
assert_eq!(
checked_extraction_chunk_count(70_770),
Ok(MAX_EXTRACTION_CHUNKS),
"the largest grid admitted sits exactly ON the ceiling"
);
assert!(
matches!(
extractor.checked_geometry(SMALLEST_REFUSED_SAMPLES - 1, FRAMES_PER_CHUNK),
Err(ExtractError::MisalignedChunkPlacement(m)) if m.chunk() == 135
),
"one sample fewer passes both chunk-axis bounds and reaches round 3's guard"
);
assert_eq!(
derived_extraction_tensor_bytes(70_770, FRAMES_PER_CHUNK),
Ok(1_217_810_160),
"round 10's ceiling was exactly this geometry's footprint at 589 frames"
);
for frames in [1usize, 588, 589, 590, 594] {
assert!(
matches!(
extractor.checked_geometry(SMALLEST_REFUSED_SAMPLES - 1, frames),
Err(ExtractError::MisalignedChunkPlacement(m)) if m.chunk() == 135
),
"frames_per_chunk={frames}: 70 770 chunks must reach round 3's guard"
);
assert_eq!(
extractor.checked_geometry(SMALLEST_REFUSED_SAMPLES, frames),
Err(ExtractError::ExtractionChunkCountTooLarge(70_771)),
"frames_per_chunk={frames}: 70 771 must not"
);
}
}
#[test]
fn max_extraction_chunks_is_the_frame_caps_own_allowance_derived_not_copied() {
use crate::audio::speaker::window;
let w = WindowOptions::new();
assert_eq!(w.step_samples(), window::DEFAULT_STEP_SAMPLES);
let chunks_sw = window::chunk_sliding_window(&w);
let frames_sw = window::frame_sliding_window();
let (mut lo, mut hi) = (1usize, 1usize << 24);
assert!(
checked_output_frame_count(hi, chunks_sw, frames_sw).is_err(),
"the search needs a refused upper bound"
);
while lo < hi {
let mid = lo + (hi - lo) / 2;
if checked_output_frame_count(mid, chunks_sw, frames_sw).is_err() {
hi = mid;
} else {
lo = mid + 1;
}
}
let first_refused = lo;
assert_eq!(first_refused, 70_770);
assert_eq!(
MAX_EXTRACTION_CHUNKS, first_refused,
"the chunk ceiling IS the allowance MAX_OUTPUT_FRAMES already gives at the \
shipped stride, admitted inclusively"
);
assert_eq!(
checked_extraction_chunk_count(first_refused - 1),
Ok(70_769)
);
assert_eq!(
checked_geometry_first_refusal_at_default_stride(),
ExtractError::OutputFrameCountTooLarge(4_194_312),
"at the shipped stride the FRAME cap is still the one that speaks"
);
assert_eq!(2 * MAX_EXTRACTION_CHUNKS, 141_540, "Extractor::extract");
assert_eq!(
3 * MAX_EXTRACTION_CHUNKS
.div_ceil(crate::audio::speaker::source::argmax::ARGMAX_WINDOWS_PER_CHUNK),
10_110,
"ArgmaxSource::extract"
);
let fine =
Extractor::with_options(Options::new().with_window(WindowOptions::new().with_step_samples(1)));
for frames in [1usize, 2, 589, 594, 51_095_812] {
assert_eq!(
fine.checked_geometry(SEG_CHUNK_SAMPLES + MAX_EXTRACTION_CHUNKS, frames),
Err(ExtractError::ExtractionChunkCountTooLarge(70_771)),
"frames_per_chunk={frames} must not buy a single extra chunk"
);
}
assert!(
checked_extraction_tensor_bytes(393_349, 1).is_ok(),
"the memory ceiling admits 393 349 one-frame chunks — 786 698 model calls"
);
assert!(
checked_extraction_chunk_count(393_349).is_err(),
"...and the compute ceiling is what refuses them"
);
}
fn checked_geometry_first_refusal_at_default_stride() -> ExtractError {
const SMALLEST_OVER_CAP_SAMPLES: usize = 1_132_448_001;
Extractor::with_options(Options::new())
.checked_geometry(SMALLEST_OVER_CAP_SAMPLES, 589)
.expect_err("a 19.66 h clip derives a grid past MAX_OUTPUT_FRAMES")
}
#[test]
fn max_extraction_tensor_bytes_is_the_addressable_grids_own_footprint_derived_not_copied() {
use crate::audio::speaker::window;
let w = WindowOptions::new();
let chunks_sw = window::chunk_sliding_window(&w);
let frames_sw = window::frame_sliding_window();
let addressable = derived_output_frame_count(1, chunks_sw, frames_sw)
.expect("the one-chunk output grid is well defined");
assert_eq!(addressable, 594);
assert_eq!(
(window::CHUNK_DURATION_S / window::FRAME_STEP_S).ceil() as usize,
593,
"594 is that span plus the endpoint frame, as try_num_output_frames counts"
);
assert_eq!(
derived_extraction_tensor_bytes(MAX_EXTRACTION_CHUNKS, addressable),
Ok(MAX_EXTRACTION_TENSOR_BYTES),
"the ceiling IS the footprint of the largest grid both other bounds admit"
);
assert_eq!(MAX_EXTRACTION_TENSOR_BYTES, 70_770 * 17_328);
let largest_admitted = MAX_EXTRACTION_CHUNKS - 1;
for frames in 1..=addressable {
assert!(
checked_extraction_tensor_bytes(largest_admitted, frames).is_ok(),
"frames_per_chunk={frames}: the frame cap's own largest grid must fit"
);
}
let widest_admitted = derived_extraction_tensor_bytes(largest_admitted, addressable)
.expect("the frame cap's largest grid at the addressable width");
assert_eq!(widest_admitted, 1_226_285_232);
assert!(
widest_admitted < MAX_EXTRACTION_TENSOR_BYTES,
"the widest geometry the other two bounds admit must fit under this one"
);
assert_eq!(
derived_extraction_tensor_bytes(70_672, 590),
Ok(1_217_819_904),
"round 10 refused this; the frame cap admits it"
);
assert!(checked_extraction_tensor_bytes(70_672, 590).is_ok());
let shipped = derived_extraction_tensor_bytes(MAX_EXTRACTION_CHUNKS, 589)
.expect("the shipped grid at the chunk ceiling");
assert_eq!(shipped, 1_217_810_160);
assert!(
shipped < MAX_EXTRACTION_TENSOR_BYTES,
"round 10's ceiling was this exact figure, so the shipped grid's largest \
admitted geometry sat ON it rather than under it"
);
let per_chunk = |f: usize| {
f * SEG_NUM_SLOTS * size_of::<f64>() + SEG_NUM_SLOTS * EMBEDDING_DIM * size_of::<f32>()
};
let widest = (MAX_EXTRACTION_TENSOR_BYTES - per_chunk(0)) / (SEG_NUM_SLOTS * size_of::<f64>());
assert_eq!(widest, 51_095_812);
assert_eq!(
checked_extraction_tensor_bytes(1, widest),
Ok(MAX_EXTRACTION_TENSOR_BYTES),
"the widest single chunk sits exactly ON the ceiling"
);
assert_eq!(
checked_extraction_tensor_bytes(1, widest + 1),
Err(ExtractError::ExtractionTensorBytesTooLarge(1_226_302_584)),
"one frame more is this bound's own smallest refusal"
);
assert!(
checked_extraction_chunk_count(1).is_ok(),
"...and the chunk bound has nothing to say about a single chunk"
);
}
#[test]
fn the_chunk_axis_cap_tracks_the_chunk_grid_across_every_legal_stride() {
use crate::audio::speaker::window;
const SAMPLES: usize = 9_600_000;
let refuses = |num_chunks: usize, frames: usize| {
checked_extraction_chunk_count(num_chunks).is_err()
|| checked_extraction_tensor_bytes(num_chunks, frames).is_err()
};
for frames in [1usize, 589, 594] {
let mut boundary = None;
for step in (1u32..=SEG_CHUNK_SAMPLES as u32).step_by(7) {
let w = WindowOptions::new().with_step_samples(step);
let num_chunks = window::num_chunks(SAMPLES, &w);
let refused = refuses(num_chunks, frames);
assert_eq!(
refused,
num_chunks > MAX_EXTRACTION_CHUNKS,
"frames={frames} step_samples={step}: {num_chunks} chunks, refused={refused}"
);
if !refused && boundary.is_none() {
boundary = Some(step);
}
}
assert_eq!(
boundary,
Some(134),
"frames={frames}: the ten-minute clip turns at step 134"
);
for (step, expected) in [(133u32, true), (134, false)] {
let w = WindowOptions::new().with_step_samples(step);
let num_chunks = window::num_chunks(SAMPLES, &w);
assert_eq!(
refuses(num_chunks, frames),
expected,
"frames={frames} step_samples={step} -> {num_chunks} chunks"
);
}
}
}
#[test]
fn derived_extraction_tensor_bytes_overflow_arms_match_check_threes_diagnosis() {
use crate::audio::speaker::error::{ExtractionGeometryOverflow, ExtractionPart};
assert_eq!(
derived_extraction_tensor_bytes(usize::MAX, 1),
Err(ExtractError::ExtractionGeometryOverflow(
ExtractionGeometryOverflow::new(ExtractionPart::RawEmbeddings, usize::MAX, 1)
))
);
let n = usize::MAX / 8_192;
assert!(
derived_extraction_tensor_bytes(n, 1).is_ok(),
"the embeddings product must still fit, or this case tests the wrong arm"
);
let m = ExtractionGeometryOverflow::new(ExtractionPart::Segmentations, n, usize::MAX);
assert_eq!(
derived_extraction_tensor_bytes(n, usize::MAX),
Err(ExtractError::ExtractionGeometryOverflow(m))
);
assert_eq!(
(m.part(), m.num_chunks(), m.num_frames_per_chunk()),
(ExtractionPart::Segmentations, n, usize::MAX)
);
let n = usize::MAX / 4_096;
let frames = 86usize;
let raw = n * SEG_NUM_SLOTS * EMBEDDING_DIM * size_of::<f32>();
let seg = n * frames * SEG_NUM_SLOTS * size_of::<f64>();
assert!(
raw.checked_add(seg).is_none(),
"this case needs two products that fit and a sum that does not"
);
assert_eq!(
derived_extraction_tensor_bytes(n, frames),
Err(ExtractError::ExtractionGeometryOverflow(
ExtractionGeometryOverflow::new(ExtractionPart::Segmentations, n, frames)
))
);
let elems = usize::MAX / (SEG_NUM_SLOTS * EMBEDDING_DIM);
assert!(
elems
.checked_mul(SEG_NUM_SLOTS)
.and_then(|v| v.checked_mul(EMBEDDING_DIM))
.is_some(),
"check 3's element product fits"
);
assert!(
derived_extraction_tensor_bytes(elems, 1).is_err(),
"...while its byte size does not"
);
}
#[test]
fn a_step_samples_floor_at_one_frame_step_would_not_have_bounded_the_tensors() {
use crate::audio::speaker::window;
const FRAMES_PER_CHUNK: usize = 589;
let one_frame_step_in_samples = window::FRAME_STEP_S * f64::from(window::SAMPLE_RATE_HZ);
assert_eq!(one_frame_step_in_samples, 270.0);
let floor = 270u32;
assert!(
window::DEFAULT_STEP_SAMPLES > floor,
"the shipped stride is 59.26 frame steps, far above any such floor"
);
let w = WindowOptions::new().with_step_samples(floor);
let chunks_sw = window::chunk_sliding_window(&w);
let frames_sw = window::frame_sliding_window();
let (mut lo, mut hi) = (1usize, 1usize << 25);
assert!(checked_output_frame_count(hi, chunks_sw, frames_sw).is_err());
while lo < hi {
let mid = lo + (hi - lo) / 2;
if checked_output_frame_count(mid, chunks_sw, frames_sw).is_err() {
hi = mid;
} else {
lo = mid + 1;
}
}
let largest_admitted = lo - 1;
assert_eq!(largest_admitted, 4_193_711);
assert_eq!(
derived_extraction_tensor_bytes(largest_admitted, FRAMES_PER_CHUNK),
Ok(72_165_378_888),
"67.21 GiB, still reachable with a 270-sample floor in force"
);
assert_eq!(
checked_extraction_tensor_bytes(largest_admitted, FRAMES_PER_CHUNK),
Err(ExtractError::ExtractionTensorBytesTooLarge(72_165_378_888)),
"only the byte ceiling refuses it"
);
}
#[test]
fn the_byte_ceiling_alone_does_not_bound_the_model_call_count() {
use crate::audio::speaker::window;
const FRAMES_PER_CHUNK: usize = 1;
const STEP: u32 = 2;
const SAMPLES: usize = 946_695;
const NUM_CHUNKS: usize = 393_349;
const TENSOR_BYTES: usize = 1_217_808_504;
const MODEL_CALLS: usize = 2 * NUM_CHUNKS;
let w = WindowOptions::new().with_step_samples(STEP);
let extractor = Extractor::with_options(Options::new().with_window(w));
let chunks_sw = window::chunk_sliding_window(&w);
let frames_sw = window::frame_sliding_window();
let num_chunks = window::num_chunks(SAMPLES, &w);
assert_eq!(num_chunks, NUM_CHUNKS);
assert_eq!(
checked_output_frame_count(num_chunks, chunks_sw, frames_sw),
Ok(3_507),
"the output-frame cap sees a 59.17 s clip and passes it"
);
assert_eq!(
window::first_misaligned_chunk(num_chunks, chunks_sw, frames_sw),
None,
"and the even stride clears the placement scan"
);
assert_eq!(
derived_extraction_tensor_bytes(num_chunks, FRAMES_PER_CHUNK),
Ok(TENSOR_BYTES),
"1 217 808 504 bytes — under round 10's 1 217 810 160 ceiling by 1 656"
);
assert_eq!(MODEL_CALLS, 786_698);
assert_eq!(
extractor.checked_geometry(SAMPLES, FRAMES_PER_CHUNK),
Err(ExtractError::ExtractionChunkCountTooLarge(NUM_CHUNKS)),
"786 698 model invocations for 59.17 s of audio must be refused from \
`samples.len()`, `step_samples` and the segmenter's frame count alone"
);
}
#[test]
fn the_byte_ceiling_admits_every_frame_grid_the_frame_cap_admits() {
use crate::audio::speaker::window;
const FRAMES_PER_CHUNK: usize = 590;
const SAMPLES: usize = 1_130_880_001;
let w = WindowOptions::new();
assert_eq!(w.step_samples(), window::DEFAULT_STEP_SAMPLES);
let extractor = Extractor::with_options(Options::new().with_window(w));
let chunks_sw = window::chunk_sliding_window(&w);
let frames_sw = window::frame_sliding_window();
let num_chunks = window::num_chunks(SAMPLES, &w);
assert_eq!(num_chunks, 70_672);
assert_eq!(
checked_output_frame_count(num_chunks, chunks_sw, frames_sw),
Ok(4_188_505),
"the output grid stays below MAX_OUTPUT_FRAMES"
);
const { assert!(4_188_505 < MAX_OUTPUT_FRAMES) };
assert_eq!(
derived_extraction_tensor_bytes(num_chunks, FRAMES_PER_CHUNK),
Ok(1_217_819_904)
);
assert_eq!(
extractor.checked_geometry(SAMPLES, FRAMES_PER_CHUNK),
Ok((num_chunks, chunks_sw, frames_sw)),
"a 590-frame segmenter at the SHIPPED stride must not be refused where the \
frame cap accepts"
);
let addressable = derived_output_frame_count(1, chunks_sw, frames_sw)
.expect("the one-chunk output grid is well defined");
assert_eq!(addressable, 594);
let largest_admitted = window::num_chunks(1_132_448_001 - 1, &w);
assert_eq!(largest_admitted, 70_769);
for frames in 1..=addressable {
assert!(
checked_extraction_tensor_bytes(largest_admitted, frames).is_ok(),
"frames_per_chunk={frames}: the frame cap's own largest grid must fit"
);
}
}
fn largest_admissible_frames_per_chunk(num_chunks: usize) -> usize {
let w = WindowOptions::new();
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let derived = derived_output_frame_count(num_chunks, chunks_sw, frames_sw)
.expect("the default grid derives a count that fits usize");
let start = crate::audio::speaker::window::reconstruct_chunk_start_frame(
num_chunks - 1,
chunks_sw,
frames_sw,
);
derived - usize::try_from(start).expect("the default grid places every chunk at or after 0")
}
#[test]
fn try_from_parts_refuses_a_grid_that_stops_short_of_the_last_chunk() {
let w = WindowOptions::new();
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
const NUM_CHUNKS: usize = 3;
const FRAMES: usize = 594;
let derived = derived_output_frame_count(NUM_CHUNKS, chunks_sw, frames_sw).expect("well defined");
assert_eq!(
derived, 712,
"three default chunks derive 712 output frames"
);
let parts = || ExtractionParts {
raw_embeddings: vec![0.0f32; NUM_CHUNKS * SEG_NUM_SLOTS * EMBEDDING_DIM],
segmentations: vec![0.0f64; NUM_CHUNKS * FRAMES * SEG_NUM_SLOTS],
count: vec![0u8; derived],
num_chunks: NUM_CHUNKS,
num_frames_per_chunk: FRAMES,
chunks_sw,
frames_sw,
};
let ExtractError::UncoveredLastChunk(u) = refused(parts()) else {
panic!("expected UncoveredLastChunk, got {:?}", refused(parts()))
};
assert_eq!((u.start_frame(), u.required(), u.got()), (119, 713, 712));
let p = parts();
let e = Extraction::from_parts_unchecked(
p.raw_embeddings,
p.segmentations,
p.count,
p.num_chunks,
p.num_frames_per_chunk,
p.chunks_sw,
p.frames_sw,
);
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
for (route, got) in [
("offline", e.diarize_with(&plda, ClusterBackend::default())),
("online", e.diarize_online(OnlineOptions::default())),
] {
let Err(diaric::offline::Error::Reconstruct(diaric::reconstruct::Error::Shape(
diaric::reconstruct::ShapeError::OutputFrameCountTooSmall { got, required },
))) = got
else {
panic!("{route} must refuse the uncovered grid, got {got:?}")
};
assert_eq!(
(got, required),
(u.got(), u.required()),
"{route} reported different numbers than check 14 predicted"
);
}
}
#[test]
fn the_producer_preflight_refuses_a_segmenter_whose_frames_overrun_the_grid() {
let w = WindowOptions::new();
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let extractor = Extractor::new();
let num_chunks = crate::audio::speaker::window::num_chunks(SEG_CHUNK_SAMPLES, &w);
assert_eq!(num_chunks, 1);
assert_eq!(largest_admissible_frames_per_chunk(1), 594);
assert_eq!(
checked_output_frame_count(num_chunks, chunks_sw, frames_sw),
Ok(594)
);
assert_eq!(checked_extraction_chunk_count(num_chunks), Ok(1));
assert_eq!(checked_extraction_tensor_bytes(num_chunks, 595), Ok(17_352));
assert_eq!(
crate::audio::speaker::window::first_misaligned_chunk(num_chunks, chunks_sw, frames_sw),
None
);
match extractor.checked_geometry(SEG_CHUNK_SAMPLES, 595) {
Err(ExtractError::UncoveredLastChunk(u)) => {
assert_eq!((u.start_frame(), u.required(), u.got()), (0, 595, 594));
}
other => panic!("a 595-frame segmenter must be refused pre-inference, got {other:?}"),
}
assert_eq!(
extractor.checked_geometry(SEG_CHUNK_SAMPLES, 594),
Ok((num_chunks, chunks_sw, frames_sw)),
"the largest admissible frame count must not be refused"
);
match Extraction::assemble_checked(
vec![0.0f32; SEG_NUM_SLOTS * EMBEDDING_DIM],
vec![0.0f64; 595 * SEG_NUM_SLOTS],
num_chunks,
595,
w.onset(),
chunks_sw,
frames_sw,
) {
Err(ExtractError::UncoveredLastChunk(u)) => {
assert_eq!((u.start_frame(), u.required(), u.got()), (0, 595, 594));
}
other => panic!("the assembly door must reach the same verdict, got {other:?}"),
}
}
#[test]
fn the_shipped_589_frame_grid_keeps_four_frames_of_headroom_everywhere() {
const SHIPPED_FRAMES: usize = 589;
assert_eq!(
crate::audio::speaker::source::argmax::ARGMAX_FRAMES_PER_WINDOW,
SHIPPED_FRAMES
);
let mut min_margin = usize::MAX;
let mut max_margin = 0usize;
for num_chunks in 1..=MAX_EXTRACTION_CHUNKS {
let admissible = largest_admissible_frames_per_chunk(num_chunks);
assert!(
admissible >= SHIPPED_FRAMES,
"num_chunks={num_chunks}: the shipped grid must fit, admissible={admissible}"
);
let margin = admissible - SHIPPED_FRAMES;
min_margin = min_margin.min(margin);
max_margin = max_margin.max(margin);
}
assert_eq!(
(min_margin, max_margin),
(4, 5),
"the shipped 589-frame grid's headroom over the whole admitted chunk range"
);
}
#[test]
fn check_14s_boundary_is_the_grids_own_allowance_at_every_chunk_count() {
let w = WindowOptions::new();
let chunks_sw = crate::audio::speaker::window::chunk_sliding_window(&w);
let frames_sw = crate::audio::speaker::window::frame_sliding_window();
let extractor = Extractor::new();
let mut seen = std::collections::BTreeSet::new();
for num_chunks in 1..=MAX_EXTRACTION_CHUNKS {
seen.insert(largest_admissible_frames_per_chunk(num_chunks));
}
assert_eq!(
seen.into_iter().collect::<Vec<_>>(),
vec![593, 594],
"the allowance at the default stride takes exactly these two values"
);
assert_eq!(
(1..=MAX_EXTRACTION_CHUNKS).find(|n| largest_admissible_frames_per_chunk(*n) == 593),
Some(3),
"594 frames per chunk first becomes unusable at three chunks"
);
for num_chunks in 1..=8usize {
let samples_len = SEG_CHUNK_SAMPLES + (num_chunks - 1) * w.step_samples() as usize;
assert_eq!(
crate::audio::speaker::window::num_chunks(samples_len, &w),
num_chunks
);
let admissible = largest_admissible_frames_per_chunk(num_chunks);
assert_eq!(
extractor.checked_geometry(samples_len, admissible),
Ok((num_chunks, chunks_sw, frames_sw)),
"num_chunks={num_chunks}: the allowance itself must be accepted"
);
match extractor.checked_geometry(samples_len, admissible + 1) {
Err(ExtractError::UncoveredLastChunk(u)) => assert_eq!(
(u.required(), u.got()),
(
u.start_frame() as usize + admissible + 1,
derived_output_frame_count(num_chunks, chunks_sw, frames_sw).expect("well defined")
),
"num_chunks={num_chunks}"
),
other => {
panic!("num_chunks={num_chunks}: one past the allowance must be refused, got {other:?}")
}
}
}
for num_chunks in 1..=8usize {
let derived =
derived_output_frame_count(num_chunks, chunks_sw, frames_sw).expect("well defined");
let admissible = largest_admissible_frames_per_chunk(num_chunks);
for (frames, want_ok) in [(admissible, true), (admissible + 1, false)] {
let parts = ExtractionParts {
raw_embeddings: vec![0.0f32; num_chunks * SEG_NUM_SLOTS * EMBEDDING_DIM],
segmentations: vec![0.0f64; num_chunks * frames * SEG_NUM_SLOTS],
count: vec![0u8; derived],
num_chunks,
num_frames_per_chunk: frames,
chunks_sw,
frames_sw,
};
let got = Extraction::try_from_parts(parts);
assert_eq!(
got.is_ok(),
want_ok,
"num_chunks={num_chunks} frames={frames}: got {:?}",
got.err()
);
}
}
}