burn-optim 0.22.0

Optimizer building blocks for the Burn deep learning framework
use burn_core as burn;

use burn::{
    config::Config,
    tensor::{DType, FloatDType, Tensor},
};

/// Gradient Clipping provides a way to mitigate exploding gradients
#[derive(Config, Debug)]
pub enum GradientClippingConfig {
    /// Clip the gradient by value.
    Value(f32),

    /// Clip the gradient by norm.
    Norm(f32),
}

impl GradientClippingConfig {
    /// Initialize the gradient clipping.
    ///
    /// # Returns
    ///
    /// The gradient clipping.
    pub fn init(&self) -> GradientClipping {
        match self {
            GradientClippingConfig::Value(val) => GradientClipping::Value(*val),
            GradientClippingConfig::Norm(val) => GradientClipping::Norm(*val),
        }
    }
}

/// Gradient Clipping provides a way to mitigate exploding gradients
/// by clipping every component of the gradient by value or by norm during
/// backpropagation.
#[derive(Clone)]
pub enum GradientClipping {
    /// Clip the gradient by value.
    Value(f32),

    /// Clip the gradient by norm.
    Norm(f32),
}

impl GradientClipping {
    /// Clip the gradient.
    ///
    /// # Arguments
    ///
    /// * `grad` - The gradient to clip.
    ///
    /// # Returns
    ///
    /// The clipped gradient.
    pub fn clip_gradient<const D: usize>(&self, grad: Tensor<D>) -> Tensor<D> {
        match self {
            GradientClipping::Value(threshold) => self.clip_by_value(grad, *threshold),
            GradientClipping::Norm(max_norm) => self.clip_by_norm(grad, *max_norm),
        }
    }

    fn clip_by_value<const D: usize>(&self, grad: Tensor<D>, threshold: f32) -> Tensor<D> {
        let greater_mask = grad.clone().greater_scalar(threshold);
        let lower_mask = grad.clone().lower_scalar(-threshold);

        let clipped_grad = grad.mask_fill(greater_mask, threshold);

        clipped_grad.mask_fill(lower_mask, -threshold)
    }

    fn clip_by_norm<const D: usize>(&self, grad: Tensor<D>, threshold: f32) -> Tensor<D> {
        // Compute the norm in F32: in F16 the sum of squares overflows once the norm
        // exceeds ~256, and BF16 loses too much precision accumulating it.
        let dtype = grad.dtype();
        let grad = match dtype {
            DType::F16 | DType::BF16 => grad.cast(FloatDType::F32),
            _ => grad,
        };

        let norm = Self::l2_norm(grad.clone());
        let min_positive = grad
            .dtype()
            .finfo()
            .unwrap_or(FloatDType::F32.finfo())
            .min_positive;
        let clip_coef = threshold / norm.add_scalar(min_positive);
        let clip_coef_clamped = clip_coef.clamp_max(1.0);
        grad.mul(clip_coef_clamped.unsqueeze()).cast(dtype)
    }

    fn l2_norm<const D: usize>(tensor: Tensor<D>) -> Tensor<1> {
        let squared = tensor.square();
        let sum = squared.sum();
        sum.sqrt()
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use burn::tensor::Tensor;

    #[test]
    fn test_clip_by_value() {
        let gradient: Tensor<2> = Tensor::from_floats(
            [
                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
            ],
            &Default::default(),
        );

        let clipped_gradient = GradientClipping::Value(0.5).clip_gradient(gradient);
        let clipped_gradient_data = clipped_gradient.into_data();

        for value in clipped_gradient_data.iter::<f32>() {
            assert!(value <= 0.5);
        }
    }

    #[test]
    fn test_clip_by_norm() {
        let gradient: Tensor<2> = Tensor::from_floats(
            [
                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
            ],
            &Default::default(),
        );

        let clipped_gradient = GradientClipping::Norm(2.2).clip_gradient(gradient);
        let clipped_gradient_data = clipped_gradient.into_data();

        for value in clipped_gradient_data.iter::<f32>() {
            assert!(value <= 0.88);
        }
    }
    #[test]
    fn test_clip_by_norm_no_clipping() {
        let gradient: Tensor<2> = Tensor::from_floats(
            [[0.3, 0.4, 0.5, 0.2], [0.1, 0.6, 0.3, 0.4]],
            &Default::default(),
        );

        let clipped_gradient = GradientClipping::Norm(2.2).clip_gradient(gradient.clone());

        clipped_gradient
            .into_data()
            .assert_eq(&gradient.into_data(), true);
    }

    #[test]
    fn test_clip_by_norm_f16_does_not_overflow() {
        let gradient = Tensor::<1>::from_floats([300.0], &Default::default()).cast(FloatDType::F16);

        let clipped_gradient = GradientClipping::Norm(1.0).clip_gradient(gradient);

        assert_eq!(clipped_gradient.dtype(), DType::F16);
        let actual = clipped_gradient.into_scalar::<f32>();
        assert!((actual - 1.0).abs() < 1e-3, "actual={actual}");
    }

    #[test]
    fn test_clip_by_norm_f16_norm_above_f16_max() {
        let gradient =
            Tensor::<1>::from_floats([60000.0, 60000.0], &Default::default()).cast(FloatDType::F16);

        let clipped_gradient = GradientClipping::Norm(1.0).clip_gradient(gradient);

        let expected = core::f32::consts::FRAC_1_SQRT_2;
        for value in clipped_gradient.into_data().iter::<f32>() {
            assert!((value - expected).abs() < 1e-3, "value={value}");
        }
    }
}