coremlit 0.1.1

Safe, synchronous CoreML runtime for macOS (CPU/GPU/Neural Engine) with opt-in on-device multimodal pipelines: speech (Whisper STT, forced alignment, speaker diarization, Silero VAD), AudioSet sound-event tagging, and audio/text/image embeddings (CLAP, granite, SigLIP)
use super::*;
use crate::audio::ced::{NUM_CLASSES, WINDOW_SAMPLES, window::Span};

/// Deterministic scripted scores: a fixed LCG over the full class range, with
/// deliberate duplicates injected so the tie-break is exercised.
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);
      // Map to [-6, 6) — a realistic logit range, and strictly below the 6.5
      // tie value injected next, so the tie triple is guaranteed maximal.
      ((state >> 40) as f32 / (1u64 << 24) as f32 - 0.5) * 12.0
    })
    .collect();
  // Ties: three classes share the maximum, two share another value.
  out[10] = 6.5;
  out[200] = 6.5;
  out[500] = 6.5;
  out[3] = -1.25;
  out[400] = -1.25;
  out
}

/// Reference ranking: full sort by (score desc via total_cmp, index asc) —
/// soundevents' RankedScore contract, spelled naively.
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);
  // Monotonic on a coarse grid.
  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() {
  // The three 6.5-scored classes are the maximum: they must come out first,
  // ordered 10, 200, 500 (soundevents' RankedScore contract).
  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();
  // `usize::MAX` is a natural "give me everything" sentinel: it must saturate
  // like any other oversized `k`, not panic on `BinaryHeap` capacity overflow
  // (`usize::MAX * size_of::<Entry>() > isize::MAX`) before saturation runs.
  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() {
  // Monotonicity: ranking raw logits and mapping sigmoid at extraction must
  // give the SAME order and values as pre-mapping every score (the soundevents
  // trick that avoids a 527-element sort per call).
  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() {
  // /m/09x0r "Speech" is class 0 in the released rated label set, and carries
  // permanent id 3 in `soundevents-dataset`'s ledger. Three distinct numbers
  // name this one class and this pins all three: the model output INDEX (0 —
  // this artifact's label ordering), the AudioSet MID (upstream's string), and
  // the permanent ID (the storage handle). A build that confused them — the
  // exact hazard of the 0.4 `id`/`mid` swap — fails here.
  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);
  // The id resolves back to the row it was read from — the bijection
  // `soundevents-dataset` promises, checked through coremlit's own surface.
  assert_eq!(
    RatedSoundEvent::from_id(p.id()).map(RatedSoundEvent::mid),
    Some("/m/09x0r")
  );
  // The last valid row exists; one past it is the typed defensive error.
  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() {
  // Internal invariant (pub(crate) constructor): a wrong-length vector is a
  // module bug, not a caller error — assert, never a silent truncation.
  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");
  // Read back positionally, unchanged: the constructor is a copy, not a
  // transform — no sigmoid, no renormalization.
  assert_eq!(c.as_slice(), values.as_slice());
}

#[test]
fn try_from_slice_rejects_a_wrong_length_vector_without_panicking() {
  // The caller-reachable counterpart to `Confidences::new`'s internal assert:
  // same mistake, a typed error instead of an unwind.
  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() {
  // `Confidences` documents "finite values in [0, 1]"; the model path gets
  // that from sigmoid, so this path is the only one that can break it.
  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:?}"
    );
  }
  // The boundaries themselves are inside the invariant, and so is -0.0.
  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() {
  // The public hand-build path and the internal model path produce equal
  // values when handed the same numbers — `try_from_slice` is the identity on
  // confidence space, `from_logits` the sigmoid onto it.
  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());
  // Same unbounded-k saturation guarantee as top_k_from_scores: no panic, no
  // abort, exactly NUM_CLASSES back.
  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);
}