Skip to main content

burn_optim/lr_scheduler/
noam.rs

1use burn_core as burn;
2
3use burn::config::Config;
4
5use super::{LrScheduler, LrSchedulerRecord, String};
6use crate::LearningRate;
7use crate::RecordState;
8use crate::lr_scheduler::module_lr_scheduler::ModuleLrScheduler;
9
10/// Configuration to create a [noam](NoamLrScheduler) learning rate scheduler.
11#[derive(Config, Debug)]
12pub struct NoamLrSchedulerConfig {
13    /// The overall scale factor for the learning rate decay.
14    factor: f64,
15    /// The number of steps before the exponential decay stats.
16    #[config(default = 4000)]
17    warmup_steps: usize,
18    /// The size of the model.
19    #[config(default = 512)]
20    model_size: usize,
21}
22
23/// Noam learning rate scheduler as described in [Attention Is All You Need](https://arxiv.org/abs/1706.03762).
24#[derive(Clone, Debug)]
25pub struct NoamLrScheduler {
26    warmup_steps: f64,
27    embedding_size: f64,
28    factor: f64,
29    step: f64,
30}
31
32impl NoamLrSchedulerConfig {
33    /// Initialize a new [noam](NoamLrScheduler) learning rate scheduler.
34    pub(crate) fn build(&self) -> Result<NoamLrScheduler, String> {
35        if self.warmup_steps == 0 {
36            return Err(
37                "Number of steps before exponential decay starts must be greater than 0".into(),
38            );
39        }
40        if self.model_size == 0 {
41            return Err("Model size must be greater than 0".into());
42        }
43
44        Ok(NoamLrScheduler {
45            warmup_steps: self.warmup_steps as f64,
46            embedding_size: self.model_size as f64,
47            factor: self.factor,
48            step: 0.0,
49        })
50    }
51
52    /// Initializes a [module learning rate scheduler](ModuleLrScheduler).
53    ///
54    /// # Errors
55    ///
56    /// An error will be returned if any of the following conditions is true:
57    ///
58    /// * `warmup_steps` is 0
59    /// * `model_size` is 0
60    pub fn init(&self) -> Result<ModuleLrScheduler, String> {
61        self.build().map(|s| s.into())
62    }
63}
64
65impl LrScheduler for NoamLrScheduler {
66    fn step(&mut self) -> LearningRate {
67        self.step += 1.0;
68
69        let arg1 = self.step.powf(-0.5);
70        let arg2 = self.step * self.warmup_steps.powf(-1.5);
71
72        self.factor * self.embedding_size.powf(-0.5) * f64::min(arg1, arg2)
73    }
74
75    fn to_record(&self) -> LrSchedulerRecord {
76        LrSchedulerRecord::from_state(&NoamLrSchedulerState { step: self.step })
77    }
78
79    fn load_record(&mut self, record: LrSchedulerRecord) {
80        if let Some(state) = record.into_state::<NoamLrSchedulerState>() {
81            self.step = state.step;
82        }
83    }
84}
85
86/// The serializable state of a [noam scheduler](NoamLrScheduler).
87#[derive(RecordState, Clone, Debug)]
88pub struct NoamLrSchedulerState {
89    step: f64,
90}
91
92#[cfg(test)]
93mod tests {
94    use super::*;
95
96    #[test]
97    fn test_config_warmup_steps_invalid() {
98        let r = NoamLrSchedulerConfig::new(0.1).with_warmup_steps(0).build();
99        assert!(r.is_err(), "Should return an error");
100    }
101
102    #[test]
103    fn test_config_warmup_steps_valid() {
104        let r = NoamLrSchedulerConfig::new(0.1).with_warmup_steps(1).build();
105        assert!(r.is_ok(), "Should return a success value");
106    }
107
108    #[test]
109    fn test_config_model_size_invalid() {
110        let r = NoamLrSchedulerConfig::new(0.1).with_model_size(0).build();
111        assert!(r.is_err(), "Should return an error");
112    }
113
114    #[test]
115    fn test_config_model_size_valid() {
116        let r = NoamLrSchedulerConfig::new(0.1).with_model_size(1).build();
117        assert!(r.is_ok(), "Should return a success value");
118    }
119
120    #[test]
121    fn test_function_increase_and_decrease() {
122        let warmup_steps = 100;
123        let mut scheduler = NoamLrSchedulerConfig::new(10.0)
124            .with_warmup_steps(warmup_steps)
125            .build()
126            .unwrap();
127        let mut lr_current = 0.0;
128
129        for _ in 0..warmup_steps {
130            let lr = scheduler.step();
131            assert!(
132                lr > lr_current,
133                "Learning rate should increase before the warmup_steps is reached."
134            );
135            lr_current = lr;
136        }
137
138        for _ in 0..warmup_steps {
139            let lr = scheduler.step();
140            assert!(
141                lr < lr_current,
142                "Learning rate should decrease after the warmup_steps is reached."
143            );
144            lr_current = lr;
145        }
146    }
147}