Skip to main content

ruda_optim/grad_clipping/
base.rs

1
2use ruda_model::tensor::backend::Backend;
3use ruda_model::{config::Config, tensor::Tensor};
4
5/// Gradient Clipping provides a way to mitigate exploding gradients
6#[derive(Config, Debug)]
7pub enum GradientClippingConfig {
8    /// Clip the gradient by value.
9    Value(f32),
10
11    /// Clip the gradient by norm.
12    Norm(f32),
13}
14
15impl GradientClippingConfig {
16    /// Initialize the gradient clipping.
17    ///
18    /// # Returns
19    ///
20    /// The gradient clipping.
21    pub fn init(&self) -> GradientClipping {
22        match self {
23            GradientClippingConfig::Value(val) => GradientClipping::Value(*val),
24            GradientClippingConfig::Norm(val) => GradientClipping::Norm(*val),
25        }
26    }
27}
28
29/// Gradient Clipping provides a way to mitigate exploding gradients
30/// by clipping every component of the gradient by value or by norm during
31/// backpropagation.
32#[derive(Clone)]
33pub enum GradientClipping {
34    /// Clip the gradient by value.
35    Value(f32),
36
37    /// Clip the gradient by norm.
38    Norm(f32),
39}
40
41impl GradientClipping {
42    /// Clip the gradient.
43    ///
44    /// # Arguments
45    ///
46    /// * `grad` - The gradient to clip.
47    ///
48    /// # Returns
49    ///
50    /// The clipped gradient.
51    pub fn clip_gradient<B: Backend, const D: usize>(&self, grad: Tensor<B, D>) -> Tensor<B, D> {
52        match self {
53            GradientClipping::Value(threshold) => self.clip_by_value(grad, *threshold),
54            GradientClipping::Norm(max_norm) => self.clip_by_norm(grad, *max_norm),
55        }
56    }
57
58    fn clip_by_value<B: Backend, const D: usize>(
59        &self,
60        grad: Tensor<B, D>,
61        threshold: f32,
62    ) -> Tensor<B, D> {
63        let greater_mask = grad.clone().greater_elem(threshold);
64        let lower_mask = grad.clone().lower_elem(-threshold);
65
66        let clipped_grad = grad.mask_fill(greater_mask, threshold);
67
68        clipped_grad.mask_fill(lower_mask, -threshold)
69    }
70
71    fn clip_by_norm<B: Backend, const D: usize>(
72        &self,
73        grad: Tensor<B, D>,
74        threshold: f32,
75    ) -> Tensor<B, D> {
76        let norm = Self::l2_norm(grad.clone());
77        let min_positive = grad
78            .dtype()
79            .finfo()
80            .unwrap_or(ruda_model::tensor::FloatDType::F32.finfo())
81            .min_positive;
82        let clip_coef = threshold / norm.add_scalar(min_positive);
83        let clip_coef_clamped = clip_coef.clamp_max(1.0);
84        grad.mul(clip_coef_clamped.unsqueeze())
85    }
86
87    fn l2_norm<B: Backend, const D: usize>(tensor: Tensor<B, D>) -> Tensor<B, 1> {
88        let squared = tensor.square();
89        let sum = squared.sum();
90        sum.sqrt()
91    }
92}
93
94#[cfg(test)]
95mod tests {
96    use super::*;
97    use crate::TestBackend;
98    use ruda_model::tensor::Tensor;
99
100    #[test]
101    fn test_clip_by_value() {
102        let gradient: Tensor<TestBackend, 2> = Tensor::from_floats(
103            [
104                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
105                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
106            ],
107            &Default::default(),
108        );
109
110        let clipped_gradient = GradientClipping::Value(0.5).clip_gradient(gradient);
111        let clipped_gradient_data = clipped_gradient.into_data();
112
113        for value in clipped_gradient_data.iter::<f32>() {
114            assert!(value <= 0.5);
115        }
116    }
117
118    #[test]
119    fn test_clip_by_norm() {
120        let gradient: Tensor<TestBackend, 2> = Tensor::from_floats(
121            [
122                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
123                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
124            ],
125            &Default::default(),
126        );
127
128        let clipped_gradient = GradientClipping::Norm(2.2).clip_gradient(gradient);
129        let clipped_gradient_data = clipped_gradient.into_data();
130
131        for value in clipped_gradient_data.iter::<f32>() {
132            assert!(value <= 0.88);
133        }
134    }
135    #[test]
136    fn test_clip_by_norm_no_clipping() {
137        let gradient: Tensor<TestBackend, 2> = Tensor::from_floats(
138            [[0.3, 0.4, 0.5, 0.2], [0.1, 0.6, 0.3, 0.4]],
139            &Default::default(),
140        );
141
142        let clipped_gradient = GradientClipping::Norm(2.2).clip_gradient(gradient.clone());
143
144        clipped_gradient
145            .into_data()
146            .assert_eq(&gradient.into_data(), true);
147    }
148}