use ndarray::{Array1, Array2, Axis};
use crate::error::Result;
use crate::scoring::Scorer;
use crate::splitters::CvSplitter;
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum TrainSize {
Count(usize),
Fraction(f64),
}
#[derive(Debug, Clone)]
pub struct LearningCurve {
pub train_sizes: Vec<usize>,
pub train_scores: Vec<Vec<f64>>,
pub val_scores: Vec<Vec<f64>>,
}
impl LearningCurve {
#[must_use]
pub fn mean_train_scores(&self) -> Vec<f64> {
self.train_scores.iter().map(|row| mean(row)).collect()
}
#[must_use]
pub fn mean_val_scores(&self) -> Vec<f64> {
self.val_scores.iter().map(|row| mean(row)).collect()
}
}
fn mean(xs: &[f64]) -> f64 {
xs.iter().sum::<f64>() / xs.len() as f64
}
struct Job {
size_idx: usize,
fold_idx: usize,
train_score: f64,
val_score: f64,
}
pub fn learning_curve<S, F, M>(
splitter: &S,
x: &Array2<f64>,
y: &Array1<f64>,
fit_fn: F,
scorer: &(dyn Scorer + Sync),
train_sizes: &[TrainSize],
) -> Result<LearningCurve>
where
S: CvSplitter,
F: Fn(&Array2<f64>, &Array1<f64>) -> M + Sync,
M: Fn(&Array2<f64>) -> Array1<f64>,
{
let splits = splitter.split(x.nrows())?;
let min_train = splits.iter().map(|(tr, _)| tr.len()).min().unwrap_or(0);
let mut abs_sizes: Vec<usize> = train_sizes
.iter()
.map(|ts| match ts {
TrainSize::Count(c) => (*c).clamp(1, min_train.max(1)),
TrainSize::Fraction(f) => {
((f * min_train as f64).round() as usize).clamp(1, min_train.max(1))
}
})
.collect();
abs_sizes.sort_unstable();
abs_sizes.dedup();
let eval = |size_idx: usize, fold_idx: usize| -> Job {
let (train, val) = &splits[fold_idx];
let size = abs_sizes[size_idx];
let sub = &train[..size];
let x_sub = x.select(Axis(0), sub);
let y_sub = y.select(Axis(0), sub);
let x_val = x.select(Axis(0), val);
let y_val = y.select(Axis(0), val);
let model = fit_fn(&x_sub, &y_sub);
let train_score = scorer.score(&y_sub, &model(&x_sub));
let val_score = scorer.score(&y_val, &model(&x_val));
Job {
size_idx,
fold_idx,
train_score,
val_score,
}
};
let coords: Vec<(usize, usize)> = (0..abs_sizes.len())
.flat_map(|s| (0..splits.len()).map(move |f| (s, f)))
.collect();
#[cfg(feature = "parallel")]
let jobs: Vec<Job> = {
use rayon::prelude::*;
coords.par_iter().map(|&(s, f)| eval(s, f)).collect()
};
#[cfg(not(feature = "parallel"))]
let jobs: Vec<Job> = coords.iter().map(|&(s, f)| eval(s, f)).collect();
let mut train_scores = vec![vec![0.0; splits.len()]; abs_sizes.len()];
let mut val_scores = vec![vec![0.0; splits.len()]; abs_sizes.len()];
for job in jobs {
train_scores[job.size_idx][job.fold_idx] = job.train_score;
val_scores[job.size_idx][job.fold_idx] = job.val_score;
}
Ok(LearningCurve {
train_sizes: abs_sizes,
train_scores,
val_scores,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::scoring::R2Score;
use crate::splitters::KFold;
fn ols_fit(x: &Array2<f64>, y: &Array1<f64>) -> impl Fn(&Array2<f64>) -> Array1<f64> {
let n = x.nrows() as f64;
let xs: Vec<f64> = x.column(0).to_vec();
let ys: Vec<f64> = y.to_vec();
let mx = xs.iter().sum::<f64>() / n;
let my = ys.iter().sum::<f64>() / n;
let cov: f64 = xs.iter().zip(&ys).map(|(a, b)| (a - mx) * (b - my)).sum();
let var: f64 = xs.iter().map(|a| (a - mx).powi(2)).sum();
let slope = if var == 0.0 { 0.0 } else { cov / var };
let intercept = my - slope * mx;
move |xq: &Array2<f64>| xq.column(0).mapv(|v| slope * v + intercept)
}
fn constant_zero(_x: &Array2<f64>, _y: &Array1<f64>) -> impl Fn(&Array2<f64>) -> Array1<f64> {
|xq: &Array2<f64>| Array1::zeros(xq.nrows())
}
#[test]
fn shapes_are_correct() {
let x = Array2::from_shape_fn((30, 1), |(i, _)| i as f64);
let y = x.column(0).mapv(|v| 2.0 * v + 1.0);
let kf = KFold::new(3).unwrap();
let lc = learning_curve(
&kf,
&x,
&y,
ols_fit,
&R2Score,
&[
TrainSize::Fraction(0.3),
TrainSize::Fraction(0.6),
TrainSize::Fraction(1.0),
],
)
.unwrap();
assert_eq!(lc.train_sizes.len(), 3);
assert_eq!(lc.train_scores.len(), 3);
assert_eq!(lc.train_scores[0].len(), 3); }
#[test]
fn high_bias_curves_are_both_mediocre() {
let x = Array2::from_shape_fn((40, 1), |(i, _)| i as f64);
let y = x.column(0).mapv(|v| 3.0 * v + 2.0);
let kf = KFold::new(4).unwrap();
let lc = learning_curve(
&kf,
&x,
&y,
constant_zero,
&R2Score,
&[TrainSize::Fraction(0.5), TrainSize::Fraction(1.0)],
)
.unwrap();
let train = lc.mean_train_scores();
let val = lc.mean_val_scores();
for (t, v) in train.iter().zip(&val) {
assert!(*t < 0.5, "train R2 should be poor, got {t}");
assert!(*v < 0.5, "val R2 should be poor, got {v}");
}
}
}