optirs_core/schedulers/exponential_decay.rs
1// Exponential decay learning rate scheduler
2
3use scirs2_core::ndarray::ScalarOperand;
4use scirs2_core::numeric::Float;
5use std::fmt::Debug;
6
7use crate::schedulers::LearningRateScheduler;
8
9/// Exponential decay learning rate scheduler
10///
11/// Applies exponential decay to the learning rate over time:
12/// lr = initial_lr * decay_rate^(step / decay_steps)
13///
14/// # Examples
15///
16/// ```
17/// use optirs_core::schedulers::{ExponentialDecay, LearningRateScheduler};
18///
19/// // Create a scheduler with initial learning rate 0.1, decay rate 0.95
20/// // and decay steps 1000
21/// let mut scheduler = ExponentialDecay::new(0.1f64, 0.95, 1000);
22///
23/// // Initial learning rate
24/// let initial_lr = scheduler.get_learning_rate();
25///
26/// // Run for a few steps (reduced for test)
27/// for _ in 0..3 {
28/// // Update learning rate
29/// let lr = scheduler.step();
30/// // Verify learning rate is decaying
31/// assert!(lr < initial_lr);
32/// }
33///
34/// // Verify scheduler is working
35/// let final_lr = scheduler.get_learning_rate();
36/// assert!(final_lr < 0.1);
37/// ```
38#[derive(Debug, Clone)]
39pub struct ExponentialDecay<A: Float + Debug> {
40 /// Initial learning rate
41 initial_lr: A,
42 /// Decay rate
43 decay_rate: A,
44 /// Number of steps after which the learning rate is decayed by decay_rate
45 decay_steps: usize,
46 /// Current step
47 step: usize,
48 /// Current learning rate
49 current_lr: A,
50}
51
52impl<A: Float + Debug + Send + Sync> ExponentialDecay<A> {
53 /// Create a new exponential decay scheduler
54 ///
55 /// # Arguments
56 ///
57 /// * `initial_lr` - Initial learning rate
58 /// * `decay_rate` - Rate at which learning rate decays (e.g., 0.95)
59 /// * `decay_steps` - Number of steps after which learning rate is decayed by decay_rate
60 pub fn new(initial_lr: A, decay_rate: A, decay_steps: usize) -> Self {
61 Self {
62 initial_lr,
63 decay_rate,
64 decay_steps,
65 step: 0,
66 current_lr: initial_lr,
67 }
68 }
69}
70
71impl<A: Float + Debug + ScalarOperand + Send + Sync> LearningRateScheduler<A>
72 for ExponentialDecay<A>
73{
74 fn get_learning_rate(&self) -> A {
75 self.current_lr
76 }
77
78 fn step(&mut self) -> A {
79 self.step += 1;
80
81 // Calculate learning rate decay
82 // lr = initial_lr * decay_rate^(step / decay_steps)
83 let power = A::from(self.step).expect("ExponentialDecay: step must fit in A (f32/f64)")
84 / A::from(self.decay_steps)
85 .expect("ExponentialDecay: decay_steps must fit in A (f32/f64)");
86 self.current_lr = self.initial_lr * self.decay_rate.powf(power);
87
88 self.current_lr
89 }
90
91 fn reset(&mut self) {
92 self.step = 0;
93 self.current_lr = self.initial_lr;
94 }
95}