use rand::rngs::StdRng;
use rand::seq::SliceRandom;
use rand::SeedableRng;
use super::CvSplitter;
use crate::error::{ModelSelectionError, Result};
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum SubsetSize {
Count(usize),
Fraction(f64),
}
impl SubsetSize {
pub(crate) fn resolve(self, n_samples: usize) -> usize {
match self {
SubsetSize::Count(c) => c,
SubsetSize::Fraction(f) => (f * n_samples as f64).round() as usize,
}
}
}
#[derive(Debug, Clone)]
pub struct ShuffleSplit {
n_splits: usize,
test_size: SubsetSize,
train_size: Option<SubsetSize>,
seed: u64,
}
impl ShuffleSplit {
#[must_use]
pub fn new(n_splits: usize) -> Self {
Self {
n_splits,
test_size: SubsetSize::Fraction(0.1),
train_size: None,
seed: 0,
}
}
#[must_use]
pub fn with_test_size(mut self, test_size: SubsetSize) -> Self {
self.test_size = test_size;
self
}
#[must_use]
pub fn with_train_size(mut self, train_size: SubsetSize) -> Self {
self.train_size = Some(train_size);
self
}
#[must_use]
pub fn with_seed(mut self, seed: u64) -> Self {
self.seed = seed;
self
}
pub(crate) fn resolve_sizes(&self, n_samples: usize) -> Result<(usize, usize)> {
let n_test = self.test_size.resolve(n_samples);
let n_train = match self.train_size {
Some(ts) => ts.resolve(n_samples),
None => n_samples.saturating_sub(n_test),
};
if n_test == 0 || n_train == 0 {
return Err(ModelSelectionError::InvalidSplitCount {
msg: format!(
"resolved train={n_train}, test={n_test}; both must be >= 1 \
(n_samples={n_samples})"
),
});
}
if n_train + n_test > n_samples {
return Err(ModelSelectionError::NotEnoughSamples {
needed: n_train + n_test,
got: n_samples,
});
}
Ok((n_train, n_test))
}
}
impl CvSplitter for ShuffleSplit {
fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
let (n_train, n_test) = self.resolve_sizes(n_samples)?;
let mut splits = Vec::with_capacity(self.n_splits);
for i in 0..self.n_splits {
let mut rng = StdRng::seed_from_u64(self.seed.wrapping_add(i as u64));
let mut indices: Vec<usize> = (0..n_samples).collect();
indices.shuffle(&mut rng);
let test: Vec<usize> = indices[..n_test].to_vec();
let train: Vec<usize> = indices[n_test..n_test + n_train].to_vec();
splits.push((train, test));
}
Ok(splits)
}
fn n_splits(&self) -> usize {
self.n_splits
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashSet;
#[test]
fn sizes_honoured_and_disjoint() {
let ss = ShuffleSplit::new(4)
.with_test_size(SubsetSize::Count(5))
.with_train_size(SubsetSize::Count(10))
.with_seed(3);
for (train, test) in ss.split(30).unwrap() {
assert_eq!(train.len(), 10);
assert_eq!(test.len(), 5);
let tr: HashSet<_> = train.iter().collect();
let te: HashSet<_> = test.iter().collect();
assert!(tr.is_disjoint(&te));
}
}
#[test]
fn fraction_test_size() {
let ss = ShuffleSplit::new(2).with_test_size(SubsetSize::Fraction(0.2));
let splits = ss.split(50).unwrap();
assert!(splits.iter().all(|(_, te)| te.len() == 10));
}
#[test]
fn deterministic_for_seed() {
let a = ShuffleSplit::new(3).with_seed(11).split(20).unwrap();
let b = ShuffleSplit::new(3).with_seed(11).split(20).unwrap();
assert_eq!(a, b);
}
#[test]
fn errors_when_sizes_dont_fit() {
let ss = ShuffleSplit::new(2)
.with_test_size(SubsetSize::Count(20))
.with_train_size(SubsetSize::Count(20));
assert!(matches!(
ss.split(30),
Err(ModelSelectionError::NotEnoughSamples { .. })
));
}
}