Skip to main content

torsh_optim/
lr_scheduler_additional.rs

1//! Additional learning rate schedulers
2
3use crate::{
4    lr_scheduler::{BaseScheduler, LRScheduler, SchedulerState},
5    Optimizer, OptimizerError, OptimizerResult,
6};
7
8/// Multi-step learning rate scheduler
9pub struct MultiStepLR<O: Optimizer> {
10    base: BaseScheduler<O>,
11    milestones: Vec<i32>,
12    gamma: f32,
13}
14
15impl<O: Optimizer> MultiStepLR<O> {
16    pub fn new(optimizer: O, milestones: Vec<i32>, gamma: f32) -> Self {
17        let mut milestones = milestones;
18        milestones.sort_unstable();
19
20        Self {
21            base: BaseScheduler::new(optimizer),
22            milestones,
23            gamma,
24        }
25    }
26}
27
28impl<O: Optimizer> LRScheduler for MultiStepLR<O> {
29    fn step(&mut self) -> OptimizerResult<()> {
30        self.base.last_epoch += 1;
31
32        let num_milestones_passed = self
33            .milestones
34            .iter()
35            .filter(|&&milestone| self.base.last_epoch >= milestone)
36            .count() as i32;
37
38        let new_lrs: Vec<f32> = self
39            .base
40            .base_lrs
41            .iter()
42            .map(|&base_lr| base_lr * self.gamma.powi(num_milestones_passed))
43            .collect();
44
45        self.base.optimizer.set_lrs(&new_lrs);
46
47        self.base.last_lr = new_lrs;
48        Ok(())
49    }
50
51    fn get_last_lr(&self) -> &[f32] {
52        &self.base.last_lr
53    }
54
55    fn get_base_lrs(&self) -> &[f32] {
56        &self.base.base_lrs
57    }
58
59    fn get_last_epoch(&self) -> i32 {
60        self.base.last_epoch
61    }
62
63    fn reset(&mut self) {
64        self.base.last_epoch = 0;
65        self.base.last_lr = self.base.base_lrs.clone();
66        // Push the base rates back onto the optimizer so a reset actually
67        // restores them (per group, not broadcast).
68        let base_lrs = self.base.base_lrs.clone();
69        self.base.optimizer.set_lrs(&base_lrs);
70    }
71
72    fn state_dict(&self) -> SchedulerState {
73        let mut state = SchedulerState::new("MultiStepLR".to_string());
74        state.last_epoch = self.base.last_epoch;
75        state.base_lrs = self.base.base_lrs.clone();
76        state.last_lr = self.base.last_lr.clone();
77        state.state.insert("gamma".to_string(), self.gamma);
78        for (i, &milestone) in self.milestones.iter().enumerate() {
79            state
80                .state
81                .insert(format!("milestone_{}", i), milestone as f32);
82        }
83        state
84    }
85
86    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
87        self.base.last_epoch = state.last_epoch;
88        self.base.base_lrs = state.base_lrs;
89        self.base.last_lr = state.last_lr;
90        if let Some(&gamma) = state.state.get("gamma") {
91            self.gamma = gamma;
92        }
93        // Note: milestones are immutable after construction
94        Ok(())
95    }
96}
97
98/// Cyclic learning rate scheduler
99pub struct CyclicLR<O: Optimizer> {
100    base: BaseScheduler<O>,
101    base_lr: Vec<f32>,
102    max_lr: Vec<f32>,
103    step_size_up: i32,
104    step_size_down: Option<i32>,
105    mode: String,
106    gamma: f32,
107    scale_fn: Option<Box<dyn Fn(i32) -> f32>>,
108    scale_mode: String,
109    cycle: i32,
110    step_in_cycle: i32,
111}
112
113impl<O: Optimizer> CyclicLR<O> {
114    #[allow(clippy::too_many_arguments)]
115    pub fn new(
116        optimizer: O,
117        base_lr: Vec<f32>,
118        max_lr: Vec<f32>,
119        step_size_up: i32,
120        step_size_down: Option<i32>,
121        mode: Option<&str>,
122        gamma: Option<f32>,
123        scale_fn: Option<Box<dyn Fn(i32) -> f32>>,
124        scale_mode: Option<&str>,
125    ) -> Self {
126        let mode = mode.unwrap_or("triangular").to_string();
127        let gamma = gamma.unwrap_or(1.0);
128        let scale_mode = scale_mode.unwrap_or("cycle").to_string();
129
130        Self {
131            base: BaseScheduler::new(optimizer),
132            base_lr,
133            max_lr,
134            step_size_up,
135            step_size_down,
136            mode,
137            gamma,
138            scale_fn,
139            scale_mode,
140            cycle: 0,
141            step_in_cycle: 0,
142        }
143    }
144
145    fn get_cycle_length(&self) -> i32 {
146        self.step_size_up + self.step_size_down.unwrap_or(self.step_size_up)
147    }
148}
149
150impl<O: Optimizer> LRScheduler for CyclicLR<O> {
151    fn step(&mut self) -> OptimizerResult<()> {
152        self.step_in_cycle += 1;
153
154        if self.step_in_cycle >= self.get_cycle_length() {
155            self.step_in_cycle = 0;
156            self.cycle += 1;
157        }
158
159        let step_size_down = self.step_size_down.unwrap_or(self.step_size_up);
160
161        let scale_factor = match self.mode.as_str() {
162            "triangular" => 1.0,
163            "triangular2" => 1.0 / (2.0_f32.powi(self.cycle)),
164            "exp_range" => self.gamma.powi(self.base.last_epoch),
165            _ => {
166                if let Some(ref scale_fn) = self.scale_fn {
167                    match self.scale_mode.as_str() {
168                        "cycle" => scale_fn(self.cycle),
169                        _ => scale_fn(self.base.last_epoch),
170                    }
171                } else {
172                    1.0
173                }
174            }
175        };
176
177        let new_lrs: Vec<f32> = if self.step_in_cycle < self.step_size_up {
178            // Ascending
179            let pct = self.step_in_cycle as f32 / self.step_size_up as f32;
180            self.base_lr
181                .iter()
182                .zip(self.max_lr.iter())
183                .map(|(&base, &max)| base + (max - base) * pct * scale_factor)
184                .collect()
185        } else {
186            // Descending
187            let down_step = self.step_in_cycle - self.step_size_up;
188            let pct = 1.0 - (down_step as f32 / step_size_down as f32);
189            self.base_lr
190                .iter()
191                .zip(self.max_lr.iter())
192                .map(|(&base, &max)| base + (max - base) * pct * scale_factor)
193                .collect()
194        };
195
196        self.base.optimizer.set_lrs(&new_lrs);
197
198        self.base.last_lr = new_lrs;
199        self.base.last_epoch += 1;
200        Ok(())
201    }
202
203    fn get_last_lr(&self) -> &[f32] {
204        &self.base.last_lr
205    }
206
207    fn get_base_lrs(&self) -> &[f32] {
208        &self.base.base_lrs
209    }
210
211    fn get_last_epoch(&self) -> i32 {
212        self.base.last_epoch
213    }
214
215    fn reset(&mut self) {
216        self.base.last_epoch = 0;
217        self.base.last_lr = self.base.base_lrs.clone();
218        // Push the base rates back onto the optimizer so a reset actually
219        // restores them (per group, not broadcast).
220        let base_lrs = self.base.base_lrs.clone();
221        self.base.optimizer.set_lrs(&base_lrs);
222        self.cycle = 0;
223        self.step_in_cycle = 0;
224    }
225
226    fn state_dict(&self) -> SchedulerState {
227        let mut state = SchedulerState::new("CyclicLR".to_string());
228        state.last_epoch = self.base.last_epoch;
229        state.base_lrs = self.base.base_lrs.clone();
230        state.last_lr = self.base.last_lr.clone();
231        state
232            .state
233            .insert("step_size_up".to_string(), self.step_size_up as f32);
234        if let Some(step_size_down) = self.step_size_down {
235            state
236                .state
237                .insert("step_size_down".to_string(), step_size_down as f32);
238        }
239        state.state.insert("gamma".to_string(), self.gamma);
240        state.state.insert("cycle".to_string(), self.cycle as f32);
241        state
242            .state
243            .insert("step_in_cycle".to_string(), self.step_in_cycle as f32);
244        state
245    }
246
247    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
248        self.base.last_epoch = state.last_epoch;
249        self.base.base_lrs = state.base_lrs;
250        self.base.last_lr = state.last_lr;
251        if let Some(&step_size_up) = state.state.get("step_size_up") {
252            self.step_size_up = step_size_up as i32;
253        }
254        if let Some(&step_size_down) = state.state.get("step_size_down") {
255            self.step_size_down = Some(step_size_down as i32);
256        }
257        if let Some(&gamma) = state.state.get("gamma") {
258            self.gamma = gamma;
259        }
260        if let Some(&cycle) = state.state.get("cycle") {
261            self.cycle = cycle as i32;
262        }
263        if let Some(&step_in_cycle) = state.state.get("step_in_cycle") {
264            self.step_in_cycle = step_in_cycle as i32;
265        }
266        Ok(())
267    }
268}
269
270/// Polynomial learning rate scheduler
271pub struct PolynomialLR<O: Optimizer> {
272    base: BaseScheduler<O>,
273    total_iters: i32,
274    power: f32,
275}
276
277impl<O: Optimizer> PolynomialLR<O> {
278    pub fn new(optimizer: O, total_iters: i32, power: f32) -> Self {
279        Self {
280            base: BaseScheduler::new(optimizer),
281            total_iters,
282            power,
283        }
284    }
285}
286
287impl<O: Optimizer> LRScheduler for PolynomialLR<O> {
288    fn step(&mut self) -> OptimizerResult<()> {
289        self.base.last_epoch += 1;
290
291        let factor = if self.base.last_epoch > self.total_iters {
292            0.0
293        } else {
294            (1.0 - self.base.last_epoch as f32 / self.total_iters as f32).powf(self.power)
295        };
296
297        let new_lrs: Vec<f32> = self
298            .base
299            .base_lrs
300            .iter()
301            .map(|&base_lr| base_lr * factor)
302            .collect();
303
304        self.base.optimizer.set_lrs(&new_lrs);
305
306        self.base.last_lr = new_lrs;
307        Ok(())
308    }
309
310    fn get_last_lr(&self) -> &[f32] {
311        &self.base.last_lr
312    }
313
314    fn get_base_lrs(&self) -> &[f32] {
315        &self.base.base_lrs
316    }
317
318    fn get_last_epoch(&self) -> i32 {
319        self.base.last_epoch
320    }
321
322    fn reset(&mut self) {
323        self.base.last_epoch = 0;
324        self.base.last_lr = self.base.base_lrs.clone();
325        // Push the base rates back onto the optimizer so a reset actually
326        // restores them (per group, not broadcast).
327        let base_lrs = self.base.base_lrs.clone();
328        self.base.optimizer.set_lrs(&base_lrs);
329    }
330
331    fn state_dict(&self) -> SchedulerState {
332        let mut state = SchedulerState::new("PolynomialLR".to_string());
333        state.last_epoch = self.base.last_epoch;
334        state.base_lrs = self.base.base_lrs.clone();
335        state.last_lr = self.base.last_lr.clone();
336        state
337            .state
338            .insert("total_iters".to_string(), self.total_iters as f32);
339        state.state.insert("power".to_string(), self.power);
340        state
341    }
342
343    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
344        self.base.last_epoch = state.last_epoch;
345        self.base.base_lrs = state.base_lrs;
346        self.base.last_lr = state.last_lr;
347        if let Some(&total_iters) = state.state.get("total_iters") {
348            self.total_iters = total_iters as i32;
349        }
350        if let Some(&power) = state.state.get("power") {
351            self.power = power;
352        }
353        Ok(())
354    }
355}
356
357/// Linear learning rate scheduler
358pub struct LinearLR<O: Optimizer> {
359    base: BaseScheduler<O>,
360    start_factor: f32,
361    end_factor: f32,
362    total_iters: i32,
363}
364
365impl<O: Optimizer> LinearLR<O> {
366    pub fn new(optimizer: O, start_factor: f32, end_factor: f32, total_iters: i32) -> Self {
367        Self {
368            base: BaseScheduler::new(optimizer),
369            start_factor,
370            end_factor,
371            total_iters,
372        }
373    }
374}
375
376impl<O: Optimizer> LRScheduler for LinearLR<O> {
377    fn step(&mut self) -> OptimizerResult<()> {
378        self.base.last_epoch += 1;
379
380        let factor = if self.base.last_epoch >= self.total_iters {
381            self.end_factor
382        } else {
383            self.start_factor
384                + (self.end_factor - self.start_factor)
385                    * (self.base.last_epoch as f32 / self.total_iters as f32)
386        };
387
388        let new_lrs: Vec<f32> = self
389            .base
390            .base_lrs
391            .iter()
392            .map(|&base_lr| base_lr * factor)
393            .collect();
394
395        self.base.optimizer.set_lrs(&new_lrs);
396
397        self.base.last_lr = new_lrs;
398        Ok(())
399    }
400
401    fn get_last_lr(&self) -> &[f32] {
402        &self.base.last_lr
403    }
404
405    fn get_base_lrs(&self) -> &[f32] {
406        &self.base.base_lrs
407    }
408
409    fn get_last_epoch(&self) -> i32 {
410        self.base.last_epoch
411    }
412
413    fn reset(&mut self) {
414        self.base.last_epoch = 0;
415        self.base.last_lr = self.base.base_lrs.clone();
416        // Push the base rates back onto the optimizer so a reset actually
417        // restores them (per group, not broadcast).
418        let base_lrs = self.base.base_lrs.clone();
419        self.base.optimizer.set_lrs(&base_lrs);
420    }
421
422    fn state_dict(&self) -> SchedulerState {
423        let mut state = SchedulerState::new("LinearLR".to_string());
424        state.last_epoch = self.base.last_epoch;
425        state.base_lrs = self.base.base_lrs.clone();
426        state.last_lr = self.base.last_lr.clone();
427        state
428            .state
429            .insert("start_factor".to_string(), self.start_factor);
430        state
431            .state
432            .insert("end_factor".to_string(), self.end_factor);
433        state
434            .state
435            .insert("total_iters".to_string(), self.total_iters as f32);
436        state
437    }
438
439    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
440        self.base.last_epoch = state.last_epoch;
441        self.base.base_lrs = state.base_lrs;
442        self.base.last_lr = state.last_lr;
443        if let Some(&start_factor) = state.state.get("start_factor") {
444            self.start_factor = start_factor;
445        }
446        if let Some(&end_factor) = state.state.get("end_factor") {
447            self.end_factor = end_factor;
448        }
449        if let Some(&total_iters) = state.state.get("total_iters") {
450            self.total_iters = total_iters as i32;
451        }
452        Ok(())
453    }
454}
455
456/// Constant learning rate scheduler
457pub struct ConstantLR<O: Optimizer> {
458    base: BaseScheduler<O>,
459    factor: f32,
460    total_iters: i32,
461}
462
463impl<O: Optimizer> ConstantLR<O> {
464    pub fn new(optimizer: O, factor: f32, total_iters: i32) -> Self {
465        Self {
466            base: BaseScheduler::new(optimizer),
467            factor,
468            total_iters,
469        }
470    }
471}
472
473impl<O: Optimizer> LRScheduler for ConstantLR<O> {
474    fn step(&mut self) -> OptimizerResult<()> {
475        self.base.last_epoch += 1;
476
477        let factor = if self.base.last_epoch < self.total_iters {
478            self.factor
479        } else {
480            1.0
481        };
482
483        let new_lrs: Vec<f32> = self
484            .base
485            .base_lrs
486            .iter()
487            .map(|&base_lr| base_lr * factor)
488            .collect();
489
490        self.base.optimizer.set_lrs(&new_lrs);
491
492        self.base.last_lr = new_lrs;
493        Ok(())
494    }
495
496    fn get_last_lr(&self) -> &[f32] {
497        &self.base.last_lr
498    }
499
500    fn get_base_lrs(&self) -> &[f32] {
501        &self.base.base_lrs
502    }
503
504    fn get_last_epoch(&self) -> i32 {
505        self.base.last_epoch
506    }
507
508    fn reset(&mut self) {
509        self.base.last_epoch = 0;
510        self.base.last_lr = self.base.base_lrs.clone();
511        // Push the base rates back onto the optimizer so a reset actually
512        // restores them (per group, not broadcast).
513        let base_lrs = self.base.base_lrs.clone();
514        self.base.optimizer.set_lrs(&base_lrs);
515    }
516
517    fn state_dict(&self) -> SchedulerState {
518        let mut state = SchedulerState::new("ConstantLR".to_string());
519        state.last_epoch = self.base.last_epoch;
520        state.base_lrs = self.base.base_lrs.clone();
521        state.last_lr = self.base.last_lr.clone();
522        state.state.insert("factor".to_string(), self.factor);
523        state
524            .state
525            .insert("total_iters".to_string(), self.total_iters as f32);
526        state
527    }
528
529    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
530        self.base.last_epoch = state.last_epoch;
531        self.base.base_lrs = state.base_lrs;
532        self.base.last_lr = state.last_lr;
533        if let Some(&factor) = state.state.get("factor") {
534            self.factor = factor;
535        }
536        if let Some(&total_iters) = state.state.get("total_iters") {
537            self.total_iters = total_iters as i32;
538        }
539        Ok(())
540    }
541}
542
543/// Cosine annealing with warm restarts
544pub struct CosineAnnealingWarmRestarts<O: Optimizer> {
545    base: BaseScheduler<O>,
546    t_0: i32,
547    t_mult: i32,
548    eta_min: f32,
549    t_cur: i32,
550}
551
552impl<O: Optimizer> CosineAnnealingWarmRestarts<O> {
553    pub fn new(optimizer: O, t_0: i32, t_mult: i32, eta_min: f32) -> Self {
554        Self {
555            base: BaseScheduler::new(optimizer),
556            t_0,
557            t_mult,
558            eta_min,
559            t_cur: -1,
560        }
561    }
562}
563
564impl<O: Optimizer> LRScheduler for CosineAnnealingWarmRestarts<O> {
565    fn step(&mut self) -> OptimizerResult<()> {
566        self.t_cur += 1;
567
568        if self.t_cur >= self.t_0 {
569            self.t_cur = 0;
570            self.t_0 *= self.t_mult;
571        }
572
573        let new_lrs: Vec<f32> = self
574            .base
575            .base_lrs
576            .iter()
577            .map(|&base_lr| {
578                self.eta_min
579                    + (base_lr - self.eta_min)
580                        * (1.0 + (std::f32::consts::PI * self.t_cur as f32 / self.t_0 as f32).cos())
581                        / 2.0
582            })
583            .collect();
584
585        self.base.optimizer.set_lrs(&new_lrs);
586
587        self.base.last_lr = new_lrs;
588        self.base.last_epoch += 1;
589        Ok(())
590    }
591
592    fn get_last_lr(&self) -> &[f32] {
593        &self.base.last_lr
594    }
595
596    fn get_base_lrs(&self) -> &[f32] {
597        &self.base.base_lrs
598    }
599
600    fn get_last_epoch(&self) -> i32 {
601        self.base.last_epoch
602    }
603
604    fn reset(&mut self) {
605        self.base.last_epoch = 0;
606        self.base.last_lr = self.base.base_lrs.clone();
607        // Push the base rates back onto the optimizer so a reset actually
608        // restores them (per group, not broadcast).
609        let base_lrs = self.base.base_lrs.clone();
610        self.base.optimizer.set_lrs(&base_lrs);
611        self.t_cur = -1;
612    }
613
614    fn state_dict(&self) -> SchedulerState {
615        let mut state = SchedulerState::new("CosineAnnealingWarmRestarts".to_string());
616        state.last_epoch = self.base.last_epoch;
617        state.base_lrs = self.base.base_lrs.clone();
618        state.last_lr = self.base.last_lr.clone();
619        state.state.insert("t_0".to_string(), self.t_0 as f32);
620        state.state.insert("t_mult".to_string(), self.t_mult as f32);
621        state.state.insert("eta_min".to_string(), self.eta_min);
622        state.state.insert("t_cur".to_string(), self.t_cur as f32);
623        state
624    }
625
626    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
627        self.base.last_epoch = state.last_epoch;
628        self.base.base_lrs = state.base_lrs;
629        self.base.last_lr = state.last_lr;
630        if let Some(&t_0) = state.state.get("t_0") {
631            self.t_0 = t_0 as i32;
632        }
633        if let Some(&t_mult) = state.state.get("t_mult") {
634            self.t_mult = t_mult as i32;
635        }
636        if let Some(&eta_min) = state.state.get("eta_min") {
637            self.eta_min = eta_min;
638        }
639        if let Some(&t_cur) = state.state.get("t_cur") {
640            self.t_cur = t_cur as i32;
641        }
642        Ok(())
643    }
644}
645
646#[cfg(test)]
647mod tests {
648    use super::*;
649    use crate::sgd::SGD;
650    use parking_lot::RwLock;
651    use std::sync::Arc;
652    use torsh_tensor::creation::ones;
653
654    #[test]
655    fn test_multi_step_lr() {
656        let param = Arc::new(RwLock::new(ones(&[10]).unwrap()));
657        let optimizer = SGD::new(vec![param], 0.1, None, None, None, false);
658        let mut scheduler = MultiStepLR::new(optimizer, vec![10, 20, 30], 0.5);
659
660        // Initial LR
661        assert_eq!(scheduler.get_last_lr(), vec![0.1]);
662
663        // Step through epochs
664        for _ in 0..9 {
665            let _ = scheduler.step();
666        }
667        assert_eq!(scheduler.get_last_lr(), vec![0.1]);
668
669        // After milestone 10
670        let _ = scheduler.step();
671        assert_eq!(scheduler.get_last_lr(), vec![0.05]);
672
673        // After milestone 20
674        for _ in 0..10 {
675            let _ = scheduler.step();
676        }
677        assert_eq!(scheduler.get_last_lr(), vec![0.025]);
678
679        // After milestone 30
680        for _ in 0..10 {
681            let _ = scheduler.step();
682        }
683        assert_eq!(scheduler.get_last_lr(), vec![0.0125]);
684    }
685
686    #[test]
687    fn test_linear_lr() {
688        let param = Arc::new(RwLock::new(ones(&[10]).unwrap()));
689        let optimizer = SGD::new(vec![param], 0.1, None, None, None, false);
690        let mut scheduler = LinearLR::new(optimizer, 0.1, 1.0, 10);
691
692        // Initial LR should be base_lr (no step taken yet)
693        assert_eq!(scheduler.get_last_lr(), vec![0.1]);
694
695        // After first step
696        let _ = scheduler.step();
697        // Now it should be base_lr * (start_factor + (end_factor - start_factor) * (1/10))
698        let expected_step1 = 0.1 * (0.1 + (1.0 - 0.1) * (1.0 / 10.0));
699        assert!((scheduler.get_last_lr()[0] - expected_step1).abs() < 1e-6);
700
701        // Halfway through (4 more steps to reach step 5)
702        for _ in 0..4 {
703            let _ = scheduler.step();
704        }
705        let expected = 0.1 * (0.1 + (1.0 - 0.1) * (5.0 / 10.0));
706        assert!((scheduler.get_last_lr()[0] - expected).abs() < 1e-6);
707
708        // At the end (5 more steps to reach step 10)
709        for _ in 0..5 {
710            let _ = scheduler.step();
711        }
712        assert_eq!(scheduler.get_last_lr(), vec![0.1]);
713    }
714
715    #[test]
716    fn test_cosine_annealing_warm_restarts() {
717        let param = Arc::new(RwLock::new(ones(&[10]).unwrap()));
718        let optimizer = SGD::new(vec![param], 0.1, None, None, None, false);
719        let mut scheduler = CosineAnnealingWarmRestarts::new(optimizer, 10, 2, 0.0);
720
721        // Initial step
722        let _ = scheduler.step();
723        let lr0 = scheduler.get_last_lr()[0];
724
725        // Should decrease then restart
726        for _ in 1..10 {
727            let _ = scheduler.step();
728            let lr = scheduler.get_last_lr()[0];
729            assert!(lr <= lr0);
730        }
731
732        // After restart, LR should be back to base
733        let _ = scheduler.step();
734        assert!((scheduler.get_last_lr()[0] - 0.1).abs() < 1e-5);
735    }
736}