1use crate::{Optimizer, OptimizerError, OptimizerResult};
4use torsh_core::error::{Result, TorshError};
5
6pub trait LRScheduler {
8 fn step(&mut self) -> OptimizerResult<()>;
10
11 fn step_with_metric(&mut self, metric: Option<f32>) -> OptimizerResult<()> {
13 self.step()
15 }
16
17 fn get_last_lr(&self) -> &[f32];
19
20 fn get_base_lrs(&self) -> &[f32];
22
23 fn get_last_epoch(&self) -> i32;
25
26 fn reset(&mut self);
28
29 fn state_dict(&self) -> SchedulerState;
31
32 fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()>;
34}
35
36#[macro_export]
38macro_rules! impl_base_scheduler_methods {
39 ($scheduler_type:ty, $scheduler_name:expr) => {
40 impl<O: Optimizer> LRScheduler for $scheduler_type {
41 fn step(&mut self) -> OptimizerResult<()> {
42 Ok(())
44 }
45
46 fn get_last_lr(&self) -> &[f32] {
47 &self.base.last_lr
48 }
49
50 fn get_base_lrs(&self) -> &[f32] {
51 &self.base.base_lrs
52 }
53
54 fn get_last_epoch(&self) -> i32 {
55 self.base.last_epoch
56 }
57
58 fn reset(&mut self) {
59 self.base.last_epoch = 0;
60 self.base.last_lr = self.base.base_lrs.clone();
61 let base_lrs = self.base.base_lrs.clone();
64 self.base.optimizer.set_lrs(&base_lrs);
65 }
66
67 fn state_dict(&self) -> SchedulerState {
68 let mut state = SchedulerState::new($scheduler_name.to_string());
69 state.last_epoch = self.base.last_epoch;
70 state.base_lrs = self.base.base_lrs.clone();
71 state.last_lr = self.base.last_lr.clone();
72 state
73 }
74
75 fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
76 self.base.last_epoch = state.last_epoch;
77 self.base.base_lrs = state.base_lrs;
78 self.base.last_lr = state.last_lr;
79 Ok(())
80 }
81 }
82 };
83}
84
85#[macro_export]
87macro_rules! impl_scheduler_with_state {
88 ($scheduler_type:ty, $scheduler_name:expr, $state_fields:expr, $load_state_fields:expr) => {
89 impl<O: Optimizer> LRScheduler for $scheduler_type {
90 fn step(&mut self) -> OptimizerResult<()> {
91 Ok(())
93 }
94
95 fn get_last_lr(&self) -> &[f32] {
96 &self.base.last_lr
97 }
98
99 fn get_base_lrs(&self) -> &[f32] {
100 &self.base.base_lrs
101 }
102
103 fn get_last_epoch(&self) -> i32 {
104 self.base.last_epoch
105 }
106
107 fn reset(&mut self) {
108 self.base.last_epoch = 0;
109 self.base.last_lr = self.base.base_lrs.clone();
110 let base_lrs = self.base.base_lrs.clone();
113 self.base.optimizer.set_lrs(&base_lrs);
114 }
115
116 fn state_dict(&self) -> SchedulerState {
117 let mut state = SchedulerState::new($scheduler_name.to_string());
118 state.last_epoch = self.base.last_epoch;
119 state.base_lrs = self.base.base_lrs.clone();
120 state.last_lr = self.base.last_lr.clone();
121
122 $state_fields(&self, &mut state);
124
125 state
126 }
127
128 fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
129 self.base.last_epoch = state.last_epoch;
130 self.base.base_lrs = state.base_lrs;
131 self.base.last_lr = state.last_lr;
132
133 $load_state_fields(self, &state)?;
135
136 Ok(())
137 }
138 }
139 };
140}
141
142#[derive(Debug, Clone)]
144pub struct SchedulerState {
145 pub scheduler_type: String,
146 pub last_epoch: i32,
147 pub base_lrs: Vec<f32>,
148 pub last_lr: Vec<f32>,
149 pub state: std::collections::HashMap<String, f32>,
150}
151
152impl SchedulerState {
153 pub fn new(scheduler_type: String) -> Self {
154 Self {
155 scheduler_type,
156 last_epoch: 0,
157 base_lrs: Vec::new(),
158 last_lr: Vec::new(),
159 state: std::collections::HashMap::new(),
160 }
161 }
162}
163
164pub struct BaseScheduler<O: Optimizer> {
166 pub optimizer: O,
167 pub base_lrs: Vec<f32>,
168 pub last_lr: Vec<f32>,
169 pub last_epoch: i32,
170}
171
172impl<O: Optimizer> BaseScheduler<O> {
173 pub fn new(optimizer: O) -> Self {
174 let base_lrs = optimizer.get_lr();
175 let last_lr = base_lrs.clone();
176
177 Self {
178 optimizer,
179 base_lrs,
180 last_lr,
181 last_epoch: 0,
182 }
183 }
184
185 pub fn set_learning_rates(&mut self, lrs: &[f32]) {
191 if !lrs.is_empty() {
192 if lrs.len() == 1 {
193 self.optimizer.set_lr(lrs[0]);
195 } else {
196 self.optimizer.set_lrs(lrs);
197 }
198 }
199 self.last_lr = lrs.to_vec();
200 }
201
202 pub fn optimizer_mut(&mut self) -> &mut O {
204 &mut self.optimizer
205 }
206
207 pub fn optimizer(&self) -> &O {
209 &self.optimizer
210 }
211
212 pub fn increment_epoch(&mut self) {
214 self.last_epoch += 1;
215 }
216
217 pub fn set_epoch(&mut self, epoch: i32) {
219 self.last_epoch = epoch;
220 }
221}
222
223impl<O: Optimizer> LRScheduler for BaseScheduler<O> {
224 fn step(&mut self) -> OptimizerResult<()> {
225 self.increment_epoch();
226 Ok(())
228 }
229
230 fn get_last_lr(&self) -> &[f32] {
231 &self.last_lr
232 }
233
234 fn get_base_lrs(&self) -> &[f32] {
235 &self.base_lrs
236 }
237
238 fn get_last_epoch(&self) -> i32 {
239 self.last_epoch
240 }
241
242 fn reset(&mut self) {
243 self.last_epoch = 0;
244 self.last_lr = self.base_lrs.clone();
245 self.set_learning_rates(&self.base_lrs.clone());
246 }
247
248 fn state_dict(&self) -> SchedulerState {
249 let mut state = SchedulerState::new("BaseScheduler".to_string());
250 state.last_epoch = self.last_epoch;
251 state.base_lrs = self.base_lrs.clone();
252 state.last_lr = self.last_lr.clone();
253 state
254 }
255
256 fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
257 self.last_epoch = state.last_epoch;
258 self.base_lrs = state.base_lrs;
259 self.last_lr = state.last_lr.clone();
260 self.set_learning_rates(&state.last_lr);
261 Ok(())
262 }
263}
264
265pub struct StepLR<O: Optimizer> {
267 base: BaseScheduler<O>,
268 step_size: i32,
269 gamma: f32,
270}
271
272impl<O: Optimizer> StepLR<O> {
273 pub fn new(optimizer: O, step_size: i32, gamma: f32) -> Self {
274 Self {
275 base: BaseScheduler::new(optimizer),
276 step_size,
277 gamma,
278 }
279 }
280}
281
282impl<O: Optimizer> LRScheduler for StepLR<O> {
283 fn step(&mut self) -> OptimizerResult<()> {
284 self.base.increment_epoch();
285
286 let new_lrs: Vec<f32> = self
287 .base
288 .base_lrs
289 .iter()
290 .map(|&base_lr| {
291 let num_steps = self.base.last_epoch / self.step_size;
292 base_lr * self.gamma.powi(num_steps)
293 })
294 .collect();
295
296 self.base.set_learning_rates(&new_lrs);
297 Ok(())
298 }
299
300 fn get_last_lr(&self) -> &[f32] {
301 self.base.get_last_lr()
302 }
303
304 fn get_base_lrs(&self) -> &[f32] {
305 self.base.get_base_lrs()
306 }
307
308 fn get_last_epoch(&self) -> i32 {
309 self.base.get_last_epoch()
310 }
311
312 fn reset(&mut self) {
313 self.base.reset()
314 }
315
316 fn state_dict(&self) -> SchedulerState {
317 let mut state = self.base.state_dict();
318 state.scheduler_type = "StepLR".to_string();
319 state
320 .state
321 .insert("step_size".to_string(), self.step_size as f32);
322 state.state.insert("gamma".to_string(), self.gamma);
323 state
324 }
325
326 fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
327 self.base.load_state_dict(state.clone())?;
328
329 if let Some(&step_size) = state.state.get("step_size") {
330 self.step_size = step_size as i32;
331 }
332 if let Some(&gamma) = state.state.get("gamma") {
333 self.gamma = gamma;
334 }
335
336 Ok(())
337 }
338}
339
340pub struct ExponentialLR<O: Optimizer> {
342 base: BaseScheduler<O>,
343 gamma: f32,
344}
345
346impl<O: Optimizer> ExponentialLR<O> {
347 pub fn new(optimizer: O, gamma: f32) -> Self {
348 Self {
349 base: BaseScheduler::new(optimizer),
350 gamma,
351 }
352 }
353
354 pub fn optimizer(&self) -> &O {
356 &self.base.optimizer
357 }
358
359 pub fn optimizer_mut(&mut self) -> &mut O {
361 &mut self.base.optimizer
362 }
363}
364
365impl<O: Optimizer> LRScheduler for ExponentialLR<O> {
366 fn step(&mut self) -> OptimizerResult<()> {
367 self.base.last_epoch += 1;
368
369 let new_lrs: Vec<f32> = self
370 .base
371 .base_lrs
372 .iter()
373 .map(|&base_lr| base_lr * self.gamma.powi(self.base.last_epoch))
374 .collect();
375
376 self.base.optimizer.set_lrs(&new_lrs);
377
378 self.base.last_lr = new_lrs;
379 Ok(())
380 }
381
382 fn get_last_lr(&self) -> &[f32] {
383 &self.base.last_lr
384 }
385
386 fn get_base_lrs(&self) -> &[f32] {
387 &self.base.base_lrs
388 }
389
390 fn get_last_epoch(&self) -> i32 {
391 self.base.last_epoch
392 }
393
394 fn reset(&mut self) {
395 self.base.last_epoch = 0;
396 self.base.last_lr = self.base.base_lrs.clone();
397 let base_lrs = self.base.base_lrs.clone();
400 self.base.optimizer.set_lrs(&base_lrs);
401 }
402
403 fn state_dict(&self) -> SchedulerState {
404 let mut state = SchedulerState::new("ExponentialLR".to_string());
405 state.last_epoch = self.base.last_epoch;
406 state.base_lrs = self.base.base_lrs.clone();
407 state.last_lr = self.base.last_lr.clone();
408 state.state.insert("gamma".to_string(), self.gamma);
409 state
410 }
411
412 fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
413 self.base.last_epoch = state.last_epoch;
414 self.base.base_lrs = state.base_lrs;
415 self.base.last_lr = state.last_lr;
416 if let Some(&gamma) = state.state.get("gamma") {
417 self.gamma = gamma;
418 }
419 Ok(())
420 }
421}
422
423pub struct CosineAnnealingLR<O: Optimizer> {
425 base: BaseScheduler<O>,
426 t_max: i32,
427 eta_min: f32,
428}
429
430impl<O: Optimizer> CosineAnnealingLR<O> {
431 pub fn new(optimizer: O, t_max: i32, eta_min: f32) -> Self {
432 Self {
433 base: BaseScheduler::new(optimizer),
434 t_max,
435 eta_min,
436 }
437 }
438}
439
440impl<O: Optimizer> LRScheduler for CosineAnnealingLR<O> {
441 fn step(&mut self) -> OptimizerResult<()> {
442 self.base.last_epoch += 1;
443
444 let new_lrs: Vec<f32> = self
445 .base
446 .base_lrs
447 .iter()
448 .map(|&base_lr| {
449 self.eta_min
450 + (base_lr - self.eta_min)
451 * (1.0
452 + (std::f32::consts::PI * self.base.last_epoch as f32
453 / self.t_max as f32)
454 .cos())
455 / 2.0
456 })
457 .collect();
458
459 self.base.optimizer.set_lrs(&new_lrs);
460
461 self.base.last_lr = new_lrs;
462 Ok(())
463 }
464
465 fn get_last_lr(&self) -> &[f32] {
466 &self.base.last_lr
467 }
468
469 fn get_base_lrs(&self) -> &[f32] {
470 &self.base.base_lrs
471 }
472
473 fn get_last_epoch(&self) -> i32 {
474 self.base.last_epoch
475 }
476
477 fn reset(&mut self) {
478 self.base.last_epoch = 0;
479 self.base.last_lr = self.base.base_lrs.clone();
480 let base_lrs = self.base.base_lrs.clone();
483 self.base.optimizer.set_lrs(&base_lrs);
484 }
485
486 fn state_dict(&self) -> SchedulerState {
487 let mut state = SchedulerState::new("CosineAnnealingLR".to_string());
488 state.last_epoch = self.base.last_epoch;
489 state.base_lrs = self.base.base_lrs.clone();
490 state.last_lr = self.base.last_lr.clone();
491 state.state.insert("t_max".to_string(), self.t_max as f32);
492 state.state.insert("eta_min".to_string(), self.eta_min);
493 state
494 }
495
496 fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
497 self.base.last_epoch = state.last_epoch;
498 self.base.base_lrs = state.base_lrs;
499 self.base.last_lr = state.last_lr;
500 if let Some(&t_max) = state.state.get("t_max") {
501 self.t_max = t_max as i32;
502 }
503 if let Some(&eta_min) = state.state.get("eta_min") {
504 self.eta_min = eta_min;
505 }
506 Ok(())
507 }
508}
509
510pub struct ReduceLROnPlateau<O: Optimizer> {
512 optimizer: O,
513 mode: String,
514 factor: f32,
515 patience: i32,
516 threshold: f32,
517 threshold_mode: String,
518 cooldown: i32,
519 min_lr: f32,
520 eps: f32,
521 best: Option<f32>,
522 num_bad_epochs: i32,
523 cooldown_counter: i32,
524}
525
526impl<O: Optimizer> ReduceLROnPlateau<O> {
527 #[allow(clippy::too_many_arguments)]
528 pub fn new(
529 optimizer: O,
530 mode: &str,
531 factor: f32,
532 patience: i32,
533 threshold: f32,
534 threshold_mode: &str,
535 cooldown: i32,
536 min_lr: f32,
537 eps: f32,
538 ) -> Result<Self> {
539 if factor >= 1.0 {
540 return Err(TorshError::Other("Factor should be < 1.0".to_string()));
541 }
542
543 Ok(Self {
544 optimizer,
545 mode: mode.to_string(),
546 factor,
547 patience,
548 threshold,
549 threshold_mode: threshold_mode.to_string(),
550 cooldown,
551 min_lr,
552 eps,
553 best: None,
554 num_bad_epochs: 0,
555 cooldown_counter: 0,
556 })
557 }
558
559 pub fn step(&mut self, metrics: f32) {
560 let current = metrics;
561
562 if self.best.is_none() {
563 self.best = Some(current);
564 } else {
565 let best_value = self.best.expect("best should exist after is_none check");
566 let is_better = match self.mode.as_str() {
567 "min" => match self.threshold_mode.as_str() {
568 "rel" => current < best_value * (1.0 - self.threshold),
569 "abs" => current < best_value - self.threshold,
570 _ => false,
571 },
572 "max" => match self.threshold_mode.as_str() {
573 "rel" => current > best_value * (1.0 + self.threshold),
574 "abs" => current > best_value + self.threshold,
575 _ => false,
576 },
577 _ => false,
578 };
579
580 if is_better {
581 self.best = Some(current);
582 self.num_bad_epochs = 0;
583 } else {
584 self.num_bad_epochs += 1;
585 }
586
587 if self.cooldown_counter > 0 {
588 self.cooldown_counter -= 1;
589 self.num_bad_epochs = 0;
590 }
591
592 if self.num_bad_epochs > self.patience {
593 self.reduce_lr();
594 self.cooldown_counter = self.cooldown;
595 self.num_bad_epochs = 0;
596 }
597 }
598 }
599
600 fn reduce_lr(&mut self) {
601 let old_lrs = self.optimizer.get_lr();
602 let new_lrs: Vec<f32> = old_lrs
603 .iter()
604 .map(|&lr| (lr * self.factor).max(self.min_lr))
605 .collect();
606
607 let mut applied: Vec<f32> = old_lrs.clone();
610 let mut any_reduced = false;
611 for (idx, (old_lr, new_lr)) in old_lrs.iter().zip(new_lrs.iter()).enumerate() {
612 if old_lr - new_lr > self.eps {
613 applied[idx] = *new_lr;
614 any_reduced = true;
615 log::debug!("Reducing learning rate from {old_lr} to {new_lr}");
616 }
617 }
618 if any_reduced {
619 self.optimizer.set_lrs(&applied);
620 }
621 }
622}
623
624pub struct OneCycleLR<O: Optimizer> {
626 base: BaseScheduler<O>,
627 max_lr: Vec<f32>,
628 total_steps: i32,
629 pct_start: f32,
630 anneal_strategy: String,
631 #[allow(dead_code)]
632 cycle_momentum: bool,
633 #[allow(dead_code)]
634 base_momentum: f32,
635 #[allow(dead_code)]
636 max_momentum: f32,
637 #[allow(dead_code)]
638 div_factor: f32,
639 final_div_factor: f32,
640 step_count: i32,
641}
642
643impl<O: Optimizer> OneCycleLR<O> {
644 #[allow(clippy::too_many_arguments)]
645 pub fn new(
646 optimizer: O,
647 max_lr: Vec<f32>,
648 total_steps: i32,
649 pct_start: Option<f32>,
650 anneal_strategy: Option<&str>,
651 cycle_momentum: Option<bool>,
652 base_momentum: Option<f32>,
653 max_momentum: Option<f32>,
654 div_factor: Option<f32>,
655 final_div_factor: Option<f32>,
656 ) -> Self {
657 let pct_start = pct_start.unwrap_or(0.3);
658 let anneal_strategy = anneal_strategy.unwrap_or("cos").to_string();
659 let cycle_momentum = cycle_momentum.unwrap_or(true);
660 let base_momentum = base_momentum.unwrap_or(0.85);
661 let max_momentum = max_momentum.unwrap_or(0.95);
662 let div_factor = div_factor.unwrap_or(25.0);
663 let final_div_factor = final_div_factor.unwrap_or(10000.0);
664
665 let mut base = BaseScheduler::new(optimizer);
666
667 base.base_lrs = max_lr.iter().map(|&lr| lr / div_factor).collect();
669
670 Self {
671 base,
672 max_lr,
673 total_steps,
674 pct_start,
675 anneal_strategy,
676 cycle_momentum,
677 base_momentum,
678 max_momentum,
679 div_factor,
680 final_div_factor,
681 step_count: 0,
682 }
683 }
684}
685
686impl<O: Optimizer> LRScheduler for OneCycleLR<O> {
687 fn step(&mut self) -> OptimizerResult<()> {
688 self.step_count += 1;
689
690 let step_num = self.step_count as f32;
691 let step_size_up = (self.pct_start * self.total_steps as f32).floor();
692 let step_size_down = self.total_steps as f32 - step_size_up;
693
694 let new_lrs: Vec<f32> = if step_num <= step_size_up {
695 let computed_lr =
697 |base_lr: f32, max_lr: f32| base_lr + (max_lr - base_lr) * step_num / step_size_up;
698
699 self.base
700 .base_lrs
701 .iter()
702 .zip(self.max_lr.iter())
703 .map(|(&base, &max)| computed_lr(base, max))
704 .collect()
705 } else {
706 let down_step_num = step_num - step_size_up;
708 match self.anneal_strategy.as_str() {
709 "cos" => {
710 let computed_lr = |max_lr: f32, base_lr: f32| {
711 let min_lr = base_lr / self.final_div_factor;
712 min_lr
713 + (max_lr - min_lr)
714 * (1.0
715 + (std::f32::consts::PI * down_step_num / step_size_down).cos())
716 / 2.0
717 };
718
719 self.max_lr
720 .iter()
721 .zip(self.base.base_lrs.iter())
722 .map(|(&max, &base)| computed_lr(max, base))
723 .collect()
724 }
725 "linear" => {
726 let computed_lr = |max_lr: f32, base_lr: f32| {
727 let min_lr = base_lr / self.final_div_factor;
728 max_lr - (max_lr - min_lr) * down_step_num / step_size_down
729 };
730
731 self.max_lr
732 .iter()
733 .zip(self.base.base_lrs.iter())
734 .map(|(&max, &base)| computed_lr(max, base))
735 .collect()
736 }
737 _ => {
738 return Err(OptimizerError::InvalidParameter(format!(
739 "Unknown anneal strategy: {}",
740 self.anneal_strategy
741 )))
742 }
743 }
744 };
745
746 self.base.optimizer.set_lrs(&new_lrs);
748
749 self.base.last_lr = new_lrs;
750 Ok(())
751 }
752
753 fn get_last_lr(&self) -> &[f32] {
754 &self.base.last_lr
755 }
756
757 fn get_base_lrs(&self) -> &[f32] {
758 &self.base.base_lrs
759 }
760
761 fn get_last_epoch(&self) -> i32 {
762 self.step_count
763 }
764
765 fn reset(&mut self) {
766 self.step_count = 0;
767 self.base.last_lr = self.base.base_lrs.clone();
768 let base_lrs = self.base.base_lrs.clone();
771 self.base.optimizer.set_lrs(&base_lrs);
772 }
773
774 fn state_dict(&self) -> SchedulerState {
775 let mut state = SchedulerState::new("OneCycleLR".to_string());
776 state.last_epoch = self.step_count;
777 state.base_lrs = self.base.base_lrs.clone();
778 state.last_lr = self.base.last_lr.clone();
779 state
780 .state
781 .insert("total_steps".to_string(), self.total_steps as f32);
782 state.state.insert("pct_start".to_string(), self.pct_start);
783 state
784 .state
785 .insert("final_div_factor".to_string(), self.final_div_factor);
786 state
787 }
788
789 fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
790 self.step_count = state.last_epoch;
791 self.base.base_lrs = state.base_lrs;
792 self.base.last_lr = state.last_lr;
793 if let Some(&total_steps) = state.state.get("total_steps") {
794 self.total_steps = total_steps as i32;
795 }
796 if let Some(&pct_start) = state.state.get("pct_start") {
797 self.pct_start = pct_start;
798 }
799 if let Some(&final_div_factor) = state.state.get("final_div_factor") {
800 self.final_div_factor = final_div_factor;
801 }
802 Ok(())
803 }
804}