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(crate) fn global_sumsq(grads: &GradStore) -> Result<f64> {
let parts: Vec<Tensor> = grads
.get_ids()
.filter_map(|id| grads.get_id(*id))
.map(|g| g.sqr()?.sum_all()?.to_dtype(candle_core::DType::F32))
.collect::<Result<_>>()?;
if parts.is_empty() {
return Ok(0.0);
}
let mut sums: Vec<f32> = Tensor::stack(&parts, 0)?.to_vec1()?;
sums.sort_by(f32::total_cmp);
Ok(sums.iter().map(|&s| f64::from(s)).sum())
}
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 norm = global_sumsq(grads)?.sqrt();
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)
}