Skip to main content

optirs_core/optimizers/
adabound.rs

1// OptiRS - AdaBound Optimizer
2// Adaptive Gradient Methods with Dynamic Bound of Learning Rate
3// Reference: "Adaptive Gradient Methods with Dynamic Bound of Learning Rate" (ICLR 2019)
4//
5// Algorithm:
6//   AdaBound employs dynamic bounds on learning rates to achieve smooth transition
7//   from adaptive methods to SGD. This prevents the generalization gap observed
8//   in pure adaptive methods.
9//
10//   Lower bound: α_l(t) = α_final * (1 - 1/(γ*t + 1))
11//   Upper bound: α_u(t) = α_final * (1 + 1/(γ*t))
12//   Clipped learning rate: η_t(i) = Clip(α / √(v_t(i) + ε), α_l(t), α_u(t))
13
14use crate::error::{OptimError, Result};
15use crate::optimizers::Optimizer;
16use scirs2_core::ndarray::{Ix1, ScalarOperand};
17use scirs2_core::ndarray_ext::{Array1, ArrayView1};
18use scirs2_core::numeric::Float;
19use serde::{Deserialize, Serialize};
20use std::fmt::Debug;
21
22/// AdaBound optimizer configuration
23///
24/// AdaBound combines the benefits of adaptive learning rate methods (like Adam)
25/// with the strong generalization of SGD by dynamically bounding the learning rates.
26///
27/// # Key Features
28/// - Smooth transition from Adam to SGD during training
29/// - Dynamic bounds prevent learning rates from becoming too large or too small
30/// - Better generalization than pure Adam
31/// - Maintains fast convergence of adaptive methods
32///
33/// # Type Parameters
34/// - `T`: Floating-point type (f32 or f64)
35#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct AdaBound<T: Float> {
37    /// Initial learning rate (α)
38    learning_rate: T,
39
40    /// Final learning rate for SGD convergence
41    /// Typically 0.1 * learning_rate
42    final_lr: T,
43
44    /// First moment decay rate (β₁) - typically 0.9
45    beta1: T,
46
47    /// Second moment decay rate (β₂) - typically 0.999
48    beta2: T,
49
50    /// Small constant for numerical stability (ε) - typically 1e-8
51    epsilon: T,
52
53    /// Convergence speed parameter (γ) - typically 1e-3
54    /// Controls how fast bounds converge to final_lr
55    gamma: T,
56
57    /// Weight decay coefficient (L2 regularization)
58    weight_decay: T,
59
60    /// Whether to use AMSBound variant (max of v_t)
61    amsbound: bool,
62
63    /// First moment vector (m_t)
64    momentum: Option<Array1<T>>,
65
66    /// Second moment vector (v_t)
67    velocity: Option<Array1<T>>,
68
69    /// Max of second moment (v̂_t) - only for AMSBound
70    max_velocity: Option<Array1<T>>,
71
72    /// Number of optimization steps performed
73    step_count: usize,
74}
75
76impl<T: Float + ScalarOperand> Default for AdaBound<T> {
77    fn default() -> Self {
78        Self::new(
79            T::from(0.001).expect("AdaBound: default learning_rate (0.001) must fit in T"),
80            T::from(0.1).expect("AdaBound: default final_lr (0.1) must fit in T"),
81            T::from(0.9).expect("AdaBound: default beta1 (0.9) must fit in T"),
82            T::from(0.999).expect("AdaBound: default beta2 (0.999) must fit in T"),
83            T::from(1e-8).expect("AdaBound: default epsilon (1e-8) must fit in T"),
84            T::from(1e-3).expect("AdaBound: default gamma (1e-3) must fit in T"),
85            T::zero(),
86            false,
87        )
88        .expect("AdaBound: default hyperparameters always satisfy validation")
89    }
90}
91
92impl<T: Float + ScalarOperand> AdaBound<T> {
93    /// Create a new AdaBound optimizer
94    ///
95    /// # Arguments
96    /// - `learning_rate`: Initial learning rate (typically 0.001)
97    /// - `final_lr`: Final learning rate for SGD convergence (typically 0.1)
98    /// - `beta1`: First moment decay rate (typically 0.9)
99    /// - `beta2`: Second moment decay rate (typically 0.999)
100    /// - `epsilon`: Small constant for numerical stability (typically 1e-8)
101    /// - `gamma`: Convergence speed parameter (typically 1e-3)
102    /// - `weight_decay`: L2 regularization coefficient (typically 0.0)
103    /// - `amsbound`: Use AMSBound variant if true
104    ///
105    /// # Example
106    /// ```
107    /// use optirs_core::optimizers::AdaBound;
108    ///
109    /// let optimizer = AdaBound::<f32>::new(
110    ///     0.001,  // learning_rate
111    ///     0.1,    // final_lr
112    ///     0.9,    // beta1
113    ///     0.999,  // beta2
114    ///     1e-8,   // epsilon
115    ///     1e-3,   // gamma
116    ///     0.0,    // weight_decay
117    ///     false   // amsbound
118    /// ).expect("AdaBound::new succeeds for finite, in-range default hyperparameters");
119    /// ```
120    // AdaBound's full-configuration constructor mirrors the paper's 8 named
121    // hyperparameters (Luo et al., 2019); grouping them into a config struct
122    // would be a breaking change to this crate's public API for no gain in
123    // clarity at the (single, non-hot-path) call site.
124    #[allow(clippy::too_many_arguments)]
125    pub fn new(
126        learning_rate: T,
127        final_lr: T,
128        beta1: T,
129        beta2: T,
130        epsilon: T,
131        gamma: T,
132        weight_decay: T,
133        amsbound: bool,
134    ) -> Result<Self> {
135        let lr_f64 = crate::optimizers::scalar_to_f64(learning_rate)?;
136        let final_f64 = crate::optimizers::scalar_to_f64(final_lr)?;
137        let beta1_f64 = crate::optimizers::scalar_to_f64(beta1)?;
138        let beta2_f64 = crate::optimizers::scalar_to_f64(beta2)?;
139        let eps_f64 = crate::optimizers::scalar_to_f64(epsilon)?;
140        let gamma_f64 = crate::optimizers::scalar_to_f64(gamma)?;
141        let wd_f64 = crate::optimizers::scalar_to_f64(weight_decay)?;
142
143        if lr_f64 <= 0.0 {
144            return Err(OptimError::InvalidParameter(format!(
145                "learning_rate must be positive, got {}",
146                lr_f64
147            )));
148        }
149        if final_f64 <= 0.0 {
150            return Err(OptimError::InvalidParameter(format!(
151                "final_lr must be positive, got {}",
152                final_f64
153            )));
154        }
155        if beta1_f64 <= 0.0 || beta1_f64 >= 1.0 {
156            return Err(OptimError::InvalidParameter(format!(
157                "beta1 must be in (0, 1), got {}",
158                beta1_f64
159            )));
160        }
161        if beta2_f64 <= 0.0 || beta2_f64 >= 1.0 {
162            return Err(OptimError::InvalidParameter(format!(
163                "beta2 must be in (0, 1), got {}",
164                beta2_f64
165            )));
166        }
167        if eps_f64 <= 0.0 {
168            return Err(OptimError::InvalidParameter(format!(
169                "epsilon must be positive, got {}",
170                eps_f64
171            )));
172        }
173        if gamma_f64 <= 0.0 {
174            return Err(OptimError::InvalidParameter(format!(
175                "gamma must be positive, got {}",
176                gamma_f64
177            )));
178        }
179        if wd_f64 < 0.0 {
180            return Err(OptimError::InvalidParameter(format!(
181                "weight_decay must be non-negative, got {}",
182                wd_f64
183            )));
184        }
185
186        Ok(Self {
187            learning_rate,
188            final_lr,
189            beta1,
190            beta2,
191            epsilon,
192            gamma,
193            weight_decay,
194            amsbound,
195            momentum: None,
196            velocity: None,
197            max_velocity: None,
198            step_count: 0,
199        })
200    }
201
202    /// Perform a single optimization step
203    ///
204    /// # Arguments
205    /// - `params`: Current parameter values
206    /// - `grads`: Gradient values
207    ///
208    /// # Returns
209    /// Result containing updated parameters or error
210    ///
211    /// # Algorithm
212    /// 1. Initialize moments on first step
213    /// 2. Apply weight decay if configured
214    /// 3. Update biased first moment: m_t = β₁ * m_{t-1} + (1 - β₁) * g_t
215    /// 4. Update biased second moment: v_t = β₂ * v_{t-1} + (1 - β₂) * g_t²
216    /// 5. Compute bias-corrected moments
217    /// 6. Compute dynamic bounds: [α_l(t), α_u(t)]
218    /// 7. Compute clipped learning rate per parameter
219    /// 8. Apply parameter update: θ_{t+1} = θ_t - η_t * m̂_t
220    ///
221    /// # Example
222    /// ```
223    /// use optirs_core::optimizers::AdaBound;
224    /// use scirs2_core::ndarray_ext::array;
225    ///
226    /// let mut optimizer = AdaBound::<f32>::default();
227    /// let params = array![1.0, 2.0, 3.0];
228    /// let grads = array![0.1, 0.2, 0.3];
229    ///
230    /// let updated_params = optimizer.step(params.view(), grads.view()).expect("optimizer.step succeeds");
231    /// ```
232    pub fn step<'a, P, G>(&mut self, params: P, grads: G) -> Result<Array1<T>>
233    where
234        P: Into<ArrayView1<'a, T>>,
235        G: Into<ArrayView1<'a, T>>,
236        T: 'a,
237    {
238        self.step_view(params.into(), grads.into())
239    }
240
241    /// Perform a single optimization step on borrowed views
242    ///
243    /// This is the concrete implementation behind [`AdaBound::step`].
244    pub fn step_view(&mut self, params: ArrayView1<T>, grads: ArrayView1<T>) -> Result<Array1<T>> {
245        let n = params.len();
246
247        if grads.len() != n {
248            return Err(OptimError::DimensionMismatch(format!(
249                "Expected gradient size {}, got {}",
250                n,
251                grads.len()
252            )));
253        }
254
255        // Initialize moments on first step
256        if self.amsbound && self.max_velocity.is_none() {
257            self.max_velocity = Some(Array1::zeros(n));
258        }
259
260        self.step_count += 1;
261        let t: T = crate::optimizers::cast_scalar(self.step_count)?;
262
263        let momentum = self.momentum.get_or_insert_with(|| Array1::zeros(n));
264        let velocity = self.velocity.get_or_insert_with(|| Array1::zeros(n));
265
266        let one = T::one();
267
268        // Apply weight decay if configured
269        let effective_grads = if self.weight_decay > T::zero() {
270            grads.to_owned() + &(params.to_owned() * self.weight_decay)
271        } else {
272            grads.to_owned()
273        };
274
275        // Update biased first moment: m_t = β₁ * m_{t-1} + (1 - β₁) * g_t
276        for i in 0..n {
277            momentum[i] = self.beta1 * momentum[i] + (one - self.beta1) * effective_grads[i];
278        }
279
280        // Update biased second moment: v_t = β₂ * v_{t-1} + (1 - β₂) * g_t²
281        for i in 0..n {
282            let grad_sq = effective_grads[i] * effective_grads[i];
283            velocity[i] = self.beta2 * velocity[i] + (one - self.beta2) * grad_sq;
284        }
285
286        // For AMSBound: v̂_t = max(v̂_{t-1}, v_t)
287        if self.amsbound {
288            let max_vel = self.max_velocity.get_or_insert_with(|| Array1::zeros(n));
289            for i in 0..n {
290                if velocity[i] > max_vel[i] {
291                    max_vel[i] = velocity[i];
292                }
293            }
294        }
295
296        // Compute bias correction terms
297        let bias_correction1 = one - self.beta1.powf(t);
298        let bias_correction2 = one - self.beta2.powf(t);
299
300        // Compute dynamic bounds
301        // Lower bound: α_l(t) = α_final * (1 - 1/(γ*t + 1))
302        let lower_bound = self.final_lr * (one - one / (self.gamma * t + one));
303
304        // Upper bound: α_u(t) = α_final * (1 + 1/(γ*t))
305        let upper_bound = self.final_lr * (one + one / (self.gamma * t));
306
307        // Apply parameter updates with clipped learning rates
308        let mut updated_params = params.to_owned();
309
310        for i in 0..n {
311            // Bias-corrected first moment
312            let m_hat = momentum[i] / bias_correction1;
313
314            // Bias-corrected second moment (or max for AMSBound)
315            let v_hat = if self.amsbound {
316                // Structural invariant: `max_velocity` is initialized to `Some` at the
317                // top of this function whenever `self.amsbound` is true, so this can
318                // never actually be `None`.
319                self.max_velocity
320                    .as_ref()
321                    .expect("AdaBound: max_velocity is Some whenever amsbound is enabled")[i]
322                    / bias_correction2
323            } else {
324                velocity[i] / bias_correction2
325            };
326
327            // Compute adaptive learning rate: α / √(v_t + ε)
328            let step_size = self.learning_rate / (v_hat.sqrt() + self.epsilon);
329
330            // Clip learning rate to dynamic bounds
331            let clipped_step_size = if step_size < lower_bound {
332                lower_bound
333            } else if step_size > upper_bound {
334                upper_bound
335            } else {
336                step_size
337            };
338
339            // Apply update: θ_{t+1} = θ_t - η_clipped * m̂_t
340            updated_params[i] = updated_params[i] - clipped_step_size * m_hat;
341        }
342
343        Ok(updated_params)
344    }
345
346    /// Get the number of optimization steps performed
347    pub fn step_count(&self) -> usize {
348        self.step_count
349    }
350
351    /// Reset the optimizer state
352    pub fn reset(&mut self) {
353        self.momentum = None;
354        self.velocity = None;
355        self.max_velocity = None;
356        self.step_count = 0;
357    }
358
359    /// Get current dynamic bounds [lower, upper]
360    pub fn current_bounds(&self) -> (T, T) {
361        if self.step_count == 0 {
362            return (self.final_lr, self.final_lr);
363        }
364
365        let t = T::from(self.step_count)
366            .expect("AdaBound: step_count must be representable in T (f32/f64)");
367        let one = T::one();
368
369        let lower_bound = self.final_lr * (one - one / (self.gamma * t + one));
370        let upper_bound = self.final_lr * (one + one / (self.gamma * t));
371
372        (lower_bound, upper_bound)
373    }
374}
375
376impl<T> Optimizer<T, Ix1> for AdaBound<T>
377where
378    T: Float + ScalarOperand + Debug + Send + Sync,
379{
380    fn step(&mut self, params: &Array1<T>, gradients: &Array1<T>) -> Result<Array1<T>> {
381        self.step_view(params.view(), gradients.view())
382    }
383
384    fn get_learning_rate(&self) -> T {
385        self.learning_rate
386    }
387
388    fn set_learning_rate(&mut self, learning_rate: T) {
389        self.learning_rate = learning_rate;
390    }
391}
392
393#[cfg(test)]
394mod tests {
395    use super::*;
396    use approx::assert_relative_eq;
397    use scirs2_core::ndarray_ext::array;
398
399    #[test]
400    fn test_adabound_creation() {
401        let optimizer = AdaBound::<f32>::default();
402        assert_eq!(optimizer.step_count(), 0);
403    }
404
405    #[test]
406    fn test_adabound_single_step() {
407        let mut optimizer = AdaBound::<f32>::default();
408        let params = array![1.0, 2.0, 3.0];
409        let grads = array![0.1, 0.2, 0.3];
410
411        let updated_params = optimizer
412            .step(params.view(), grads.view())
413            .expect("step succeeds in test_adabound_single_step");
414
415        assert_eq!(updated_params.len(), 3);
416        assert_eq!(optimizer.step_count(), 1);
417
418        // Parameters should decrease (gradient descent)
419        for i in 0..3 {
420            assert!(updated_params[i] < params[i]);
421        }
422    }
423
424    #[test]
425    fn test_adabound_multiple_steps() {
426        let mut optimizer = AdaBound::<f32>::default();
427        let mut params = array![1.0, 2.0, 3.0];
428
429        for _ in 0..10 {
430            let grads = array![0.1, 0.2, 0.3];
431            params = optimizer
432                .step(params.view(), grads.view())
433                .expect("step succeeds in test_adabound_multiple_steps");
434        }
435
436        assert_eq!(optimizer.step_count(), 10);
437    }
438
439    #[test]
440    fn test_adabound_dynamic_bounds() {
441        let mut optimizer = AdaBound::<f32>::default();
442        let params = array![1.0, 2.0, 3.0];
443        let grads = array![0.1, 0.2, 0.3];
444
445        // Before any steps, bounds should be equal to final_lr
446        let (lower0, upper0) = optimizer.current_bounds();
447        assert_relative_eq!(lower0, 0.1, epsilon = 1e-6);
448        assert_relative_eq!(upper0, 0.1, epsilon = 1e-6);
449
450        // After first step, bounds should widen
451        optimizer
452            .step(params.view(), grads.view())
453            .expect("step succeeds in test_adabound_dynamic_bounds");
454        let (lower1, upper1) = optimizer.current_bounds();
455        assert!(lower1 < upper1);
456        assert!(lower1 >= 0.0);
457
458        // After many steps, bounds should converge to final_lr
459        for _ in 0..10000 {
460            // Need many more steps for bound convergence
461            optimizer
462                .step(params.view(), grads.view())
463                .expect("step succeeds in test_adabound_dynamic_bounds");
464        }
465        let (lower_final, upper_final) = optimizer.current_bounds();
466        assert_relative_eq!(lower_final, 0.1, epsilon = 0.01);
467        assert_relative_eq!(upper_final, 0.1, epsilon = 0.01);
468    }
469
470    #[test]
471    fn test_amsbound() {
472        let mut optimizer = AdaBound::<f32>::new(0.001, 0.1, 0.9, 0.999, 1e-8, 1e-3, 0.0, true)
473            .expect("AdaBound::<f32>::new succeeds in test_amsbound");
474
475        let params = array![1.0, 2.0, 3.0];
476        let grads = array![0.1, 0.2, 0.3];
477
478        let updated_params = optimizer
479            .step(params.view(), grads.view())
480            .expect("step succeeds in test_amsbound");
481        assert_eq!(updated_params.len(), 3);
482        assert!(optimizer.max_velocity.is_some());
483    }
484
485    #[test]
486    fn test_adabound_weight_decay() {
487        let mut optimizer = AdaBound::<f32>::new(0.001, 0.1, 0.9, 0.999, 1e-8, 1e-3, 0.01, false)
488            .expect("AdaBound::<f32>::new succeeds in test_adabound_weight_decay");
489
490        let params = array![1.0, 2.0, 3.0];
491        let grads = array![0.1, 0.2, 0.3];
492
493        let updated_params = optimizer
494            .step(params.view(), grads.view())
495            .expect("step succeeds in test_adabound_weight_decay");
496
497        // With weight decay, updates should be larger
498        for i in 0..3 {
499            assert!(updated_params[i] < params[i]);
500        }
501    }
502
503    #[test]
504    fn test_adabound_convergence() {
505        // Test convergence on quadratic function f(x) = x²
506        let mut optimizer = AdaBound::<f64>::default();
507        let mut params = array![5.0];
508
509        for _ in 0..500 {
510            // AdaBound needs more iterations for tight convergence
511            let grads = params.mapv(|x| 2.0 * x);
512            params = optimizer
513                .step(params.view(), grads.view())
514                .expect("step succeeds in test_adabound_convergence");
515        }
516
517        // Should converge close to zero
518        assert!(
519            params[0].abs() < 0.1,
520            "Failed to converge, got {}",
521            params[0]
522        );
523    }
524
525    #[test]
526    fn test_adabound_reset() {
527        let mut optimizer = AdaBound::<f32>::default();
528        let params = array![1.0, 2.0, 3.0];
529        let grads = array![0.1, 0.2, 0.3];
530
531        optimizer
532            .step(params.view(), grads.view())
533            .expect("step succeeds in test_adabound_reset");
534        assert_eq!(optimizer.step_count(), 1);
535
536        optimizer.reset();
537        assert_eq!(optimizer.step_count(), 0);
538        assert!(optimizer.momentum.is_none());
539        assert!(optimizer.velocity.is_none());
540    }
541
542    /// AdaBound must be usable through the generic `Optimizer` trait.
543    #[test]
544    fn test_adabound_optimizer_trait() {
545        let mut optimizer = AdaBound::<f64>::default();
546        let params = scirs2_core::ndarray_ext::array![1.0f64, 2.0, 3.0];
547        let grads = scirs2_core::ndarray_ext::array![0.1f64, 0.2, 0.3];
548
549        let updated =
550            Optimizer::<f64, scirs2_core::ndarray::Ix1>::step(&mut optimizer, &params, &grads)
551                .expect("trait step failed");
552        assert_eq!(updated.len(), 3);
553
554        let lr = Optimizer::<f64, scirs2_core::ndarray::Ix1>::get_learning_rate(&optimizer);
555        Optimizer::<f64, scirs2_core::ndarray::Ix1>::set_learning_rate(&mut optimizer, lr * 2.0);
556        assert!(
557            (Optimizer::<f64, scirs2_core::ndarray::Ix1>::get_learning_rate(&optimizer) - lr * 2.0)
558                .abs()
559                < 1e-12
560        );
561
562        // The generic inherent `step` also accepts plain references.
563        let again = optimizer.step(&params, &grads).expect("ref step failed");
564        assert_eq!(again.len(), 3);
565    }
566}