1pub trait LRScheduler: Send + Sync {
70 fn get_lr(&self, step: usize) -> f32;
72 fn step(&mut self);
74}
75
76#[derive(Debug)]
81pub struct LinearScheduler {
82 base_lr: f32,
83 warmup_steps: usize,
84 total_steps: usize,
85 current_step: usize,
86}
87
88impl LinearScheduler {
89 pub fn new(base_lr: f32, warmup_steps: usize, total_steps: usize) -> Self {
90 Self {
91 base_lr,
92 warmup_steps,
93 total_steps,
94 current_step: 0,
95 }
96 }
97}
98
99impl LRScheduler for LinearScheduler {
100 fn get_lr(&self, step: usize) -> f32 {
101 if self.warmup_steps > 0 && step < self.warmup_steps {
102 return self.base_lr * (step as f32) / (self.warmup_steps as f32);
103 }
104
105 let decay_steps = self.total_steps.saturating_sub(self.warmup_steps);
107 if decay_steps == 0 {
108 return 0.0;
109 }
110
111 let progress = (step - self.warmup_steps) as f32 / decay_steps as f32;
112 self.base_lr * (1.0 - progress).max(0.0)
113 }
114
115 fn step(&mut self) {
116 self.current_step += 1;
117 }
118}
119
120#[derive(Debug)]
126pub struct CosineScheduler {
127 base_lr: f32,
128 warmup_steps: usize,
129 total_steps: usize,
130 current_step: usize,
131 min_lr: f32,
132}
133
134impl CosineScheduler {
135 pub fn new(base_lr: f32, warmup_steps: usize, total_steps: usize, min_lr: f32) -> Self {
136 Self {
137 base_lr,
138 warmup_steps,
139 total_steps,
140 current_step: 0,
141 min_lr,
142 }
143 }
144}
145
146impl LRScheduler for CosineScheduler {
147 fn get_lr(&self, step: usize) -> f32 {
152 use std::f32::consts::PI;
153
154 if self.warmup_steps > 0 && step < self.warmup_steps {
155 return self.base_lr * (step as f32) / (self.warmup_steps as f32);
156 }
157
158 let decay_steps = self.total_steps.saturating_sub(self.warmup_steps);
160 if decay_steps == 0 {
161 return self.min_lr;
162 }
163
164 let progress = ((step - self.warmup_steps) as f32 / decay_steps as f32).clamp(0.0, 1.0);
165 let cosine_decay = 0.5 * (1.0 + (PI * progress).cos());
166 self.min_lr + (self.base_lr - self.min_lr) * cosine_decay
167 }
168
169 fn step(&mut self) {
170 self.current_step += 1;
171 }
172}
173
174#[derive(Debug)]
182pub struct PolynomialScheduler {
183 base_lr: f32,
184 warmup_steps: usize,
185 total_steps: usize,
186 current_step: usize,
187 min_lr: f32,
188 power: f32,
189}
190
191impl PolynomialScheduler {
192 pub fn new(
193 base_lr: f32,
194 warmup_steps: usize,
195 total_steps: usize,
196 min_lr: f32,
197 power: f32,
198 ) -> Self {
199 Self {
200 base_lr,
201 warmup_steps,
202 total_steps,
203 current_step: 0,
204 min_lr,
205 power,
206 }
207 }
208}
209
210impl LRScheduler for PolynomialScheduler {
211 fn get_lr(&self, step: usize) -> f32 {
212 if step < self.warmup_steps {
213 self.base_lr * (step as f32) / (self.warmup_steps as f32)
214 } else {
215 let progress =
216 (step - self.warmup_steps) as f32 / (self.total_steps - self.warmup_steps) as f32;
217 let decay_factor = (1.0 - progress.min(1.0)).powf(self.power);
218 self.min_lr + (self.base_lr - self.min_lr) * decay_factor
219 }
220 }
221
222 fn step(&mut self) {
223 self.current_step += 1;
224 }
225}
226
227#[derive(Debug)]
229pub struct ConstantWithWarmupScheduler {
230 base_lr: f32,
231 warmup_steps: usize,
232 current_step: usize,
233}
234
235impl ConstantWithWarmupScheduler {
236 pub fn new(base_lr: f32, warmup_steps: usize) -> Self {
237 Self {
238 base_lr,
239 warmup_steps,
240 current_step: 0,
241 }
242 }
243}
244
245impl LRScheduler for ConstantWithWarmupScheduler {
246 fn get_lr(&self, step: usize) -> f32 {
247 if step < self.warmup_steps {
248 self.base_lr * (step as f32) / (self.warmup_steps as f32)
249 } else {
250 self.base_lr
251 }
252 }
253
254 fn step(&mut self) {
255 self.current_step += 1;
256 }
257}
258
259#[derive(Debug)]
261pub struct ExponentialScheduler {
262 base_lr: f32,
263 warmup_steps: usize,
264 current_step: usize,
265 decay_rate: f32,
266 decay_steps: usize,
267}
268
269impl ExponentialScheduler {
270 pub fn new(base_lr: f32, warmup_steps: usize, decay_rate: f32, decay_steps: usize) -> Self {
271 Self {
272 base_lr,
273 warmup_steps,
274 current_step: 0,
275 decay_rate,
276 decay_steps,
277 }
278 }
279}
280
281impl LRScheduler for ExponentialScheduler {
282 fn get_lr(&self, step: usize) -> f32 {
283 if step < self.warmup_steps {
284 self.base_lr * (step as f32) / (self.warmup_steps as f32)
285 } else {
286 let decay_step = (step - self.warmup_steps) / self.decay_steps;
287 self.base_lr * self.decay_rate.powf(decay_step as f32)
288 }
289 }
290
291 fn step(&mut self) {
292 self.current_step += 1;
293 }
294}
295
296#[derive(Debug)]
298pub struct StepScheduler {
299 base_lr: f32,
300 warmup_steps: usize,
301 current_step: usize,
302 step_size: usize,
303 gamma: f32,
304}
305
306impl StepScheduler {
307 pub fn new(base_lr: f32, warmup_steps: usize, step_size: usize, gamma: f32) -> Self {
308 Self {
309 base_lr,
310 warmup_steps,
311 current_step: 0,
312 step_size,
313 gamma,
314 }
315 }
316}
317
318impl LRScheduler for StepScheduler {
319 fn get_lr(&self, step: usize) -> f32 {
320 if step < self.warmup_steps {
321 self.base_lr * (step as f32) / (self.warmup_steps as f32)
322 } else {
323 let decay_step = (step - self.warmup_steps) / self.step_size;
324 self.base_lr * self.gamma.powf(decay_step as f32)
325 }
326 }
327
328 fn step(&mut self) {
329 self.current_step += 1;
330 }
331}
332
333#[derive(Debug)]
339pub struct OneCycleScheduler {
340 max_lr: f32,
341 final_lr: f32,
342 total_steps: usize,
343 pct_start: f32,
344 current_step: usize,
345}
346
347impl OneCycleScheduler {
348 pub fn new(max_lr: f32, total_steps: usize, pct_start: f32, final_lr: f32) -> Self {
354 const MIN_PCT: f32 = 1e-3;
355 Self {
356 max_lr,
357 final_lr,
358 total_steps,
359 pct_start: if pct_start.is_finite() {
360 pct_start.clamp(MIN_PCT, 1.0 - MIN_PCT)
361 } else {
362 0.3
363 },
364 current_step: 0,
365 }
366 }
367}
368
369impl LRScheduler for OneCycleScheduler {
370 fn get_lr(&self, step: usize) -> f32 {
371 use std::f32::consts::PI;
372
373 if self.total_steps == 0 {
374 return self.final_lr;
375 }
376 let step = step.min(self.total_steps);
377 let pct = step as f32 / self.total_steps as f32;
378
379 if pct <= self.pct_start {
380 let phase_pct = pct / self.pct_start;
382 let cosine_term = 0.5 * (1.0 - (PI * phase_pct).cos());
383 self.final_lr + (self.max_lr - self.final_lr) * cosine_term
384 } else {
385 let remaining_pct = (pct - self.pct_start) / (1.0 - self.pct_start);
387 let cosine_term = 0.5 * (1.0 + (PI * remaining_pct).cos());
388 self.final_lr + (self.max_lr - self.final_lr) * cosine_term
389 }
390 }
391
392 fn step(&mut self) {
393 self.current_step += 1;
394 }
395}
396
397#[derive(Debug)]
402pub struct CosineWithRestartsScheduler {
403 base_lr: f32,
404 min_lr: f32,
405 t_0: usize,
406 t_mult: f32,
407 current_step: usize,
408 next_restart: usize,
409 current_t: usize,
410}
411
412impl CosineWithRestartsScheduler {
413 pub fn new(base_lr: f32, min_lr: f32, t_0: usize, t_mult: f32) -> Self {
423 let t_0 = t_0.max(1);
424 let t_mult = if t_mult.is_finite() { t_mult.max(1.0) } else { 1.0 };
425 Self {
426 base_lr,
427 min_lr,
428 t_0,
429 t_mult,
430 current_step: 0,
431 next_restart: t_0,
432 current_t: t_0,
433 }
434 }
435
436 pub fn try_new(
442 base_lr: f32,
443 min_lr: f32,
444 t_0: usize,
445 t_mult: f32,
446 ) -> trustformers_core::errors::Result<Self> {
447 if t_0 == 0 {
448 return Err(
449 trustformers_core::errors::TrustformersError::invalid_config(
450 "CosineWithRestartsScheduler requires t_0 >= 1".to_string(),
451 ),
452 );
453 }
454 if !t_mult.is_finite() || t_mult < 1.0 {
455 return Err(
456 trustformers_core::errors::TrustformersError::invalid_config(format!(
457 "CosineWithRestartsScheduler requires a finite t_mult >= 1.0, got {t_mult}"
458 )),
459 );
460 }
461 Ok(Self::new(base_lr, min_lr, t_0, t_mult))
462 }
463}
464
465impl LRScheduler for CosineWithRestartsScheduler {
466 fn get_lr(&self, step: usize) -> f32 {
467 use std::f32::consts::PI;
468
469 let mut step_in_cycle = step;
470 let mut cycle_length = self.t_0.max(1);
473
474 while step_in_cycle >= cycle_length {
476 step_in_cycle -= cycle_length;
477 cycle_length = ((cycle_length as f32 * self.t_mult) as usize).max(1);
478 }
479
480 let progress = step_in_cycle as f32 / cycle_length as f32;
481 let cosine_decay = 0.5 * (1.0 + (PI * progress).cos());
482
483 self.min_lr + (self.base_lr - self.min_lr) * cosine_decay
484 }
485
486 fn step(&mut self) {
487 self.current_step += 1;
488
489 if self.current_step >= self.next_restart {
490 self.current_t = (self.current_t as f32 * self.t_mult) as usize;
491 self.next_restart += self.current_t;
492 }
493 }
494}
495
496#[derive(Debug)]
501pub struct CyclicalScheduler {
502 base_lr: f32,
503 max_lr: f32,
504 step_size_up: usize,
505 step_size_down: usize,
506 current_step: usize,
507 mode: CyclicalMode,
508}
509
510#[derive(Debug, Clone)]
511pub enum CyclicalMode {
512 Triangular,
513 Triangular2,
514 ExpRange(f32), }
516
517impl CyclicalScheduler {
518 pub fn new(
519 base_lr: f32,
520 max_lr: f32,
521 step_size_up: usize,
522 step_size_down: usize,
523 mode: CyclicalMode,
524 ) -> Self {
525 Self {
526 base_lr,
527 max_lr,
528 step_size_up,
529 step_size_down,
530 current_step: 0,
531 mode,
532 }
533 }
534}
535
536impl LRScheduler for CyclicalScheduler {
537 fn get_lr(&self, step: usize) -> f32 {
538 let cycle_length = self.step_size_up + self.step_size_down;
539 let cycle = (step / cycle_length) + 1;
540 let x = (step % cycle_length) as f32;
541
542 let (amplitude, _phase) = if x <= self.step_size_up as f32 {
543 (x / self.step_size_up as f32, 1.0)
545 } else {
546 (
548 (self.step_size_down as f32 - (x - self.step_size_up as f32))
549 / self.step_size_down as f32,
550 1.0,
551 )
552 };
553
554 let scale_factor = match &self.mode {
555 CyclicalMode::Triangular => 1.0,
556 CyclicalMode::Triangular2 => 1.0 / (2.0_f32.powi((cycle - 1) as i32)),
557 CyclicalMode::ExpRange(gamma) => gamma.powi(step as i32),
558 };
559
560 self.base_lr + (self.max_lr - self.base_lr) * amplitude * scale_factor
561 }
562
563 fn step(&mut self) {
564 self.current_step += 1;
565 }
566}
567
568#[cfg(test)]
569mod tests {
570 use super::*;
571
572 #[test]
573 fn test_linear_scheduler() {
574 let scheduler = LinearScheduler::new(1e-3, 100, 1000);
575
576 assert_eq!(scheduler.get_lr(0), 0.0);
578 assert_eq!(scheduler.get_lr(50), 5e-4);
579 assert_eq!(scheduler.get_lr(100), 1e-3);
580
581 assert_eq!(scheduler.get_lr(550), 5e-4);
583 assert_eq!(scheduler.get_lr(1000), 0.0);
584 }
585
586 #[test]
587 fn test_cosine_scheduler() {
588 let scheduler = CosineScheduler::new(1e-3, 100, 1000, 1e-5);
589
590 assert_eq!(scheduler.get_lr(0), 0.0);
592 assert_eq!(scheduler.get_lr(50), 5e-4);
593 assert_eq!(scheduler.get_lr(100), 1e-3);
594
595 let mid_lr = scheduler.get_lr(550);
597 assert!(mid_lr > 1e-5 && mid_lr < 1e-3);
598
599 let end_lr = scheduler.get_lr(1000);
601 assert!((end_lr - 1e-5).abs() < 1e-6);
602 }
603
604 #[test]
605 fn test_polynomial_scheduler() {
606 let scheduler = PolynomialScheduler::new(1e-3, 100, 1000, 1e-5, 2.0);
607
608 assert_eq!(scheduler.get_lr(0), 0.0);
610 assert_eq!(scheduler.get_lr(100), 1e-3);
611
612 let mid_lr = scheduler.get_lr(550);
614 assert!(mid_lr > 1e-5 && mid_lr < 1e-3);
615 }
616
617 #[test]
618 fn test_constant_with_warmup_scheduler() {
619 let scheduler = ConstantWithWarmupScheduler::new(1e-3, 100);
620
621 assert_eq!(scheduler.get_lr(0), 0.0);
623 assert_eq!(scheduler.get_lr(50), 5e-4);
624 assert_eq!(scheduler.get_lr(100), 1e-3);
625
626 assert_eq!(scheduler.get_lr(200), 1e-3);
628 assert_eq!(scheduler.get_lr(1000), 1e-3);
629 }
630
631 #[test]
632 fn test_exponential_scheduler() {
633 let scheduler = ExponentialScheduler::new(1e-3, 100, 0.9, 100);
634
635 assert_eq!(scheduler.get_lr(0), 0.0);
637 assert_eq!(scheduler.get_lr(100), 1e-3);
638
639 assert_eq!(scheduler.get_lr(200), 1e-3 * 0.9);
641 assert_eq!(scheduler.get_lr(300), 1e-3 * 0.9 * 0.9);
642 }
643
644 #[test]
645 fn test_step_scheduler() {
646 let scheduler = StepScheduler::new(1e-3, 100, 200, 0.5);
647
648 assert_eq!(scheduler.get_lr(0), 0.0);
650 assert_eq!(scheduler.get_lr(100), 1e-3);
651
652 assert_eq!(scheduler.get_lr(250), 1e-3); assert_eq!(scheduler.get_lr(300), 1e-3 * 0.5); assert_eq!(scheduler.get_lr(500), 1e-3 * 0.5 * 0.5); }
657
658 #[test]
659 fn test_onecycle_scheduler() {
660 let scheduler = OneCycleScheduler::new(1e-2, 1000, 0.3, 1e-5);
661
662 assert_eq!(scheduler.get_lr(0), 1e-5);
664
665 let peak_lr = scheduler.get_lr(150);
667 assert!(peak_lr > 5e-3);
668
669 let end_lr = scheduler.get_lr(1000);
671 assert!((end_lr - 1e-5).abs() < 1e-6);
672 }
673
674 #[test]
675 fn test_cosine_with_restarts_scheduler() {
676 let scheduler = CosineWithRestartsScheduler::new(1e-3, 1e-5, 100, 2.0);
677
678 assert!((scheduler.get_lr(0) - 1e-3).abs() < 1e-6);
680
681 let mid_lr = scheduler.get_lr(50);
683 assert!(mid_lr > 1e-5 && mid_lr < 1e-3);
684
685 let near_end_lr = scheduler.get_lr(99);
687 assert!(near_end_lr < 2e-4);
688
689 let restart_lr = scheduler.get_lr(100);
691 assert!(restart_lr > 5e-4);
692 }
693
694 #[test]
695 fn test_cyclical_scheduler() {
696 let scheduler = CyclicalScheduler::new(1e-4, 1e-3, 50, 50, CyclicalMode::Triangular);
697
698 assert!((scheduler.get_lr(0) - 1e-4).abs() < 1e-6);
700
701 assert!((scheduler.get_lr(50) - 1e-3).abs() < 1e-6);
703
704 assert!((scheduler.get_lr(100) - 1e-4).abs() < 1e-6);
706
707 assert!((scheduler.get_lr(150) - 1e-3).abs() < 1e-6);
709 }
710}
711
712#[derive(Debug, Clone)]
718pub struct AdaptiveScheduler {
719 current_lr: f32,
721 factor: f32,
723 patience: usize,
725 threshold: f32,
727 min_lr: f32,
729 mode: String,
731 epochs_since_improvement: usize,
733 best_metric: Option<f32>,
735 current_step: usize,
737}
738
739impl AdaptiveScheduler {
740 pub fn new(
765 initial_lr: f32,
766 factor: f32,
767 patience: usize,
768 threshold: f32,
769 min_lr: f32,
770 mode: &str,
771 ) -> Self {
772 match Self::try_new(initial_lr, factor, patience, threshold, min_lr, mode) {
773 Ok(scheduler) => scheduler,
774 Err(error) => panic!("invalid AdaptiveScheduler configuration: {error}"),
775 }
776 }
777
778 pub fn try_new(
789 initial_lr: f32,
790 factor: f32,
791 patience: usize,
792 threshold: f32,
793 min_lr: f32,
794 mode: &str,
795 ) -> trustformers_core::errors::Result<Self> {
796 use trustformers_core::errors::TrustformersError;
797
798 if !(factor > 0.0 && factor < 1.0) {
799 return Err(TrustformersError::invalid_config(format!(
800 "AdaptiveScheduler factor must be in (0, 1), got {factor}"
801 )));
802 }
803 if patience == 0 {
804 return Err(TrustformersError::invalid_config(
805 "AdaptiveScheduler patience must be positive".to_string(),
806 ));
807 }
808 if threshold < 0.0 {
809 return Err(TrustformersError::invalid_config(format!(
810 "AdaptiveScheduler threshold must be non-negative, got {threshold}"
811 )));
812 }
813 if min_lr < 0.0 {
814 return Err(TrustformersError::invalid_config(format!(
815 "AdaptiveScheduler min_lr must be non-negative, got {min_lr}"
816 )));
817 }
818 if mode != "min" && mode != "max" {
819 return Err(TrustformersError::invalid_config(format!(
820 "AdaptiveScheduler mode must be \"min\" or \"max\", got \"{mode}\""
821 )));
822 }
823
824 Ok(Self {
825 current_lr: initial_lr,
826 factor,
827 patience,
828 threshold,
829 min_lr,
830 mode: mode.to_string(),
831 epochs_since_improvement: 0,
832 best_metric: None,
833 current_step: 0,
834 })
835 }
836
837 pub fn step_with_metric(&mut self, metric: f32) -> (f32, bool) {
840 self.current_step += 1;
841 let mut lr_reduced = false;
842
843 let is_improvement = match self.best_metric {
844 None => {
845 self.best_metric = Some(metric);
847 true
848 },
849 Some(best) => {
850 let improvement = if self.mode == "min" {
851 (best - metric) / best.abs().max(1e-8) > self.threshold
853 } else {
854 (metric - best) / best.abs().max(1e-8) > self.threshold
856 };
857
858 if improvement {
859 self.best_metric = Some(metric);
860 }
861
862 improvement
863 },
864 };
865
866 if is_improvement {
867 self.epochs_since_improvement = 0;
868 } else {
869 self.epochs_since_improvement += 1;
870
871 if self.epochs_since_improvement >= self.patience {
872 let new_lr = (self.current_lr * self.factor).max(self.min_lr);
874 if new_lr < self.current_lr {
875 self.current_lr = new_lr;
876 lr_reduced = true;
877 self.epochs_since_improvement = 0; }
879 }
880 }
881
882 (self.current_lr, lr_reduced)
883 }
884
885 pub fn get_current_lr(&self) -> f32 {
887 self.current_lr
888 }
889
890 pub fn get_best_metric(&self) -> Option<f32> {
892 self.best_metric
893 }
894
895 pub fn get_epochs_since_improvement(&self) -> usize {
897 self.epochs_since_improvement
898 }
899
900 pub fn reset(&mut self) {
902 self.epochs_since_improvement = 0;
903 self.best_metric = None;
904 self.current_step = 0;
905 }
906
907 pub fn set_lr(&mut self, lr: f32) {
909 self.current_lr = lr;
910 }
911}
912
913impl LRScheduler for AdaptiveScheduler {
914 fn get_lr(&self, _step: usize) -> f32 {
915 self.current_lr
916 }
917
918 fn step(&mut self) {
919 }
922}
923
924pub struct CompositeScheduler {
929 schedulers: Vec<Box<dyn LRScheduler>>,
930 step_boundaries: Vec<usize>,
931 current_step: usize,
932 #[allow(dead_code)]
935 global_step_offset: usize,
936}
937
938impl CompositeScheduler {
939 pub fn new(schedulers: Vec<Box<dyn LRScheduler>>, step_boundaries: Vec<usize>) -> Self {
962 match Self::try_new(schedulers, step_boundaries) {
963 Ok(scheduler) => scheduler,
964 Err(error) => panic!("invalid CompositeScheduler configuration: {error}"),
965 }
966 }
967
968 pub fn try_new(
975 schedulers: Vec<Box<dyn LRScheduler>>,
976 step_boundaries: Vec<usize>,
977 ) -> trustformers_core::errors::Result<Self> {
978 use trustformers_core::errors::TrustformersError;
979
980 if schedulers.is_empty() {
981 return Err(TrustformersError::invalid_config(
982 "CompositeScheduler needs at least one scheduler".to_string(),
983 ));
984 }
985 if schedulers.len() != step_boundaries.len() {
986 return Err(TrustformersError::invalid_config(format!(
987 "CompositeScheduler has {} schedulers but {} boundaries",
988 schedulers.len(),
989 step_boundaries.len()
990 )));
991 }
992
993 Ok(Self {
994 schedulers,
995 step_boundaries,
996 current_step: 0,
997 global_step_offset: 0,
998 })
999 }
1000
1001 fn get_active_scheduler_index(&self, step: usize) -> usize {
1002 for (i, &boundary) in self.step_boundaries.iter().enumerate() {
1003 if step < boundary {
1004 return i;
1005 }
1006 }
1007 self.schedulers.len() - 1
1008 }
1009
1010 fn get_local_step(&self, global_step: usize, scheduler_index: usize) -> usize {
1011 if scheduler_index == 0 {
1012 global_step
1013 } else {
1014 global_step - self.step_boundaries[scheduler_index - 1]
1015 }
1016 }
1017}
1018
1019impl LRScheduler for CompositeScheduler {
1020 fn get_lr(&self, step: usize) -> f32 {
1021 let scheduler_idx = self.get_active_scheduler_index(step);
1022 let local_step = self.get_local_step(step, scheduler_idx);
1023 self.schedulers[scheduler_idx].get_lr(local_step)
1024 }
1025
1026 fn step(&mut self) {
1027 self.current_step += 1;
1028 let _scheduler_idx = self.get_active_scheduler_index(self.current_step);
1029 }
1031}
1032
1033pub struct PhaseBasedScheduler {
1037 phases: Vec<Phase>,
1038 current_phase: usize,
1039 current_step: usize,
1040 phase_start_step: usize,
1041}
1042
1043pub struct Phase {
1044 pub name: String,
1045 pub scheduler: Box<dyn LRScheduler>,
1046 pub duration_steps: usize,
1047 pub lr_multiplier: f32,
1048}
1049
1050impl PhaseBasedScheduler {
1051 pub fn new(phases: Vec<Phase>) -> Self {
1084 match Self::try_new(phases) {
1085 Ok(scheduler) => scheduler,
1086 Err(error) => panic!("invalid PhaseBasedScheduler configuration: {error}"),
1087 }
1088 }
1089
1090 pub fn try_new(phases: Vec<Phase>) -> trustformers_core::errors::Result<Self> {
1096 if phases.is_empty() {
1097 return Err(
1098 trustformers_core::errors::TrustformersError::invalid_config(
1099 "PhaseBasedScheduler needs at least one phase".to_string(),
1100 ),
1101 );
1102 }
1103
1104 Ok(Self {
1105 phases,
1106 current_phase: 0,
1107 current_step: 0,
1108 phase_start_step: 0,
1109 })
1110 }
1111
1112 pub fn get_current_phase(&self) -> &str {
1114 &self.phases[self.current_phase].name
1115 }
1116
1117 pub fn get_current_phase_index(&self) -> usize {
1119 self.current_phase
1120 }
1121
1122 pub fn is_complete(&self) -> bool {
1124 self.current_phase >= self.phases.len()
1125 }
1126
1127 fn update_phase(&mut self, step: usize) {
1128 while self.current_phase < self.phases.len() {
1129 let phase_end = self.phase_start_step + self.phases[self.current_phase].duration_steps;
1130
1131 if step < phase_end {
1132 break; }
1134
1135 self.current_phase += 1;
1137 self.phase_start_step = phase_end;
1138 }
1139 }
1140}
1141
1142impl LRScheduler for PhaseBasedScheduler {
1143 fn get_lr(&self, step: usize) -> f32 {
1144 if self.current_phase >= self.phases.len() {
1145 return 0.0; }
1147
1148 let phase = &self.phases[self.current_phase];
1149 let phase_step = step - self.phase_start_step;
1150 let base_lr = phase.scheduler.get_lr(phase_step);
1151
1152 base_lr * phase.lr_multiplier
1153 }
1154
1155 fn step(&mut self) {
1156 self.current_step += 1;
1157 self.update_phase(self.current_step);
1158 }
1159}
1160
1161pub struct DynamicScheduler {
1166 primary_scheduler: Box<dyn LRScheduler>,
1167 fallback_scheduler: Box<dyn LRScheduler>,
1168 current_scheduler: usize, switch_condition: SwitchCondition,
1170 metrics_window: Vec<f32>,
1171 window_size: usize,
1172 current_step: usize,
1173}
1174
1175#[derive(Debug)]
1176pub enum SwitchCondition {
1177 LossPlateauSteps(usize),
1179 GradientNormThreshold(f32),
1181 StepThreshold(usize),
1183 LossIncreaseFactor(f32),
1185}
1186
1187impl DynamicScheduler {
1188 pub fn new(
1190 primary_scheduler: Box<dyn LRScheduler>,
1191 fallback_scheduler: Box<dyn LRScheduler>,
1192 switch_condition: SwitchCondition,
1193 window_size: usize,
1194 ) -> Self {
1195 Self {
1196 primary_scheduler,
1197 fallback_scheduler,
1198 current_scheduler: 0,
1199 switch_condition,
1200 metrics_window: Vec::with_capacity(window_size),
1201 window_size,
1202 current_step: 0,
1203 }
1204 }
1205
1206 pub fn update_metric(&mut self, metric: f32) {
1208 self.metrics_window.push(metric);
1209 if self.metrics_window.len() > self.window_size {
1210 self.metrics_window.remove(0);
1211 }
1212
1213 if self.current_scheduler == 0 && self.should_switch() {
1215 self.current_scheduler = 1;
1216 }
1217 }
1218
1219 fn should_switch(&self) -> bool {
1220 match &self.switch_condition {
1221 SwitchCondition::LossPlateauSteps(steps) => {
1222 if self.metrics_window.len() < *steps {
1223 return false;
1224 }
1225
1226 let recent_avg =
1227 self.metrics_window.iter().rev().take(*steps).sum::<f32>() / *steps as f32;
1228 let older_avg =
1229 self.metrics_window.iter().take(self.metrics_window.len() - steps).sum::<f32>()
1230 / (self.metrics_window.len() - steps) as f32;
1231
1232 recent_avg >= older_avg * 0.995 },
1234 SwitchCondition::StepThreshold(step) => self.current_step >= *step,
1235 SwitchCondition::LossIncreaseFactor(factor) => {
1236 if self.metrics_window.len() < 2 {
1237 return false;
1238 }
1239 let latest = self.metrics_window[self.metrics_window.len() - 1];
1240 let previous = self.metrics_window[self.metrics_window.len() - 2];
1241 latest > previous * factor
1242 },
1243 SwitchCondition::GradientNormThreshold(_) => false, }
1245 }
1246
1247 pub fn get_active_scheduler(&self) -> &str {
1249 if self.current_scheduler == 0 {
1250 "primary"
1251 } else {
1252 "fallback"
1253 }
1254 }
1255}
1256
1257impl LRScheduler for DynamicScheduler {
1258 fn get_lr(&self, step: usize) -> f32 {
1259 if self.current_scheduler == 0 {
1260 self.primary_scheduler.get_lr(step)
1261 } else {
1262 self.fallback_scheduler.get_lr(step)
1263 }
1264 }
1265
1266 fn step(&mut self) {
1267 self.current_step += 1;
1268 if self.current_scheduler == 0 {
1269 self.primary_scheduler.step();
1270 } else {
1271 self.fallback_scheduler.step();
1272 }
1273 }
1274}
1275
1276pub struct TaskSpecificScheduler {
1278 scheduler: Box<dyn LRScheduler>,
1279 task_type: TaskType,
1280 current_step: usize,
1281}
1282
1283#[derive(Debug)]
1284pub enum TaskType {
1285 LanguageModelPretraining,
1287 FineTuning,
1289 ComputerVision,
1291 ReinforcementLearning,
1293 GANTraining,
1295}
1296
1297impl TaskSpecificScheduler {
1298 pub fn new(task_type: TaskType, base_lr: f32, total_steps: usize) -> Self {
1300 let scheduler: Box<dyn LRScheduler> = match task_type {
1301 TaskType::LanguageModelPretraining => {
1302 Box::new(CosineScheduler::new(
1303 base_lr,
1304 (total_steps as f32 * 0.06) as usize, total_steps,
1306 base_lr * 0.1, ))
1308 },
1309 TaskType::FineTuning => {
1310 Box::new(LinearScheduler::new(
1311 base_lr * 0.1, (total_steps as f32 * 0.1) as usize, total_steps,
1314 ))
1315 },
1316 TaskType::ComputerVision => {
1317 Box::new(StepScheduler::new(
1318 base_lr,
1319 (total_steps as f32 * 0.05) as usize, total_steps / 3, 0.1, ))
1323 },
1324 TaskType::ReinforcementLearning => {
1325 Box::new(AdaptiveScheduler::new(
1326 base_lr,
1327 0.5, 10, 1e-4, base_lr * 1e-3, "max", ))
1333 },
1334 TaskType::GANTraining => {
1335 Box::new(ConstantWithWarmupScheduler::new(
1336 base_lr,
1337 (total_steps as f32 * 0.02) as usize, ))
1339 },
1340 };
1341
1342 Self {
1343 scheduler,
1344 task_type,
1345 current_step: 0,
1346 }
1347 }
1348
1349 pub fn get_task_type(&self) -> &TaskType {
1351 &self.task_type
1352 }
1353}
1354
1355impl LRScheduler for TaskSpecificScheduler {
1356 fn get_lr(&self, step: usize) -> f32 {
1357 self.scheduler.get_lr(step)
1358 }
1359
1360 fn step(&mut self) {
1361 self.current_step += 1;
1362 self.scheduler.step();
1363 }
1364}
1365
1366#[cfg(test)]
1367mod boundary_tests {
1368 use super::*;
1369
1370 #[test]
1373 fn cosine_with_restarts_never_hangs() {
1374 for (t_0, t_mult) in [(0_usize, 2.0_f32), (10, 0.5), (0, 0.0), (5, f32::NAN)] {
1375 let scheduler = CosineWithRestartsScheduler::new(1e-3, 1e-5, t_0, t_mult);
1376 let lr = scheduler.get_lr(1_000);
1378 assert!(lr.is_finite(), "t_0={t_0}, t_mult={t_mult} produced {lr}");
1379 }
1380 }
1381
1382 #[test]
1384 fn cosine_with_restarts_try_new_validates() {
1385 assert!(CosineWithRestartsScheduler::try_new(1e-3, 1e-5, 10, 2.0).is_ok());
1386 assert!(CosineWithRestartsScheduler::try_new(1e-3, 1e-5, 0, 2.0).is_err());
1387 assert!(CosineWithRestartsScheduler::try_new(1e-3, 1e-5, 10, 0.5).is_err());
1388 assert!(CosineWithRestartsScheduler::try_new(1e-3, 1e-5, 10, f32::NAN).is_err());
1389 }
1390
1391 #[test]
1394 fn cosine_lr_never_rises_after_the_schedule_ends() {
1395 let scheduler = CosineScheduler::new(1e-3, 100, 1000, 1e-5);
1396
1397 assert!((scheduler.get_lr(0) - 0.0).abs() < 1e-9);
1398 assert!((scheduler.get_lr(100) - 1e-3).abs() < 1e-9);
1399 let at_end = scheduler.get_lr(1000);
1400 assert!((at_end - 1e-5).abs() < 1e-7, "at total_steps: {at_end}");
1401
1402 for step in [1001_usize, 1500, 2000, 10_000] {
1403 let lr = scheduler.get_lr(step);
1404 assert!(
1405 (lr - 1e-5).abs() < 1e-7,
1406 "step {step} must stay at min_lr, got {lr}"
1407 );
1408 }
1409 }
1410
1411 #[test]
1413 fn degenerate_step_counts_do_not_underflow() {
1414 let cosine = CosineScheduler::new(1e-3, 1000, 100, 1e-5);
1415 assert!(cosine.get_lr(2000).is_finite());
1416
1417 let linear = LinearScheduler::new(1e-3, 1000, 100);
1418 assert!(linear.get_lr(2000).is_finite());
1419 }
1420
1421 #[test]
1423 fn one_cycle_endpoints_stay_finite() {
1424 for pct_start in [0.0_f32, 1.0, -1.0, 2.0, f32::NAN] {
1425 let scheduler = OneCycleScheduler::new(1e-2, 1000, pct_start, 1e-5);
1426 for step in [0_usize, 1, 500, 999, 1000, 5000] {
1427 let lr = scheduler.get_lr(step);
1428 assert!(lr.is_finite(), "pct_start={pct_start}, step={step} -> {lr}");
1429 }
1430 }
1431
1432 let empty = OneCycleScheduler::new(1e-2, 0, 0.3, 1e-5);
1434 assert!(empty.get_lr(0).is_finite());
1435 }
1436
1437 #[test]
1439 fn constructors_report_invalid_configuration() {
1440 assert!(AdaptiveScheduler::try_new(1e-3, 0.1, 5, 1e-4, 1e-8, "min").is_ok());
1441 assert!(AdaptiveScheduler::try_new(1e-3, 1.5, 5, 1e-4, 1e-8, "min").is_err());
1442 assert!(AdaptiveScheduler::try_new(1e-3, 0.1, 0, 1e-4, 1e-8, "min").is_err());
1443 assert!(AdaptiveScheduler::try_new(1e-3, 0.1, 5, -1.0, 1e-8, "min").is_err());
1444 assert!(AdaptiveScheduler::try_new(1e-3, 0.1, 5, 1e-4, -1.0, "min").is_err());
1445 assert!(AdaptiveScheduler::try_new(1e-3, 0.1, 5, 1e-4, 1e-8, "Min").is_err());
1447
1448 assert!(CompositeScheduler::try_new(Vec::new(), Vec::new()).is_err());
1449 assert!(CompositeScheduler::try_new(
1450 vec![Box::new(LinearScheduler::new(1e-3, 10, 100)) as Box<dyn LRScheduler>],
1451 vec![10, 20],
1452 )
1453 .is_err());
1454
1455 assert!(PhaseBasedScheduler::try_new(Vec::new()).is_err());
1456 }
1457}