nato-opt 0.1.0

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

pub fn fourier_spectral_penalty(
    named_parameters: &[(String, Tensor)],
    lambda_fsp: f64,
    tau: Option<f64>,
    mask_mode: &str, // "hypercube" or "radial"
    include_conv: bool,
    include_linear: bool,
    module_whitelist: Option<&HashSet<String>>,
    module_blacklist: Option<&HashSet<String>>,
) -> Tensor {
    let mut device = tch::Device::Cpu;
    let mut has_param = false;

    for (_, p) in named_parameters.iter() {
        if p.requires_grad() {
            device = p.device();
            has_param = true;
            break;
        }
    }

    let mut penalty = Tensor::zeros([], (Kind::Float, device));
    if !has_param {
        return penalty;
    }

    for (pname, p) in named_parameters.iter() {
        if !p.requires_grad() {
            continue;
        }

        let is_weight = pname.ends_with("weight");
        if !is_weight {
            continue;
        }

        let module_name = if let Some(idx) = pname.rfind('.') {
            &pname[0..idx]
        } else {
            ""
        };

        if let Some(wl) = module_whitelist {
            if !wl.contains(module_name) {
                continue;
            }
        }
        if let Some(bl) = module_blacklist {
            if bl.contains(module_name) {
                continue;
            }
        }

        let mut applies = false;
        let dim = p.dim();
        if dim == 2 && include_linear {
            applies = true;
        } else if dim > 2 && include_conv {
            applies = true;
        }

        if !applies {
            continue;
        }

        let w = if p.device() == device {
            p.shallow_clone()
        } else {
            p.to_device(device)
        };

        let w_fft = w.to_kind(Kind::ComplexFloat).fft_fftn(
            None::<&[i64]>,
            None::<&[i64]>,
            "ortho",
        );

        if let Some(t) = tau {
            let shape = w.size();
            let mut mask_low = Tensor::ones(shape.as_slice(), (Kind::Bool, device));

            if mask_mode == "hypercube" {
                for (axis, &dim_size) in shape.iter().enumerate() {
                    let mut idx_mask = vec![false; dim_size as usize];
                    for i in 0..dim_size {
                        let f_centered = if i <= dim_size / 2 { i } else { i - dim_size };
                        if (f_centered as f64).abs() <= t {
                            idx_mask[i as usize] = true;
                        }
                    }
                    let axis_low = Tensor::from_slice(&idx_mask).to_device(device);
                    let mut view_shape = vec![1; shape.len()];
                    view_shape[axis] = dim_size;
                    let axis_low = axis_low.view(view_shape.as_slice());
                    mask_low = mask_low.logical_and(&axis_low);
                }
            } else if mask_mode == "radial" {
                let mut r_squared_data = Tensor::zeros(shape.as_slice(), (Kind::Float, device));
                for (axis, &dim_size) in shape.iter().enumerate() {
                    let mut idx_vals = vec![0.0f32; dim_size as usize];
                    for i in 0..dim_size {
                        let f_centered = if i <= dim_size / 2 { i } else { i - dim_size };
                        idx_vals[i as usize] = f_centered as f32;
                    }
                    let axis_vals = Tensor::from_slice(&idx_vals).to_device(device);
                    let mut view_shape = vec![1; shape.len()];
                    view_shape[axis] = dim_size;
                    let axis_grid = axis_vals.view(view_shape.as_slice());
                    r_squared_data = r_squared_data + axis_grid.pow_tensor_scalar(2.0);
                }
                mask_low = r_squared_data.le_tensor(&Tensor::from_slice(&[(t * t) as f32]).to_device(device));
            }

            let mask_hf = mask_low.logical_not();
            let hf_energy = w_fft.masked_select(&mask_hf).abs().pow_tensor_scalar(2.0).sum(Kind::Float);
            penalty = penalty + hf_energy;
        } else {
            let total_energy = w_fft.abs().pow_tensor_scalar(2.0).sum(Kind::Float);
            penalty = penalty + total_energy;
        }
    }

    penalty * lambda_fsp
}