Skip to main content

ruda_optim/lr_scheduler/
constant.rs

1
2use ruda_model::tensor::backend::Backend;
3
4use super::LrScheduler;
5use crate::LearningRate;
6
7/// Constant learning rate implementing [learning rate scheduler](LrScheduler).
8///
9/// # Notes
10///
11/// You can also use [learning rate](LearningRate) which the same effect.
12#[derive(new, Clone, Debug)]
13pub struct ConstantLr {
14    lr: LearningRate,
15}
16
17impl From<LearningRate> for ConstantLr {
18    fn from(lr: LearningRate) -> Self {
19        Self { lr }
20    }
21}
22
23impl LrScheduler for ConstantLr {
24    type Record<B: Backend> = ();
25
26    fn step(&mut self) -> LearningRate {
27        self.lr
28    }
29
30    fn to_record<B: Backend>(&self) -> Self::Record<B> {}
31
32    fn load_record<B: Backend>(self, _record: Self::Record<B>) -> Self {
33        self
34    }
35}
36
37impl LrScheduler for LearningRate {
38    type Record<B: Backend> = ();
39
40    fn step(&mut self) -> LearningRate {
41        *self
42    }
43
44    fn to_record<B: Backend>(&self) -> Self::Record<B> {}
45
46    fn load_record<B: Backend>(self, _record: Self::Record<B>) -> Self {
47        self
48    }
49}