use std::collections::HashMap;
use tch::Tensor;
pub struct NatoOptimizerState {
pub m: Tensor,
pub v: Tensor,
pub prev: Tensor,
pub step: i64,
}
pub struct NatoOptimizer {
pub lr: f64,
pub beta1: f64,
pub beta2: f64,
pub epsilon: f64,
pub gamma: f64,
pub tether_interval: i64,
pub states: HashMap<String, NatoOptimizerState>,
}
impl NatoOptimizer {
pub fn new(
lr: f64,
beta1: f64,
beta2: f64,
epsilon: f64,
gamma: f64,
tether_interval: i64,
) -> Self {
Self {
lr,
beta1,
beta2,
epsilon,
gamma,
tether_interval,
states: HashMap::new(),
}
}
pub fn step(&mut self, named_parameters: &[(String, Tensor)]) {
tch::no_grad(|| {
for (name, p) in named_parameters.iter() {
if !p.requires_grad() || !p.grad().defined() {
continue;
}
let grad = p.grad();
if !self.states.contains_key(name) {
self.states.insert(
name.clone(),
NatoOptimizerState {
m: Tensor::zeros_like(p),
v: Tensor::zeros_like(p),
prev: p.shallow_clone(),
step: 0,
},
);
}
let state = self.states.get_mut(name).unwrap();
state.step += 1;
let t = state.step;
let next_m = &state.m * self.beta1 + &grad * (1.0 - self.beta1);
let _ = state.m.copy_(&next_m);
let next_v = &state.v * self.beta2 + (&grad * &grad) * (1.0 - self.beta2);
let _ = state.v.copy_(&next_v);
let m_hat = &state.m / (1.0 - self.beta1.powi(t as i32));
let v_hat = &state.v / (1.0 - self.beta2.powi(t as i32));
let update = (m_hat * self.lr) / (v_hat.sqrt() + self.epsilon);
let update = update + (p.shallow_clone() - &state.prev) * self.gamma;
let mut p_mut = p.shallow_clone();
let _ = p_mut.copy_(&(p.shallow_clone() - update));
if t % self.tether_interval == 0 {
state.prev = p.shallow_clone();
}
}
});
}
}