use crate::error::Result;
mod group_kfold;
mod kfold;
mod leave_one_out;
mod repeated;
mod shuffle_split;
mod stratified_group_kfold;
mod stratified_kfold;
mod stratified_shuffle_split;
mod time_series_split;
pub use group_kfold::GroupKFold;
pub use kfold::KFold;
pub use leave_one_out::LeaveOneOut;
pub use repeated::{RepeatedKFold, RepeatedStratifiedKFold};
pub use shuffle_split::{ShuffleSplit, SubsetSize};
pub use stratified_group_kfold::StratifiedGroupKFold;
pub use stratified_kfold::StratifiedKFold;
pub use stratified_shuffle_split::StratifiedShuffleSplit;
pub use time_series_split::TimeSeriesSplit;
pub trait CvSplitter {
fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>>;
fn n_splits(&self) -> usize;
}
pub(crate) fn fold_bounds(n_samples: usize, k: usize) -> Vec<(usize, usize)> {
let base = n_samples / k;
let remainder = n_samples % k;
let mut bounds = Vec::with_capacity(k);
let mut start = 0;
for fold in 0..k {
let size = base + usize::from(fold < remainder);
bounds.push((start, start + size));
start += size;
}
bounds
}
#[cfg(test)]
mod tests {
use super::*;
struct PassthroughSplitter;
impl CvSplitter for PassthroughSplitter {
fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
Ok(vec![((0..n_samples).collect(), Vec::new())])
}
fn n_splits(&self) -> usize {
1
}
}
#[test]
fn passthrough_works_as_trait_object() {
let s: &dyn CvSplitter = &PassthroughSplitter;
let splits = s.split(5).unwrap();
assert_eq!(s.n_splits(), 1);
assert_eq!(splits[0].0, vec![0, 1, 2, 3, 4]);
assert!(splits[0].1.is_empty());
}
#[test]
fn fold_bounds_distributes_remainder_to_leading_folds() {
assert_eq!(fold_bounds(7, 3), vec![(0, 3), (3, 5), (5, 7)]);
assert_eq!(fold_bounds(6, 3), vec![(0, 2), (2, 4), (4, 6)]);
}
}