use super::*;
use crate::audio::ced::{NUM_CLASSES, WINDOW_SAMPLES, window::Span};
fn scripted_scores() -> Vec<f32> {
let mut state = 0x2545F4914F6CDD1Du64;
let mut out: Vec<f32> = (0..NUM_CLASSES)
.map(|_| {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 40) as f32 / (1u64 << 24) as f32 - 0.5) * 12.0
})
.collect();
out[10] = 6.5;
out[200] = 6.5;
out[500] = 6.5;
out[3] = -1.25;
out[400] = -1.25;
out
}
fn sorted_reference(scores: &[f32]) -> Vec<(usize, f32)> {
let mut pairs: Vec<(usize, f32)> = scores.iter().copied().enumerate().collect();
pairs.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
pairs
}
#[test]
fn sigmoid_matches_the_soundevents_form() {
assert_eq!(sigmoid(0.0), 0.5);
assert!((sigmoid(2.0) - 0.880_797).abs() < 1e-6);
assert!((sigmoid(-2.0) - 0.119_203).abs() < 1e-6);
let mut prev = f32::NEG_INFINITY;
for i in -50..=50 {
let v = sigmoid(i as f32 / 5.0);
assert!(v > prev);
prev = v;
}
}
#[test]
fn top_k_matches_a_full_sort_reference() {
let scores = scripted_scores();
let reference = sorted_reference(&scores);
for k in [1usize, 5, 50, NUM_CLASSES] {
let preds = top_k_from_scores(scores.iter().copied().enumerate(), k, sigmoid).unwrap();
assert_eq!(preds.len(), k.min(NUM_CLASSES));
for (p, &(ref_index, ref_score)) in preds.iter().zip(reference.iter()) {
assert_eq!(p.index(), ref_index, "k={k}");
assert_eq!(p.confidence(), sigmoid(ref_score), "k={k}");
}
}
}
#[test]
fn ties_break_by_ascending_class_index() {
let scores = scripted_scores();
let preds = top_k_from_scores(scores.iter().copied().enumerate(), 3, sigmoid).unwrap();
let indices: Vec<usize> = preds.iter().map(|p| p.index()).collect();
assert_eq!(indices, vec![10, 200, 500]);
}
#[test]
fn zero_k_is_empty_not_an_error() {
let scores = scripted_scores();
let preds = top_k_from_scores(scores.iter().copied().enumerate(), 0, sigmoid).unwrap();
assert!(preds.is_empty());
}
#[test]
fn oversized_k_saturates_at_num_classes() {
let scores = scripted_scores();
for k in [NUM_CLASSES + 100, usize::MAX] {
let preds = top_k_from_scores(scores.iter().copied().enumerate(), k, sigmoid).unwrap();
assert_eq!(preds.len(), NUM_CLASSES, "k={k}");
}
}
#[test]
fn sigmoid_at_extraction_equals_sigmoid_then_rank() {
let scores = scripted_scores();
let at_extraction = top_k_from_scores(scores.iter().copied().enumerate(), 20, sigmoid).unwrap();
let confidences: Vec<f32> = scores.iter().copied().map(sigmoid).collect();
let pre_mapped = top_k_from_scores(confidences.iter().copied().enumerate(), 20, |c| c).unwrap();
for (a, b) in at_extraction.iter().zip(pre_mapped.iter()) {
assert_eq!(a.index(), b.index());
assert_eq!(a.confidence(), b.confidence());
}
}
#[test]
fn event_prediction_round_trips_known_rows() {
let p = EventPrediction::from_confidence(0, 0.75).unwrap();
assert_eq!(p.index(), 0);
assert_eq!(p.id(), SoundEventId::new(3));
assert_eq!(p.id().get(), 3u16);
assert_eq!(p.mid(), "/m/09x0r");
assert_eq!(p.name(), "Speech");
assert_eq!(p.confidence(), 0.75);
assert_eq!(p.event().index(), 0);
assert_eq!(
RatedSoundEvent::from_id(p.id()).map(RatedSoundEvent::mid),
Some("/m/09x0r")
);
assert!(EventPrediction::from_confidence(NUM_CLASSES - 1, 0.5).is_ok());
let err = EventPrediction::from_confidence(NUM_CLASSES, 0.5).unwrap_err();
assert!(
matches!(err, crate::audio::ced::Error::UnknownClassIndex(index) if index == NUM_CLASSES),
"got {err:?}"
);
}
#[test]
fn confidences_hold_the_class_count_invariant() {
let c = Confidences::new(vec![0.5; NUM_CLASSES]);
assert_eq!(c.as_slice().len(), NUM_CLASSES);
}
#[test]
#[should_panic(expected = "NUM_CLASSES")]
fn confidences_reject_a_wrong_length_vector() {
let _ = Confidences::new(vec![0.5; NUM_CLASSES - 1]);
}
#[test]
fn try_from_slice_accepts_a_hand_built_confidence_vector() {
let mut values = vec![0.0f32; NUM_CLASSES];
values[74] = 0.86;
values[NUM_CLASSES - 1] = 1.0;
let c = Confidences::try_from_slice(&values).expect("valid confidences");
assert_eq!(c.as_slice(), values.as_slice());
}
#[test]
fn try_from_slice_rejects_a_wrong_length_vector_without_panicking() {
for got in [NUM_CLASSES - 1, NUM_CLASSES + 1, 0] {
let err = Confidences::try_from_slice(&vec![0.5f32; got]).unwrap_err();
assert!(
matches!(
&err,
crate::audio::ced::Error::ClassCountMismatch(e)
if e.expected() == NUM_CLASSES && e.got() == got
),
"len {got} gave {err:?}"
);
}
}
#[test]
fn try_from_slice_rejects_values_outside_the_stated_invariant() {
for (index, bad) in [
(0usize, -0.001f32),
(74, 1.001),
(300, f32::NAN),
(NUM_CLASSES - 1, f32::INFINITY),
(5, f32::NEG_INFINITY),
] {
let mut values = vec![0.5f32; NUM_CLASSES];
values[index] = bad;
let err = Confidences::try_from_slice(&values).unwrap_err();
assert!(
matches!(
&err,
crate::audio::ced::Error::InvalidConfidence(e)
if e.index() == index && (e.value() == bad || (e.value().is_nan() && bad.is_nan()))
),
"{bad} at {index} gave {err:?}"
);
}
let mut edges = vec![0.5f32; NUM_CLASSES];
edges[0] = 0.0;
edges[1] = 1.0;
edges[2] = -0.0;
assert!(Confidences::try_from_slice(&edges).is_ok());
}
#[test]
fn try_from_slice_and_from_logits_agree_on_the_same_confidences() {
let logits = scripted_scores();
let via_logits = Confidences::from_logits(&logits);
let via_slice = Confidences::try_from_slice(via_logits.as_slice()).expect("sigmoid is in range");
assert_eq!(via_slice, via_logits);
}
#[test]
fn from_logits_maps_sigmoid_elementwise() {
let mut logits = vec![0.0f32; NUM_CLASSES];
logits[0] = 2.0;
logits[526] = -2.0;
let c = Confidences::from_logits(&logits);
assert_eq!(c.as_slice()[0], sigmoid(2.0));
assert_eq!(c.as_slice()[1], 0.5);
assert_eq!(c.as_slice()[526], sigmoid(-2.0));
}
#[test]
fn confidences_top_k_ranks_in_confidence_space() {
let mut values = vec![0.1f32; NUM_CLASSES];
values[7] = 0.9;
values[300] = 0.8;
let c = Confidences::new(values);
let preds = c.top_k(2).unwrap();
assert_eq!(preds[0].index(), 7);
assert_eq!(preds[0].confidence(), 0.9);
assert_eq!(preds[1].index(), 300);
assert!(c.top_k(0).unwrap().is_empty());
assert_eq!(c.top_k(usize::MAX).unwrap().len(), NUM_CLASSES);
}
#[test]
fn window_confidences_pair_values_with_spans() {
let span = Span::new(0, WINDOW_SAMPLES, WINDOW_SAMPLES);
let w = WindowConfidences::new(Confidences::new(vec![0.5; NUM_CLASSES]), span);
assert_eq!(w.span().start(), 0);
assert_eq!(w.value().as_slice().len(), NUM_CLASSES);
}