Skip to main content

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}