use std::collections::HashMap;
use tch::{Tensor, Kind};
pub struct KakeyaState {
pub prev_grads: HashMap<String, Tensor>,
}
impl KakeyaState {
pub fn new() -> Self {
Self {
prev_grads: HashMap::new(),
}
}
}
pub fn kakeya_directional_penalty(
named_parameters: &[(String, Tensor)],
state: &mut KakeyaState,
lambda_k: f64,
) -> Tensor {
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();
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
}