pub fn argmax(v: &[f32]) -> usize {
let mut best = 0;
let mut bv = f32::NEG_INFINITY;
for (i, &x) in v.iter().enumerate() {
if x > bv {
bv = x;
best = i;
}
}
best
}
pub fn softmax(logits: &[f32]) -> Vec<f32> {
let m = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = logits.iter().map(|&x| (x - m).exp()).collect();
let s: f32 = exps.iter().sum();
if s <= 0.0 {
return vec![1.0 / logits.len() as f32; logits.len()];
}
exps.iter().map(|&e| e / s).collect()
}
pub fn confusion(preds: &[usize], labels: &[usize], k: usize) -> Vec<Vec<u64>> {
let mut m = vec![vec![0u64; k]; k];
for (&p, &t) in preds.iter().zip(labels) {
if t < k && p < k {
m[t][p] += 1;
}
}
m
}
pub fn accuracy(preds: &[usize], labels: &[usize]) -> f64 {
if preds.is_empty() {
return f64::NAN;
}
let correct = preds.iter().zip(labels).filter(|(p, t)| p == t).count();
correct as f64 / preds.len() as f64
}
pub fn balanced_accuracy(preds: &[usize], labels: &[usize], k: usize) -> f64 {
let cm = confusion(preds, labels, k);
let mut recalls = Vec::new();
for c in 0..k {
let support: u64 = cm[c].iter().sum();
if support > 0 {
recalls.push(cm[c][c] as f64 / support as f64);
}
}
if recalls.is_empty() {
return f64::NAN;
}
recalls.iter().sum::<f64>() / recalls.len() as f64
}
fn average_ranks(scores: &[f64]) -> Vec<f64> {
let n = scores.len();
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a, &b| scores[a].partial_cmp(&scores[b]).unwrap());
let mut ranks = vec![0f64; n];
let mut i = 0;
while i < n {
let mut j = i + 1;
while j < n && scores[order[j]] == scores[order[i]] {
j += 1;
}
let mean_rank = ((i + 1 + j) as f64) / 2.0;
for &idx in &order[i..j] {
ranks[idx] = mean_rank;
}
i = j;
}
ranks
}
pub fn auroc_binary(scores: &[f64], labels: &[u8]) -> f64 {
let n_pos = labels.iter().filter(|&&l| l == 1).count();
let n_neg = labels.len() - n_pos;
if n_pos == 0 || n_neg == 0 {
return f64::NAN;
}
let ranks = average_ranks(scores);
let sum_pos_ranks: f64 = ranks
.iter()
.zip(labels)
.filter(|(_, &l)| l == 1)
.map(|(r, _)| *r)
.sum();
let u = sum_pos_ranks - (n_pos as f64) * (n_pos as f64 + 1.0) / 2.0;
u / (n_pos as f64 * n_neg as f64)
}
pub fn auroc_macro_ovr(probs: &[Vec<f32>], labels: &[usize], k: usize) -> f64 {
let mut per_class = Vec::new();
for c in 0..k {
let present = labels.contains(&c);
let absent = labels.iter().any(|&l| l != c);
if !(present && absent) {
continue;
}
let scores: Vec<f64> = probs.iter().map(|p| p[c] as f64).collect();
let bin: Vec<u8> = labels.iter().map(|&l| (l == c) as u8).collect();
let a = auroc_binary(&scores, &bin);
if a.is_finite() {
per_class.push(a);
}
}
if per_class.is_empty() {
return f64::NAN;
}
per_class.iter().sum::<f64>() / per_class.len() as f64
}
pub fn average_precision(scores: &[f64], labels: &[u8]) -> f64 {
let n_pos = labels.iter().filter(|&&l| l == 1).count();
if n_pos == 0 {
return f64::NAN;
}
let mut order: Vec<usize> = (0..scores.len()).collect();
order.sort_by(|&a, &b| scores[b].partial_cmp(&scores[a]).unwrap());
let mut ap = 0.0;
let mut tp = 0u64;
let mut fp = 0u64;
let mut prev_recall = 0.0;
let mut i = 0;
while i < order.len() {
let s = scores[order[i]];
let mut j = i;
while j < order.len() && scores[order[j]] == s {
if labels[order[j]] == 1 {
tp += 1;
} else {
fp += 1;
}
j += 1;
}
let recall = tp as f64 / n_pos as f64;
let precision = if tp + fp > 0 { tp as f64 / (tp + fp) as f64 } else { 1.0 };
ap += (recall - prev_recall) * precision;
prev_recall = recall;
i = j;
}
ap
}
#[cfg(test)]
mod tests {
use super::*;
fn close(a: f64, b: f64) -> bool {
(a - b).abs() < 1e-9
}
#[test]
fn auroc_matches_sklearn() {
let s = [0.1, 0.4, 0.35, 0.8];
let y = [0u8, 0, 1, 1];
assert!(close(auroc_binary(&s, &y), 0.75), "{}", auroc_binary(&s, &y));
}
#[test]
fn auroc_ties() {
assert!(close(auroc_binary(&[0.5, 0.5], &[0, 1]), 0.5));
}
#[test]
fn average_precision_matches_sklearn() {
let s = [0.1, 0.4, 0.35, 0.8];
let y = [0u8, 0, 1, 1];
let ap = average_precision(&s, &y);
assert!((ap - 0.8333333333).abs() < 1e-6, "{ap}");
}
#[test]
fn balanced_accuracy_basic() {
let preds = [0usize, 0, 1, 0];
let labels = [0usize, 0, 1, 1];
assert!(close(balanced_accuracy(&preds, &labels, 2), 0.75));
}
#[test]
fn balanced_accuracy_imbalanced_beats_plain() {
let mut preds = vec![0usize; 10];
let mut labels = vec![0usize; 10];
labels[9] = 1; let _ = &mut preds;
assert!(close(accuracy(&preds, &labels), 0.9));
assert!(close(balanced_accuracy(&preds, &labels, 2), 0.5));
}
#[test]
fn macro_auroc_perfect() {
let probs = vec![
vec![0.9f32, 0.05, 0.05],
vec![0.05, 0.9, 0.05],
vec![0.05, 0.05, 0.9],
vec![0.8, 0.1, 0.1],
];
let labels = [0usize, 1, 2, 0];
assert!(close(auroc_macro_ovr(&probs, &labels, 3), 1.0));
}
}