diff --git a/crates/ml-alpha/src/gate/auc.rs b/crates/ml-alpha/src/gate/auc.rs new file mode 100644 index 000000000..04c40aa3e --- /dev/null +++ b/crates/ml-alpha/src/gate/auc.rs @@ -0,0 +1,133 @@ +//! Per-horizon AUC computation (Mann-Whitney U statistic). +//! +//! Slow-path CPU operation — one shot per epoch per horizon, on data +//! already downloaded from GPU (probs + labels). Not a hot-path kernel. + +use anyhow::Result; + +pub struct AucInput { + pub probs: Vec, + pub labels: Vec, +} + +/// Computes ROC AUC via the Mann-Whitney U statistic. +/// +/// `labels` entries in {0.0, 1.0}; NaN entries are filtered (mirrors +/// `multi_horizon_labels::generate_labels` masking semantics). +/// Returns 0.5 if either class is empty after filtering. +pub fn compute_auc(input: &AucInput) -> Result { + anyhow::ensure!( + input.probs.len() == input.labels.len(), + "probs/labels size mismatch: {} vs {}", + input.probs.len(), + input.labels.len() + ); + + // Filter out NaN labels (right-edge / tied-price drops). + let paired: Vec<(f32, f32)> = input + .probs + .iter() + .zip(&input.labels) + .filter(|(p, y)| p.is_finite() && y.is_finite() && !y.is_nan()) + .map(|(&p, &y)| (p, y)) + .collect(); + if paired.is_empty() { + return Ok(0.5); + } + + let n_pos = paired.iter().filter(|(_, y)| *y > 0.5).count(); + let n_neg = paired.len() - n_pos; + if n_pos == 0 || n_neg == 0 { + return Ok(0.5); + } + + // Sort ascending by score; positives at higher ranks contribute to AUC. + let mut sorted = paired; + sorted.sort_by(|a, b| { + a.0.partial_cmp(&b.0) + .unwrap_or(std::cmp::Ordering::Equal) + }); + + // Sum of ranks of positives (rank starts at 1 to match Mann-Whitney + // formula). Average ranks for ties. + let mut rank_sum: f64 = 0.0; + let mut i = 0; + while i < sorted.len() { + let mut j = i + 1; + while j < sorted.len() && (sorted[j].0 - sorted[i].0).abs() < f32::EPSILON { + j += 1; + } + // Group [i, j) shares score → assign their average rank. + let avg_rank = (i as f64 + j as f64 + 1.0) / 2.0; + for k in i..j { + if sorted[k].1 > 0.5 { + rank_sum += avg_rank; + } + } + i = j; + } + + let u = rank_sum - (n_pos as f64) * (n_pos as f64 + 1.0) / 2.0; + Ok((u / (n_pos as f64 * n_neg as f64)) as f32) +} + +#[cfg(test)] +mod tests { + use super::*; + use approx::assert_relative_eq; + + #[test] + fn perfect_separation_yields_one() { + let probs = vec![0.1, 0.2, 0.3, 0.4, 0.8, 0.85, 0.9, 0.95]; + let labels = vec![0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0]; + let a = compute_auc(&AucInput { probs, labels }).unwrap(); + assert_relative_eq!(a, 1.0, epsilon = 1e-6); + } + + #[test] + fn perfect_anti_separation_yields_zero() { + let probs = vec![0.1, 0.2, 0.3, 0.4, 0.8, 0.85, 0.9, 0.95]; + let labels = vec![1.0, 1.0, 1.0, 1.0, 0.0, 0.0, 0.0, 0.0]; + let a = compute_auc(&AucInput { probs, labels }).unwrap(); + assert_relative_eq!(a, 0.0, epsilon = 1e-6); + } + + #[test] + fn random_yields_around_half() { + // 1000 random pairs — uniform; AUC should be close to 0.5. + use rand::{Rng, SeedableRng}; + use rand_chacha::ChaCha8Rng; + let mut r = ChaCha8Rng::seed_from_u64(0xA0C_5EED); + let n = 1000usize; + let probs: Vec = (0..n).map(|_| r.gen_range(0.0..1.0)).collect(); + let labels: Vec = (0..n).map(|_| if r.gen_bool(0.5) { 1.0 } else { 0.0 }).collect(); + let a = compute_auc(&AucInput { probs, labels }).unwrap(); + assert!((a - 0.5).abs() < 0.05, "random AUC should be ≈0.5, got {a}"); + } + + #[test] + fn nan_labels_are_filtered() { + let probs = vec![0.1, 0.2, 0.3, 0.4, 0.8, 0.9]; + let labels = vec![f32::NAN, 0.0, f32::NAN, 0.0, 1.0, 1.0]; + let a = compute_auc(&AucInput { probs, labels }).unwrap(); + // Effective pairs: (0.2, 0), (0.4, 0), (0.8, 1), (0.9, 1) → perfect → 1.0 + assert_relative_eq!(a, 1.0, epsilon = 1e-6); + } + + #[test] + fn empty_class_yields_half() { + let probs = vec![0.1, 0.5, 0.9]; + let labels = vec![1.0, 1.0, 1.0]; // no negatives + let a = compute_auc(&AucInput { probs, labels }).unwrap(); + assert_relative_eq!(a, 0.5, epsilon = 1e-6); + } + + #[test] + fn tied_scores_use_average_rank() { + // All probs equal → AUC = 0.5 regardless of labels. + let probs = vec![0.5; 8]; + let labels = vec![1.0, 1.0, 1.0, 1.0, 0.0, 0.0, 0.0, 0.0]; + let a = compute_auc(&AucInput { probs, labels }).unwrap(); + assert_relative_eq!(a, 0.5, epsilon = 1e-6); + } +} diff --git a/crates/ml-alpha/src/gate/mod.rs b/crates/ml-alpha/src/gate/mod.rs new file mode 100644 index 000000000..4fd0a4399 --- /dev/null +++ b/crates/ml-alpha/src/gate/mod.rs @@ -0,0 +1,8 @@ +//! Validation gate — CfC vs Mamba2 multi-horizon AUC comparison. +//! +//! Per the design spec Section 4 gate criterion: CfC AUC ≥ Mamba2 AUC +//! at every horizon (with the 2026-05-16 stacked amendment shifting the +//! framing to "stacked AUC ≥ Mamba2-only baseline AUC"). + +pub mod auc; +// pub mod cfc_vs_mamba2; — Task 17 will land this diff --git a/crates/ml-alpha/src/lib.rs b/crates/ml-alpha/src/lib.rs index dd2c27012..130dd9827 100644 --- a/crates/ml-alpha/src/lib.rs +++ b/crates/ml-alpha/src/lib.rs @@ -27,6 +27,7 @@ // Task 16+: pub mod gate; pub mod cfc; pub mod data; +pub mod gate; pub mod heads; pub mod isv; pub mod pinned;