model-selection-rs 0.1.0

Cross-validation and model-selection utilities for Rust: stratified / group-aware / time-series splitting, nested CV, and learning & validation curves. Dependency-light, composes with any modeling crate.
Documentation
//! Cross-validate a hand-rolled least-squares model with several metrics at once.
//!
//! Run with:
//!   `cargo run --example cross_validate`
//!   `cargo run --example cross_validate --features parallel`  (same numbers, threaded)

use model_selection_rs::evaluate::{cross_validate, BoxedScorer};
use model_selection_rs::scoring::{MeanAbsoluteError, R2Score, RootMeanSquaredError};
use model_selection_rs::splitters::KFold;
use ndarray::{Array1, Array2};

/// Ordinary least squares y = a*x + b on a single feature.
fn ols_fit(x: &Array2<f64>, y: &Array1<f64>) -> impl Fn(&Array2<f64>) -> Array1<f64> {
    let n = x.nrows() as f64;
    let xs = x.column(0).to_owned();
    let mx = xs.sum() / n;
    let my = y.sum() / n;
    let cov: f64 = xs
        .iter()
        .zip(y.iter())
        .map(|(a, b)| (a - mx) * (b - my))
        .sum();
    let var: f64 = xs.iter().map(|a| (a - mx).powi(2)).sum();
    let slope = cov / var;
    let intercept = my - slope * mx;
    move |xq: &Array2<f64>| xq.column(0).mapv(|v| slope * v + intercept)
}

fn main() {
    // Noisy-ish linear data.
    let x = Array2::from_shape_fn((60, 1), |(i, _)| i as f64 / 10.0);
    let y = x
        .column(0)
        .mapv(|v| 3.0 * v + 2.0 + ((v * 7.0).sin()) * 0.3);

    let kf = KFold::new(5).unwrap().with_shuffle(0);
    let scorers: Vec<BoxedScorer> = vec![
        Box::new(R2Score),
        Box::new(MeanAbsoluteError),
        Box::new(RootMeanSquaredError),
    ];

    let res = cross_validate(&kf, &x, &y, ols_fit, &scorers, true).unwrap();

    println!("5-fold cross-validation of an OLS line:\n");
    for (i, name) in res.scorer_names.iter().enumerate() {
        println!(
            "  {name:>5}: test {:.4} +/- {:.4}   (train {:.4})",
            res.mean_test_score(i),
            res.std_test_score(i),
            res.train_scores.as_ref().unwrap()[i].iter().sum::<f64>() / res.n_splits() as f64,
        );
    }
    println!(
        "\n  total fit time across folds: {:?}",
        res.total_fit_time()
    );

    #[cfg(feature = "parallel")]
    println!("  (folds were fit in parallel via rayon)");
}