use std::collections::HashMap;
use std::hash::Hash;
use ndarray::Array1;
use rand::rngs::StdRng;
use rand::seq::SliceRandom;
use rand::SeedableRng;
use super::shuffle_split::SubsetSize;
use super::CvSplitter;
use crate::error::{ModelSelectionError, Result};
#[derive(Debug, Clone)]
pub struct StratifiedShuffleSplit<L> {
n_splits: usize,
test_size: SubsetSize,
train_size: Option<SubsetSize>,
seed: u64,
labels: Vec<L>,
}
impl<L: Eq + Hash + Clone> StratifiedShuffleSplit<L> {
#[must_use]
pub fn new(n_splits: usize, y: &Array1<L>) -> Self {
Self {
n_splits,
test_size: SubsetSize::Fraction(0.1),
train_size: None,
seed: 0,
labels: y.to_vec(),
}
}
#[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
}
fn class_indices(&self) -> Vec<Vec<usize>> {
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 StratifiedShuffleSplit<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,
});
}
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"),
});
}
if n_train + n_test > n_samples {
return Err(ModelSelectionError::NotEnoughSamples {
needed: n_train + n_test,
got: n_samples,
});
}
let class_indices = self.class_indices();
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 train = Vec::new();
let mut test = Vec::new();
for members in &class_indices {
let n_c = members.len();
let test_c = ((n_test as f64) * (n_c as f64) / (n_samples as f64)).round() as usize;
let train_c =
((n_train as f64) * (n_c as f64) / (n_samples as f64)).round() as usize;
let (test_c, train_c) = if test_c + train_c > n_c {
(test_c.min(n_c), n_c.saturating_sub(test_c).min(train_c))
} else {
(test_c, train_c)
};
let mut shuffled = members.clone();
shuffled.shuffle(&mut rng);
test.extend_from_slice(&shuffled[..test_c]);
train.extend_from_slice(&shuffled[test_c..test_c + train_c]);
}
train.sort_unstable();
test.sort_unstable();
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 preserves_class_proportion_per_split() {
let mut v = vec![0; 80];
v.extend(std::iter::repeat(1).take(20));
let y = Array1::from(v);
let sss = StratifiedShuffleSplit::new(5, &y)
.with_test_size(SubsetSize::Fraction(0.2))
.with_seed(7);
for (_, test) in sss.split(100).unwrap() {
let ones = test.iter().filter(|&&i| y[i] == 1).count();
let frac = ones as f64 / test.len() as f64;
assert!((frac - 0.2).abs() < 0.1, "test class-1 share {frac}");
}
}
#[test]
fn train_and_test_disjoint() {
let y = Array1::from(vec![0, 1, 0, 1, 0, 1, 0, 1, 0, 1]);
let sss = StratifiedShuffleSplit::new(3, &y).with_test_size(SubsetSize::Fraction(0.4));
for (train, test) in sss.split(10).unwrap() {
let tr: HashSet<_> = train.iter().collect();
let te: HashSet<_> = test.iter().collect();
assert!(tr.is_disjoint(&te));
}
}
}