nato-opt 0.1.0

NATO Optimizer and Spectral Penalties (Rust Port)
use std::collections::HashMap;
use tch::{Tensor, Kind};

/// State for the Kakeya Directional Penalty.
/// Maintains the previous gradient for each parameter.
pub struct KakeyaState {
    pub prev_grads: HashMap<String, Tensor>,
}

impl KakeyaState {
    pub fn new() -> Self {
        Self {
            prev_grads: HashMap::new(),
        }
    }
}

/// Compute Kakeya directional penalty based on gradient direction consistency.
pub fn kakeya_directional_penalty(
    named_parameters: &[(String, Tensor)],
    state: &mut KakeyaState,
    lambda_k: f64,
) -> Tensor {
    // If no device can be inferred from gradients, we just return a scalar 0
    let mut device = tch::Device::Cpu;
    let mut has_grad = false;
    
    for (_, p) in named_parameters.iter() {
        if p.requires_grad() && p.grad().defined() {
            device = p.device();
            has_grad = true;
            break;
        }
    }

    let mut penalty = Tensor::zeros([], (Kind::Float, device));

    if !has_grad {
        return penalty;
    }

    for (name, p) in named_parameters.iter() {
        if !p.requires_grad() || !p.grad().defined() {
            continue;
        }

        let grad = p.grad().view([-1]);

        if !state.prev_grads.contains_key(name) {
            state.prev_grads.insert(name.clone(), grad.detach().copy());
            continue;
        }

        let prev_grad = state.prev_grads.get(name).unwrap();

        let grad_norm = grad.norm();
        let prev_norm = prev_grad.norm();

        // Need to convert to f64 to compare with 1e-8
        let g_norm_val = f64::try_from(&grad_norm).unwrap_or(0.0);
        let p_norm_val = f64::try_from(&prev_norm).unwrap_or(0.0);

        if g_norm_val > 1e-8 && p_norm_val > 1e-8 {
            let cos_sim = grad.dot(prev_grad) / (grad_norm * prev_norm + 1e-8);
            penalty = penalty + cos_sim.pow_tensor_scalar(2.0);
        }

        state.prev_grads.insert(name.clone(), grad.detach().copy());
    }

    penalty * lambda_k
}