1use crate::{
4 lr_scheduler::{BaseScheduler, LRScheduler, SchedulerState},
5 Optimizer, OptimizerError, OptimizerResult,
6};
7
8pub 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 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 Ok(())
95 }
96}
97
98pub 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 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 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 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
270pub 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 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
357pub 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 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
456pub 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 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
543pub 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 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 assert_eq!(scheduler.get_last_lr(), vec![0.1]);
662
663 for _ in 0..9 {
665 let _ = scheduler.step();
666 }
667 assert_eq!(scheduler.get_last_lr(), vec![0.1]);
668
669 let _ = scheduler.step();
671 assert_eq!(scheduler.get_last_lr(), vec![0.05]);
672
673 for _ in 0..10 {
675 let _ = scheduler.step();
676 }
677 assert_eq!(scheduler.get_last_lr(), vec![0.025]);
678
679 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 assert_eq!(scheduler.get_last_lr(), vec![0.1]);
694
695 let _ = scheduler.step();
697 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 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 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 let _ = scheduler.step();
723 let lr0 = scheduler.get_last_lr()[0];
724
725 for _ in 1..10 {
727 let _ = scheduler.step();
728 let lr = scheduler.get_last_lr()[0];
729 assert!(lr <= lr0);
730 }
731
732 let _ = scheduler.step();
734 assert!((scheduler.get_last_lr()[0] - 0.1).abs() < 1e-5);
735 }
736}