Skip to main content

model_selection_rs/splitters/
stratified_shuffle_split.rs

1//! Class-proportion-preserving random-permutation splitting.
2
3use std::collections::HashMap;
4use std::hash::Hash;
5
6use ndarray::Array1;
7use rand::rngs::StdRng;
8use rand::seq::SliceRandom;
9use rand::SeedableRng;
10
11use super::shuffle_split::SubsetSize;
12use super::CvSplitter;
13use crate::error::{ModelSelectionError, Result};
14
15/// Stratified random-permutation cross-validation.
16///
17/// Like [`ShuffleSplit`](super::ShuffleSplit), but every split keeps
18/// approximately the dataset's class proportions in both its train and test
19/// subsets. Class labels are supplied at construction (generic
20/// `L: Eq + Hash + Clone`) and stored, so the splitter still satisfies the plain
21/// [`CvSplitter`](super::CvSplitter) interface.
22///
23/// Per-class subset sizes are allocated proportionally and rounded, so realised
24/// sizes may differ from the requested totals by a sample or two — the guarantee
25/// is proportional balance, not an exact global count.
26///
27/// ```
28/// use ndarray::Array1;
29/// use model_selection_rs::splitters::{CvSplitter, StratifiedShuffleSplit, SubsetSize};
30///
31/// let mut v = vec![0; 80];
32/// v.extend(std::iter::repeat(1).take(20));
33/// let y = Array1::from(v);
34/// let sss = StratifiedShuffleSplit::new(3, &y)
35///     .with_test_size(SubsetSize::Fraction(0.2))
36///     .with_seed(0);
37/// let splits = sss.split(100).unwrap();
38/// assert_eq!(splits.len(), 3);
39/// ```
40#[derive(Debug, Clone)]
41pub struct StratifiedShuffleSplit<L> {
42    n_splits: usize,
43    test_size: SubsetSize,
44    train_size: Option<SubsetSize>,
45    seed: u64,
46    labels: Vec<L>,
47}
48
49impl<L: Eq + Hash + Clone> StratifiedShuffleSplit<L> {
50    /// Create a `StratifiedShuffleSplit` over class labels `y`, with a default
51    /// test size of 10%.
52    #[must_use]
53    pub fn new(n_splits: usize, y: &Array1<L>) -> Self {
54        Self {
55            n_splits,
56            test_size: SubsetSize::Fraction(0.1),
57            train_size: None,
58            seed: 0,
59            labels: y.to_vec(),
60        }
61    }
62
63    /// Set the (approximate) total test-subset size.
64    #[must_use]
65    pub fn with_test_size(mut self, test_size: SubsetSize) -> Self {
66        self.test_size = test_size;
67        self
68    }
69
70    /// Set the (approximate) total train-subset size.
71    #[must_use]
72    pub fn with_train_size(mut self, train_size: SubsetSize) -> Self {
73        self.train_size = Some(train_size);
74        self
75    }
76
77    /// Set the base RNG seed.
78    #[must_use]
79    pub fn with_seed(mut self, seed: u64) -> Self {
80        self.seed = seed;
81        self
82    }
83
84    fn class_indices(&self) -> Vec<Vec<usize>> {
85        let mut order: Vec<L> = Vec::new();
86        let mut map: HashMap<L, Vec<usize>> = HashMap::new();
87        for (i, label) in self.labels.iter().enumerate() {
88            map.entry(label.clone()).or_insert_with(|| {
89                order.push(label.clone());
90                Vec::new()
91            });
92            map.get_mut(label).unwrap().push(i);
93        }
94        order.into_iter().map(|c| map.remove(&c).unwrap()).collect()
95    }
96}
97
98impl<L: Eq + Hash + Clone> CvSplitter for StratifiedShuffleSplit<L> {
99    fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
100        if n_samples != self.labels.len() {
101            return Err(ModelSelectionError::ShapeMismatch {
102                expected: self.labels.len(),
103                got: n_samples,
104            });
105        }
106        let n_test = self.test_size.resolve(n_samples);
107        let n_train = match self.train_size {
108            Some(ts) => ts.resolve(n_samples),
109            None => n_samples.saturating_sub(n_test),
110        };
111        if n_test == 0 || n_train == 0 {
112            return Err(ModelSelectionError::InvalidSplitCount {
113                msg: format!("resolved train={n_train}, test={n_test}; both must be >= 1"),
114            });
115        }
116        if n_train + n_test > n_samples {
117            return Err(ModelSelectionError::NotEnoughSamples {
118                needed: n_train + n_test,
119                got: n_samples,
120            });
121        }
122
123        let class_indices = self.class_indices();
124        let mut splits = Vec::with_capacity(self.n_splits);
125
126        for i in 0..self.n_splits {
127            let mut rng = StdRng::seed_from_u64(self.seed.wrapping_add(i as u64));
128            let mut train = Vec::new();
129            let mut test = Vec::new();
130
131            for members in &class_indices {
132                let n_c = members.len();
133                // Proportional per-class allocation.
134                let test_c = ((n_test as f64) * (n_c as f64) / (n_samples as f64)).round() as usize;
135                let train_c =
136                    ((n_train as f64) * (n_c as f64) / (n_samples as f64)).round() as usize;
137                // Never over-draw a class.
138                let (test_c, train_c) = if test_c + train_c > n_c {
139                    (test_c.min(n_c), n_c.saturating_sub(test_c).min(train_c))
140                } else {
141                    (test_c, train_c)
142                };
143
144                let mut shuffled = members.clone();
145                shuffled.shuffle(&mut rng);
146                test.extend_from_slice(&shuffled[..test_c]);
147                train.extend_from_slice(&shuffled[test_c..test_c + train_c]);
148            }
149
150            train.sort_unstable();
151            test.sort_unstable();
152            splits.push((train, test));
153        }
154        Ok(splits)
155    }
156
157    fn n_splits(&self) -> usize {
158        self.n_splits
159    }
160}
161
162#[cfg(test)]
163mod tests {
164    use super::*;
165    use std::collections::HashSet;
166
167    #[test]
168    fn preserves_class_proportion_per_split() {
169        let mut v = vec![0; 80];
170        v.extend(std::iter::repeat(1).take(20));
171        let y = Array1::from(v);
172        let sss = StratifiedShuffleSplit::new(5, &y)
173            .with_test_size(SubsetSize::Fraction(0.2))
174            .with_seed(7);
175        for (_, test) in sss.split(100).unwrap() {
176            let ones = test.iter().filter(|&&i| y[i] == 1).count();
177            let frac = ones as f64 / test.len() as f64;
178            assert!((frac - 0.2).abs() < 0.1, "test class-1 share {frac}");
179        }
180    }
181
182    #[test]
183    fn train_and_test_disjoint() {
184        let y = Array1::from(vec![0, 1, 0, 1, 0, 1, 0, 1, 0, 1]);
185        let sss = StratifiedShuffleSplit::new(3, &y).with_test_size(SubsetSize::Fraction(0.4));
186        for (train, test) in sss.split(10).unwrap() {
187            let tr: HashSet<_> = train.iter().collect();
188            let te: HashSet<_> = test.iter().collect();
189            assert!(tr.is_disjoint(&te));
190        }
191    }
192}