use candle_core::backprop::GradStore;
use candle_core::{Result, Tensor};
use candle_nn::optim::Optimizer;
pub fn clipped_backward_step<O: Optimizer>(
opt: &mut O,
loss: &Tensor,
max_norm: f64,
) -> Result<f64> {
let mut grads = loss.backward()?;
let norm = clip_grad_global_norm(&mut grads, max_norm)?;
opt.step(&grads)?;
Ok(norm)
}
pub fn clip_grad_global_norm(grads: &mut GradStore, max_norm: f64) -> Result<f64> {
if max_norm <= 0.0 {
return Ok(0.0);
}
let ids: Vec<_> = grads.get_ids().copied().collect();
let mut sumsq: Option<Tensor> = None;
for id in &ids {
if let Some(g) = grads.get_id(*id) {
let s = g.sqr()?.sum_all()?;
sumsq = Some(match sumsq {
None => s,
Some(t) => (t + s)?,
});
}
}
let norm = match sumsq {
Some(t) => (t.to_scalar::<f32>()? as f64).sqrt(),
None => 0.0,
};
if norm > max_norm && norm > 0.0 {
let scale = max_norm / norm;
for id in &ids {
if let Some(g) = grads.get_id(*id) {
let scaled = g.affine(scale, 0.0)?;
grads.insert_id(*id, scaled);
}
}
}
Ok(norm)
}