Skip to main content

ruda_optim/optim/
adam.rs

1
2use ruda_model::{module::AutodiffModule, record::Record};
3
4use ruda_model::config::Config;
5use ruda_model::tensor::{Tensor, backend::AutodiffBackend};
6use ruda_model::tensor::{backend::Backend, ops::Device};
7
8use super::{
9    SimpleOptimizer,
10    adaptor::OptimizerAdaptor,
11    decay::{WeightDecay, WeightDecayConfig},
12};
13use crate::{LearningRate, grad_clipping::GradientClippingConfig};
14
15#[cfg(not(feature = "std"))]
16#[allow(unused_imports)]
17use num_traits::Float as _;
18
19/// Adam configuration.
20#[derive(Config, Debug)]
21pub struct AdamConfig {
22    /// Parameter for Adam.
23    #[config(default = 0.9)]
24    beta_1: f32,
25    /// Parameter for Adam.
26    #[config(default = 0.999)]
27    beta_2: f32,
28    /// A value required for numerical stability.
29    #[config(default = 1e-5)]
30    epsilon: f32,
31    /// Whether to use AMSGrad algorithm
32    #[config(default = false)]
33    amsgrad: bool,
34    /// [Weight decay](WeightDecayConfig) config.
35    weight_decay: Option<WeightDecayConfig>,
36    /// [Gradient Clipping](GradientClippingConfig) config.
37    grad_clipping: Option<GradientClippingConfig>,
38}
39
40/// Adam optimizer.
41///
42/// See:
43/// - [Adam: A Method for Stochastic Optimization](https://arxiv.org/pdf/1412.6980.pdf).
44/// - [On the Convergence of Adam and Beyond](https://openreview.net/forum?id=ryQu7f-RZ)
45#[derive(Clone)]
46pub struct Adam {
47    momentum: AdaptiveMomentum,
48    weight_decay: Option<WeightDecay>,
49}
50
51/// Adam state.
52#[derive(Record, Clone, new)]
53pub struct AdamState<B: Backend, const D: usize> {
54    /// The current adaptive momentum.
55    pub momentum: AdaptiveMomentumState<B, D>,
56}
57
58impl<B: Backend> SimpleOptimizer<B> for Adam {
59    type State<const D: usize> = AdamState<B, D>;
60
61    fn step<const D: usize>(
62        &self,
63        lr: LearningRate,
64        tensor: Tensor<B, D>,
65        mut grad: Tensor<B, D>,
66        state: Option<Self::State<D>>,
67    ) -> (Tensor<B, D>, Option<Self::State<D>>) {
68        let mut state_momentum = None;
69
70        if let Some(state) = state {
71            state_momentum = Some(state.momentum);
72        }
73
74        if let Some(weight_decay) = &self.weight_decay {
75            grad = weight_decay.transform(grad, tensor.clone());
76        }
77
78        let (grad, state_momentum) = self.momentum.transform(grad, state_momentum);
79
80        let state = AdamState::new(state_momentum);
81        let delta = grad.mul_scalar(lr);
82
83        (tensor - delta, Some(state))
84    }
85
86    fn to_device<const D: usize>(mut state: Self::State<D>, device: &Device<B>) -> Self::State<D> {
87        state.momentum = state.momentum.to_device(device);
88        state
89    }
90}
91
92impl AdamConfig {
93    /// Build an [`Adam`] from the config.
94    pub fn build(&self) -> Adam {
95        Adam {
96            momentum: AdaptiveMomentum {
97                beta_1: self.beta_1,
98                beta_2: self.beta_2,
99                epsilon: self.epsilon,
100                amsgrad: self.amsgrad,
101            },
102            weight_decay: self.weight_decay.as_ref().map(WeightDecay::new),
103        }
104    }
105
106    /// Initialize Adam optimizer.
107    ///
108    /// # Returns
109    ///
110    /// Returns an optimizer that can be used to optimize a module.
111    pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(&self) -> OptimizerAdaptor<Adam, M, B> {
112        let mut optim = OptimizerAdaptor::from(self.build());
113        if let Some(config) = &self.grad_clipping {
114            optim = optim.with_grad_clipping(config.init());
115        }
116        optim
117    }
118}
119
120/// Adaptive momentum state.
121#[derive(Record, new, Clone)]
122pub struct AdaptiveMomentumState<B: Backend, const D: usize> {
123    /// The number of iterations aggregated.
124    pub time: usize,
125    /// The first order momentum.
126    pub moment_1: Tensor<B, D>,
127    /// The second order momentum.
128    pub moment_2: Tensor<B, D>,
129    /// Max of second  order momentum (for AMSGrad)
130    #[new(default)]
131    pub max_moment_2: Option<Tensor<B, D>>,
132}
133
134#[derive(Clone)]
135struct AdaptiveMomentum {
136    beta_1: f32,
137    beta_2: f32,
138    epsilon: f32,
139    amsgrad: bool,
140}
141
142impl AdaptiveMomentum {
143    pub fn transform<B: Backend, const D: usize>(
144        &self,
145        grad: Tensor<B, D>,
146        momentum_state: Option<AdaptiveMomentumState<B, D>>,
147    ) -> (Tensor<B, D>, AdaptiveMomentumState<B, D>) {
148        let state = if let Some(mut state) = momentum_state {
149            let factor = 1.0 - self.beta_1;
150            state.moment_1 = state
151                .moment_1
152                .mul_scalar(self.beta_1)
153                .add(grad.clone().mul_scalar(factor));
154
155            let factor = 1.0 - self.beta_2;
156            state.moment_2 = state
157                .moment_2
158                .mul_scalar(self.beta_2)
159                .add(grad.square().mul_scalar(factor));
160            if self.amsgrad {
161                let max_v = state
162                    .max_moment_2
163                    .take()
164                    .unwrap_or_else(|| state.moment_2.clone());
165
166                let new_max = max_v.max_pair(state.moment_2.clone());
167                state.max_moment_2 = Some(new_max);
168            }
169
170            state.time += 1;
171
172            state
173        } else {
174            let factor = 1.0 - self.beta_1;
175            let moment_1 = grad.clone().mul_scalar(factor);
176
177            let factor = 1.0 - self.beta_2;
178            let moment_2 = grad.square().mul_scalar(factor);
179            let max_moment_2 = self.amsgrad.then(|| moment_2.clone());
180            AdaptiveMomentumState {
181                time: 1,
182                moment_1,
183                moment_2,
184                max_moment_2,
185            }
186        };
187
188        let time = state.time as i32;
189        let bias_correction2_sqrt = (1.0 - self.beta_2.powi(time)).sqrt();
190        let combined_factor = bias_correction2_sqrt / (1.0 - self.beta_1.powi(time));
191
192        let v_to_use = if self.amsgrad {
193            state.max_moment_2.as_ref().unwrap_or(&state.moment_2)
194        } else {
195            &state.moment_2
196        };
197
198        let grad = state.moment_1.clone().mul_scalar(combined_factor).div(
199            v_to_use
200                .clone()
201                .sqrt()
202                .add_scalar(self.epsilon * bias_correction2_sqrt),
203        );
204        (grad, state)
205    }
206}
207
208impl<B: Backend, const D: usize> AdaptiveMomentumState<B, D> {
209    /// Move state to device.
210    ///
211    /// # Arguments
212    ///
213    /// * `device` - Device to move state to.
214    ///
215    /// # Returns
216    ///
217    /// Returns state moved to device.
218    pub fn to_device(mut self, device: &B::Device) -> Self {
219        self.moment_1 = self.moment_1.to_device(device);
220        self.moment_2 = self.moment_2.to_device(device);
221        self.max_moment_2 = self.max_moment_2.map(|tensor| tensor.to_device(device));
222        self
223    }
224}
225
226#[cfg(test)]
227mod tests {
228    use ruda_model::tensor::Tolerance;
229    use ruda_model::tensor::ops::FloatElem;
230
231    use super::*;
232    use crate::TestAutodiffBackend;
233    use crate::{GradientsParams, Optimizer};
234    use ruda_model::module::{Module, Param};
235    use ruda_model::tensor::{Distribution, Tensor, TensorData};
236    use ruda_nn::{Linear, LinearConfig, LinearRecord};
237
238    const LEARNING_RATE: LearningRate = 0.01;
239
240    #[test]
241    fn test_adam_optimizer_save_load_state() {
242        let device = Default::default();
243        let linear = LinearConfig::new(6, 6).init(&device);
244        let x = Tensor::<TestAutodiffBackend, 2>::random([2, 6], Distribution::Default, &device);
245        let mut optimizer = create_adam();
246        let grads = linear.forward(x).backward();
247        let grads = GradientsParams::from_grads(grads, &linear);
248        let _linear = optimizer.step(LEARNING_RATE, linear, grads);
249
250        #[cfg(feature = "std")]
251        {
252            use ruda_model::record::{BinFileRecorder, FullPrecisionSettings, Recorder};
253
254            BinFileRecorder::<FullPrecisionSettings>::default()
255                .record(
256                    optimizer.to_record(),
257                    std::env::temp_dir().as_path().join("test_optim_adam"),
258                )
259                .unwrap();
260        }
261        #[cfg(not(feature = "std"))]
262        {
263            use ruda_model::record::{BinBytesRecorder, FullPrecisionSettings, Recorder};
264
265            let result = BinBytesRecorder::<FullPrecisionSettings>::default()
266                .record(optimizer.to_record(), ())
267                .unwrap();
268            assert!(!result.is_empty());
269        }
270
271        let state_optim_before = optimizer.to_record();
272        let state_optim_before_copy = optimizer.to_record();
273        let optimizer = create_adam();
274        let optimizer = optimizer.load_record(state_optim_before_copy);
275        let state_optim_after = optimizer.to_record();
276
277        assert_eq!(state_optim_before.len(), state_optim_after.len());
278    }
279    #[test]
280    fn test_adam_optimizer_with_amsgrad_50_steps() {
281        let device = Default::default();
282        let mut linear = given_linear_layer(
283            TensorData::from([
284                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
285                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
286                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
287                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
288                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
289                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
290            ]),
291            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
292        );
293
294        let mut optimizer = AdamConfig::new()
295            .with_epsilon(1e-8)
296            .with_beta_1(0.9)
297            .with_beta_2(0.999)
298            .with_amsgrad(true)
299            .with_weight_decay(Some(WeightDecayConfig::new(0.5)))
300            .init();
301
302        for i in 1..=50 {
303            let x = Tensor::<TestAutodiffBackend, 2>::ones([2, 6], &device)
304                .mul_scalar(i as f32 * 0.1)
305                .require_grad();
306
307            let grads = linear.forward(x).backward();
308            let grads = GradientsParams::from_grads(grads, &linear);
309            linear = optimizer.step(LEARNING_RATE, linear, grads);
310        }
311
312        let state_updated = linear.into_record();
313        let weight_updated = state_updated.weight.to_data();
314        let bias_updated = state_updated.bias.unwrap().to_data();
315
316        let weights_expected = TensorData::from([
317            [
318                -0.9125810265541077,
319                -0.45855265855789185,
320                -0.1915993094444275,
321                -0.2759990692138672,
322                -0.5099529027938843,
323                -0.5287043452262878,
324            ],
325            [
326                -0.5181325674057007,
327                -0.6139854788780212,
328                -0.9574727416038513,
329                -0.34102925658226013,
330                -0.400514155626297,
331                -0.8847861886024475,
332            ],
333            [
334                -0.614483118057251,
335                -0.5611032247543335,
336                -0.8887064456939697,
337                -0.34762972593307495,
338                -0.8708556890487671,
339                -0.2830044627189636,
340            ],
341            [
342                -0.8904699683189392,
343                -0.8151527643203735,
344                -0.9621278643608093,
345                -0.8905676603317261,
346                -0.671261191368103,
347                -0.4333854615688324,
348            ],
349            [
350                -0.26599061489105225,
351                -0.8119961023330688,
352                -0.22424538433551788,
353                -0.7672406435012817,
354                -0.2163349837064743,
355                -0.6258266568183899,
356            ],
357            [
358                -0.611397922039032,
359                -0.6075160503387451,
360                -0.4701341986656189,
361                -0.4039117991924286,
362                -0.5663845539093018,
363                -0.21262989938259125,
364            ],
365        ]);
366        let bias_expected = TensorData::from([
367            -0.8817203044891357,
368            -0.4038999378681183,
369            -0.5889149308204651,
370            -0.37475723028182983,
371            -0.3557940721511841,
372            -0.47914788126945496,
373        ]);
374
375        type FT = FloatElem<TestAutodiffBackend>;
376        let tolerance = Tolerance::absolute(1e-5);
377        weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
378        bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
379    }
380    #[test]
381    fn test_adam_optimizer_with_numbers() {
382        let device = Default::default();
383        let linear = given_linear_layer(
384            TensorData::from([
385                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
386                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
387                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
388                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
389                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
390                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
391            ]),
392            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
393        );
394        let x_1 = Tensor::<TestAutodiffBackend, 2>::from_floats(
395            [
396                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
397                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
398            ],
399            &device,
400        )
401        .require_grad();
402        let x_2 = Tensor::<TestAutodiffBackend, 2>::from_floats(
403            [
404                [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
405                [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, 0.9085],
406            ],
407            &device,
408        )
409        .require_grad();
410
411        let mut optimizer = AdamConfig::new()
412            .with_epsilon(1e-8)
413            .with_beta_1(0.9)
414            .with_beta_2(0.999)
415            .with_weight_decay(Some(WeightDecayConfig::new(0.5)))
416            .init();
417
418        let grads = linear.forward(x_1).backward();
419        let grads = GradientsParams::from_grads(grads, &linear);
420        let linear = optimizer.step(LEARNING_RATE, linear, grads);
421
422        let grads = linear.forward(x_2).backward();
423        let grads = GradientsParams::from_grads(grads, &linear);
424        let linear = optimizer.step(LEARNING_RATE, linear, grads);
425
426        let state_updated = linear.into_record();
427        let weights_expected = TensorData::from([
428            [-0.340528, 0.118929, 0.384336, 0.300010, 0.066034, 0.047154],
429            [
430                0.057757, -0.036690, -0.386649, 0.235010, 0.175624, -0.312133,
431            ],
432            [
433                -0.038940, 0.016306, -0.316151, 0.228410, -0.297819, 0.293047,
434            ],
435            [
436                -0.317929, -0.239100, -0.391449, -0.318087, -0.095948, 0.142651,
437            ],
438            [
439                0.310050, -0.235909, 0.351736, -0.192888, 0.359710, -0.050343,
440            ],
441            [-0.035840, -0.030203, 0.105840, 0.172110, 0.009440, 0.363346],
442        ]);
443        let bias_expected = TensorData::from([
444            -0.410499, 0.068401, -0.116999, 0.097601, 0.116601, -0.006999,
445        ]);
446
447        let (weight_updated, bias_updated) = (
448            state_updated.weight.to_data(),
449            state_updated.bias.unwrap().to_data(),
450        );
451
452        type FT = FloatElem<TestAutodiffBackend>;
453        let tolerance = Tolerance::absolute(1e-2);
454        bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
455        weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
456    }
457
458    #[test]
459    fn test_adam_optimizer_no_nan() {
460        let linear = given_linear_layer(
461            TensorData::from([
462                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
463                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
464                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
465                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
466                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
467                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
468            ]),
469            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
470        );
471
472        let x = Tensor::<TestAutodiffBackend, 2>::from_floats(
473            [
474                [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
475                [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, 0.9085],
476            ],
477            &Default::default(),
478        )
479        .require_grad();
480
481        let mut optimizer = AdamConfig::new()
482            .with_epsilon(1e-8)
483            .with_beta_1(0.9)
484            .with_beta_2(0.999)
485            .with_weight_decay(Some(WeightDecayConfig::new(0.5)))
486            .init();
487
488        let grads = linear.forward(x.clone()).backward();
489        let grads = GradientsParams::from_grads(grads, &linear);
490        let linear = optimizer.step(LEARNING_RATE, linear, grads);
491
492        let grads = linear.forward(x).backward();
493        let grads = GradientsParams::from_grads(grads, &linear);
494        let linear = optimizer.step(LEARNING_RATE, linear, grads);
495
496        let state_updated = linear.into_record();
497        assert!(!state_updated.weight.to_data().as_slice::<f32>().unwrap()[0].is_nan());
498    }
499
500    fn given_linear_layer(weight: TensorData, bias: TensorData) -> Linear<TestAutodiffBackend> {
501        let device = Default::default();
502        let record = LinearRecord {
503            weight: Param::from_data(weight, &device),
504            bias: Some(Param::from_data(bias, &device)),
505        };
506
507        LinearConfig::new(6, 6).init(&device).load_record(record)
508    }
509
510    fn create_adam() -> OptimizerAdaptor<Adam, Linear<TestAutodiffBackend>, TestAutodiffBackend> {
511        let config = AdamConfig::new();
512        Adam {
513            momentum: AdaptiveMomentum {
514                beta_1: config.beta_1,
515                beta_2: config.beta_2,
516                epsilon: config.epsilon,
517                amsgrad: config.amsgrad,
518            },
519            weight_decay: config.weight_decay.as_ref().map(WeightDecay::new),
520        }
521        .into()
522    }
523}