1use scirs2_core::ndarray::ScalarOperand;
7use scirs2_core::numeric::Float;
8use std::cell::RefCell;
9use std::fmt::Debug;
10use std::marker::PhantomData;
11use std::rc::Rc;
12
13use super::LearningRateScheduler;
14
15pub struct CustomScheduler<A, F>
17where
18 A: Float + Debug + ScalarOperand,
19 F: FnMut(usize) -> A,
20{
21 lr_func: Rc<RefCell<F>>,
23 step_count: usize,
25 _phantom: PhantomData<A>,
27}
28
29impl<A, F> CustomScheduler<A, F>
30where
31 A: Float + Debug + ScalarOperand,
32 F: FnMut(usize) -> A,
33{
34 pub fn new(_initial_lr: A, lrfunc: F) -> Self {
55 Self {
56 lr_func: Rc::new(RefCell::new(lrfunc)),
57 step_count: 0,
58 _phantom: PhantomData,
59 }
60 }
61
62 pub fn get_step_count(&self) -> usize {
64 self.step_count
65 }
66}
67
68impl<A, F> LearningRateScheduler<A> for CustomScheduler<A, F>
69where
70 A: Float + Debug + ScalarOperand,
71 F: FnMut(usize) -> A,
72{
73 fn get_learning_rate(&self) -> A {
74 let mut func = self.lr_func.borrow_mut();
76 func(self.step_count)
77 }
78
79 fn step(&mut self) -> A {
80 self.step_count += 1;
81 self.get_learning_rate()
82 }
83
84 fn reset(&mut self) {
85 self.step_count = 0;
86 }
87}
88
89pub struct CombinedScheduler<A, F1, F2, C>
91where
92 A: Float + Debug + ScalarOperand,
93 F1: FnMut(usize) -> A,
94 F2: FnMut(usize) -> A,
95 C: FnMut(A, A) -> A,
96{
97 scheduler1: CustomScheduler<A, F1>,
99 scheduler2: CustomScheduler<A, F2>,
101 combinator: Rc<RefCell<C>>,
103}
104
105impl<A, F1, F2, C> CombinedScheduler<A, F1, F2, C>
106where
107 A: Float + Debug + ScalarOperand,
108 F1: FnMut(usize) -> A,
109 F2: FnMut(usize) -> A,
110 C: FnMut(A, A) -> A,
111{
112 pub fn new(
147 scheduler1: CustomScheduler<A, F1>,
148 scheduler2: CustomScheduler<A, F2>,
149 combinator: C,
150 ) -> Self {
151 Self {
152 scheduler1,
153 scheduler2,
154 combinator: Rc::new(RefCell::new(combinator)),
155 }
156 }
157}
158
159impl<A, F1, F2, C> LearningRateScheduler<A> for CombinedScheduler<A, F1, F2, C>
160where
161 A: Float + Debug + ScalarOperand,
162 F1: FnMut(usize) -> A,
163 F2: FnMut(usize) -> A,
164 C: FnMut(A, A) -> A,
165{
166 fn get_learning_rate(&self) -> A {
167 let lr1 = self.scheduler1.get_learning_rate();
168 let lr2 = self.scheduler2.get_learning_rate();
169
170 let mut combinator = self.combinator.borrow_mut();
172 combinator(lr1, lr2)
173 }
174
175 fn step(&mut self) -> A {
176 self.scheduler1.step();
177 self.scheduler2.step();
178 self.get_learning_rate()
179 }
180
181 fn reset(&mut self) {
182 self.scheduler1.reset();
183 self.scheduler2.reset();
184 }
185}
186
187pub struct SchedulerBuilder<A>
189where
190 A: Float + Debug + ScalarOperand,
191{
192 initial_lr: A,
193}
194
195impl<A> SchedulerBuilder<A>
196where
197 A: Float + Debug + ScalarOperand,
198{
199 pub fn new(initiallr: A) -> Self {
201 Self {
202 initial_lr: initiallr,
203 }
204 }
205
206 pub fn step_decay(
213 self,
214 step_size: usize,
215 gamma: A,
216 ) -> CustomScheduler<A, impl FnMut(usize) -> A> {
217 let initial_lr = self.initial_lr;
218 CustomScheduler::new(initial_lr, move |step| {
219 let decay_factor = gamma.powi((step / step_size) as i32);
220 initial_lr * decay_factor
221 })
222 }
223
224 pub fn exponential_decay(self, gamma: A) -> CustomScheduler<A, impl FnMut(usize) -> A> {
230 let initial_lr = self.initial_lr;
231 CustomScheduler::new(initial_lr, move |step| initial_lr * gamma.powi(step as i32))
232 }
233
234 pub fn linear_decay(
241 self,
242 total_steps: usize,
243 final_lr: A,
244 ) -> CustomScheduler<A, impl FnMut(usize) -> A> {
245 let initial_lr = self.initial_lr;
246 let total_steps =
247 A::from(total_steps).expect("CustomScheduler: total_steps must fit in A (f32/f64)");
248 CustomScheduler::new(initial_lr, move |step| {
249 let step = A::from(step).expect("CustomScheduler: step must fit in A (f32/f64)");
250 if step >= total_steps {
251 final_lr
252 } else {
253 let progress = step / total_steps;
254 initial_lr + progress * (final_lr - initial_lr)
255 }
256 })
257 }
258
259 pub fn cosine_annealing(
266 self,
267 total_steps: usize,
268 min_lr: A,
269 ) -> CustomScheduler<A, impl FnMut(usize) -> A> {
270 let initial_lr = self.initial_lr;
271 let total_steps =
272 A::from(total_steps).expect("CustomScheduler: total_steps must fit in A (f32/f64)");
273 let pi = A::from(std::f64::consts::PI)
274 .expect("CustomScheduler: pi constant must fit in A (f32/f64)");
275 CustomScheduler::new(initial_lr, move |step| {
276 let step = A::from(step).expect("CustomScheduler: step must fit in A (f32/f64)");
277 if step >= total_steps {
278 min_lr
279 } else {
280 let progress = pi * step / total_steps;
281 min_lr + (initial_lr - min_lr) * (A::one() + progress.cos()) / (A::one() + A::one())
282 }
283 })
284 }
285
286 pub fn cyclic_lr(
294 self,
295 step_size: usize,
296 max_lr: A,
297 mode: CyclicMode<A>,
298 ) -> CustomScheduler<A, impl FnMut(usize) -> A> {
299 let min_lr = self.initial_lr;
300 let step_size =
301 A::from(step_size).expect("CustomScheduler: step_size must fit in A (f32/f64)");
302 let two = A::one() + A::one();
303
304 let mode_inner = mode;
306
307 CustomScheduler::new(min_lr, move |step| {
308 let step = A::from(step).expect("CustomScheduler: step must fit in A (f32/f64)");
309 let cycle = (step / (two * step_size)).floor();
310 let x = (step / step_size - two * cycle).abs();
311
312 let scale = match mode_inner {
313 CyclicMode::Triangular => A::one(),
314 CyclicMode::Triangular2 => A::one() / (two.powi(cycle.to_i32().unwrap_or(0))),
315 CyclicMode::ExpRange(gamma) => gamma.powi(step.to_i32().unwrap_or(0)),
316 };
317
318 min_lr + scale * (max_lr - min_lr) * (A::one() - x).max(A::zero())
319 })
320 }
321
322 pub fn custom<F>(self, func: F) -> CustomScheduler<A, F>
328 where
329 F: FnMut(usize) -> A,
330 {
331 CustomScheduler::new(self.initial_lr, func)
332 }
333}
334
335#[derive(Debug, Clone, Copy)]
337pub enum CyclicMode<A: Float> {
338 Triangular,
340 Triangular2,
342 ExpRange(A),
344}
345
346#[cfg(test)]
347mod tests {
348 use super::*;
349 use approx::assert_relative_eq;
350
351 #[test]
352 fn test_custom_scheduler() {
353 let mut scheduler =
354 CustomScheduler::new(0.1f64, |step| 0.1 * 0.9f64.powi((step / 10) as i32));
355
356 assert_eq!(scheduler.get_learning_rate(), 0.1);
357 assert_eq!(scheduler.step(), 0.1);
358 assert_eq!(scheduler.step(), 0.1);
359
360 for _ in 0..8 {
362 scheduler.step();
363 }
364 assert_relative_eq!(scheduler.get_learning_rate(), 0.09, epsilon = 1e-10);
365 }
366
367 #[test]
368 fn test_combined_scheduler() {
369 let scheduler1 = CustomScheduler::new(0.1f64, |step| 0.1 * 0.9f64.powi((step / 10) as i32));
370
371 let scheduler2 = CustomScheduler::new(0.2f64, |step| 0.2 * 0.8f64.powi((step / 5) as i32));
372
373 let mut combined =
374 CombinedScheduler::new(scheduler1, scheduler2, |lr1, lr2| lr1 * 0.3 + lr2 * 0.7);
375
376 assert_relative_eq!(
377 combined.get_learning_rate(),
378 0.1 * 0.3 + 0.2 * 0.7,
379 epsilon = 1e-10
380 );
381 combined.step();
382 assert_relative_eq!(
383 combined.get_learning_rate(),
384 0.1 * 0.3 + 0.2 * 0.7,
385 epsilon = 1e-10
386 );
387
388 for _ in 0..4 {
390 combined.step();
391 }
392 assert_relative_eq!(
393 combined.get_learning_rate(),
394 0.1 * 0.3 + 0.2 * 0.8 * 0.7,
395 epsilon = 1e-10
396 );
397 }
398
399 #[test]
400 fn test_scheduler_builder() {
401 let mut step_scheduler = SchedulerBuilder::new(0.1f64).step_decay(10, 0.5);
403 assert_eq!(step_scheduler.get_learning_rate(), 0.1);
404 for _ in 0..10 {
405 step_scheduler.step();
406 }
407 assert_relative_eq!(step_scheduler.get_learning_rate(), 0.05, epsilon = 1e-10);
408
409 let mut exp_scheduler = SchedulerBuilder::new(0.1f64).exponential_decay(0.95);
411 assert_eq!(exp_scheduler.get_learning_rate(), 0.1);
412 exp_scheduler.step();
413 assert_relative_eq!(
414 exp_scheduler.get_learning_rate(),
415 0.1 * 0.95,
416 epsilon = 1e-10
417 );
418
419 let mut linear_scheduler = SchedulerBuilder::new(0.1f64).linear_decay(100, 0.01);
421 assert_eq!(linear_scheduler.get_learning_rate(), 0.1);
422 linear_scheduler.step();
423 assert_relative_eq!(
424 linear_scheduler.get_learning_rate(),
425 0.1 - 0.0009, epsilon = 1e-10
427 );
428
429 let mut cosine_scheduler = SchedulerBuilder::new(0.1f64).cosine_annealing(100, 0.01);
431 assert_eq!(cosine_scheduler.get_learning_rate(), 0.1);
432 cosine_scheduler.step();
433 assert!(cosine_scheduler.get_learning_rate() < 0.1);
435 assert!(cosine_scheduler.get_learning_rate() > 0.01);
437 }
438}