ruda_optim/grad_clipping/
base.rs1
2use ruda_model::tensor::backend::Backend;
3use ruda_model::{config::Config, tensor::Tensor};
4
5#[derive(Config, Debug)]
7pub enum GradientClippingConfig {
8 Value(f32),
10
11 Norm(f32),
13}
14
15impl GradientClippingConfig {
16 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#[derive(Clone)]
33pub enum GradientClipping {
34 Value(f32),
36
37 Norm(f32),
39}
40
41impl GradientClipping {
42 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}