burn_optim/lr_scheduler/
sequential.rs1use 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#[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 pub fn init(&self) -> Result<ModuleLrScheduler, String> {
85 self.build().map(Into::into)
86 }
87}
88
89#[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}