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, 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
}