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
//! Leave-one-out cross-validation.

use super::CvSplitter;
use crate::error::{ModelSelectionError, Result};

/// Leave-one-out cross-validation (LOO).
///
/// Each sample serves as the sole test point exactly once, with every other
/// sample forming the training set — equivalent to [`KFold`](super::KFold) with
/// `n_splits == n_samples`. Included for API completeness.
///
/// **Cost warning:** LOO fits the model `n_samples` times. On anything but small
/// datasets this is very expensive; prefer a modest `KFold` (5 or 10) unless you
/// specifically need LOO's near-unbiased (but high-variance) estimate.
///
/// ```
/// use model_selection_rs::splitters::{CvSplitter, LeaveOneOut};
///
/// let loo = LeaveOneOut;
/// let splits = loo.split(4).unwrap();
/// assert_eq!(splits.len(), 4);
/// assert_eq!(splits[0], (vec![1, 2, 3], vec![0]));
/// ```
#[derive(Debug, Clone, Copy, Default)]
pub struct LeaveOneOut;

impl CvSplitter for LeaveOneOut {
    fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
        if n_samples < 2 {
            return Err(ModelSelectionError::NotEnoughSamples {
                needed: 2,
                got: n_samples,
            });
        }
        let splits = (0..n_samples)
            .map(|test_idx| {
                let train: Vec<usize> = (0..n_samples).filter(|&i| i != test_idx).collect();
                (train, vec![test_idx])
            })
            .collect();
        Ok(splits)
    }

    fn n_splits(&self) -> usize {
        // Unknown without n_samples; LOO's split count equals n_samples, which
        // is only known at `split` time. Report 0 as "depends on data".
        0
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn each_sample_is_test_once() {
        let splits = LeaveOneOut.split(5).unwrap();
        assert_eq!(splits.len(), 5);
        for (i, (train, test)) in splits.iter().enumerate() {
            assert_eq!(test, &vec![i]);
            assert_eq!(train.len(), 4);
            assert!(!train.contains(&i));
        }
    }

    #[test]
    fn errors_on_tiny_input() {
        assert!(LeaveOneOut.split(1).is_err());
    }
}