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
//! Repeated K-fold variants for more robust performance estimates.

use std::hash::Hash;

use ndarray::Array1;

use super::stratified_kfold::{stratified_test_folds, test_folds_to_splits};
use super::{CvSplitter, KFold};
use crate::error::{ModelSelectionError, Result};

/// Repeated plain K-fold.
///
/// Runs [`KFold`] `n_repeats` times, each with a different shuffle seed, and
/// concatenates the resulting splits. Repeated evaluation smooths out the
/// variance that a single random fold assignment introduces, giving a more
/// robust performance estimate. Total splits = `n_repeats * n_splits`.
///
/// ```
/// use model_selection_rs::splitters::{CvSplitter, RepeatedKFold};
///
/// let rkf = RepeatedKFold::new(5, 3, 0).unwrap();
/// assert_eq!(rkf.n_splits(), 15);
/// assert_eq!(rkf.split(50).unwrap().len(), 15);
/// ```
#[derive(Debug, Clone)]
pub struct RepeatedKFold {
    n_splits: usize,
    n_repeats: usize,
    base_seed: u64,
}

impl RepeatedKFold {
    /// Create a `RepeatedKFold` with `n_splits` folds repeated `n_repeats` times.
    ///
    /// # Errors
    ///
    /// Returns [`ModelSelectionError::InvalidSplitCount`] if `n_splits < 2` or
    /// `n_repeats < 1`.
    pub fn new(n_splits: usize, n_repeats: usize, base_seed: u64) -> Result<Self> {
        if n_splits < 2 {
            return Err(ModelSelectionError::InvalidSplitCount {
                msg: format!("n_splits must be >= 2, got {n_splits}"),
            });
        }
        if n_repeats < 1 {
            return Err(ModelSelectionError::InvalidSplitCount {
                msg: format!("n_repeats must be >= 1, got {n_repeats}"),
            });
        }
        Ok(Self {
            n_splits,
            n_repeats,
            base_seed,
        })
    }
}

impl CvSplitter for RepeatedKFold {
    fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
        let mut all = Vec::with_capacity(self.n_splits * self.n_repeats);
        for r in 0..self.n_repeats {
            let kf = KFold::new(self.n_splits)?.with_shuffle(self.base_seed.wrapping_add(r as u64));
            all.extend(kf.split(n_samples)?);
        }
        Ok(all)
    }

    fn n_splits(&self) -> usize {
        self.n_splits * self.n_repeats
    }
}

/// Repeated stratified K-fold.
///
/// The stratified analogue of [`RepeatedKFold`]: runs stratified K-fold
/// `n_repeats` times with different seeds, preserving class proportions in every
/// fold of every repeat. Total splits = `n_repeats * n_splits`.
///
/// ```
/// use ndarray::Array1;
/// use model_selection_rs::splitters::{CvSplitter, RepeatedStratifiedKFold};
///
/// let y = Array1::from(vec![0, 1, 0, 1, 0, 1, 0, 1, 0, 1]);
/// let rskf = RepeatedStratifiedKFold::new(2, 3, 0, &y).unwrap();
/// assert_eq!(rskf.n_splits(), 6);
/// ```
#[derive(Debug, Clone)]
pub struct RepeatedStratifiedKFold<L> {
    n_splits: usize,
    n_repeats: usize,
    base_seed: u64,
    labels: Vec<L>,
}

impl<L: Eq + Hash + Clone> RepeatedStratifiedKFold<L> {
    /// Create a `RepeatedStratifiedKFold` over labels `y`.
    ///
    /// # Errors
    ///
    /// Returns [`ModelSelectionError::InvalidSplitCount`] if `n_splits < 2` or
    /// `n_repeats < 1`.
    pub fn new(n_splits: usize, n_repeats: usize, base_seed: u64, y: &Array1<L>) -> Result<Self> {
        if n_splits < 2 {
            return Err(ModelSelectionError::InvalidSplitCount {
                msg: format!("n_splits must be >= 2, got {n_splits}"),
            });
        }
        if n_repeats < 1 {
            return Err(ModelSelectionError::InvalidSplitCount {
                msg: format!("n_repeats must be >= 1, got {n_repeats}"),
            });
        }
        Ok(Self {
            n_splits,
            n_repeats,
            base_seed,
            labels: y.to_vec(),
        })
    }

    fn class_indices(&self) -> Vec<Vec<usize>> {
        use std::collections::HashMap;
        let mut order: Vec<L> = Vec::new();
        let mut map: HashMap<L, Vec<usize>> = HashMap::new();
        for (i, label) in self.labels.iter().enumerate() {
            map.entry(label.clone()).or_insert_with(|| {
                order.push(label.clone());
                Vec::new()
            });
            map.get_mut(label).unwrap().push(i);
        }
        order.into_iter().map(|c| map.remove(&c).unwrap()).collect()
    }
}

impl<L: Eq + Hash + Clone> CvSplitter for RepeatedStratifiedKFold<L> {
    fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
        if n_samples != self.labels.len() {
            return Err(ModelSelectionError::ShapeMismatch {
                expected: self.labels.len(),
                got: n_samples,
            });
        }
        if self.n_splits > n_samples {
            return Err(ModelSelectionError::NotEnoughSamples {
                needed: self.n_splits,
                got: n_samples,
            });
        }
        let class_indices = self.class_indices();
        let mut all = Vec::with_capacity(self.n_splits * self.n_repeats);
        for r in 0..self.n_repeats {
            let test_folds = stratified_test_folds(
                &class_indices,
                self.n_splits,
                true,
                self.base_seed.wrapping_add(r as u64),
            );
            all.extend(test_folds_to_splits(test_folds, n_samples));
        }
        Ok(all)
    }

    fn n_splits(&self) -> usize {
        self.n_splits * self.n_repeats
    }
}

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

    #[test]
    fn repeated_kfold_yields_product_of_splits() {
        let rkf = RepeatedKFold::new(5, 3, 42).unwrap();
        assert_eq!(rkf.n_splits(), 15);
        assert_eq!(rkf.split(50).unwrap().len(), 15);
    }

    #[test]
    fn repeats_differ_from_each_other() {
        let rkf = RepeatedKFold::new(2, 2, 1).unwrap();
        let splits = rkf.split(20).unwrap();
        // The first repeat's first-fold test set should differ from the second
        // repeat's first-fold test set (different seed).
        assert_ne!(splits[0].1, splits[2].1);
    }

    #[test]
    fn repeated_stratified_counts() {
        let y = Array1::from(vec![0, 1, 0, 1, 0, 1, 0, 1, 0, 1]);
        let rskf = RepeatedStratifiedKFold::new(2, 4, 0, &y).unwrap();
        assert_eq!(rskf.n_splits(), 8);
        assert_eq!(rskf.split(10).unwrap().len(), 8);
    }
}