Skip to main content

burn_optim/lr_scheduler/
module_lr_scheduler.rs

1use burn_core::{self as burn, module::ParamId};
2
3use burn::config::Config;
4use burn_core::module::ParamGroup;
5
6use crate::{
7    LearningRate,
8    lr_scheduler::{DynLrScheduler, LrScheduler, LrSchedulerConfig, LrSchedulerRecord},
9};
10
11#[derive(Clone)]
12struct LrGroup {
13    group: ParamGroup,
14    lr: f64,
15}
16
17/// Determines what learning rate to use for a given trainable parameter.
18#[derive(Clone, Default)]
19pub struct ModuleLearningRate {
20    groups: Vec<LrGroup>,
21}
22
23impl From<LearningRate> for ModuleLearningRate {
24    fn from(value: LearningRate) -> Self {
25        Self {
26            groups: vec![LrGroup {
27                group: ParamGroup::all(),
28                lr: value,
29            }],
30        }
31    }
32}
33
34impl ModuleLearningRate {
35    /// Get the effective learning rate for the given parameter.
36    pub fn lr_from_param(&self, id: ParamId, path: Option<&str>) -> LearningRate {
37        self.groups
38            .iter()
39            .filter_map(|val| val.group.matches(&id, path).then_some(val.lr))
40            .next_back()
41            .expect("Should match at least one parameter group.")
42    }
43
44    /// Get the base learning rate value which's group matches all parameters.
45    pub fn base(&self) -> LearningRate {
46        self.groups
47            .first()
48            .expect("Should have at least one learning rate.")
49            .lr
50    }
51}
52
53#[derive(Config, Debug)]
54struct LrSchedulerGroupConfig {
55    group: ParamGroup,
56    scheduler: LrSchedulerConfig,
57}
58
59#[derive(new, Clone)]
60struct LrSchedulerGroup {
61    group: ParamGroup,
62    scheduler: DynLrScheduler,
63}
64
65/// Configuration for a [ModuleLrScheduler].
66#[derive(Config, Debug)]
67pub struct ModuleLrSchedulerConfig {
68    base: LrSchedulerConfig,
69    #[config(default = "Vec::new()")]
70    scheduler_groups: Vec<LrSchedulerGroupConfig>,
71}
72
73/// A learning rate scheduler that maps specific parameter groups to dedicated sub-schedulers.
74///
75/// This allows heterogeneous learning rate schedules across different layers of a model
76/// (e.g., discriminative layer training or fine-tuning).
77#[derive(Clone)]
78pub struct ModuleLrScheduler {
79    groups: Vec<LrSchedulerGroup>,
80}
81
82impl ModuleLrSchedulerConfig {
83    /// Initialize a new learning rate policy scheduler.
84    pub fn init(&self) -> Result<ModuleLrScheduler, String> {
85        let mut groups = Vec::with_capacity(self.scheduler_groups.len());
86
87        let base = self.base.build()?;
88        groups.push(LrSchedulerGroup {
89            group: ParamGroup::all(),
90            scheduler: base,
91        });
92
93        for group in self.scheduler_groups.iter() {
94            let scheduler = group.scheduler.build()?;
95            groups.push(LrSchedulerGroup::new(group.group.clone(), scheduler));
96        }
97
98        Ok(ModuleLrScheduler { groups })
99    }
100
101    /// Add a new parameter group to the scheduler's policy.
102    pub fn with_group(
103        mut self,
104        group: ParamGroup,
105        scheduler: impl Into<LrSchedulerConfig>,
106    ) -> Self {
107        self.scheduler_groups.push(LrSchedulerGroupConfig {
108            group,
109            scheduler: scheduler.into(),
110        });
111        self
112    }
113}
114
115impl ModuleLrScheduler {
116    /// Create a [ModuleLrScheduler].
117    ///
118    /// # Arguments
119    ///
120    /// * `scheduler` - The policy's default learning rate scheduler.
121    pub fn new<S: LrScheduler + 'static>(scheduler: S) -> Self {
122        Self {
123            groups: vec![LrSchedulerGroup {
124                group: ParamGroup::all(),
125                scheduler: scheduler.into(),
126            }],
127        }
128    }
129
130    /// Perform the scheduler step of every scheduler and returns the effective learning rate policy.
131    pub fn step(&mut self) -> ModuleLearningRate {
132        let groups = self
133            .groups
134            .iter_mut()
135            .map(|s| {
136                let lr = s.scheduler.step();
137
138                LrGroup {
139                    group: s.group.clone(),
140                    lr,
141                }
142            })
143            .collect();
144
145        ModuleLearningRate { groups }
146    }
147
148    /// Get the current state of the schedulers as a [record](LrSchedulerRecord).
149    pub fn to_record(&self) -> super::LrSchedulerRecord {
150        let mut record = LrSchedulerRecord::new();
151        for (index, item) in self.groups.iter().enumerate() {
152            let sub = item.scheduler.to_record();
153            record = record.with_record(&index.to_string(), sub);
154        }
155
156        record
157    }
158
159    /// Load the state of the schedulers from a [record](LrSchedulerRecord).
160    pub fn load_record(mut self, record: super::LrSchedulerRecord) -> Self {
161        self.groups = self
162            .groups
163            .into_iter()
164            .enumerate()
165            .map(|(index, item)| {
166                let sub = record.record(&index.to_string());
167                let scheduler = item.scheduler.load_record(sub);
168
169                LrSchedulerGroup {
170                    group: item.group,
171                    scheduler,
172                }
173            })
174            .collect();
175
176        self
177    }
178
179    /// Add a new parameter group to the scheduler's policy.
180    pub fn with_group(mut self, group: ParamGroup, scheduler: impl Into<DynLrScheduler>) -> Self {
181        self.groups.push(LrSchedulerGroup {
182            group,
183            scheduler: scheduler.into(),
184        });
185        self
186    }
187}
188
189impl<S> From<S> for ModuleLrScheduler
190where
191    S: LrScheduler + 'static,
192{
193    fn from(value: S) -> Self {
194        Self::new(value)
195    }
196}
197
198#[cfg(test)]
199mod tests {
200    use super::*;
201    use crate::lr_scheduler::linear::LinearLrSchedulerConfig;
202    use burn_core::module::ParamGroup;
203
204    const EPSILON: f64 = 1e-10;
205
206    fn check_approx(actual: f64, expected: f64) {
207        assert!(
208            (actual - expected).abs() < EPSILON,
209            "expected {expected}, got {actual}",
210        );
211    }
212
213    #[test]
214    fn step_yields_constant_default_lr() {
215        let mut scheduler = ModuleLrScheduler::new(0.01_f64);
216        for _ in 0..3 {
217            check_approx(scheduler.step().base(), 0.01);
218        }
219    }
220
221    #[test]
222    fn step_advances_linear_default_scheduler() {
223        let linear = LinearLrSchedulerConfig::new(0.9, 0.5, 4).build().unwrap();
224        let mut scheduler = ModuleLrScheduler::new(linear);
225
226        let expected = [0.9, 0.8, 0.7, 0.6, 0.5, 0.5];
227        for expected_lr in expected {
228            check_approx(scheduler.step().base(), expected_lr);
229        }
230    }
231
232    #[test]
233    fn save_load_preserves_default_scheduler_state() {
234        let make =
235            || ModuleLrScheduler::new(LinearLrSchedulerConfig::new(1.0, 0.1, 9).build().unwrap());
236
237        let mut original = make();
238        let mut truth = make();
239
240        for _ in 0..5 {
241            original.step();
242            truth.step();
243        }
244
245        let record = original.to_record();
246        let mut restored = make().load_record(record);
247
248        for _ in 0..4 {
249            check_approx(restored.step().base(), truth.step().base());
250        }
251    }
252
253    #[test]
254    fn group_param_gets_group_lr() {
255        let id_group = ParamId::new();
256        let id_default = ParamId::new();
257
258        let mut scheduler = ModuleLrSchedulerConfig::new(0.001.into())
259            .with_group(ParamGroup::from_ids(vec![id_group.clone()]), 0.1)
260            .init()
261            .unwrap();
262
263        let policy = scheduler.step();
264        // id_group is in the explicit group = group LR
265        check_approx(policy.lr_from_param(id_group, None), 0.1);
266        // id_default is not in any group = default LR
267        check_approx(policy.lr_from_param(id_default, None), 0.001);
268    }
269
270    #[test]
271    fn path_group_matches_param_by_path_substring() {
272        let mut scheduler = ModuleLrSchedulerConfig::new(0.001.into())
273            .with_group(ParamGroup::from_predicate("backbone"), 0.1)
274            .init()
275            .unwrap();
276
277        let policy = scheduler.step();
278        let id = ParamId::new();
279
280        check_approx(
281            policy.lr_from_param(id.clone(), Some("model.backbone.layer.weight")),
282            0.1,
283        );
284        check_approx(
285            policy.lr_from_param(id, Some("model.head.layer.weight")),
286            0.001,
287        );
288    }
289
290    #[test]
291    fn multiple_groups_are_independent() {
292        let id_a = ParamId::new();
293        let id_b = ParamId::new();
294        let id_default = ParamId::new();
295
296        let mut scheduler =
297            ModuleLrSchedulerConfig::new(LinearLrSchedulerConfig::new(0.001, 0.0001, 4).into())
298                .with_group(
299                    ParamGroup::from_ids(vec![id_a.clone()]),
300                    LinearLrSchedulerConfig::new(0.1, 0.01, 4),
301                )
302                .with_group(
303                    ParamGroup::from_ids(vec![id_b.clone()]),
304                    LinearLrSchedulerConfig::new(0.5, 0.05, 4),
305                )
306                .init()
307                .unwrap();
308
309        // Each group returns its own initial LR
310        let policy = scheduler.step();
311        check_approx(policy.lr_from_param(id_a.clone(), None), 0.1);
312        check_approx(policy.lr_from_param(id_b.clone(), None), 0.5);
313        check_approx(policy.lr_from_param(id_default.clone(), None), 0.001);
314
315        // All three schedulers advanced; LRs are strictly between initial and final
316        let policy = scheduler.step();
317        let lr_a = policy.lr_from_param(id_a, None);
318        let lr_b = policy.lr_from_param(id_b, None);
319        let lr_default = policy.lr_from_param(id_default, None);
320        assert!(
321            lr_a < 0.1 && lr_a > 0.01,
322            "group-a LR should have decayed: {lr_a}"
323        );
324        assert!(
325            lr_b < 0.5 && lr_b > 0.05,
326            "group-b LR should have decayed: {lr_b}"
327        );
328        assert!(
329            lr_default < 0.001 && lr_default > 0.0001,
330            "default LR should have decayed: {lr_default}"
331        );
332    }
333
334    #[test]
335    fn save_load_with_groups_preserves_state() {
336        let id_group = ParamId::new();
337
338        let make = || {
339            ModuleLrSchedulerConfig::new(LinearLrSchedulerConfig::new(0.01, 0.001, 9).into())
340                .with_group(
341                    ParamGroup::from_ids(vec![id_group.clone()]),
342                    LinearLrSchedulerConfig::new(0.1, 0.01, 9),
343                )
344                .init()
345                .unwrap()
346        };
347
348        let mut original = make();
349        let mut truth = make();
350
351        for _ in 0..5 {
352            original.step();
353            truth.step();
354        }
355
356        let record = original.to_record();
357        let mut restored = make().load_record(record);
358
359        for _ in 0..4 {
360            let lr_restored = restored.step().lr_from_param(id_group.clone(), None);
361            let lr_truth = truth.step().lr_from_param(id_group.clone(), None);
362            check_approx(lr_restored, lr_truth);
363        }
364    }
365}