nato-opt 0.1.0

NATO Optimizer and Spectral Penalties (Rust Port)
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;

                // m = beta1 * m + (1 - beta1) * grad
                let next_m = &state.m * self.beta1 + &grad * (1.0 - self.beta1);
                let _ = state.m.copy_(&next_m);

                // v = beta2 * v + (1 - beta2) * grad^2
                let next_v = &state.v * self.beta2 + (&grad * &grad) * (1.0 - self.beta2);
                let _ = state.v.copy_(&next_v);

                // Proper bias correction
                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);

                // Tether term
                let update = update + (p.shallow_clone() - &state.prev) * self.gamma;

                // Apply update
                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();
                }
            }
        });
    }
}