use zyx::Tensor;
use zyx_derive::Module;
#[derive(Module)]
#[cfg_attr(feature = "py", pyo3::pyclass)]
pub struct Adam {
pub learning_rate: f32,
pub betas: (f32, f32),
pub eps: f32,
pub weight_decay: f32,
pub amsgrad: bool,
pub m: Vec<Tensor>,
pub v: Vec<Tensor>,
pub vm: Vec<Tensor>,
pub t: usize,
}
impl Default for Adam {
fn default() -> Self {
Self {
learning_rate: 0.001,
betas: (0.9, 0.999),
eps: 1e-8,
weight_decay: 0.0,
amsgrad: false,
m: Vec::new(),
v: Vec::new(),
vm: Vec::new(),
t: 0,
}
}
}
impl Adam {
pub fn update<'a>(
&mut self,
parameters: impl IntoIterator<Item = &'a mut Tensor>,
gradients: impl IntoIterator<Item = Tensor>,
) {
use zyx::Scalar;
self.t += 1;
for (i, (param, mut grad)) in parameters.into_iter().zip(gradients).enumerate() {
if self.weight_decay != 0.0 {
grad = grad + &*param * self.weight_decay;
}
if let Some(m) = self.m.get_mut(i) {
*m = &*m * self.betas.0 + &grad * (1.0 - self.betas.0);
} else {
self.m.push(&grad * (1.0 - self.betas.0));
}
if let Some(v) = self.v.get_mut(i) {
*v = &*v * self.betas.1 + &grad * &grad * (1.0 - self.betas.1);
} else {
self.v.push(&grad * &grad * (1.0 - self.betas.1));
}
let mh = &self.m[i] / (1.0 - self.betas.0.pow(self.t as f32));
let vh = &self.v[i] / (1.0 - self.betas.1.pow(self.t as f32));
if self.amsgrad {
if let Some(vm) = self.vm.get_mut(i) {
*vm = vm.cmplt(&vh).unwrap().where_(vh, &*vm).unwrap();
} else {
self.vm.push(vh);
}
*param = (&*param - self.learning_rate * mh / (self.vm[i].sqrt() + self.eps))
.cast(param.dtype());
} else {
*param = (&*param - self.learning_rate * mh / (vh.sqrt() + self.eps))
.cast(param.dtype());
}
}
}
}