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#[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 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 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#[derive(Config, Debug)]
67pub struct ModuleLrSchedulerConfig {
68 base: LrSchedulerConfig,
69 #[config(default = "Vec::new()")]
70 scheduler_groups: Vec<LrSchedulerGroupConfig>,
71}
72
73#[derive(Clone)]
78pub struct ModuleLrScheduler {
79 groups: Vec<LrSchedulerGroup>,
80}
81
82impl ModuleLrSchedulerConfig {
83 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 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 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 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 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 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 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 check_approx(policy.lr_from_param(id_group, None), 0.1);
266 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 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 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}