Skip to main content

ruda_optim/optim/
adan.rs

1
2use ruda_model::config::Config;
3use ruda_model::tensor::{Tensor, backend::AutodiffBackend};
4use ruda_model::tensor::{backend::Backend, ops::Device};
5use ruda_model::{module::AutodiffModule, record::Record};
6
7use super::{SimpleOptimizer, adaptor::OptimizerAdaptor};
8use crate::{LearningRate, grad_clipping::GradientClippingConfig};
9
10#[cfg(not(feature = "std"))]
11#[allow(unused_imports)]
12use num_traits::Float as _;
13
14/// [`Adan`] Configuration.
15///
16/// See:
17/// - [Adan: Adaptive Nesterov Momentum Algorithm for Faster Optimizing Deep Models](https://arxiv.org/abs/2208.06677).
18#[derive(Config, Debug)]
19pub struct AdanConfig {
20    /// Parameter for the first moment.
21    #[config(default = 0.98)]
22    beta_1: f32,
23    /// Parameter for the gradient-difference momentum.
24    #[config(default = 0.92)]
25    beta_2: f32,
26    /// Parameter for the second moment.
27    #[config(default = 0.99)]
28    beta_3: f32,
29    /// A value required for numerical stability.
30    #[config(default = 1e-8)]
31    epsilon: f32,
32    /// Weight decay factor.
33    #[config(default = 0.0)]
34    weight_decay: f32,
35    /// Disable proximal weight decay and use the decoupled update instead.
36    #[config(default = false)]
37    no_prox: bool,
38    /// [Gradient Clipping](GradientClippingConfig) config.
39    grad_clipping: Option<GradientClippingConfig>,
40}
41
42/// Adan optimizer.
43///
44/// See:
45/// - [Adan: Adaptive Nesterov Momentum Algorithm for Faster Optimizing Deep Models](https://arxiv.org/abs/2208.06677).
46///
47/// Configured by [`AdanConfig`].
48#[derive(Clone)]
49pub struct Adan {
50    momentum: AdaptiveNesterovMomentum,
51    weight_decay: f32,
52    no_prox: bool,
53}
54
55/// Adan state.
56#[derive(Record, Clone, new)]
57pub struct AdanState<B: Backend, const D: usize> {
58    /// The current adaptive Nesterov momentum state.
59    pub momentum: AdaptiveNesterovMomentumState<B, D>,
60}
61
62impl<B: Backend> SimpleOptimizer<B> for Adan {
63    type State<const D: usize> = AdanState<B, D>;
64
65    fn step<const D: usize>(
66        &self,
67        lr: LearningRate,
68        tensor: Tensor<B, D>,
69        grad: Tensor<B, D>,
70        state: Option<Self::State<D>>,
71    ) -> (Tensor<B, D>, Option<Self::State<D>>) {
72        let (raw_delta, momentum_state) = self.momentum.transform(grad, state.map(|s| s.momentum));
73
74        let decay_rate = lr * (self.weight_decay as f64);
75        let delta = raw_delta.mul_scalar(lr);
76
77        let tensor_updated = if self.no_prox {
78            if decay_rate == 0.0 {
79                tensor - delta
80            } else {
81                tensor.mul_scalar(1.0 - decay_rate) - delta
82            }
83        } else {
84            let updated = tensor - delta;
85            if decay_rate == 0.0 {
86                updated
87            } else {
88                updated.div_scalar(1.0 + decay_rate)
89            }
90        };
91
92        (tensor_updated, Some(AdanState::new(momentum_state)))
93    }
94
95    fn to_device<const D: usize>(mut state: Self::State<D>, device: &Device<B>) -> Self::State<D> {
96        state.momentum = state.momentum.to_device(device);
97        state
98    }
99}
100
101impl AdanConfig {
102    /// Build an [`Adan`] from the config.
103    pub fn build(&self) -> Adan {
104        Adan {
105            momentum: AdaptiveNesterovMomentum {
106                beta_1: self.beta_1,
107                beta_2: self.beta_2,
108                beta_3: self.beta_3,
109                epsilon: self.epsilon,
110            },
111            weight_decay: self.weight_decay,
112            no_prox: self.no_prox,
113        }
114    }
115
116    /// Initialize Adan optimizer.
117    ///
118    /// # Returns
119    ///
120    /// Returns an optimizer that can be used to optimize a module.
121    pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(&self) -> OptimizerAdaptor<Adan, M, B> {
122        let mut optim = OptimizerAdaptor::from(self.build());
123        if let Some(config) = &self.grad_clipping {
124            optim = optim.with_grad_clipping(config.init());
125        }
126        optim
127    }
128}
129
130/// Adaptive Nesterov momentum state.
131#[derive(Record, Clone, new)]
132pub struct AdaptiveNesterovMomentumState<B: Backend, const D: usize> {
133    /// The number of iterations aggregated.
134    pub time: usize,
135    /// The first order momentum.
136    pub exp_avg: Tensor<B, D>,
137    /// The gradient-difference weighted second order momentum.
138    pub exp_avg_sq: Tensor<B, D>,
139    /// The gradient-difference momentum.
140    pub exp_avg_diff: Tensor<B, D>,
141    /// The negated previous gradient.
142    pub neg_pre_grad: Tensor<B, D>,
143}
144
145#[derive(Clone)]
146struct AdaptiveNesterovMomentum {
147    beta_1: f32,
148    beta_2: f32,
149    beta_3: f32,
150    epsilon: f32,
151}
152
153impl AdaptiveNesterovMomentum {
154    pub fn transform<B: Backend, const D: usize>(
155        &self,
156        grad: Tensor<B, D>,
157        state: Option<AdaptiveNesterovMomentumState<B, D>>,
158    ) -> (Tensor<B, D>, AdaptiveNesterovMomentumState<B, D>) {
159        let state = if let Some(mut state) = state {
160            let grad_diff = state.neg_pre_grad.clone().add(grad.clone());
161            let grad_diff_sq = grad_diff
162                .clone()
163                .mul_scalar(self.beta_2)
164                .add(grad.clone())
165                .square();
166
167            state.exp_avg = state
168                .exp_avg
169                .mul_scalar(self.beta_1)
170                .add(grad.clone().mul_scalar(1.0 - self.beta_1));
171            state.exp_avg_diff = state
172                .exp_avg_diff
173                .mul_scalar(self.beta_2)
174                .add(grad_diff.mul_scalar(1.0 - self.beta_2));
175            state.exp_avg_sq = state
176                .exp_avg_sq
177                .mul_scalar(self.beta_3)
178                .add(grad_diff_sq.mul_scalar(1.0 - self.beta_3));
179            state.neg_pre_grad = grad.mul_scalar(-1.0);
180            state.time += 1;
181            state
182        } else {
183            AdaptiveNesterovMomentumState::new(
184                1,
185                grad.clone().mul_scalar(1.0 - self.beta_1),
186                grad.clone().square().mul_scalar(1.0 - self.beta_3),
187                grad.zeros_like(),
188                grad.clone().mul_scalar(-1.0),
189            )
190        };
191
192        let time = state.time as i32;
193        let denom = state
194            .exp_avg_sq
195            .clone()
196            .sqrt()
197            .div_scalar((1.0 - self.beta_3.powi(time)).sqrt())
198            .add_scalar(self.epsilon);
199        let update = state
200            .exp_avg
201            .clone()
202            .div_scalar(1.0 - self.beta_1.powi(time))
203            .div(denom.clone())
204            .add(
205                state
206                    .exp_avg_diff
207                    .clone()
208                    .mul_scalar(self.beta_2)
209                    .div_scalar(1.0 - self.beta_2.powi(time))
210                    .div(denom),
211            );
212
213        (update, state)
214    }
215}
216
217impl<B: Backend, const D: usize> AdaptiveNesterovMomentumState<B, D> {
218    #[allow(clippy::wrong_self_convention)]
219    fn to_device(mut self, device: &B::Device) -> Self {
220        self.exp_avg = self.exp_avg.to_device(device);
221        self.exp_avg_sq = self.exp_avg_sq.to_device(device);
222        self.exp_avg_diff = self.exp_avg_diff.to_device(device);
223        self.neg_pre_grad = self.neg_pre_grad.to_device(device);
224        self
225    }
226}
227
228#[cfg(test)]
229mod tests {
230    use super::*;
231    use crate::TestAutodiffBackend;
232    use crate::{GradientsParams, Optimizer};
233    use ruda_model::module::{Module, Param};
234    use ruda_model::tensor::{Distribution, Tensor, TensorData};
235    use ruda_model::tensor::{Tolerance, ops::FloatElem};
236    use ruda_nn::{Linear, LinearConfig, LinearRecord};
237
238    type FT = FloatElem<TestAutodiffBackend>;
239
240    const LEARNING_RATE: LearningRate = 0.01;
241
242    #[test]
243    fn test_adan_optimizer_save_load_state() {
244        let device = Default::default();
245        let linear = LinearConfig::new(6, 6).init(&device);
246        let x = Tensor::<TestAutodiffBackend, 2>::random([2, 6], Distribution::Default, &device);
247        let mut optimizer = create_adan();
248        let grads = linear.forward(x).backward();
249        let grads = GradientsParams::from_grads(grads, &linear);
250        let _linear = optimizer.step(LEARNING_RATE, linear, grads);
251
252        #[cfg(feature = "std")]
253        {
254            use ruda_model::record::{BinFileRecorder, FullPrecisionSettings, Recorder};
255
256            BinFileRecorder::<FullPrecisionSettings>::default()
257                .record(
258                    optimizer.to_record(),
259                    std::env::temp_dir().as_path().join("test_optim_adan"),
260                )
261                .unwrap();
262        }
263        #[cfg(not(feature = "std"))]
264        {
265            use ruda_model::record::{BinBytesRecorder, FullPrecisionSettings, Recorder};
266
267            let result = BinBytesRecorder::<FullPrecisionSettings>::default()
268                .record(optimizer.to_record(), ())
269                .unwrap();
270            assert!(!result.is_empty());
271        }
272
273        let state_optim_before = optimizer.to_record();
274        let state_optim_before_copy = optimizer.to_record();
275        let optimizer = create_adan();
276        let optimizer = optimizer.load_record(state_optim_before_copy);
277        let state_optim_after = optimizer.to_record();
278
279        assert_eq!(state_optim_before.len(), state_optim_after.len());
280    }
281
282    #[test]
283    fn test_adan_optimizer_with_numbers() {
284        let linear = given_linear_layer(
285            TensorData::from([
286                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
287                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
288                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
289                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
290                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
291                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
292            ]),
293            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
294        );
295        let device = Default::default();
296        let x_1 = Tensor::<TestAutodiffBackend, 2>::from_floats(
297            [
298                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
299                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
300            ],
301            &device,
302        )
303        .require_grad();
304        let x_2 = Tensor::<TestAutodiffBackend, 2>::from_floats(
305            [
306                [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
307                [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, 0.9085],
308            ],
309            &device,
310        )
311        .require_grad();
312
313        let mut optimizer = AdanConfig::new()
314            .with_beta_1(0.98)
315            .with_beta_2(0.92)
316            .with_beta_3(0.99)
317            .with_epsilon(1e-8)
318            .with_weight_decay(0.02)
319            .init();
320
321        let grads = linear.forward(x_1).backward();
322        let grads = GradientsParams::from_grads(grads, &linear);
323        let linear = optimizer.step(LEARNING_RATE, linear, grads);
324
325        let grads = linear.forward(x_2).backward();
326        let grads = GradientsParams::from_grads(grads, &linear);
327        let linear = optimizer.step(LEARNING_RATE, linear, grads);
328
329        let state_updated = linear.into_record();
330        let weights_expected = TensorData::from([
331            [
332                -0.34034607,
333                0.11747075,
334                0.38426402,
335                0.29999772,
336                0.06599136,
337                0.04719888,
338            ],
339            [
340                0.0644293,
341                -0.031732224,
342                -0.37979296,
343                0.24165839,
344                0.18218218,
345                -0.30532277,
346            ],
347            [
348                -0.038910445,
349                0.01466812,
350                -0.31599957,
351                0.2283826,
352                -0.29780683,
353                0.2929568,
354            ],
355            [
356                -0.3178632,
357                -0.24129382,
358                -0.39133376,
359                -0.31796312,
360                -0.09605193,
361                0.14255258,
362            ],
363            [
364                0.31026322,
365                -0.23771758,
366                0.3519465,
367                -0.19243571,
368                0.35984334,
369                -0.049992695,
370            ],
371            [
372                -0.03577819,
373                -0.031879753,
374                0.10586514,
375                0.17213862,
376                0.009403733,
377                0.36326218,
378            ],
379        ]);
380        let bias_expected = TensorData::from([
381            -0.4103378,
382            0.06837065,
383            -0.116955206,
384            0.097558975,
385            0.11655137,
386            -0.006999196,
387        ]);
388
389        let (weight_updated, bias_updated) = (
390            state_updated.weight.to_data(),
391            state_updated.bias.unwrap().to_data(),
392        );
393
394        let tolerance = Tolerance::absolute(1e-5);
395        bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
396        weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
397    }
398
399    #[test]
400    fn test_adan_optimizer_no_nan() {
401        let linear = given_linear_layer(
402            TensorData::from([
403                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
404                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
405                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
406                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
407                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
408                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
409            ]),
410            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
411        );
412
413        let x = Tensor::<TestAutodiffBackend, 2>::from_floats(
414            [
415                [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
416                [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, 0.9085],
417            ],
418            &Default::default(),
419        )
420        .require_grad();
421
422        let mut optimizer = AdanConfig::new()
423            .with_epsilon(1e-8)
424            .with_weight_decay(0.02)
425            .init();
426
427        let grads = linear.forward(x.clone()).backward();
428        let grads = GradientsParams::from_grads(grads, &linear);
429        let linear = optimizer.step(LEARNING_RATE, linear, grads);
430
431        let grads = linear.forward(x).backward();
432        let grads = GradientsParams::from_grads(grads, &linear);
433        let linear = optimizer.step(LEARNING_RATE, linear, grads);
434
435        let state_updated = linear.into_record();
436        assert!(!state_updated.weight.to_data().as_slice::<f32>().unwrap()[0].is_nan());
437    }
438
439    fn given_linear_layer(weight: TensorData, bias: TensorData) -> Linear<TestAutodiffBackend> {
440        let device = Default::default();
441        let record = LinearRecord {
442            weight: Param::from_data(weight, &device),
443            bias: Some(Param::from_data(bias, &device)),
444        };
445
446        LinearConfig::new(6, 6).init(&device).load_record(record)
447    }
448
449    fn create_adan() -> OptimizerAdaptor<Adan, Linear<TestAutodiffBackend>, TestAutodiffBackend> {
450        let config = AdanConfig::new();
451        Adan {
452            momentum: AdaptiveNesterovMomentum {
453                beta_1: config.beta_1,
454                beta_2: config.beta_2,
455                beta_3: config.beta_3,
456                epsilon: config.epsilon,
457            },
458            weight_decay: config.weight_decay,
459            no_prox: config.no_prox,
460        }
461        .into()
462    }
463}