#![allow(dead_code)]
#[path = "../../support/models_lock_manifest.rs"]
#[allow(dead_code)]
pub mod models_lock_manifest;
#[path = "../../support/workspace_root.rs"]
#[allow(dead_code)]
mod workspace_root;
#[allow(unused_imports)]
pub use workspace_root::{checkout_parent, models_root, workspace_root};
use std::{collections::BTreeMap, path::PathBuf};
use coremlit::audio::speaker::{
extract::EXCLUDE_OVERLAP_MIN_FRAMES,
segment::{POWERSET_CLASSES, SEG_CHUNK_SAMPLES, SEG_NUM_SLOTS, multilabel},
window::DEFAULT_ONSET,
};
use sha2::{Digest, Sha256};
#[path = "../../support/coremlit_dir.rs"]
#[allow(dead_code)]
mod coremlit_dir;
use coremlit_dir::coremlit_path;
#[allow(dead_code)]
pub const VENDOR_DIR: &str = "speakerkit";
#[allow(dead_code)]
pub const OVERLAY_LOCK_REVISION: &str = "3db69988bf2de12bab250614d6ac2b03d35132a2";
#[allow(dead_code)]
pub const BASE_LOCK_REVISION: &str = "1ed7a662fdc7109e36d822db793ee6eebdaf8594";
#[allow(dead_code)]
pub fn overlay_sha256(bundle: &str) -> Vec<(String, String)> {
models_lock_manifest::bundle_manifest(
&workspace_root::workspace_root(),
VENDOR_DIR,
OVERLAY_LOCK_REVISION,
bundle,
)
}
#[allow(dead_code)]
pub fn base_sha256(bundle: &str) -> Vec<(String, String)> {
models_lock_manifest::bundle_manifest(
&workspace_root::workspace_root(),
VENDOR_DIR,
BASE_LOCK_REVISION,
bundle,
)
}
pub fn models_dir() -> PathBuf {
std::env::var_os("SPEAKERKIT_TEST_MODELS").map_or_else(
|| workspace_root::models_root().join("speakerkit"),
PathBuf::from,
)
}
pub fn seg_path() -> PathBuf {
models_dir().join("pyannote_segmentation.mlmodelc")
}
pub fn argmax_models_dir() -> PathBuf {
std::env::var_os("ARGMAX_TEST_MODELS").map_or_else(
|| workspace_root::models_root().join("argmax-speakerkit"),
PathBuf::from,
)
}
pub fn embed_path() -> PathBuf {
models_dir().join("wespeaker_v2.mlmodelc")
}
pub fn embed_fp32_path() -> PathBuf {
models_dir().join("wespeaker.mlmodelc")
}
pub struct Fixture {
pub name: &'static str,
pub source: &'static str,
pub sha256: &'static str,
pub note: &'static str,
}
pub const FIXTURES: &[Fixture] = &[
Fixture {
name: "02_pyannote_sample",
source: "diarization/tests/parity/fixtures/02_pyannote_sample/clip_16k.wav",
sha256: "c319b4abca767b124e41432d364fd7df006cb26bb79d09326c487d606a134e6e",
note: "pyannote's canonical 30.0 s sample → exactly 3 full 10 s chunks (no padding)",
},
Fixture {
name: "07_yuhewei_dongbei_english",
source: "diarization/tests/parity/fixtures/07_yuhewei_dongbei_english/clip_16k.wav",
sha256: "096890ba8ffbaf10ca770c5373bf6c6664777f9421595c2cb7780af8cb2e46ff",
note: "25.26 s clip → 2 full chunks + 1 partial (exercises final-chunk zero-padding)",
},
];
pub const SEG_MODEL_LABEL: &str =
"segmentation-3.0.onnx (dia bundled, ort CPU EP, raw powerset logits)";
pub fn fixtures_dir() -> PathBuf {
coremlit_path("tests/speaker/fixtures")
}
pub fn audio_path(name: &str) -> PathBuf {
fixtures_dir().join("audio").join(format!("{name}.wav"))
}
pub fn golden_path(name: &str) -> PathBuf {
fixtures_dir().join("golden").join(format!("{name}.json"))
}
pub fn load_wav_16k_mono(path: &std::path::Path) -> Vec<f32> {
let mut reader =
hound::WavReader::open(path).unwrap_or_else(|e| panic!("open {}: {e}", path.display()));
let spec = reader.spec();
assert_eq!(spec.sample_rate, 16_000, "{}: not 16 kHz", path.display());
assert_eq!(spec.channels, 1, "{}: not mono", path.display());
match spec.sample_format {
hound::SampleFormat::Int => {
assert_eq!(
spec.bits_per_sample,
16,
"{}: only 16-bit int PCM supported",
path.display()
);
reader
.samples::<i16>()
.map(|s| f32::from(s.expect("read i16 sample")) / 32_768.0)
.collect()
}
hound::SampleFormat::Float => reader
.samples::<f32>()
.map(|s| s.expect("read f32 sample"))
.collect(),
}
}
pub fn chunk_and_pad(samples: &[f32]) -> Vec<Vec<f32>> {
let n = samples.len().div_ceil(SEG_CHUNK_SAMPLES).max(1);
(0..n)
.map(|c| {
let start = c * SEG_CHUNK_SAMPLES;
let mut chunk = vec![0.0f32; SEG_CHUNK_SAMPLES];
if start < samples.len() {
let end = (start + SEG_CHUNK_SAMPLES).min(samples.len());
chunk[..end - start].copy_from_slice(&samples[start..end]);
}
chunk
})
.collect()
}
pub fn fnv1a_f32(samples: &[f32]) -> u64 {
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
for &s in samples {
for b in s.to_le_bytes() {
h ^= u64::from(b);
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}
}
h
}
pub fn fnv_hex(h: u64) -> String {
format!("{h:016x}")
}
pub fn sha256_hex(data: &[u8]) -> String {
Sha256::digest(data)
.iter()
.map(|byte| format!("{byte:02x}"))
.collect()
}
const OVERLAY_MODEL_MIL_PINS: &[(&str, &str)] = &[
(
"pyannote_segmentation.mlmodelc",
"ded0d1ee11d77976b5c706ce667d0c8cb49977d3fe4367cccbd7b582bdb86dec",
),
(
"wespeaker.mlmodelc",
"cff0cfe914078e9336754a9b38a68c2cdd88ca7b6bf97568ad551ab03ae1b666",
),
];
#[allow(dead_code)]
pub fn skipped_for_stale_overlay(gate: &str, bundle_root: &std::path::Path, bundle: &str) -> bool {
use std::io::Write;
let expected = OVERLAY_MODEL_MIL_PINS
.iter()
.find_map(|(name, hash)| (*name == bundle).then_some(*hash))
.unwrap_or_else(|| {
panic!("skipped_for_stale_overlay: {bundle:?} is not a pinned overlay bundle")
});
let mil = bundle_root.join("model.mil");
let Ok(bytes) = std::fs::read(&mil) else {
return false;
};
let actual = sha256_hex(&bytes);
if actual == expected {
return false;
}
let line = format!(
"model-gates | SKIPPED {gate}: {} sha256 {actual} != MODELS_LOCK overlay pin {expected} — \
this host staged the FluidInference pre-repair build, not the fp16-guard-repaired \
FinDIT-Studio overlay; re-run the download in MODELS_LOCK's lock order before trusting this \
gate\n",
mil.display()
);
let mut fd2 = std::mem::ManuallyDrop::new(unsafe {
<std::fs::File as std::os::fd::FromRawFd>::from_raw_fd(2)
});
let _ = fd2.write_all(line.as_bytes());
true
}
pub fn cosine(a: &[f32], b: &[f32]) -> f64 {
assert_eq!(a.len(), b.len(), "cosine: length mismatch");
assert!(
a.iter().all(|v| v.is_finite()),
"cosine: vector `a` contains a non-finite element"
);
assert!(
b.iter().all(|v| v.is_finite()),
"cosine: vector `b` contains a non-finite element"
);
let (mut dot, mut na, mut nb) = (0.0f64, 0.0f64, 0.0f64);
for (&x, &y) in a.iter().zip(b) {
let (x, y) = (f64::from(x), f64::from(y));
dot += x * y;
na += x * x;
nb += y * y;
}
assert!(na > 0.0, "cosine: vector `a` has zero norm");
assert!(nb > 0.0, "cosine: vector `b` has zero norm");
dot / (na.sqrt() * nb.sqrt())
}
pub fn softmax_row(row: &[f32]) -> [f32; POWERSET_CLASSES] {
assert_eq!(row.len(), POWERSET_CLASSES, "softmax_row: bad row length");
let max = row.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut out = [0f32; POWERSET_CLASSES];
let mut sum = 0f32;
for (o, &l) in out.iter_mut().zip(row) {
*o = (l - max).exp();
sum += *o;
}
for o in &mut out {
*o /= sum;
}
out
}
pub fn max_abs_diff(a: &[f32], b: &[f32]) -> f64 {
assert_eq!(a.len(), b.len(), "max_abs_diff: length mismatch");
a.iter()
.zip(b)
.map(|(&x, &y)| (f64::from(x) - f64::from(y)).abs())
.fold(0.0, f64::max)
}
pub const SEG_ROW_SUM_EXP_TOL: f64 = 1e-4;
pub fn check_seg_log_probs(seg_logits: &[f32], num_frames: usize) -> Result<(), String> {
let expected = num_frames * POWERSET_CLASSES;
if seg_logits.len() != expected {
return Err(format!(
"seg_logits length {} != num_frames*POWERSET_CLASSES ({num_frames}*{POWERSET_CLASSES}={expected})",
seg_logits.len()
));
}
for (f, row) in seg_logits
.as_chunks::<POWERSET_CLASSES>()
.0
.iter()
.enumerate()
{
let mut sum_exp = 0.0_f64;
for (k, &v) in row.iter().enumerate() {
if !v.is_finite() {
return Err(format!("frame {f} class {k}: non-finite log-prob {v}"));
}
if v > 0.0 {
return Err(format!(
"frame {f} class {k}: log-prob {v} > 0 — a probability's log is ≤ 0; this looks like a \
raw logit"
));
}
sum_exp += f64::from(v).exp();
}
let dev = (sum_exp - 1.0).abs();
if dev > SEG_ROW_SUM_EXP_TOL {
return Err(format!(
"frame {f}: Σexp(row) = {sum_exp:.9}, off 1.0 by {dev:.3e} (> {SEG_ROW_SUM_EXP_TOL:.0e}) — \
the row is not a normalized log-softmax (raw logits, or a broken normalization)"
));
}
}
Ok(())
}
pub fn powerset_argmax(row: &[f32]) -> usize {
assert_eq!(
row.len(),
POWERSET_CLASSES,
"powerset_argmax: bad row length"
);
let mut argmax = 0usize;
let mut max = row[0];
for (k, &v) in row.iter().enumerate().skip(1) {
if v > max {
max = v;
argmax = k;
}
}
argmax
}
pub fn mask_to_string(mask: &[bool]) -> String {
mask.iter().map(|&b| if b { '1' } else { '0' }).collect()
}
pub fn parse_bit_mask(context: &str, expected_len: usize, raw: &str) -> Vec<bool> {
let bits: Vec<bool> = raw
.chars()
.map(|c| match c {
'0' => false,
'1' => true,
other => panic!(
"{context} contains {other:?} — a golden bit mask is hard-binary, so only '0'/'1' are \
valid; an unknown character is a malformed golden, not an inactive frame."
),
})
.collect();
assert_eq!(
bits.len(),
expected_len,
"{context} has {} characters, expected exactly {expected_len}. A short, empty, or over-long \
mask makes the comparison vacuous (or is truncated back by the model's frame padding) while \
the gate still reports full coverage.",
bits.len()
);
bits
}
pub struct GoldenSlot {
pub slot: usize,
pub mask: Vec<bool>,
pub embedding: Vec<f32>,
}
pub struct GoldenChunk {
pub input_len: usize,
pub input_fnv1a: u64,
pub seg_logits: Vec<f32>,
pub slots: Vec<GoldenSlot>,
}
pub struct Golden {
pub fixture: String,
pub num_chunks: usize,
pub num_frames: usize,
pub chunks: Vec<GoldenChunk>,
}
pub fn load_golden(name: &str) -> Golden {
let path = golden_path(name);
let bytes =
std::fs::read(&path).unwrap_or_else(|e| panic!("read golden {}: {e}", path.display()));
let v: serde_json::Value =
serde_json::from_slice(&bytes).unwrap_or_else(|e| panic!("parse golden {name}: {e}"));
let num_frames = v["num_frames"].as_u64().expect("num_frames") as usize;
let chunks = v["chunks"]
.as_array()
.expect("chunks array")
.iter()
.enumerate()
.map(|(c_idx, c)| {
let seg_logits = c["seg_logits"]
.as_array()
.expect("seg_logits array")
.iter()
.map(|x| x.as_f64().expect("logit f64") as f32)
.collect();
let slots = c["slots"]
.as_array()
.expect("slots array")
.iter()
.map(|s| {
let slot = s["slot"].as_u64().expect("slot") as usize;
GoldenSlot {
slot,
mask: parse_bit_mask(
&format!("{name}: chunk {c_idx} slot {slot}: mask"),
num_frames,
s["mask"].as_str().expect("mask string"),
),
embedding: s["embedding"]
.as_array()
.expect("embedding array")
.iter()
.map(|x| x.as_f64().expect("embed f64") as f32)
.collect(),
}
})
.collect();
let hex = c["input_fnv1a"].as_str().expect("input_fnv1a hex");
GoldenChunk {
input_len: c["input_len"].as_u64().expect("input_len") as usize,
input_fnv1a: u64::from_str_radix(hex, 16).expect("parse fnv hex"),
seg_logits,
slots,
}
})
.collect();
Golden {
fixture: v["fixture"].as_str().expect("fixture").to_string(),
num_chunks: v["num_chunks"].as_u64().expect("num_chunks") as usize,
num_frames,
chunks,
}
}
pub fn derive_expected_slot_masks(
seg_logits: &[f32],
num_frames: usize,
) -> [Option<Vec<bool>>; SEG_NUM_SLOTS] {
let slab = multilabel(seg_logits, num_frames);
let onset = f64::from(DEFAULT_ONSET);
let mut clean = vec![false; num_frames];
for (f, clean_f) in clean.iter_mut().enumerate() {
let active = (0..SEG_NUM_SLOTS)
.filter(|&s| slab[f * SEG_NUM_SLOTS + s] >= onset)
.count();
*clean_f = active < 2;
}
core::array::from_fn(|s| {
let mut frame_mask = vec![false; num_frames];
let mut any = false;
for (f, m) in frame_mask.iter_mut().enumerate() {
*m = slab[f * SEG_NUM_SLOTS + s] >= onset;
any |= *m;
}
if !any {
return None; }
let mut used = vec![false; num_frames];
let mut clean_count = 0usize;
for (f, u) in used.iter_mut().enumerate() {
*u = frame_mask[f] && clean[f];
if *u {
clean_count += 1;
}
}
if clean_count <= EXCLUDE_OVERLAP_MIN_FRAMES {
used = frame_mask;
}
Some(used)
})
}
pub fn assert_golden_roster(golden: &Golden) -> usize {
let mut total = 0usize;
for (c_idx, chunk) in golden.chunks.iter().enumerate() {
let derived = derive_expected_slot_masks(&chunk.seg_logits, golden.num_frames);
let expected: BTreeMap<usize, &Vec<bool>> = derived
.iter()
.enumerate()
.filter_map(|(s, m)| m.as_ref().map(|mask| (s, mask)))
.collect();
let mut stored: BTreeMap<usize, &Vec<bool>> = BTreeMap::new();
for slot in &chunk.slots {
assert!(
stored.insert(slot.slot, &slot.mask).is_none(),
"{}: chunk {c_idx}: slot {} appears more than once in the golden",
golden.fixture,
slot.slot
);
}
let stored_roster: Vec<usize> = stored.keys().copied().collect();
let expected_roster: Vec<usize> = expected.keys().copied().collect();
assert_eq!(
stored_roster, expected_roster,
"{}: chunk {c_idx}: stored (chunk, slot) roster {stored_roster:?} != roster \
{expected_roster:?} independently derived from seg_logits — a golden slot was added or \
dropped",
golden.fixture
);
for (s, exp_mask) in &expected {
let got = stored
.get(s)
.expect("roster equality checked directly above");
assert_eq!(
got,
exp_mask,
"{}: chunk {c_idx} slot {s}: stored mask ({} active frame(s)) != mask independently \
derived from seg_logits ({} active frame(s))",
golden.fixture,
got.iter().filter(|&&b| b).count(),
exp_mask.iter().filter(|&&b| b).count()
);
}
total += expected.len();
}
total
}
#[path = "../../support/host_class.rs"]
#[allow(dead_code)]
mod host_class;
#[allow(unused_imports)]
pub use host_class::{HostClass, HostVerdict, RecordedHost, check_host_class, legacy_failure_note};
#[path = "../../support/model_gate_report.rs"]
mod model_gate_report;
#[test]
fn model_gate_report() {
model_gate_report::report(&[
("SPEAKERKIT_TEST_MODELS", models_dir()),
("ARGMAX_TEST_MODELS", argmax_models_dir()),
]);
}