1use 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#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct AdaBound<T: Float> {
37 learning_rate: T,
39
40 final_lr: T,
43
44 beta1: T,
46
47 beta2: T,
49
50 epsilon: T,
52
53 gamma: T,
56
57 weight_decay: T,
59
60 amsbound: bool,
62
63 momentum: Option<Array1<T>>,
65
66 velocity: Option<Array1<T>>,
68
69 max_velocity: Option<Array1<T>>,
71
72 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 #[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 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 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 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 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 for i in 0..n {
277 momentum[i] = self.beta1 * momentum[i] + (one - self.beta1) * effective_grads[i];
278 }
279
280 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 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 let bias_correction1 = one - self.beta1.powf(t);
298 let bias_correction2 = one - self.beta2.powf(t);
299
300 let lower_bound = self.final_lr * (one - one / (self.gamma * t + one));
303
304 let upper_bound = self.final_lr * (one + one / (self.gamma * t));
306
307 let mut updated_params = params.to_owned();
309
310 for i in 0..n {
311 let m_hat = momentum[i] / bias_correction1;
313
314 let v_hat = if self.amsbound {
316 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 let step_size = self.learning_rate / (v_hat.sqrt() + self.epsilon);
329
330 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 updated_params[i] = updated_params[i] - clipped_step_size * m_hat;
341 }
342
343 Ok(updated_params)
344 }
345
346 pub fn step_count(&self) -> usize {
348 self.step_count
349 }
350
351 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 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 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 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 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 for _ in 0..10000 {
460 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 for i in 0..3 {
499 assert!(updated_params[i] < params[i]);
500 }
501 }
502
503 #[test]
504 fn test_adabound_convergence() {
505 let mut optimizer = AdaBound::<f64>::default();
507 let mut params = array![5.0];
508
509 for _ in 0..500 {
510 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 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 #[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, ¶ms, &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 let again = optimizer.step(¶ms, &grads).expect("ref step failed");
564 assert_eq!(again.len(), 3);
565 }
566}