use crate::compat::*;
use crate::optimizer::Optimizer;
use hodu_core::{error::HoduResult, scalar::Scalar, tensor::Tensor};
#[derive(Optimizer, Clone)]
pub struct Adam {
learning_rate: Scalar,
beta1: Scalar,
beta2: Scalar,
epsilon: Scalar,
t: usize,
m: Vec<Tensor>,
v: Vec<Tensor>,
}
impl Adam {
pub fn new(
learning_rate: impl Into<Scalar>,
beta1: impl Into<Scalar>,
beta2: impl Into<Scalar>,
epsilon: impl Into<Scalar>,
) -> Self {
Self {
learning_rate: learning_rate.into(),
beta1: beta1.into(),
beta2: beta2.into(),
epsilon: epsilon.into(),
t: 0,
m: Vec::new(),
v: Vec::new(),
}
}
pub fn step(&mut self, parameters: &mut [&mut Tensor]) -> HoduResult<()> {
if self.m.is_empty() || self.v.is_empty() {
self.m = parameters
.iter()
.map(|param| Tensor::zeros_like(param))
.collect::<HoduResult<Vec<_>>>()?;
self.v = parameters
.iter()
.map(|param| Tensor::zeros_like(param))
.collect::<HoduResult<Vec<_>>>()?;
}
self.t += 1;
for ((param, m), v) in parameters.iter_mut().zip(self.m.iter_mut()).zip(self.v.iter_mut()) {
let grad = param.grad()?;
let lr = self.learning_rate.to_dtype(grad.get_dtype());
let beta1 = self.beta1.to_dtype(grad.get_dtype());
let beta2 = self.beta2.to_dtype(grad.get_dtype());
let epsilon = self.epsilon.to_dtype(grad.get_dtype());
let one = Scalar::one(grad.get_dtype());
let mut beta1_t = beta1;
let mut beta2_t = beta2;
for _ in 0..self.t {
beta1_t = beta1_t * beta1;
beta2_t = beta2_t * beta2;
}
let one_minus_beta1_t = one - beta1_t;
let one_minus_beta2_t = one - beta2_t;
let lr_t = lr * (one_minus_beta2_t.sqrt()) / one_minus_beta1_t;
let one_minus_beta1 = one - beta1;
let one_minus_beta2 = one - beta2;
*m = m.mul_scalar(beta1)?.add(&grad.mul_scalar(one_minus_beta1)?)?;
*v = v.mul_scalar(beta2)?.add(&grad.square()?.mul_scalar(one_minus_beta2)?)?;
let m_hat = m.div_scalar(one_minus_beta1_t)?;
let v_hat = v.div_scalar(one_minus_beta2_t)?;
let step = m_hat.div(&v_hat.sqrt()?.add_scalar(epsilon)?)?;
param.set_(¶m.sub(&step.mul_scalar(lr_t)?)?)?;
}
Ok(())
}
pub fn set_learning_rate(&mut self, learning_rate: impl Into<Scalar>) {
self.learning_rate = learning_rate.into();
}
}
#[derive(Optimizer, Clone)]
pub struct AdamW {
learning_rate: Scalar,
beta1: Scalar,
beta2: Scalar,
epsilon: Scalar,
weight_decay: Scalar,
t: usize,
m: Vec<Tensor>,
v: Vec<Tensor>,
}
impl AdamW {
pub fn new(
learning_rate: impl Into<Scalar>,
beta1: impl Into<Scalar>,
beta2: impl Into<Scalar>,
epsilon: impl Into<Scalar>,
weight_decay: impl Into<Scalar>,
) -> Self {
Self {
learning_rate: learning_rate.into(),
beta1: beta1.into(),
beta2: beta2.into(),
epsilon: epsilon.into(),
weight_decay: weight_decay.into(),
t: 0,
m: Vec::new(),
v: Vec::new(),
}
}
pub fn step(&mut self, parameters: &mut [&mut Tensor]) -> HoduResult<()> {
if self.m.is_empty() || self.v.is_empty() {
self.m = parameters
.iter()
.map(|param| Tensor::zeros_like(param))
.collect::<HoduResult<Vec<_>>>()?;
self.v = parameters
.iter()
.map(|param| Tensor::zeros_like(param))
.collect::<HoduResult<Vec<_>>>()?;
}
self.t += 1;
for ((param, m), v) in parameters.iter_mut().zip(self.m.iter_mut()).zip(self.v.iter_mut()) {
let grad = param.grad()?;
let lr = self.learning_rate.to_dtype(grad.get_dtype());
let beta1 = self.beta1.to_dtype(grad.get_dtype());
let beta2 = self.beta2.to_dtype(grad.get_dtype());
let epsilon = self.epsilon.to_dtype(grad.get_dtype());
let weight_decay = self.weight_decay.to_dtype(grad.get_dtype());
let one = Scalar::one(grad.get_dtype());
let mut beta1_t = beta1;
let mut beta2_t = beta2;
for _ in 1..self.t {
beta1_t = beta1_t * beta1;
beta2_t = beta2_t * beta2;
}
let one_minus_beta1 = one - beta1;
let one_minus_beta2 = one - beta2;
let one_minus_beta1_t = one - beta1_t;
let one_minus_beta2_t = one - beta2_t;
*m = m.mul_scalar(beta1)?.add(&grad.mul_scalar(one_minus_beta1)?)?;
*v = v.mul_scalar(beta2)?.add(&grad.square()?.mul_scalar(one_minus_beta2)?)?;
let m_hat = m.div_scalar(one_minus_beta1_t)?;
let v_hat = v.div_scalar(one_minus_beta2_t)?;
let adaptive_step = m_hat.div(&v_hat.sqrt()?.add_scalar(epsilon)?)?;
let decay_factor = one - lr * weight_decay;
let decayed_param = param.mul_scalar(decay_factor)?;
param.set_(&decayed_param.sub(&adaptive_step.mul_scalar(lr)?)?)?;
}
Ok(())
}
pub fn set_learning_rate(&mut self, learning_rate: impl Into<Scalar>) {
self.learning_rate = learning_rate.into();
}
pub fn set_weight_decay(&mut self, weight_decay: impl Into<Scalar>) {
self.weight_decay = weight_decay.into();
}
pub fn get_learning_rate(&self) -> Scalar {
self.learning_rate
}
pub fn get_weight_decay(&self) -> Scalar {
self.weight_decay
}
}