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};
#[derive(Debug, Clone)]
pub struct RepeatedKFold {
n_splits: usize,
n_repeats: usize,
base_seed: u64,
}
impl RepeatedKFold {
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
}
}
#[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> {
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();
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);
}
}