use crate::{Element, Parameters, Symbol, Tensor};
use super::Optimizer;
use super::optimizer::assert_single_value;
#[derive(Debug, Clone)]
pub struct Adam<E> {
beta1: Tensor<E>,
beta2: Tensor<E>,
epsilon: Tensor<E>,
first_share: Tensor<E>,
second_share: Tensor<E>,
first: Option<Parameters<E>>,
second: Option<Parameters<E>>,
beta1_power: Tensor<E>,
beta2_power: Tensor<E>,
}
impl<E: Element> Adam<E> {
pub fn new(beta1: Tensor<E>, beta2: Tensor<E>, epsilon: Tensor<E>) -> Self {
assert_single_value(&beta1, "beta1");
assert_single_value(&beta2, "beta2");
assert_single_value(&epsilon, "epsilon");
Self {
first_share: beta1.one_like() - beta1.clone(),
second_share: beta2.one_like() - beta2.clone(),
beta1_power: beta1.one_like(),
beta2_power: beta2.one_like(),
beta1,
beta2,
epsilon,
first: None,
second: None,
}
}
fn direction(&mut self, gradients: &Parameters<E>) -> Parameters<E> {
let zeros = || gradients.map(|gradient| gradient.zero_like());
let first = self.first.take().unwrap_or_else(zeros);
let second = self.second.take().unwrap_or_else(zeros);
let first = first.scale(&self.beta1) + gradients.scale(&self.first_share);
let squared = gradients.map(|gradient| gradient.clone() * gradient.clone());
let second = second.scale(&self.beta2) + squared.scale(&self.second_share);
self.beta1_power = self.beta1_power.clone() * self.beta1.clone();
self.beta2_power = self.beta2_power.clone() * self.beta2.clone();
let first_correction = self.beta1_power.one_like() - self.beta1_power.clone();
let second_correction = self.beta2_power.one_like() - self.beta2_power.clone();
let direction = first.zip(&second, |first, second| {
let corrected_first = first.clone() / first_correction.broadcast_like(first);
let corrected_second = second.clone() / second_correction.broadcast_like(second);
corrected_first
/ (corrected_second.sqrt() + self.epsilon.broadcast_like(&corrected_second))
});
self.first = Some(first);
self.second = Some(second);
direction
}
}
impl<E: Element> Optimizer<E> for Adam<E> {
fn step(
&mut self,
parameters: &Parameters<E>,
gradients: &Parameters<E>,
learning_rate: &Tensor<E>,
) -> Parameters<E> {
let direction = self.direction(gradients);
parameters.step(&direction, |parameter, direction| {
parameter.clone() - direction.clone() * learning_rate.broadcast_like(direction)
})
}
}
#[derive(Debug, Clone)]
pub struct AdamW<E> {
adam: Adam<E>,
decay: Tensor<E>,
}
impl<E: Element> AdamW<E> {
pub fn new(beta1: Tensor<E>, beta2: Tensor<E>, epsilon: Tensor<E>, decay: Tensor<E>) -> Self {
assert_single_value(&decay, "decay");
Self {
adam: Adam::new(beta1, beta2, epsilon),
decay,
}
}
pub fn step_where(
&mut self,
parameters: &Parameters<E>,
gradients: &Parameters<E>,
learning_rate: &Tensor<E>,
mut policy: impl FnMut(Symbol, &Tensor<E>) -> bool,
) -> Parameters<E> {
let direction = self.adam.direction(gradients);
parameters.step_each(&direction, |symbol, current, direction| {
let stepped =
current.clone() - direction.clone() * learning_rate.broadcast_like(direction);
if policy(symbol, current) {
let decayed = current.clone()
* self.decay.broadcast_like(current)
* learning_rate.broadcast_like(current);
stepped - decayed
} else {
stepped
}
})
}
}
impl<E: Element> Optimizer<E> for AdamW<E> {
fn step(
&mut self,
parameters: &Parameters<E>,
gradients: &Parameters<E>,
learning_rate: &Tensor<E>,
) -> Parameters<E> {
self.step_where(parameters, gradients, learning_rate, |_, parameter| {
parameter.shape().rank() >= 2
})
}
}
#[cfg(test)]
#[path = "tests/adam_tests.rs"]
mod adam_tests;