1#[cfg(not(feature = "std"))]
7use alloc::vec::Vec;
8use std::collections::HashMap;
9
10use scirs2_core::rand_prelude::SliceRandom;
12
13use super::core::{rng_utils, Sampler, SamplerIterator};
14
15#[derive(Debug, Clone)]
34pub struct StratifiedSampler {
35 strata: Vec<Vec<usize>>,
36 num_samples: usize,
37 replacement: bool,
38 generator: Option<u64>,
39}
40
41impl StratifiedSampler {
42 pub fn new(labels: &[usize], num_samples: usize, replacement: bool) -> Self {
60 let mut strata: HashMap<usize, Vec<usize>> = HashMap::new();
62
63 for (idx, &label) in labels.iter().enumerate() {
64 strata.entry(label).or_default().push(idx);
65 }
66
67 let mut strata_pairs: Vec<(usize, Vec<usize>)> = strata.into_iter().collect();
69 strata_pairs.sort_unstable_by_key(|(label, _)| *label);
70 let strata: Vec<Vec<usize>> = strata_pairs
71 .into_iter()
72 .map(|(_, indices)| indices)
73 .collect();
74
75 Self {
76 strata,
77 num_samples,
78 replacement,
79 generator: None,
80 }
81 }
82
83 pub fn from_strata(strata: Vec<Vec<usize>>, num_samples: usize, replacement: bool) -> Self {
91 Self {
92 strata,
93 num_samples,
94 replacement,
95 generator: None,
96 }
97 }
98
99 pub fn with_generator(mut self, seed: u64) -> Self {
105 self.generator = Some(seed);
106 self
107 }
108
109 pub fn num_strata(&self) -> usize {
111 self.strata.len()
112 }
113
114 pub fn strata(&self) -> &[Vec<usize>] {
116 &self.strata
117 }
118
119 pub fn num_samples(&self) -> usize {
121 self.num_samples
122 }
123
124 pub fn replacement(&self) -> bool {
126 self.replacement
127 }
128
129 pub fn generator(&self) -> Option<u64> {
131 self.generator
132 }
133
134 pub fn get_stratum_sample_counts(&self) -> Vec<usize> {
139 let total_population: usize = self.strata.iter().map(|s| s.len()).sum();
140
141 if total_population == 0 {
142 return vec![0; self.strata.len()];
143 }
144
145 let mut counts = Vec::with_capacity(self.strata.len());
146 let mut allocated = 0;
147
148 for (i, stratum) in self.strata.iter().enumerate() {
150 let count = if i == self.strata.len() - 1 {
151 self.num_samples.saturating_sub(allocated)
153 } else {
154 let proportion = stratum.len() as f64 / total_population as f64;
155 (self.num_samples as f64 * proportion).round() as usize
156 };
157
158 counts.push(count);
159 allocated += count;
160 }
161
162 counts
163 }
164
165 pub fn stratum_sizes(&self) -> Vec<usize> {
167 self.strata.iter().map(|s| s.len()).collect()
168 }
169
170 pub fn total_population(&self) -> usize {
172 self.strata.iter().map(|s| s.len()).sum()
173 }
174
175 pub fn is_valid(&self) -> bool {
177 !self.strata.is_empty() && self.total_population() > 0
178 }
179}
180
181impl Sampler for StratifiedSampler {
182 type Iter = SamplerIterator;
183
184 fn iter(&self) -> Self::Iter {
185 if !self.is_valid() {
186 return SamplerIterator::new(vec![]);
187 }
188
189 let mut rng = rng_utils::create_rng(self.generator);
191 let stratum_counts = self.get_stratum_sample_counts();
192 let mut all_indices = Vec::with_capacity(self.num_samples);
193
194 for (stratum, &count) in self.strata.iter().zip(stratum_counts.iter()) {
196 if count == 0 || stratum.is_empty() {
197 continue;
198 }
199
200 let stratum_samples: Vec<usize> = if self.replacement || count <= stratum.len() {
201 if self.replacement {
202 (0..count)
204 .map(|_| stratum[rng_utils::gen_range(&mut rng, 0..stratum.len())])
205 .collect()
206 } else {
207 let mut shuffled = stratum.clone();
209 shuffled.shuffle(&mut rng);
210 shuffled.into_iter().take(count).collect()
211 }
212 } else {
213 (0..count)
215 .map(|_| stratum[rng_utils::gen_range(&mut rng, 0..stratum.len())])
216 .collect()
217 };
218
219 all_indices.extend(stratum_samples);
220 }
221
222 all_indices.shuffle(&mut rng);
224
225 SamplerIterator::new(all_indices)
226 }
227
228 fn len(&self) -> usize {
229 self.num_samples
230 }
231}
232
233pub fn stratified(
244 labels: &[usize],
245 num_samples: usize,
246 replacement: bool,
247 seed: Option<u64>,
248) -> StratifiedSampler {
249 let mut sampler = StratifiedSampler::new(labels, num_samples, replacement);
250 if let Some(s) = seed {
251 sampler = sampler.with_generator(s);
252 }
253 sampler
254}
255
256pub fn balanced_stratified(
268 labels: &[usize],
269 samples_per_stratum: usize,
270 replacement: bool,
271 seed: Option<u64>,
272) -> StratifiedSampler {
273 let mut strata: HashMap<usize, Vec<usize>> = HashMap::new();
275 for (idx, &label) in labels.iter().enumerate() {
276 strata.entry(label).or_default().push(idx);
277 }
278
279 let strata: Vec<Vec<usize>> = strata.into_values().collect();
280 let num_samples = strata.len() * samples_per_stratum;
281
282 let mut sampler = StratifiedSampler::from_strata(strata, num_samples, replacement);
283 if let Some(s) = seed {
284 sampler = sampler.with_generator(s);
285 }
286 sampler
287}
288
289pub fn stratified_train_test_split(
304 labels: &[usize],
305 test_ratio: f64,
306 seed: Option<u64>,
307) -> (StratifiedSampler, StratifiedSampler) {
308 assert!(
309 (0.0..=1.0).contains(&test_ratio),
310 "test_ratio must be between 0.0 and 1.0"
311 );
312
313 let mut strata: HashMap<usize, Vec<usize>> = HashMap::new();
315 for (idx, &label) in labels.iter().enumerate() {
316 strata.entry(label).or_default().push(idx);
317 }
318
319 let mut train_strata = Vec::new();
320 let mut test_strata = Vec::new();
321
322 let mut rng = rng_utils::create_rng(seed);
324
325 for (_, mut stratum) in strata {
326 stratum.shuffle(&mut rng);
328
329 let test_size = ((stratum.len() as f64) * test_ratio).round() as usize;
331 let test_size = test_size.min(stratum.len());
332
333 let (train_indices, test_indices) = stratum.split_at(stratum.len() - test_size);
334
335 if !train_indices.is_empty() {
336 train_strata.push(train_indices.to_vec());
337 }
338 if !test_indices.is_empty() {
339 test_strata.push(test_indices.to_vec());
340 }
341 }
342
343 let train_size = train_strata.iter().map(|s| s.len()).sum();
344 let test_size = test_strata.iter().map(|s| s.len()).sum();
345
346 let train_sampler = StratifiedSampler::from_strata(train_strata, train_size, false);
347 let test_sampler = StratifiedSampler::from_strata(test_strata, test_size, false);
348
349 (train_sampler, test_sampler)
350}
351
352#[cfg(test)]
353mod tests {
354 use super::*;
355
356 #[test]
357 fn test_stratified_sampler_basic() {
358 let labels = vec![0, 0, 0, 1, 1, 1, 2, 2, 2];
360 let sampler = StratifiedSampler::new(&labels, 6, false).with_generator(42);
361
362 assert_eq!(sampler.len(), 6);
363 assert_eq!(sampler.num_strata(), 3);
364 assert_eq!(sampler.num_samples(), 6);
365 assert!(!sampler.replacement());
366 assert_eq!(sampler.generator(), Some(42));
367 assert!(sampler.is_valid());
368
369 let indices: Vec<usize> = sampler.iter().collect();
370 assert_eq!(indices.len(), 6);
371
372 for &idx in &indices {
374 assert!(idx < labels.len());
375 }
376
377 let mut class_counts = [0; 3];
379 for &idx in &indices {
380 class_counts[labels[idx]] += 1;
381 }
382
383 assert_eq!(class_counts[0], 2);
385 assert_eq!(class_counts[1], 2);
386 assert_eq!(class_counts[2], 2);
387 }
388
389 #[test]
390 fn test_stratified_sampler_imbalanced() {
391 let labels = vec![0, 0, 0, 0, 0, 1, 1, 2];
393 let sampler = StratifiedSampler::new(&labels, 8, false).with_generator(42);
394
395 assert_eq!(sampler.len(), 8);
396 assert_eq!(sampler.num_strata(), 3);
397
398 let indices: Vec<usize> = sampler.iter().collect();
399 assert_eq!(indices.len(), 8);
400
401 let mut class_counts = [0; 3];
403 for &idx in &indices {
404 class_counts[labels[idx]] += 1;
405 }
406
407 assert_eq!(class_counts[0], 5);
410 assert_eq!(class_counts[1], 2);
411 assert_eq!(class_counts[2], 1);
412 }
413
414 #[test]
415 fn test_stratified_sampler_with_replacement() {
416 let labels = vec![0, 1, 2];
417 let sampler = StratifiedSampler::new(&labels, 9, true).with_generator(42);
418
419 assert_eq!(sampler.len(), 9);
420 assert!(sampler.replacement());
421
422 let indices: Vec<usize> = sampler.iter().collect();
423 assert_eq!(indices.len(), 9);
424
425 let mut class_counts = [0; 3];
427 for &idx in &indices {
428 class_counts[labels[idx]] += 1;
429 }
430
431 assert_eq!(class_counts[0], 3);
433 assert_eq!(class_counts[1], 3);
434 assert_eq!(class_counts[2], 3);
435 }
436
437 #[test]
438 fn test_stratified_sampler_empty() {
439 let labels: Vec<usize> = vec![];
440 let sampler = StratifiedSampler::new(&labels, 5, false);
441
442 assert_eq!(sampler.len(), 5);
443 assert_eq!(sampler.num_strata(), 0);
444 assert!(!sampler.is_valid());
445
446 let indices: Vec<usize> = sampler.iter().collect();
447 assert_eq!(indices.len(), 0);
448 }
449
450 #[test]
451 fn test_stratified_sampler_single_stratum() {
452 let labels = vec![0, 0, 0, 0, 0];
453 let sampler = StratifiedSampler::new(&labels, 3, false).with_generator(42);
454
455 assert_eq!(sampler.len(), 3);
456 assert_eq!(sampler.num_strata(), 1);
457
458 let indices: Vec<usize> = sampler.iter().collect();
459 assert_eq!(indices.len(), 3);
460
461 for &idx in &indices {
463 assert!(idx < 5);
464 assert_eq!(labels[idx], 0);
465 }
466 }
467
468 #[test]
469 fn test_stratified_sampler_oversample() {
470 let labels = vec![0, 1];
472 let sampler = StratifiedSampler::new(&labels, 10, true).with_generator(42);
473
474 let indices: Vec<usize> = sampler.iter().collect();
475 assert_eq!(indices.len(), 10);
476
477 let mut class_counts = [0; 2];
479 for &idx in &indices {
480 class_counts[labels[idx]] += 1;
481 }
482
483 assert!(class_counts[0] > 0);
484 assert!(class_counts[1] > 0);
485 assert_eq!(class_counts[0] + class_counts[1], 10);
486 }
487
488 #[test]
489 fn test_stratified_sampler_from_strata() {
490 let strata = vec![
491 vec![0, 1, 2], vec![3, 4], vec![5, 6, 7, 8], ];
495 let sampler = StratifiedSampler::from_strata(strata.clone(), 6, false).with_generator(42);
496
497 assert_eq!(sampler.len(), 6);
498 assert_eq!(sampler.num_strata(), 3);
499 assert_eq!(sampler.strata(), &strata);
500
501 let indices: Vec<usize> = sampler.iter().collect();
502 assert_eq!(indices.len(), 6);
503
504 for &idx in &indices {
506 let found = strata.iter().any(|stratum| stratum.contains(&idx));
507 assert!(found);
508 }
509 }
510
511 #[test]
512 fn test_stratified_sampler_properties() {
513 let labels = vec![0, 1, 2, 0, 1, 2];
514 let sampler = StratifiedSampler::new(&labels, 4, false);
515
516 assert_eq!(sampler.stratum_sizes(), vec![2, 2, 2]);
517 assert_eq!(sampler.total_population(), 6);
518
519 let counts = sampler.get_stratum_sample_counts();
520 assert_eq!(counts.iter().sum::<usize>(), 4); }
522
523 #[test]
524 fn test_convenience_functions() {
525 let labels = vec![0, 0, 1, 1, 2, 2];
526
527 let sampler = stratified(&labels, 4, false, Some(42));
529 assert_eq!(sampler.len(), 4);
530 assert_eq!(sampler.generator(), Some(42));
531
532 let balanced = balanced_stratified(&labels, 2, false, Some(42));
534 assert_eq!(balanced.len(), 6); assert_eq!(balanced.generator(), Some(42));
536
537 let indices: Vec<usize> = balanced.iter().collect();
538 assert_eq!(indices.len(), 6);
539
540 let mut class_counts = [0; 3];
542 for &idx in &indices {
543 class_counts[labels[idx]] += 1;
544 }
545 assert_eq!(class_counts[0], 2);
546 assert_eq!(class_counts[1], 2);
547 assert_eq!(class_counts[2], 2);
548 }
549
550 #[test]
551 fn test_stratified_train_test_split() {
552 let labels = vec![0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2];
553 let (train_sampler, test_sampler) = stratified_train_test_split(&labels, 0.25, Some(42));
554
555 assert_eq!(train_sampler.len() + test_sampler.len(), labels.len());
557
558 let train_indices: Vec<usize> = train_sampler.iter().collect();
560 let test_indices: Vec<usize> = test_sampler.iter().collect();
561
562 let mut train_class_counts = [0; 3];
564 for &idx in &train_indices {
565 train_class_counts[labels[idx]] += 1;
566 }
567
568 let mut test_class_counts = [0; 3];
570 for &idx in &test_indices {
571 test_class_counts[labels[idx]] += 1;
572 }
573
574 for i in 0..3 {
576 assert!(train_class_counts[i] > 0);
577 assert!(test_class_counts[i] > 0);
578 assert_eq!(train_class_counts[i] + test_class_counts[i], 4); }
580 }
581
582 #[test]
583 #[should_panic(expected = "test_ratio must be between 0.0 and 1.0")]
584 fn test_stratified_train_test_split_invalid_ratio() {
585 let labels = vec![0, 1, 2];
586 stratified_train_test_split(&labels, 1.5, None);
587 }
588
589 #[test]
590 fn test_stratified_sampler_clone() {
591 let labels = vec![0, 1, 2, 0, 1, 2];
592 let sampler = StratifiedSampler::new(&labels, 4, false).with_generator(42);
593 let cloned = sampler.clone();
594
595 assert_eq!(sampler.len(), cloned.len());
596 assert_eq!(sampler.num_strata(), cloned.num_strata());
597 assert_eq!(sampler.replacement(), cloned.replacement());
598 assert_eq!(sampler.generator(), cloned.generator());
599 assert_eq!(sampler.strata(), cloned.strata());
600 }
601
602 #[test]
603 fn test_stratified_sampler_reproducible() {
604 let labels = vec![0, 0, 1, 1, 2, 2];
605 let sampler1 = StratifiedSampler::new(&labels, 4, false).with_generator(123);
606 let sampler2 = StratifiedSampler::new(&labels, 4, false).with_generator(123);
607
608 let indices1: Vec<usize> = sampler1.iter().collect();
609 let indices2: Vec<usize> = sampler2.iter().collect();
610
611 assert_eq!(indices1, indices2);
612 }
613
614 #[test]
615 fn test_edge_cases() {
616 let labels = vec![0, 1, 2];
618 let sampler = StratifiedSampler::new(&labels, 0, false);
619 let indices: Vec<usize> = sampler.iter().collect();
620 assert_eq!(indices.len(), 0);
621
622 let labels = vec![0, 1];
624 let sampler = StratifiedSampler::new(&labels, 1000, true);
625 let indices: Vec<usize> = sampler.iter().collect();
626 assert_eq!(indices.len(), 1000);
627 }
628}