model_selection_rs/splitters/
stratified_shuffle_split.rs1use 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#[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 #[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 #[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 #[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 #[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 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 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}