1#[cfg(not(feature = "std"))]
7use alloc::vec::Vec;
8
9use super::core::{rng_utils, Sampler, SamplerIterator};
10
11#[derive(Clone, Debug, PartialEq)]
16pub enum CurriculumStrategy {
17 Linear,
23
24 Exponential { base: f64 },
34
35 Step { thresholds: Vec<usize> },
44
45 AntiCurriculum,
50
51 SelfPaced { lambda: f64 },
60}
61
62impl Default for CurriculumStrategy {
63 fn default() -> Self {
64 CurriculumStrategy::Linear
65 }
66}
67
68#[derive(Clone)]
99pub struct CurriculumSampler<F> {
100 difficulties: Vec<f64>,
101 difficulty_fn: F,
102 current_epoch: usize,
103 total_epochs: usize,
104 curriculum_strategy: CurriculumStrategy,
105 generator: Option<u64>,
106}
107
108impl<F> CurriculumSampler<F>
109where
110 F: Fn(usize) -> f64 + Send + Clone,
111{
112 pub fn new(
137 dataset_size: usize,
138 difficulty_fn: F,
139 total_epochs: usize,
140 strategy: CurriculumStrategy,
141 ) -> Self {
142 let difficulties: Vec<f64> = (0..dataset_size).map(&difficulty_fn).collect();
143
144 Self {
145 difficulties,
146 difficulty_fn,
147 current_epoch: 0,
148 total_epochs,
149 curriculum_strategy: strategy,
150 generator: None,
151 }
152 }
153
154 pub fn from_difficulties(
162 difficulties: Vec<f64>,
163 total_epochs: usize,
164 strategy: CurriculumStrategy,
165 ) -> Self
166 where
167 F: Default,
168 {
169 Self {
170 difficulty_fn: F::default(),
171 difficulties,
172 current_epoch: 0,
173 total_epochs,
174 curriculum_strategy: strategy,
175 generator: None,
176 }
177 }
178
179 pub fn set_epoch(&mut self, epoch: usize) {
185 self.current_epoch = epoch;
186 }
187
188 pub fn with_generator(mut self, seed: u64) -> Self {
194 self.generator = Some(seed);
195 self
196 }
197
198 pub fn current_epoch(&self) -> usize {
200 self.current_epoch
201 }
202
203 pub fn total_epochs(&self) -> usize {
205 self.total_epochs
206 }
207
208 pub fn strategy(&self) -> &CurriculumStrategy {
210 &self.curriculum_strategy
211 }
212
213 pub fn difficulties(&self) -> &[f64] {
215 &self.difficulties
216 }
217
218 pub fn generator(&self) -> Option<u64> {
220 self.generator
221 }
222
223 pub fn progress(&self) -> f64 {
225 if self.total_epochs <= 1 {
226 1.0
227 } else {
228 (self.current_epoch as f64 / (self.total_epochs - 1) as f64).min(1.0)
229 }
230 }
231
232 pub fn get_difficulty_threshold(&self) -> f64 {
237 let progress = self.progress();
238
239 match &self.curriculum_strategy {
240 CurriculumStrategy::Linear => progress,
241 CurriculumStrategy::Exponential { base } => {
242 if *base <= 1.0 {
243 progress } else {
245 (base.powf(progress) - 1.0) / (base - 1.0)
246 }
247 }
248 CurriculumStrategy::Step { thresholds } => {
249 if thresholds.is_empty() {
250 1.0
251 } else {
252 let mut threshold = 0.0;
253 for &epoch_threshold in thresholds {
254 if self.current_epoch >= epoch_threshold {
255 threshold += 1.0 / thresholds.len() as f64;
256 }
257 }
258 threshold.min(1.0)
259 }
260 }
261 CurriculumStrategy::AntiCurriculum => {
262 1.0 }
268 CurriculumStrategy::SelfPaced { lambda } => {
269 (progress * lambda).min(1.0)
272 }
273 }
274 }
275
276 pub fn get_curriculum_indices(&self) -> Vec<usize> {
281 if self.difficulties.is_empty() {
282 return Vec::new();
283 }
284
285 let threshold = self.get_difficulty_threshold();
286 let max_difficulty = self
287 .difficulties
288 .iter()
289 .fold(f64::NEG_INFINITY, |a, &b| a.max(b));
290 let min_difficulty = self
291 .difficulties
292 .iter()
293 .fold(f64::INFINITY, |a, &b| a.min(b));
294
295 let range = max_difficulty - min_difficulty;
297 let normalized_threshold = if range > 0.0 {
298 threshold
299 } else {
300 1.0 };
302
303 self.difficulties
304 .iter()
305 .enumerate()
306 .filter_map(|(idx, &difficulty)| {
307 let normalized_difficulty = if range > 0.0 {
308 (difficulty - min_difficulty) / range
309 } else {
310 0.0
311 };
312
313 let include_sample = match &self.curriculum_strategy {
314 CurriculumStrategy::AntiCurriculum => {
315 normalized_difficulty <= normalized_threshold
319 }
320 _ => {
321 normalized_difficulty <= normalized_threshold
323 }
324 };
325
326 if include_sample {
327 Some(idx)
328 } else {
329 None
330 }
331 })
332 .collect()
333 }
334
335 pub fn update_difficulties(&mut self) {
340 let dataset_size = self.difficulties.len();
341 self.difficulties = (0..dataset_size).map(&self.difficulty_fn).collect();
342 }
343
344 pub fn set_strategy(&mut self, strategy: CurriculumStrategy) {
350 self.curriculum_strategy = strategy;
351 }
352
353 pub fn reset(&mut self) {
355 self.current_epoch = 0;
356 }
357
358 pub fn is_complete(&self) -> bool {
360 self.current_epoch >= self.total_epochs || self.get_difficulty_threshold() >= 1.0
361 }
362
363 pub fn curriculum_stats(&self) -> CurriculumStats {
365 let threshold = self.get_difficulty_threshold();
366 let indices = self.get_curriculum_indices();
367 let total_samples = self.difficulties.len();
368
369 CurriculumStats {
370 current_epoch: self.current_epoch,
371 total_epochs: self.total_epochs,
372 progress: self.progress(),
373 difficulty_threshold: threshold,
374 included_samples: indices.len(),
375 total_samples,
376 inclusion_ratio: if total_samples > 0 {
377 indices.len() as f64 / total_samples as f64
378 } else {
379 0.0
380 },
381 }
382 }
383}
384
385impl<F: Send + Clone> Sampler for CurriculumSampler<F>
386where
387 F: Fn(usize) -> f64 + Send + Clone,
388{
389 type Iter = SamplerIterator;
390
391 fn iter(&self) -> Self::Iter {
392 let mut indices = self.get_curriculum_indices();
393
394 rng_utils::shuffle_indices(&mut indices, self.generator);
396
397 SamplerIterator::new(indices)
398 }
399
400 fn len(&self) -> usize {
401 self.get_curriculum_indices().len()
402 }
403}
404
405#[derive(Debug, Clone, PartialEq)]
407pub struct CurriculumStats {
408 pub current_epoch: usize,
410 pub total_epochs: usize,
412 pub progress: f64,
414 pub difficulty_threshold: f64,
416 pub included_samples: usize,
418 pub total_samples: usize,
420 pub inclusion_ratio: f64,
422}
423
424pub fn linear_curriculum<F>(
435 dataset_size: usize,
436 difficulty_fn: F,
437 total_epochs: usize,
438 seed: Option<u64>,
439) -> CurriculumSampler<F>
440where
441 F: Fn(usize) -> f64 + Send + Clone,
442{
443 let mut sampler = CurriculumSampler::new(
444 dataset_size,
445 difficulty_fn,
446 total_epochs,
447 CurriculumStrategy::Linear,
448 );
449 if let Some(s) = seed {
450 sampler = sampler.with_generator(s);
451 }
452 sampler
453}
454
455pub fn exponential_curriculum<F>(
467 dataset_size: usize,
468 difficulty_fn: F,
469 total_epochs: usize,
470 base: f64,
471 seed: Option<u64>,
472) -> CurriculumSampler<F>
473where
474 F: Fn(usize) -> f64 + Send + Clone,
475{
476 let mut sampler = CurriculumSampler::new(
477 dataset_size,
478 difficulty_fn,
479 total_epochs,
480 CurriculumStrategy::Exponential { base },
481 );
482 if let Some(s) = seed {
483 sampler = sampler.with_generator(s);
484 }
485 sampler
486}
487
488pub fn step_curriculum<F>(
500 dataset_size: usize,
501 difficulty_fn: F,
502 total_epochs: usize,
503 thresholds: Vec<usize>,
504 seed: Option<u64>,
505) -> CurriculumSampler<F>
506where
507 F: Fn(usize) -> f64 + Send + Clone,
508{
509 let mut sampler = CurriculumSampler::new(
510 dataset_size,
511 difficulty_fn,
512 total_epochs,
513 CurriculumStrategy::Step { thresholds },
514 );
515 if let Some(s) = seed {
516 sampler = sampler.with_generator(s);
517 }
518 sampler
519}
520
521pub fn anti_curriculum<F>(
532 dataset_size: usize,
533 difficulty_fn: F,
534 total_epochs: usize,
535 seed: Option<u64>,
536) -> CurriculumSampler<F>
537where
538 F: Fn(usize) -> f64 + Send + Clone,
539{
540 let mut sampler = CurriculumSampler::new(
541 dataset_size,
542 difficulty_fn,
543 total_epochs,
544 CurriculumStrategy::AntiCurriculum,
545 );
546 if let Some(s) = seed {
547 sampler = sampler.with_generator(s);
548 }
549 sampler
550}
551
552#[cfg(test)]
553mod tests {
554 use super::*;
555
556 fn linear_difficulty(idx: usize) -> f64 {
558 idx as f64 / 100.0
559 }
560
561 fn center_distance_difficulty(idx: usize) -> f64 {
563 (idx as f64 - 50.0).abs() / 50.0
564 }
565
566 #[test]
567 fn test_curriculum_sampler_basic() {
568 let mut sampler =
569 CurriculumSampler::new(100, linear_difficulty, 10, CurriculumStrategy::Linear)
570 .with_generator(42);
571
572 assert_eq!(sampler.total_epochs(), 10);
573 assert_eq!(sampler.current_epoch(), 0);
574 assert_eq!(sampler.generator(), Some(42));
575 assert_eq!(sampler.difficulties().len(), 100);
576 assert!(!sampler.is_complete());
577
578 sampler.set_epoch(0);
580 let early_indices = sampler.get_curriculum_indices();
581 assert!(!early_indices.is_empty());
582 assert!(early_indices.len() < 100);
583
584 sampler.set_epoch(9);
586 let late_indices = sampler.get_curriculum_indices();
587 assert_eq!(late_indices.len(), 100);
588 assert!(sampler.is_complete());
589 }
590
591 #[test]
592 fn test_curriculum_strategies() {
593 let dataset_size = 100;
594 let total_epochs = 10;
595
596 let mut linear_sampler = CurriculumSampler::new(
598 dataset_size,
599 linear_difficulty,
600 total_epochs,
601 CurriculumStrategy::Linear,
602 );
603
604 linear_sampler.set_epoch(0);
605 let linear_early = linear_sampler.get_curriculum_indices().len();
606 linear_sampler.set_epoch(5);
607 let linear_mid = linear_sampler.get_curriculum_indices().len();
608 linear_sampler.set_epoch(9);
609 let linear_late = linear_sampler.get_curriculum_indices().len();
610
611 assert!(linear_early < linear_mid);
612 assert!(linear_mid < linear_late);
613 assert_eq!(linear_late, dataset_size);
614
615 let mut exp_sampler = CurriculumSampler::new(
617 dataset_size,
618 linear_difficulty,
619 total_epochs,
620 CurriculumStrategy::Exponential { base: 2.0 },
621 );
622
623 exp_sampler.set_epoch(0);
624 let exp_early = exp_sampler.get_curriculum_indices().len();
625 exp_sampler.set_epoch(5);
626 let _exp_mid = exp_sampler.get_curriculum_indices().len();
627
628 assert!(exp_early <= linear_early);
630
631 let mut step_sampler = CurriculumSampler::new(
633 dataset_size,
634 linear_difficulty,
635 total_epochs,
636 CurriculumStrategy::Step {
637 thresholds: vec![3, 6, 9],
638 },
639 );
640
641 step_sampler.set_epoch(2);
642 let step_before = step_sampler.get_curriculum_indices().len();
643 step_sampler.set_epoch(3);
644 let step_after = step_sampler.get_curriculum_indices().len();
645
646 assert!(step_after > step_before);
648
649 let mut anti_sampler = CurriculumSampler::new(
651 dataset_size,
652 linear_difficulty,
653 total_epochs,
654 CurriculumStrategy::AntiCurriculum,
655 );
656
657 anti_sampler.set_epoch(0);
658 let anti_early = anti_sampler.get_curriculum_indices().len();
659 anti_sampler.set_epoch(9);
660 let anti_late = anti_sampler.get_curriculum_indices().len();
661
662 assert!(anti_early > linear_early);
664 assert_eq!(anti_late, dataset_size);
665 }
666
667 #[test]
668 fn test_difficulty_threshold_calculation() {
669 let sampler =
670 CurriculumSampler::new(100, linear_difficulty, 10, CurriculumStrategy::Linear);
671
672 assert_eq!(sampler.progress(), 0.0);
674
675 let mut sampler = sampler;
676 sampler.set_epoch(5);
677 assert!((sampler.progress() - 5.0 / 9.0).abs() < f64::EPSILON);
678
679 sampler.set_epoch(9);
680 assert_eq!(sampler.progress(), 1.0);
681
682 assert_eq!(sampler.get_difficulty_threshold(), 1.0);
684
685 sampler.set_strategy(CurriculumStrategy::Exponential { base: 2.0 });
686 sampler.set_epoch(0);
687 assert_eq!(sampler.get_difficulty_threshold(), 0.0);
688
689 sampler.set_strategy(CurriculumStrategy::AntiCurriculum);
690 sampler.set_epoch(0);
691 assert_eq!(sampler.get_difficulty_threshold(), 1.0);
692 }
693
694 #[test]
695 fn test_curriculum_from_difficulties() {
696 let difficulties = vec![0.1, 0.3, 0.5, 0.7, 0.9];
697 let difficulty_fn = |idx: usize| difficulties.get(idx).copied().unwrap_or(0.0);
698 let sampler = CurriculumSampler::new(
699 difficulties.len(),
700 difficulty_fn,
701 5,
702 CurriculumStrategy::Linear,
703 );
704
705 assert_eq!(sampler.difficulties(), &difficulties);
706 assert_eq!(sampler.total_epochs(), 5);
707 }
708
709 #[test]
710 fn test_curriculum_indices_selection() {
711 let difficulties = vec![0.0, 0.2, 0.4, 0.6, 0.8, 1.0];
712 let difficulty_fn = |idx: usize| difficulties.get(idx).copied().unwrap_or(0.0);
713 let mut sampler = CurriculumSampler::new(
714 difficulties.len(),
715 difficulty_fn,
716 6,
717 CurriculumStrategy::Linear,
718 );
719
720 sampler.set_epoch(0);
722 let indices = sampler.get_curriculum_indices();
723 assert!(indices.contains(&0)); assert!(!indices.contains(&5)); sampler.set_epoch(5);
728 let indices = sampler.get_curriculum_indices();
729 assert_eq!(indices.len(), 6);
730 for i in 0..6 {
731 assert!(indices.contains(&i));
732 }
733 }
734
735 #[test]
736 fn test_curriculum_stats() {
737 let mut sampler =
738 CurriculumSampler::new(100, linear_difficulty, 10, CurriculumStrategy::Linear);
739
740 sampler.set_epoch(5);
741 let stats = sampler.curriculum_stats();
742
743 assert_eq!(stats.current_epoch, 5);
744 assert_eq!(stats.total_epochs, 10);
745 assert_eq!(stats.total_samples, 100);
746 assert!(stats.progress > 0.0 && stats.progress < 1.0);
747 assert!(stats.difficulty_threshold > 0.0 && stats.difficulty_threshold < 1.0);
748 assert!(stats.included_samples > 0 && stats.included_samples < 100);
749 assert!(stats.inclusion_ratio > 0.0 && stats.inclusion_ratio < 1.0);
750 }
751
752 #[test]
753 fn test_curriculum_sampler_iter() {
754 let mut sampler =
755 CurriculumSampler::new(20, linear_difficulty, 5, CurriculumStrategy::Linear)
756 .with_generator(42);
757
758 sampler.set_epoch(0);
759 let indices1: Vec<usize> = sampler.iter().collect();
760 let indices2: Vec<usize> = sampler.iter().collect();
761
762 assert_eq!(indices1.len(), sampler.len());
763 assert_eq!(indices2.len(), sampler.len());
764
765 assert_eq!(indices1, indices2);
767
768 sampler.set_epoch(4);
770 let late_indices: Vec<usize> = sampler.iter().collect();
771 assert!(late_indices.len() >= indices1.len());
772 }
773
774 #[test]
775 fn test_convenience_functions() {
776 let linear = linear_curriculum(50, linear_difficulty, 10, Some(42));
778 assert_eq!(linear.total_epochs(), 10);
779 assert_eq!(linear.generator(), Some(42));
780 assert!(matches!(linear.strategy(), CurriculumStrategy::Linear));
781
782 let exponential = exponential_curriculum(50, linear_difficulty, 10, 2.0, Some(42));
784 assert!(
785 matches!(exponential.strategy(), CurriculumStrategy::Exponential { base } if *base == 2.0)
786 );
787
788 let step = step_curriculum(50, linear_difficulty, 10, vec![2, 5, 8], Some(42));
790 assert!(
791 matches!(step.strategy(), CurriculumStrategy::Step { thresholds } if thresholds == &vec![2, 5, 8])
792 );
793
794 let anti = anti_curriculum(50, linear_difficulty, 10, Some(42));
796 assert!(matches!(
797 anti.strategy(),
798 CurriculumStrategy::AntiCurriculum
799 ));
800 }
801
802 #[test]
803 fn test_curriculum_methods() {
804 let mut sampler =
805 CurriculumSampler::new(100, linear_difficulty, 10, CurriculumStrategy::Linear);
806
807 sampler.set_epoch(5);
809 assert_eq!(sampler.current_epoch(), 5);
810 sampler.reset();
811 assert_eq!(sampler.current_epoch(), 0);
812
813 sampler.set_strategy(CurriculumStrategy::Exponential { base: 3.0 });
815 assert!(
816 matches!(sampler.strategy(), CurriculumStrategy::Exponential { base } if *base == 3.0)
817 );
818
819 let original_difficulties = sampler.difficulties().to_vec();
821 sampler.update_difficulties();
822 assert_eq!(sampler.difficulties(), &original_difficulties);
823 }
824
825 #[test]
826 fn test_edge_cases() {
827 let empty_sampler =
829 CurriculumSampler::new(0, linear_difficulty, 5, CurriculumStrategy::Linear);
830 assert_eq!(empty_sampler.len(), 0);
831 assert!(empty_sampler.get_curriculum_indices().is_empty());
832
833 let mut single_epoch =
835 CurriculumSampler::new(10, linear_difficulty, 1, CurriculumStrategy::Linear);
836 single_epoch.set_epoch(0);
837 assert_eq!(single_epoch.progress(), 1.0);
838 assert!(single_epoch.is_complete());
839
840 let same_difficulties = vec![0.5; 10];
842 let same_difficulty_fn = |idx: usize| same_difficulties.get(idx).copied().unwrap_or(0.0);
843 let mut same_sampler = CurriculumSampler::new(
844 same_difficulties.len(),
845 same_difficulty_fn,
846 5,
847 CurriculumStrategy::Linear,
848 );
849 same_sampler.set_epoch(0);
850 assert_eq!(same_sampler.get_curriculum_indices().len(), 10); let mut invalid_exp = CurriculumSampler::new(
854 10,
855 linear_difficulty,
856 5,
857 CurriculumStrategy::Exponential { base: 0.5 }, );
859 invalid_exp.set_epoch(2);
860 assert!(invalid_exp.get_difficulty_threshold() >= 0.0);
862 }
863
864 #[test]
865 fn test_curriculum_strategy_equality() {
866 assert_eq!(CurriculumStrategy::Linear, CurriculumStrategy::Linear);
867 assert_eq!(
868 CurriculumStrategy::Exponential { base: 2.0 },
869 CurriculumStrategy::Exponential { base: 2.0 }
870 );
871 assert_ne!(
872 CurriculumStrategy::Linear,
873 CurriculumStrategy::AntiCurriculum
874 );
875 }
876
877 #[test]
878 fn test_curriculum_strategy_default() {
879 assert_eq!(CurriculumStrategy::default(), CurriculumStrategy::Linear);
880 }
881
882 #[test]
883 fn test_center_distance_difficulty() {
884 let mut sampler = CurriculumSampler::new(
885 101, center_distance_difficulty,
887 10,
888 CurriculumStrategy::Linear,
889 )
890 .with_generator(42);
891
892 sampler.set_epoch(0);
894 let early_indices = sampler.get_curriculum_indices();
895 assert!(early_indices.contains(&50)); sampler.set_epoch(9);
899 let late_indices = sampler.get_curriculum_indices();
900 assert!(late_indices.len() > early_indices.len());
901 assert!(late_indices.contains(&0) || late_indices.contains(&100)); }
903}