pub fn confusion_matrix(
gt_labels: &[Option<usize>],
dt_labels: &[Option<usize>],
num_classes: usize,
) -> Vec<u64> {
let mut matrix = vec![0u64; (num_classes + 1) * (num_classes + 1)];
accumulate_confusion(&mut matrix, gt_labels, dt_labels, num_classes);
matrix
}
pub fn accumulate_confusion(
matrix: &mut [u64],
gt_labels: &[Option<usize>],
dt_labels: &[Option<usize>],
num_classes: usize,
) {
let k = num_classes + 1;
assert_eq!(
matrix.len(),
k * k,
"accumulate_confusion: matrix must be (num_classes + 1)² = {} elements \
for num_classes = {num_classes} (got {})",
k * k,
matrix.len()
);
assert_eq!(
gt_labels.len(),
dt_labels.len(),
"accumulate_confusion: gt_labels and dt_labels must be parallel arrays \
(got {} vs {})",
gt_labels.len(),
dt_labels.len()
);
for (gl, dl) in gt_labels.iter().zip(dt_labels) {
let row = match gl.filter(|&g| g < num_classes) {
Some(g) => g,
None => num_classes,
};
let col = match dl.filter(|&d| d < num_classes) {
Some(d) => d,
None => num_classes,
};
if row == num_classes && col == num_classes {
continue; }
matrix[row * k + col] += 1;
}
}
pub fn row_normalize(matrix: &[u64], num_classes: usize) -> Vec<f64> {
let k = num_classes + 1;
let mut norm = vec![0.0f64; k * k];
for row in 0..k {
let row_sum: u64 = (0..k).map(|col| matrix[row * k + col]).sum();
if row_sum > 0 {
let denom = row_sum as f64;
for col in 0..k {
norm[row * k + col] = matrix[row * k + col] as f64 / denom;
}
}
}
norm
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
#[test]
fn matched_pairs_match_sklearn() {
#[derive(serde::Deserialize)]
struct Case {
num_classes: usize,
gt: Vec<usize>,
dt: Vec<usize>,
matrix: Vec<Vec<u64>>,
}
let data = include_str!("testdata/confusion_sklearn.json");
let cases: Vec<Case> = serde_json::from_str(data).expect("parse fixture");
assert!(cases.len() > 150, "fixture looks truncated");
for (i, c) in cases.iter().enumerate() {
let gt: Vec<Option<usize>> = c.gt.iter().map(|&v| Some(v)).collect();
let dt: Vec<Option<usize>> = c.dt.iter().map(|&v| Some(v)).collect();
let ours = confusion_matrix(>, &dt, c.num_classes);
let side = c.num_classes + 1;
for g in 0..c.num_classes {
for d in 0..c.num_classes {
assert_eq!(
ours[g * side + d],
c.matrix[g][d],
"case {i}: cell [{g}][{d}] is {} but sklearn says {}",
ours[g * side + d],
c.matrix[g][d]
);
}
}
for k in 0..side {
assert_eq!(
ours[c.num_classes * side + k],
0,
"case {i}: background row"
);
assert_eq!(
ours[k * side + c.num_classes],
0,
"case {i}: background col"
);
}
}
}
#[test]
fn confusion_marginals_account_for_every_record() {
let mut rng = StdRng::seed_from_u64(0xC0F5);
for case in 0..5000 {
let num_classes = rng.random_range(1..=6);
let n = rng.random_range(0..=40);
let side = num_classes + 1;
let label = |rng: &mut StdRng| -> Option<usize> {
match rng.random_range(0..4) {
0 => None,
1 => Some(rng.random_range(num_classes..num_classes + 3)),
_ => Some(rng.random_range(0..num_classes)),
}
};
let gt: Vec<Option<usize>> = (0..n).map(|_| label(&mut rng)).collect();
let dt: Vec<Option<usize>> = (0..n).map(|_| label(&mut rng)).collect();
let m = confusion_matrix(>, &dt, num_classes);
let ctx = format!("case {case}: num_classes={num_classes} n={n}");
assert_eq!(m.len(), side * side, "{ctx}");
let eff = |l: &Option<usize>| l.filter(|&v| v < num_classes);
let counted = gt
.iter()
.zip(&dt)
.filter(|(g, d)| {
!(eff(g).is_none() && eff(d).is_none())
})
.count() as u64;
assert_eq!(
m.iter().sum::<u64>(),
counted,
"{ctx}: grand total disagrees with the number of countable records"
);
for g in 0..side {
let row: u64 = m[g * side..(g + 1) * side].iter().sum();
let want = gt
.iter()
.zip(&dt)
.filter(|(gl, dl)| match eff(gl) {
Some(v) => v == g,
None => g == num_classes && eff(dl).is_some(),
})
.count() as u64;
assert_eq!(row, want, "{ctx}: row {g} sum {row} != {want}");
}
for d in 0..side {
let col: u64 = (0..side).map(|g| m[g * side + d]).sum();
let want = gt
.iter()
.zip(&dt)
.filter(|(gl, dl)| match eff(dl) {
Some(v) => v == d,
None => d == num_classes && eff(gl).is_some(),
})
.count() as u64;
assert_eq!(col, want, "{ctx}: column {d} sum {col} != {want}");
}
if n >= 2 {
let split = rng.random_range(1..n);
let mut batched = vec![0u64; side * side];
accumulate_confusion(&mut batched, >[..split], &dt[..split], num_classes);
accumulate_confusion(&mut batched, >[split..], &dt[split..], num_classes);
assert_eq!(batched, m, "{ctx}: batched at {split} != whole");
}
}
}
#[test]
fn diagonal_counts_correct_predictions() {
let gt = [Some(0), Some(1), Some(2)];
let dt = [Some(0), Some(1), Some(2)];
let m = confusion_matrix(>, &dt, 3);
let k = 4;
for c in 0..3 {
assert_eq!(m[c * k + c], 1);
}
assert_eq!(m.iter().sum::<u64>(), 3);
}
#[test]
fn unmatched_gt_and_dt_land_in_the_background_lane() {
let gt = [Some(1), None];
let dt = [None, Some(2)];
let m = confusion_matrix(>, &dt, 3);
let at = |gt_class: usize, dt_class: usize| m[gt_class * 4 + dt_class];
assert_eq!(at(1, 3), 1, "missed GT -> background column");
assert_eq!(at(3, 2), 1, "spurious DT -> background row");
}
#[test]
fn background_to_background_is_not_counted() {
let m = confusion_matrix(&[None, None], &[None, None], 3);
assert_eq!(m.iter().sum::<u64>(), 0);
}
#[test]
fn out_of_range_partner_does_not_delete_the_record() {
let m = confusion_matrix(&[Some(99), Some(0)], &[Some(0), Some(99)], 3);
let at = |gt_class: usize, dt_class: usize| m[gt_class * 4 + dt_class];
assert_eq!(at(3, 0), 1, "valid DT must survive an invalid GT label");
assert_eq!(at(0, 3), 1, "valid GT must survive an invalid DT label");
assert_eq!(m.iter().sum::<u64>(), 2);
let m = confusion_matrix(&[Some(99)], &[Some(42)], 3);
assert_eq!(m.iter().sum::<u64>(), 0);
let m = confusion_matrix(&[Some(99), None], &[None, Some(42)], 3);
assert_eq!(m.iter().sum::<u64>(), 0);
}
#[test]
fn per_batch_sums_equal_one_whole_call() {
let gt = [Some(0), Some(1), None, Some(2)];
let dt = [Some(1), None, Some(0), Some(2)];
let whole = confusion_matrix(>, &dt, 3);
let a = confusion_matrix(>[..2], &dt[..2], 3);
let b = confusion_matrix(>[2..], &dt[2..], 3);
let summed: Vec<u64> = a.iter().zip(b.iter()).map(|(x, y)| x + y).collect();
assert_eq!(whole, summed, "batching must not change the counts");
}
#[test]
fn row_normalize_makes_rows_sum_to_one() {
let m = vec![3, 1, 0, 0];
let n = row_normalize(&m, 1);
assert!((n[0] - 0.75).abs() < 1e-12);
assert!((n[1] - 0.25).abs() < 1e-12);
assert_eq!(n[2], 0.0);
assert_eq!(n[3], 0.0);
}
#[test]
#[should_panic(expected = "parallel arrays")]
fn mismatched_lengths_panic_instead_of_truncating() {
confusion_matrix(&[Some(0), Some(1), Some(2)], &[Some(0)], 3);
}
#[test]
#[should_panic(expected = "(num_classes + 1)²")]
fn undersized_matrix_panics_instead_of_no_op() {
let mut too_small = vec![0u64; 4]; accumulate_confusion(&mut too_small, &[Some(0)], &[Some(0)], 3);
}
#[test]
#[should_panic(expected = "(num_classes + 1)²")]
fn oversized_matrix_panics_too() {
let mut too_big = vec![0u64; 25]; accumulate_confusion(&mut too_big, &[Some(0)], &[Some(0)], 3);
}
}