lumamba 0.0.1

LuMamba EEG foundation model — inference in Rust on the RLX runtime
Documentation
// lumamba-rs — LuMamba EEG foundation-model inference on the RLX runtime.
// Copyright (C) 2026 Nataliya Kosmyna.
// SPDX-License-Identifier: GPL-3.0-only

//! Classification metrics for downstream evaluation.
//!
//! Implementations match scikit-learn conventions so results line up with the
//! numbers reported in the LuMamba paper:
//! * [`balanced_accuracy`] — mean per-class recall.
//! * [`auroc_binary`] / [`auroc_macro_ovr`] — tie-aware Mann–Whitney AUROC.
//! * [`average_precision`] — step-wise AP (`Σ (Rₙ − Rₙ₋₁)·Pₙ`), as in
//!   `sklearn.metrics.average_precision_score`.

/// Argmax of a logit/probability vector.
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
}

/// Numerically stable softmax.
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()
}

/// `k × k` confusion matrix, `conf[true][pred]`.
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
}

/// Plain accuracy.
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
}

/// Balanced accuracy = mean over classes of recall (TPᵢ / support_i).
/// Classes with zero support are skipped (sklearn behaviour).
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
}

/// Average ranks (1-based) of `scores`, ties share the mean rank.
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;
        }
        // positions i..j tie; 1-based ranks are (i+1..=j), mean = (i+1+j)/2
        let mean_rank = ((i + 1 + j) as f64) / 2.0;
        for &idx in &order[i..j] {
            ranks[idx] = mean_rank;
        }
        i = j;
    }
    ranks
}

/// Binary AUROC via the tie-aware Mann–Whitney U statistic. `labels` are
/// 0/1 with `1` the positive class; `scores` is the positive-class score.
/// Returns `NaN` if only one class is present.
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)
}

/// Macro one-vs-rest AUROC over `k` classes. `probs[i]` is the length-`k`
/// probability vector for sample `i`. Classes absent from `labels` are skipped.
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
}

/// Average precision (area under the precision–recall curve), computed as
/// `Σ (Rₙ − Rₙ₋₁)·Pₙ` over score thresholds — matches
/// `sklearn.metrics.average_precision_score`.
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;
    }
    // Sort by descending score.
    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() {
        // Process all samples sharing this score together (threshold step).
        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() {
        // sklearn.roc_auc_score([0,0,1,1], [0.1,0.4,0.35,0.8]) = 0.75
        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() {
        // Tied scores between a pos and neg → 0.5 contribution.
        // sklearn.roc_auc_score([0,1], [0.5,0.5]) = 0.5
        assert!(close(auroc_binary(&[0.5, 0.5], &[0, 1]), 0.5));
    }

    #[test]
    fn average_precision_matches_sklearn() {
        // sklearn.average_precision_score([0,0,1,1], [0.1,0.4,0.35,0.8]) = 0.8333…
        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() {
        // 2 classes; class 0: 2/2 recall, class 1: 1/2 recall → 0.75
        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() {
        // 9 negatives all correct, 1 positive misclassified.
        let mut preds = vec![0usize; 10];
        let mut labels = vec![0usize; 10];
        labels[9] = 1; // one positive, predicted 0
        let _ = &mut preds;
        // plain acc = 0.9, balanced = (1.0 + 0.0)/2 = 0.5
        assert!(close(accuracy(&preds, &labels), 0.9));
        assert!(close(balanced_accuracy(&preds, &labels, 2), 0.5));
    }

    #[test]
    fn macro_auroc_perfect() {
        // 3-class, perfectly separable probs → 1.0
        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));
    }
}