Skip to main content

torsh_data/sampler/
stratified.rs

1//! Stratified sampling functionality
2//!
3//! This module provides stratified sampling strategies that ensure proportional
4//! representation of different classes or strata in the sampled data.
5
6#[cfg(not(feature = "std"))]
7use alloc::vec::Vec;
8use std::collections::HashMap;
9
10// ✅ SciRS2 Policy Compliant - Using scirs2_core for all random operations
11use scirs2_core::rand_prelude::SliceRandom;
12
13use super::core::{rng_utils, Sampler, SamplerIterator};
14
15/// Stratified sampler that samples proportionally from different strata/classes
16///
17/// This sampler ensures that samples from each stratum are drawn proportionally
18/// to their size in the population, which is useful for maintaining class balance
19/// in machine learning datasets.
20///
21/// # Examples
22///
23/// ```rust,ignore
24/// use torsh_data::sampler::{StratifiedSampler, Sampler};
25///
26/// // Dataset with 3 classes: [0,0,0,1,1,1,2,2,2]
27/// let labels = vec![0, 0, 0, 1, 1, 1, 2, 2, 2];
28/// let sampler = StratifiedSampler::new(&labels, 6, false).with_generator(42);
29///
30/// let indices: Vec<usize> = sampler.iter().collect();
31/// assert_eq!(indices.len(), 6); // 2 samples from each class
32/// ```
33#[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    /// Create a new stratified sampler
43    ///
44    /// # Arguments
45    ///
46    /// * `labels` - Labels for each sample in the dataset
47    /// * `num_samples` - Total number of samples to draw
48    /// * `replacement` - Whether to sample with replacement
49    ///
50    /// # Examples
51    ///
52    /// ```rust,ignore
53    /// use torsh_data::sampler::StratifiedSampler;
54    ///
55    /// let labels = vec![0, 0, 1, 1, 2, 2];
56    /// let sampler = StratifiedSampler::new(&labels, 3, false);
57    /// assert_eq!(sampler.num_strata(), 3);
58    /// ```
59    pub fn new(labels: &[usize], num_samples: usize, replacement: bool) -> Self {
60        // Group indices by their labels
61        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        // Convert to Vec and sort by label to ensure deterministic ordering
68        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    /// Create a stratified sampler from pre-grouped strata
84    ///
85    /// # Arguments
86    ///
87    /// * `strata` - Pre-grouped indices for each stratum
88    /// * `num_samples` - Total number of samples to draw
89    /// * `replacement` - Whether to sample with replacement
90    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    /// Set random generator seed
100    ///
101    /// # Arguments
102    ///
103    /// * `seed` - Random seed for reproducible sampling
104    pub fn with_generator(mut self, seed: u64) -> Self {
105        self.generator = Some(seed);
106        self
107    }
108
109    /// Get the number of strata
110    pub fn num_strata(&self) -> usize {
111        self.strata.len()
112    }
113
114    /// Get the strata
115    pub fn strata(&self) -> &[Vec<usize>] {
116        &self.strata
117    }
118
119    /// Get the number of samples
120    pub fn num_samples(&self) -> usize {
121        self.num_samples
122    }
123
124    /// Check if sampling with replacement
125    pub fn replacement(&self) -> bool {
126        self.replacement
127    }
128
129    /// Get the generator seed if set
130    pub fn generator(&self) -> Option<u64> {
131        self.generator
132    }
133
134    /// Get the proportional sample count for each stratum
135    ///
136    /// This method calculates how many samples should be drawn from each stratum
137    /// to maintain proportional representation in the final sample.
138    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        // Calculate proportional samples for each stratum
149        for (i, stratum) in self.strata.iter().enumerate() {
150            let count = if i == self.strata.len() - 1 {
151                // Last stratum gets the remainder to ensure exact total
152                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    /// Get stratum sizes
166    pub fn stratum_sizes(&self) -> Vec<usize> {
167        self.strata.iter().map(|s| s.len()).collect()
168    }
169
170    /// Calculate the total population size across all strata
171    pub fn total_population(&self) -> usize {
172        self.strata.iter().map(|s| s.len()).sum()
173    }
174
175    /// Check if the sampler is valid (has non-empty strata)
176    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        // ✅ SciRS2 Policy Compliant - Using scirs2_core for random operations
190        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        // Sample from each stratum
195        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                    // Sample with replacement
203                    (0..count)
204                        .map(|_| stratum[rng_utils::gen_range(&mut rng, 0..stratum.len())])
205                        .collect()
206                } else {
207                    // Sample without replacement
208                    let mut shuffled = stratum.clone();
209                    shuffled.shuffle(&mut rng);
210                    shuffled.into_iter().take(count).collect()
211                }
212            } else {
213                // Need more samples than available in stratum - sample with replacement
214                (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        // Shuffle the final combined indices to avoid grouping by stratum
223        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
233/// Create a stratified sampler from labels
234///
235/// Convenience function for creating a stratified sampler.
236///
237/// # Arguments
238///
239/// * `labels` - Labels for each sample in the dataset
240/// * `num_samples` - Total number of samples to draw
241/// * `replacement` - Whether to sample with replacement
242/// * `seed` - Optional random seed for reproducible sampling
243pub 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
256/// Create a balanced stratified sampler
257///
258/// Creates a stratified sampler that draws equal numbers of samples from each stratum,
259/// regardless of their original sizes.
260///
261/// # Arguments
262///
263/// * `labels` - Labels for each sample in the dataset
264/// * `samples_per_stratum` - Number of samples to draw from each stratum
265/// * `replacement` - Whether to sample with replacement
266/// * `seed` - Optional random seed for reproducible sampling
267pub fn balanced_stratified(
268    labels: &[usize],
269    samples_per_stratum: usize,
270    replacement: bool,
271    seed: Option<u64>,
272) -> StratifiedSampler {
273    // Group indices by labels
274    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
289/// Create a stratified train-test split
290///
291/// Splits the data into training and testing sets while maintaining
292/// the proportion of samples from each class.
293///
294/// # Arguments
295///
296/// * `labels` - Labels for each sample in the dataset
297/// * `test_ratio` - Proportion of data to use for testing (0.0 to 1.0)
298/// * `seed` - Optional random seed for reproducible splits
299///
300/// # Returns
301///
302/// A tuple of (train_sampler, test_sampler)
303pub 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    // Group indices by labels
314    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    // ✅ SciRS2 Policy Compliant - Using scirs2_core for random operations
323    let mut rng = rng_utils::create_rng(seed);
324
325    for (_, mut stratum) in strata {
326        // Shuffle the stratum
327        stratum.shuffle(&mut rng);
328
329        // Split into train and test
330        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        // Test with balanced classes: [0,0,0,1,1,1,2,2,2]
359        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        // Check that all indices are valid
373        for &idx in &indices {
374            assert!(idx < labels.len());
375        }
376
377        // Count samples per class
378        let mut class_counts = [0; 3];
379        for &idx in &indices {
380            class_counts[labels[idx]] += 1;
381        }
382
383        // Should be roughly proportional (2 samples per class for balanced classes)
384        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        // Test with imbalanced classes: [0,0,0,0,0,1,1,2]
392        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        // Count samples per class
402        let mut class_counts = [0; 3];
403        for &idx in &indices {
404            class_counts[labels[idx]] += 1;
405        }
406
407        // Should be proportional: class 0 (5/8 = 62.5%), class 1 (2/8 = 25%), class 2 (1/8 = 12.5%)
408        // Expected: class 0 → 5 samples, class 1 → 2 samples, class 2 → 1 sample
409        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        // Count samples per class
426        let mut class_counts = [0; 3];
427        for &idx in &indices {
428            class_counts[labels[idx]] += 1;
429        }
430
431        // Should be roughly equal (3 samples per class)
432        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        // All indices should be valid and from class 0
462        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        // Test when requesting more samples than available
471        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        // Should have samples from both classes
478        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],    // First stratum
492            vec![3, 4],       // Second stratum
493            vec![5, 6, 7, 8], // Third stratum
494        ];
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        // All indices should be from the original strata
505        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); // Should sum to num_samples
521    }
522
523    #[test]
524    fn test_convenience_functions() {
525        let labels = vec![0, 0, 1, 1, 2, 2];
526
527        // Test stratified convenience function
528        let sampler = stratified(&labels, 4, false, Some(42));
529        assert_eq!(sampler.len(), 4);
530        assert_eq!(sampler.generator(), Some(42));
531
532        // Test balanced_stratified convenience function
533        let balanced = balanced_stratified(&labels, 2, false, Some(42));
534        assert_eq!(balanced.len(), 6); // 3 strata * 2 samples each
535        assert_eq!(balanced.generator(), Some(42));
536
537        let indices: Vec<usize> = balanced.iter().collect();
538        assert_eq!(indices.len(), 6);
539
540        // Count samples per class - should be exactly 2 each
541        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        // Total samples should equal original dataset size
556        assert_eq!(train_sampler.len() + test_sampler.len(), labels.len());
557
558        // Should maintain proportion in both sets
559        let train_indices: Vec<usize> = train_sampler.iter().collect();
560        let test_indices: Vec<usize> = test_sampler.iter().collect();
561
562        // Count classes in train set
563        let mut train_class_counts = [0; 3];
564        for &idx in &train_indices {
565            train_class_counts[labels[idx]] += 1;
566        }
567
568        // Count classes in test set
569        let mut test_class_counts = [0; 3];
570        for &idx in &test_indices {
571            test_class_counts[labels[idx]] += 1;
572        }
573
574        // Each class should appear in both train and test sets
575        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); // 4 samples per class
579        }
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        // Test with zero samples requested
617        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        // Test with large number of samples
623        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}