use ndarray::{Array1, Array2, Axis};
use crate::error::Result;
use crate::scoring::Scorer;
use crate::splitters::CvSplitter;
#[derive(Debug, Clone)]
pub struct NestedCvResults<P> {
pub outer_scores: Vec<f64>,
pub selected_params: Vec<P>,
}
impl<P> NestedCvResults<P> {
#[must_use]
pub fn mean_score(&self) -> f64 {
self.outer_scores.iter().sum::<f64>() / self.outer_scores.len() as f64
}
#[must_use]
pub fn std_score(&self) -> f64 {
let m = self.mean_score();
let n = self.outer_scores.len();
if n < 2 {
return 0.0;
}
(self
.outer_scores
.iter()
.map(|s| (s - m).powi(2))
.sum::<f64>()
/ n as f64)
.sqrt()
}
}
pub fn nested_cross_validate<OS, IS, P, Tune, Fit, M>(
outer: &OS,
inner: &IS,
x: &Array2<f64>,
y: &Array1<f64>,
tune_fn: Tune,
fit_fn: Fit,
scorer: &(dyn Scorer + Sync),
) -> Result<NestedCvResults<P>>
where
OS: CvSplitter,
IS: CvSplitter + Sync,
Tune: Fn(&Array2<f64>, &Array1<f64>, &IS) -> P + Sync,
Fit: Fn(&P, &Array2<f64>, &Array1<f64>) -> M + Sync,
M: Fn(&Array2<f64>) -> Array1<f64>,
P: Send,
{
let outer_splits = outer.split(x.nrows())?;
let eval = |train: &[usize], test: &[usize]| -> (f64, P) {
let x_train = x.select(Axis(0), train);
let y_train = y.select(Axis(0), train);
let x_test = x.select(Axis(0), test);
let y_test = y.select(Axis(0), test);
let best = tune_fn(&x_train, &y_train, inner);
let model = fit_fn(&best, &x_train, &y_train);
let score = scorer.score(&y_test, &model(&x_test));
(score, best)
};
#[cfg(feature = "parallel")]
let results: Vec<(f64, P)> = {
use rayon::prelude::*;
outer_splits
.par_iter()
.map(|(tr, te)| eval(tr, te))
.collect()
};
#[cfg(not(feature = "parallel"))]
let results: Vec<(f64, P)> = outer_splits.iter().map(|(tr, te)| eval(tr, te)).collect();
let mut outer_scores = Vec::with_capacity(results.len());
let mut selected_params = Vec::with_capacity(results.len());
for (score, param) in results {
outer_scores.push(score);
selected_params.push(param);
}
Ok(NestedCvResults {
outer_scores,
selected_params,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::evaluate::{cross_validate, BoxedScorer};
use crate::scoring::MeanSquaredError;
use crate::splitters::KFold;
fn ridge_fit(
lambda: &f64,
x: &Array2<f64>,
y: &Array1<f64>,
) -> impl Fn(&Array2<f64>) -> Array1<f64> {
let lambda = *lambda;
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 sxx: f64 = xs.iter().map(|v| (v - mx).powi(2)).sum();
let sxy: f64 = xs.iter().zip(&ys).map(|(a, b)| (a - mx) * (b - my)).sum();
let slope = sxy / (sxx + lambda);
let intercept = my - slope * mx;
move |xq: &Array2<f64>| xq.column(0).mapv(|v| slope * v + intercept)
}
fn tune_lambda(x: &Array2<f64>, y: &Array1<f64>, inner: &KFold) -> f64 {
let candidates = [0.0f64, 0.1, 1.0, 10.0, 100.0];
let mut best = (f64::INFINITY, 0.0);
for &lam in &candidates {
let scorers: Vec<BoxedScorer> = vec![Box::new(MeanSquaredError)];
let res = cross_validate(
inner,
x,
y,
move |xt, yt| ridge_fit(&lam, xt, yt),
&scorers,
false,
)
.unwrap();
let mse = res.mean_test_score(0);
if mse < best.0 {
best = (mse, lam);
}
}
best.1
}
#[test]
fn recovers_low_regularization_on_clean_linear_data() {
let x = Array2::from_shape_fn((60, 1), |(i, _)| i as f64 / 10.0);
let y = x.column(0).mapv(|v| 2.0 * v + 1.0);
let outer = KFold::new(5).unwrap().with_shuffle(0);
let inner = KFold::new(4).unwrap().with_shuffle(1);
let res = nested_cross_validate(
&outer,
&inner,
&x,
&y,
tune_lambda,
ridge_fit,
&MeanSquaredError,
)
.unwrap();
assert_eq!(res.outer_scores.len(), 5);
assert!(
res.selected_params.iter().all(|&l| l <= 1.0),
"selected lambdas: {:?}",
res.selected_params
);
assert!(res.mean_score() < 1.0, "mean MSE {}", res.mean_score());
}
}