use std::collections::{HashMap, HashSet};
use crate::linalg::basic::arrays::ArrayView1;
use crate::numbers::realnum::RealNumber;
pub(crate) struct ConfusionCounts {
classes_set: HashSet<u64>,
predicted: HashMap<u64, usize>,
support: HashMap<u64, usize>,
tp_map: HashMap<u64, usize>,
}
impl ConfusionCounts {
pub(crate) fn new<T: RealNumber>(
y_true: &dyn ArrayView1<T>,
y_pred: &dyn ArrayView1<T>,
) -> Self {
let n = y_true.shape();
let mut classes_set: HashSet<u64> = HashSet::new();
let mut predicted: HashMap<u64, usize> = HashMap::new();
let mut support: HashMap<u64, usize> = HashMap::new();
let mut tp_map: HashMap<u64, usize> = HashMap::new();
for i in 0..n {
let t_bits = y_true.get(i).to_f64_bits();
classes_set.insert(t_bits);
*support.entry(t_bits).or_insert(0) += 1;
*predicted.entry(y_pred.get(i).to_f64_bits()).or_insert(0) += 1;
if *y_true.get(i) == *y_pred.get(i) {
*tp_map.entry(t_bits).or_insert(0) += 1;
}
}
Self {
classes_set,
predicted,
support,
tp_map,
}
}
pub(crate) fn classes_set(&self) -> &HashSet<u64> {
&self.classes_set
}
pub(crate) fn predicted(&self, bits: u64) -> usize {
*self.predicted.get(&bits).unwrap_or(&0)
}
pub(crate) fn support(&self, bits: u64) -> usize {
*self.support.get(&bits).unwrap_or(&0)
}
pub(crate) fn tp(&self, bits: u64) -> usize {
*self.tp_map.get(&bits).unwrap_or(&0)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn bits_of(v: f64) -> u64 {
v.to_f64_bits()
}
#[test]
fn confusion_counts_binary_basic() {
let y_true: Vec<f64> = vec![0., 1., 1., 0.];
let y_pred: Vec<f64> = vec![0., 0., 1., 1.];
let counts = ConfusionCounts::new(&y_true, &y_pred);
assert_eq!(counts.classes_set().len(), 2);
assert!(counts.classes_set().contains(&bits_of(0.0)));
assert!(counts.classes_set().contains(&bits_of(1.0)));
assert_eq!(counts.support(bits_of(0.0)), 2);
assert_eq!(counts.predicted(bits_of(0.0)), 2);
assert_eq!(counts.tp(bits_of(0.0)), 1);
assert_eq!(counts.support(bits_of(1.0)), 2);
assert_eq!(counts.predicted(bits_of(1.0)), 2);
assert_eq!(counts.tp(bits_of(1.0)), 1);
}
#[test]
fn confusion_counts_multiclass() {
let y_true: Vec<f64> = vec![0., 0., 1., 2., 2., 2.];
let y_pred: Vec<f64> = vec![0., 1., 1., 2., 0., 2.];
let counts = ConfusionCounts::new(&y_true, &y_pred);
assert_eq!(counts.classes_set().len(), 3);
assert_eq!(counts.support(bits_of(0.0)), 2);
assert_eq!(counts.predicted(bits_of(0.0)), 2);
assert_eq!(counts.tp(bits_of(0.0)), 1);
assert_eq!(counts.support(bits_of(1.0)), 1);
assert_eq!(counts.predicted(bits_of(1.0)), 2);
assert_eq!(counts.tp(bits_of(1.0)), 1);
assert_eq!(counts.support(bits_of(2.0)), 3);
assert_eq!(counts.predicted(bits_of(2.0)), 2);
assert_eq!(counts.tp(bits_of(2.0)), 2);
}
#[test]
fn confusion_counts_spurious_predicted_label() {
let y_true: Vec<f64> = vec![0., 0., 1., 1.];
let y_pred: Vec<f64> = vec![0., 2., 1., 1.];
let counts = ConfusionCounts::new(&y_true, &y_pred);
assert_eq!(counts.classes_set().len(), 2);
assert!(!counts.classes_set().contains(&bits_of(2.0)));
assert_eq!(counts.predicted(bits_of(2.0)), 1);
assert_eq!(counts.support(bits_of(2.0)), 0);
assert_eq!(counts.tp(bits_of(2.0)), 0);
}
#[test]
fn confusion_counts_empty_input() {
let y_true: Vec<f64> = vec![];
let y_pred: Vec<f64> = vec![];
let counts = ConfusionCounts::new(&y_true, &y_pred);
assert!(counts.classes_set().is_empty());
assert_eq!(counts.predicted(bits_of(0.0)), 0);
assert_eq!(counts.support(bits_of(0.0)), 0);
assert_eq!(counts.tp(bits_of(0.0)), 0);
}
#[test]
fn confusion_counts_perfect_predictions() {
let y_true: Vec<f64> = vec![0., 1., 2., 0., 1.];
let counts = ConfusionCounts::new(&y_true, &y_true);
for &bits in counts.classes_set() {
assert_eq!(counts.tp(bits), counts.support(bits));
assert_eq!(counts.tp(bits), counts.predicted(bits));
}
}
}