Skip to main content

burn_optim/lr_scheduler/
sequential.rs

1use alloc::vec::Vec;
2use burn_core as burn;
3
4use burn::config::Config;
5
6use super::module_lr_scheduler::ModuleLrScheduler;
7use super::{DynLrScheduler, LrScheduler, LrSchedulerConfig, LrSchedulerRecord, String};
8use crate::{LearningRate, RecordState};
9
10/// Configuration for a [sequential learning rate scheduler](SequentialLrScheduler).
11///
12/// Each milestone is the number of completed calls to [`LrScheduler::step`] at which the next
13/// scheduler takes over. Therefore, `N` schedulers require exactly `N - 1` milestones. Milestones
14/// must be greater than zero and strictly increasing.
15///
16/// # Example
17///
18/// The linear scheduler below handles the first 100 steps, after which the cosine scheduler starts
19/// from its own first step.
20///
21/// ```
22/// use burn_optim::lr_scheduler::{
23///     LrSchedulerConfig,
24///     cosine::CosineAnnealingLrSchedulerConfig, linear::LinearLrSchedulerConfig,
25///     sequential::SequentialLrSchedulerConfig,
26/// };
27///
28/// let config = SequentialLrSchedulerConfig::new(
29///     vec![
30///         LrSchedulerConfig::Linear(LinearLrSchedulerConfig::new(1e-5, 1e-2, 100)),
31///         LrSchedulerConfig::Cosine(CosineAnnealingLrSchedulerConfig::new(1e-2, 900)),
32///     ],
33///     vec![100],
34/// );
35/// let scheduler = config.init().unwrap();
36/// ```
37#[derive(Config, Debug)]
38pub struct SequentialLrSchedulerConfig {
39    schedulers: Vec<LrSchedulerConfig>,
40    milestones: Vec<usize>,
41}
42
43impl SequentialLrSchedulerConfig {
44    pub(crate) fn build(&self) -> Result<SequentialLrScheduler, String> {
45        if self.schedulers.is_empty() {
46            return Err("At least one scheduler is required".into());
47        }
48        if self.milestones.len() + 1 != self.schedulers.len() {
49            return Err(
50                "The number of milestones must be one less than the number of schedulers".into(),
51            );
52        }
53        if self
54            .milestones
55            .iter()
56            .enumerate()
57            .any(|(index, milestone)| {
58                *milestone == 0 || index > 0 && *milestone <= self.milestones[index - 1]
59            })
60        {
61            return Err("Milestones must be greater than zero and strictly increasing".into());
62        }
63
64        let schedulers = self
65            .schedulers
66            .iter()
67            .map(LrSchedulerConfig::build)
68            .collect::<Result<Vec<_>, _>>()?;
69
70        Ok(SequentialLrScheduler {
71            schedulers,
72            milestones: self.milestones.clone(),
73            step: 0,
74        })
75    }
76
77    /// Initializes a [module learning rate scheduler](ModuleLrScheduler).
78    ///
79    /// # Errors
80    ///
81    /// An error is returned when there are no schedulers, the number of milestones is not one less
82    /// than the number of schedulers, milestones are zero or not strictly increasing, or a child
83    /// scheduler configuration is invalid.
84    pub fn init(&self) -> Result<ModuleLrScheduler, String> {
85        self.build().map(Into::into)
86    }
87}
88
89/// Runs learning rate schedulers one after another at configured milestones.
90#[derive(Clone)]
91pub struct SequentialLrScheduler {
92    schedulers: Vec<DynLrScheduler>,
93    milestones: Vec<usize>,
94    step: usize,
95}
96
97impl SequentialLrScheduler {
98    fn active_scheduler(&self) -> usize {
99        self.milestones.partition_point(|&m| self.step >= m)
100    }
101}
102
103impl LrScheduler for SequentialLrScheduler {
104    fn step(&mut self) -> LearningRate {
105        let index = self.active_scheduler();
106        let lr = self.schedulers[index].step();
107        self.step = self
108            .step
109            .checked_add(1)
110            .expect("The sequential scheduler step counter overflowed");
111        lr
112    }
113
114    fn to_record(&self) -> LrSchedulerRecord {
115        let mut record =
116            LrSchedulerRecord::from_state(&SequentialLrSchedulerState { step: self.step });
117        for (index, scheduler) in self.schedulers.iter().enumerate() {
118            record = record.with_record(&index.to_string(), scheduler.to_record());
119        }
120        record
121    }
122
123    fn load_record(&mut self, record: LrSchedulerRecord) {
124        if let Some(state) = record.into_state::<SequentialLrSchedulerState>() {
125            self.step = state.step;
126        }
127
128        let schedulers = core::mem::take(&mut self.schedulers);
129        self.schedulers = schedulers
130            .into_iter()
131            .enumerate()
132            .map(|(index, scheduler)| scheduler.load_record(record.record(&index.to_string())))
133            .collect();
134    }
135}
136
137#[derive(RecordState, Clone, Debug)]
138struct SequentialLrSchedulerState {
139    step: usize,
140}
141
142#[cfg(test)]
143mod tests {
144    use super::super::cosine::CosineAnnealingLrSchedulerConfig;
145    use super::super::exponential::ExponentialLrSchedulerConfig;
146    use super::super::linear::LinearLrSchedulerConfig;
147    use super::super::test_utils;
148    use super::*;
149
150    fn config(milestones: Vec<usize>) -> SequentialLrSchedulerConfig {
151        SequentialLrSchedulerConfig::new(
152            vec![
153                LinearLrSchedulerConfig::new(0.1, 0.3, 2).into(),
154                ExponentialLrSchedulerConfig::new(0.5, 0.5).into(),
155                CosineAnnealingLrSchedulerConfig::new(0.8, 2).into(),
156            ],
157            milestones,
158        )
159    }
160
161    #[test]
162    fn switches_schedulers_at_milestones() {
163        let scheduler = config(vec![2, 5]).build().unwrap();
164        test_utils::check_lr_sequence(scheduler, [0.1, 0.2, 0.5, 0.25, 0.125, 0.8, 0.4, 0.0]);
165    }
166
167    #[test]
168    fn rejects_empty_scheduler_list() {
169        let result = SequentialLrSchedulerConfig::new(vec![], vec![]).build();
170        assert_eq!(result.err().unwrap(), "At least one scheduler is required");
171    }
172
173    #[test]
174    fn rejects_wrong_milestone_count() {
175        let result = config(vec![2]).build();
176        assert_eq!(
177            result.err().unwrap(),
178            "The number of milestones must be one less than the number of schedulers"
179        );
180    }
181
182    #[test]
183    fn rejects_zero_or_non_increasing_milestones() {
184        for milestones in [vec![0, 2], vec![2, 2], vec![3, 2]] {
185            let result = config(milestones).build();
186            assert_eq!(
187                result.err().unwrap(),
188                "Milestones must be greater than zero and strictly increasing"
189            );
190        }
191    }
192
193    #[test]
194    fn reports_invalid_child_config() {
195        let result = SequentialLrSchedulerConfig::new(
196            vec![LinearLrSchedulerConfig::new(0.1, 0.2, 0).into()],
197            vec![],
198        )
199        .build();
200        assert_eq!(
201            result.err().unwrap(),
202            "Number of iterations must be at least 1"
203        );
204    }
205
206    #[test]
207    fn saves_and_loads_before_and_after_transitions() {
208        test_utils::check_save_load(config(vec![2, 5]).build().unwrap(), 1);
209        test_utils::check_save_load(config(vec![2, 5]).build().unwrap(), 4);
210        test_utils::check_save_load(config(vec![2, 5]).build().unwrap(), 6);
211    }
212}