use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand, Zip};
use scirs2_core::numeric::Float;
use std::fmt::Debug;
use crate::error::{OptimError, Result};
use crate::optimizers::Optimizer;
#[derive(Debug, Clone)]
pub struct RAdam<A: Float + ScalarOperand + Debug> {
learning_rate: A,
beta1: A,
beta2: A,
epsilon: A,
weight_decay: A,
m: Option<Vec<Array<A, IxDyn>>>,
v: Option<Vec<Array<A, IxDyn>>>,
t: Vec<usize>,
rho_inf: A,
}
impl<A: Float + ScalarOperand + Debug + Send + Sync> RAdam<A> {
pub fn new(learning_rate: A) -> Self {
let beta2 = A::from(0.999).expect("RAdam: default beta2 (0.999) must fit in A");
Self {
learning_rate,
beta1: A::from(0.9).expect("RAdam: default beta1 (0.9) must fit in A"),
beta2,
epsilon: A::from(1e-8).expect("RAdam: default epsilon (1e-8) must fit in A"),
weight_decay: A::zero(),
m: None,
v: None,
t: Vec::new(),
rho_inf: A::from(2.0).expect("RAdam: integer literal 2.0 must fit in A")
/ (A::one() - beta2)
- A::one(),
}
}
pub fn new_with_config(
learning_rate: A,
beta1: A,
beta2: A,
epsilon: A,
weight_decay: A,
) -> Self {
Self {
learning_rate,
beta1,
beta2,
epsilon,
weight_decay,
m: None,
v: None,
t: Vec::new(),
rho_inf: A::from(2.0).expect("RAdam: integer literal 2.0 must fit in A")
/ (A::one() - beta2)
- A::one(),
}
}
pub fn set_beta1(&mut self, beta1: A) -> &mut Self {
self.beta1 = beta1;
self
}
pub fn get_beta1(&self) -> A {
self.beta1
}
pub fn set_beta2(&mut self, beta2: A) -> &mut Self {
self.beta2 = beta2;
self.rho_inf = A::from(2.0).expect("RAdam: integer literal 2.0 must fit in A")
/ (A::one() - beta2)
- A::one();
self
}
pub fn get_beta2(&self) -> A {
self.beta2
}
pub fn set_epsilon(&mut self, epsilon: A) -> &mut Self {
self.epsilon = epsilon;
self
}
pub fn get_epsilon(&self) -> A {
self.epsilon
}
pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
self.weight_decay = weight_decay;
self
}
pub fn get_weight_decay(&self) -> A {
self.weight_decay
}
pub fn learning_rate(&self) -> A {
self.learning_rate
}
pub fn set_lr(&mut self, lr: A) {
self.learning_rate = lr;
}
pub fn reset(&mut self) {
self.m = None;
self.v = None;
self.t.clear();
}
pub fn timestep(&self, index: usize) -> usize {
self.t.get(index).copied().unwrap_or(0)
}
pub fn rho_inf(&self) -> A {
self.rho_inf
}
pub fn rho_t(&self, t: usize) -> Option<A> {
if t == 0 {
return None;
}
let exp = i32::try_from(t).ok()?;
let two = A::one() + A::one();
let t_f = A::from(t)?;
let beta2_t = self.beta2.powi(exp);
Some(self.rho_inf - two * t_f * beta2_t / (A::one() - beta2_t))
}
pub fn rectification_term(&self, t: usize) -> Option<A> {
let rho_t = self.rho_t(t)?;
let two = A::one() + A::one();
let four = two + two;
if rho_t <= four {
return None;
}
let rho_inf = self.rho_inf;
let numerator = (rho_t - four) * (rho_t - two) * rho_inf;
let denominator = (rho_inf - four) * (rho_inf - two) * rho_t;
if denominator <= A::zero() {
return None;
}
Some((numerator / denominator).sqrt())
}
fn advance_state(&mut self, index: usize, dim: &IxDyn) -> Result<usize> {
let m = self.m.get_or_insert_with(Vec::new);
let v = self.v.get_or_insert_with(Vec::new);
while m.len() <= index {
m.push(Array::zeros(dim.clone()));
}
while v.len() <= index {
v.push(Array::zeros(dim.clone()));
}
while self.t.len() <= index {
self.t.push(0);
}
if m[index].raw_dim() != *dim || v[index].raw_dim() != *dim {
m[index] = Array::zeros(dim.clone());
v[index] = Array::zeros(dim.clone());
self.t[index] = 0;
}
let next = self.t[index].checked_add(1).ok_or_else(|| {
OptimError::InvalidConfig(
"Timestep counter overflow - too many optimization steps".to_string(),
)
})?;
self.t[index] = next;
Ok(next)
}
pub fn step_inplace_indexed<D: Dimension>(
&mut self,
index: usize,
params: &mut Array<A, D>,
gradients: &Array<A, D>,
) -> Result<()> {
if params.shape() != gradients.shape() {
return Err(OptimError::DimensionMismatch(format!(
"Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
params.shape(),
gradients.shape()
)));
}
let dim = params.raw_dim().into_dyn();
let t = self.advance_state(index, &dim)?;
let exp = i32::try_from(t).map_err(|_| {
OptimError::InvalidConfig(
"Timestep too large for bias correction calculation".to_string(),
)
})?;
let beta1 = self.beta1;
let beta2 = self.beta2;
let lr = self.learning_rate;
let eps = self.epsilon;
let weight_decay = self.weight_decay;
let use_weight_decay = weight_decay > A::zero();
let one = A::one();
let bias_correction1 = one - beta1.powi(exp);
let bias_correction2 = one - beta2.powi(exp);
let rect = self.rectification_term(t);
let m = self
.m
.as_mut()
.ok_or_else(|| OptimError::InvalidConfig("RAdam state not initialized".to_string()))?;
let v = self
.v
.as_mut()
.ok_or_else(|| OptimError::InvalidConfig("RAdam state not initialized".to_string()))?;
let mut params_view = params.view_mut().into_dyn();
let gradients_view = gradients.view().into_dyn();
Zip::from(&mut params_view)
.and(&gradients_view)
.and(&mut m[index])
.and(&mut v[index])
.for_each(|p, &g, m_i, v_i| {
let grad = if use_weight_decay {
g + weight_decay * *p
} else {
g
};
*m_i = *m_i * beta1 + grad * (one - beta1);
*v_i = *v_i * beta2 + grad * grad * (one - beta2);
let m_hat = *m_i / bias_correction1;
match rect {
Some(r_t) => {
let v_hat = *v_i / bias_correction2;
*p = *p - lr * r_t * m_hat / (v_hat.sqrt() + eps);
}
None => {
*p = *p - lr * m_hat;
}
}
});
Ok(())
}
pub fn step_inplace<D: Dimension>(
&mut self,
params: &mut Array<A, D>,
gradients: &Array<A, D>,
) -> Result<()> {
self.step_inplace_indexed(0, params, gradients)
}
pub fn step_indexed<D: Dimension>(
&mut self,
index: usize,
params: &Array<A, D>,
gradients: &Array<A, D>,
) -> Result<Array<A, D>> {
let mut updated = params.to_owned();
self.step_inplace_indexed(index, &mut updated, gradients)?;
Ok(updated)
}
}
impl<A, D> Optimizer<A, D> for RAdam<A>
where
A: Float + ScalarOperand + Debug + Send + Sync + std::convert::From<f64>,
D: Dimension,
{
fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
self.step_indexed(0, params, gradients)
}
fn step_list(
&mut self,
params_list: &[&Array<A, D>],
gradients_list: &[&Array<A, D>],
) -> Result<Vec<Array<A, D>>> {
if params_list.len() != gradients_list.len() {
return Err(OptimError::InvalidConfig(format!(
"Number of parameter arrays ({}) does not match number of gradient arrays ({})",
params_list.len(),
gradients_list.len()
)));
}
let mut results = Vec::with_capacity(params_list.len());
for (index, (params, grads)) in params_list.iter().zip(gradients_list.iter()).enumerate() {
results.push(self.step_indexed(index, params, grads)?);
}
Ok(results)
}
fn get_learning_rate(&self) -> A {
self.learning_rate
}
fn set_learning_rate(&mut self, learning_rate: A) {
self.learning_rate = learning_rate;
}
}
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::Array1;
#[test]
fn test_radam_step() {
let params = Array1::zeros(3);
let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let mut optimizer = RAdam::new(0.01);
let new_params = optimizer
.step(¶ms, &gradients)
.expect("optimizer.step succeeds in test_radam_step");
assert!(new_params.iter().all(|&x| x != 0.0));
for i in 1..3 {
assert!(new_params[i].abs() > new_params[i - 1].abs());
}
}
#[test]
fn test_radam_multiple_steps() {
let mut params = Array1::zeros(3);
let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let mut optimizer = RAdam::new(0.01);
for _ in 0..100 {
params = optimizer
.step(¶ms, &gradients)
.expect("optimizer.step succeeds in test_radam_multiple_steps");
}
for i in 1..3 {
assert!(params[i].abs() > params[i - 1].abs());
}
}
#[test]
fn test_radam_weight_decay() {
let params = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let gradients = Array1::from_vec(vec![0.01, 0.01, 0.01]);
let mut optimizer = RAdam::new_with_config(
0.01, 0.9, 0.999, 1e-8, 0.1, );
let new_params = optimizer
.step(¶ms, &gradients)
.expect("optimizer.step succeeds in test_radam_weight_decay");
for i in 0..3 {
assert!(new_params[i].abs() < params[i].abs());
}
}
#[test]
fn test_radam_reset() {
let params = Array1::zeros(3);
let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
let mut optimizer = RAdam::new(0.01);
optimizer
.step(¶ms, &gradients)
.expect("optimizer.step succeeds in test_radam_reset");
assert_eq!(optimizer.timestep(0), 1);
assert!(optimizer.m.is_some());
assert!(optimizer.v.is_some());
optimizer.reset();
assert_eq!(optimizer.timestep(0), 0);
assert!(optimizer.m.is_none());
assert!(optimizer.v.is_none());
}
}