mod common;
use coremlit::audio::ced::{
CedModel, ChunkAggregation, Classifier, ClassifierOptions, SoundEventId, WindowPlan,
};
const SINE_CLASS: usize = 501; const SINE_NAME: &str = "Sine wave";
const SINE_MID: &str = "/m/01v_m0";
const SINE_ID: SoundEventId = SoundEventId::new(593);
const SINE_CONF_LO: f32 = 0.80;
const SINE_CONF_HI: f32 = 0.98;
fn clip_wav(model: CedModel, file: &str) -> std::path::PathBuf {
common::fixture_path(&format!("goldens/{}", model.as_str())).join(file)
}
fn single_window_top_k(model: CedModel) {
let corpus = common::load_golden_corpus(model);
assert!(!corpus.clips.is_empty(), "goldens corpus must not be empty");
let clf = Classifier::from_file(common::model_path(model)).unwrap();
let wav = common::read_wav_16k_mono(&clip_wav(model, "../../mel/sine440_10s.wav"));
let top = clf.classify(&wav, 5).unwrap();
println!(
"[e2e] {model} sine440 top-5: {:?}",
top
.iter()
.map(|p| (p.index(), p.name(), p.confidence()))
.collect::<Vec<_>>()
);
assert_eq!(
top[0].index(),
SINE_CLASS,
"{model}: sine top-1 not Sine wave"
);
assert_eq!(top[0].name(), SINE_NAME);
assert_eq!(top[0].mid(), SINE_MID, "{model}: sine top-1 mid");
assert_eq!(top[0].id(), SINE_ID, "{model}: sine top-1 permanent id");
let c = top[0].confidence();
assert!(
(SINE_CONF_LO..=SINE_CONF_HI).contains(&c),
"{model}: sine confidence {c:.4} outside [{SINE_CONF_LO}, {SINE_CONF_HI}]"
);
}
fn long_clip_rank(model: CedModel) {
let clf = Classifier::from_file(common::model_path(model)).unwrap();
let wav = common::read_wav_16k_mono(&clip_wav(model, "../clips/long_15s.wav"));
let plan = WindowPlan::new();
let n_spans = plan.spans(wav.len()).unwrap().len();
assert_eq!(
n_spans, 2,
"{model}: 15 s clip must plan two 10 s windows, got {n_spans}"
);
let windows = clf.classify_windows(&wav, &plan).unwrap();
assert_eq!(
windows.len(),
n_spans,
"{model}: per-window count != plan spans"
);
let mean = clf
.classify_long(&wav, 5, &plan, ChunkAggregation::Mean)
.unwrap();
let max = clf
.classify_long(&wav, 5, &plan, ChunkAggregation::Max)
.unwrap();
println!(
"[e2e] {model} long Mean top1={}({}) Max top1={}({})",
mean[0].name(),
mean[0].confidence(),
max[0].name(),
max[0].confidence()
);
assert_eq!(
mean[0].index(),
SINE_CLASS,
"{model}: long Mean top-1 not Sine wave"
);
assert_eq!(
max[0].index(),
SINE_CLASS,
"{model}: long Max top-1 not Sine wave"
);
assert_eq!(mean.len(), 5);
assert_eq!(max.len(), 5);
}
fn prewarm(model: CedModel) {
let clf = Classifier::load(common::model_path(model), ClassifierOptions::new()).unwrap();
clf.prewarm().unwrap();
let wav = common::read_wav_16k_mono(&clip_wav(model, "../../mel/silence_2s.wav"));
let out = clf.classify(&wav, 1).unwrap();
assert_eq!(
out.len(),
1,
"{model}: warm classify returned no prediction"
);
}
macro_rules! per_model_gates {
($($m:ident => $v:expr),+ $(,)?) => {$(
mod $m {
use super::CedModel;
#[test]
#[ignore = "requires staged CED model + fixtures (CED_TEST_MODELS) — Wave C"]
fn single_window_clip_yields_the_pinned_top_k() {
super::single_window_top_k($v);
}
#[test]
#[ignore = "requires staged CED model + fixtures (CED_TEST_MODELS) — Wave C"]
fn long_clip_windows_aggregate_and_rank_as_pinned() {
super::long_clip_rank($v);
}
#[test]
#[ignore = "requires staged CED model (CED_TEST_MODELS) — Wave C"]
fn prewarm_smoke() {
super::prewarm($v);
}
}
)+};
}
per_model_gates!(
tiny => CedModel::Tiny,
mini => CedModel::Mini,
small => CedModel::Small,
base => CedModel::Base,
);