Skip to main content

ruda_optim/optim/
adamw.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::{AdaptiveMomentumState, 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/// [`AdamW`] Configuration.
15#[derive(Config, Debug)]
16pub struct AdamWConfig {
17    /// Parameter for AdamW.
18    #[config(default = 0.9)]
19    beta_1: f32,
20    /// Parameter for AdamW.
21    #[config(default = 0.999)]
22    beta_2: f32,
23    /// A value required for numerical stability.
24    #[config(default = 1e-5)]
25    epsilon: f32,
26    /// Weight decay config.
27    #[config(default = 1e-4)]
28    weight_decay: f32,
29
30    /// Cautious weight decay config.
31    ///
32    /// See: <https://arxiv.org/abs/2510.12402>
33    #[config(default = false)]
34    cautious_weight_decay: bool,
35
36    /// Whether to use AMSGrad algorithm
37    #[config(default = false)]
38    amsgrad: bool,
39    /// [Gradient Clipping](GradientClippingConfig) config.
40    grad_clipping: Option<GradientClippingConfig>,
41}
42
43/// AdamW optimizer.
44///
45/// See:
46/// - [Decoupled Weight Decay Regularization, Loshchilov and Hutter, 2019](https://arxiv.org/abs/1711.05101).
47/// - [Cautious Weight Decay, 2025](https://arxiv.org/abs/2510.12402)
48/// - [On the Convergence of Adam and Beyond](https://openreview.net/forum?id=ryQu7f-RZ)
49///
50/// Configured by [`AdamWConfig`].
51#[derive(Clone)]
52pub struct AdamW {
53    momentum: AdaptiveMomentumW,
54    weight_decay: f32,
55    cautious_weight_decay: bool,
56}
57
58/// AdamW state.
59#[derive(Record, Clone, new)]
60pub struct AdamWState<B: Backend, const D: usize> {
61    /// Th current adaptive momentum state.
62    pub momentum: AdaptiveMomentumState<B, D>,
63}
64
65impl<B: Backend> SimpleOptimizer<B> for AdamW {
66    type State<const D: usize> = AdamWState<B, D>;
67
68    /// A single optimization step for any tensor that represents the parameters of a model.
69    fn step<const D: usize>(
70        &self,
71        // Learning rate.
72        lr: LearningRate,
73        // Any tensor that represents the parameters of a model.
74        tensor: Tensor<B, D>,
75        // Gradient of the loss w.r.t. the parameters.
76        grad: Tensor<B, D>,
77        // State of the optimizer.
78        state: Option<Self::State<D>>,
79    ) -> (Tensor<B, D>, Option<Self::State<D>>) {
80        let (raw_delta, momentum_state) = self.momentum.transform(grad, state.map(|s| s.momentum));
81
82        let decay_rate = lr * (self.weight_decay as f64);
83
84        let decayed_tensor = if decay_rate == 0.0 {
85            tensor.clone()
86        } else if self.cautious_weight_decay {
87            // Cautious weight decay.
88            // See: https://arxiv.org/abs/2510.12402
89            let tensor_pos = tensor.clone().greater_equal_elem(0.0);
90            let grad_pos = momentum_state.moment_1.clone().greater_equal_elem(0.0);
91            let differ = tensor_pos.not_equal(grad_pos);
92
93            // Zero out the decay where the decay is counter to the update direction.
94            tensor.clone() - tensor.mul_scalar(decay_rate).mask_fill(differ, 0.0)
95        } else {
96            tensor.clone().mul_scalar(1.0 - decay_rate)
97        };
98
99        let tensor_updated = decayed_tensor - raw_delta.mul_scalar(lr);
100
101        let state = AdamWState {
102            momentum: momentum_state,
103        };
104
105        (tensor_updated, Some(state))
106    }
107
108    fn to_device<const D: usize>(mut state: Self::State<D>, device: &Device<B>) -> Self::State<D> {
109        state.momentum = state.momentum.to_device(device);
110        state
111    }
112}
113
114impl AdamWConfig {
115    /// Validation shared by explicit mixed-optimizer configuration.
116    pub(crate) fn validate_hyperparameters(&self) -> Result<(), &'static str> {
117        if !self.beta_1.is_finite() || !self.beta_2.is_finite()
118            || !(0.0..1.0).contains(&self.beta_1) || !(0.0..1.0).contains(&self.beta_2) {
119            return Err("AdamW betas must be finite in [0, 1)");
120        }
121        if !self.epsilon.is_finite() || self.epsilon <= 0.0 {
122            return Err("AdamW epsilon must be finite and positive");
123        }
124        if !self.weight_decay.is_finite() || self.weight_decay < 0.0 {
125            return Err("AdamW weight decay must be finite and nonnegative");
126        }
127        Ok(())
128    }
129    /// Build an [`AdamW`] from the config.
130    pub fn build(&self) -> AdamW {
131        AdamW {
132            momentum: AdaptiveMomentumW {
133                beta_1: self.beta_1,
134                beta_2: self.beta_2,
135                epsilon: self.epsilon,
136                amsgrad: self.amsgrad,
137            },
138            weight_decay: self.weight_decay,
139            cautious_weight_decay: self.cautious_weight_decay,
140        }
141    }
142
143    /// Initialize AdamW optimizer.
144    ///
145    /// # Returns
146    ///
147    /// Returns an optimizer that can be used to optimize a module.
148    pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(&self) -> OptimizerAdaptor<AdamW, M, B> {
149        let mut optim = OptimizerAdaptor::from(self.build());
150        if let Some(config) = &self.grad_clipping {
151            optim = optim.with_grad_clipping(config.init());
152        }
153        optim
154    }
155}
156
157#[derive(Clone)]
158struct AdaptiveMomentumW {
159    beta_1: f32,
160    beta_2: f32,
161    epsilon: f32,
162    amsgrad: bool,
163}
164
165impl AdaptiveMomentumW {
166    pub fn transform<B: Backend, const D: usize>(
167        &self,
168        grad: Tensor<B, D>,
169        state: Option<AdaptiveMomentumState<B, D>>,
170    ) -> (Tensor<B, D>, AdaptiveMomentumState<B, D>) {
171        let factor_1 = 1.0 - self.beta_1;
172        let factor_2 = 1.0 - self.beta_2;
173
174        let state = if let Some(mut state) = state {
175            // Update first moment estimate.
176            state.moment_1 = state
177                .moment_1
178                .mul_scalar(self.beta_1)
179                .add(grad.clone().mul_scalar(factor_1));
180
181            // Update second moment estimate.
182            state.moment_2 = state
183                .moment_2
184                .mul_scalar(self.beta_2)
185                .add(grad.square().mul_scalar(factor_2));
186
187            if self.amsgrad {
188                let max_v = state
189                    .max_moment_2
190                    .take()
191                    .unwrap_or_else(|| state.moment_2.clone());
192                state.max_moment_2 = Some(max_v.max_pair(state.moment_2.clone()));
193            }
194
195            // Update time.
196            state.time += 1;
197
198            state
199        } else {
200            // Initialize first moment estimate.
201            let moment_1 = grad.clone().mul_scalar(factor_1);
202
203            // Initialize second moment estimate.
204            let moment_2 = grad.square().mul_scalar(factor_2);
205            let max_moment_2 = self.amsgrad.then(|| moment_2.clone());
206            AdaptiveMomentumState {
207                time: 1,
208                moment_1,
209                moment_2,
210                max_moment_2,
211            }
212        };
213
214        let time: i32 = state.time as i32;
215
216        // Compute bias-corrected first and second moment estimates.
217        let moment_1_corrected = state
218            .moment_1
219            .clone()
220            .div_scalar(1f32 - self.beta_1.powi(time));
221
222        let v_to_use = if self.amsgrad {
223            state.max_moment_2.as_ref().unwrap_or(&state.moment_2)
224        } else {
225            &state.moment_2
226        };
227
228        let moment_2_corrected = v_to_use.clone().div_scalar(1f32 - self.beta_2.powi(time));
229
230        let update_delta =
231            moment_1_corrected.div(moment_2_corrected.sqrt().add_scalar(self.epsilon));
232
233        (update_delta, state)
234    }
235}
236
237#[cfg(test)]
238mod tests {
239    use super::*;
240    use crate::TestAutodiffBackend;
241    use crate::{GradientsParams, Optimizer};
242    use ruda_model::module::{Module, Param};
243    use ruda_model::tensor::{Distribution, Tensor, TensorData};
244    use ruda_model::tensor::{Tolerance, ops::FloatElem};
245    use ruda_nn::{Linear, LinearConfig, LinearRecord};
246
247    type FT = FloatElem<TestAutodiffBackend>;
248
249    const LEARNING_RATE: LearningRate = 0.01;
250
251    #[test]
252    fn test_adamw_optimizer_save_load_state() {
253        let device = Default::default();
254        let linear = LinearConfig::new(6, 6).init(&device);
255        let x = Tensor::<TestAutodiffBackend, 2>::random([2, 6], Distribution::Default, &device);
256        let mut optimizer = create_adamw();
257        let grads = linear.forward(x).backward();
258        let grads = GradientsParams::from_grads(grads, &linear);
259        let _linear = optimizer.step(LEARNING_RATE, linear, grads);
260
261        #[cfg(feature = "std")]
262        {
263            use ruda_model::record::{BinFileRecorder, FullPrecisionSettings, Recorder};
264
265            BinFileRecorder::<FullPrecisionSettings>::default()
266                .record(
267                    optimizer.to_record(),
268                    std::env::temp_dir().as_path().join("test_optim_adamw"),
269                )
270                .unwrap();
271        }
272        #[cfg(not(feature = "std"))]
273        {
274            use ruda_model::record::{BinBytesRecorder, FullPrecisionSettings, Recorder};
275
276            let result = BinBytesRecorder::<FullPrecisionSettings>::default()
277                .record(optimizer.to_record(), ())
278                .unwrap();
279            assert!(!result.is_empty());
280        }
281
282        let state_optim_before = optimizer.to_record();
283        let state_optim_before_copy = optimizer.to_record();
284        let optimizer = create_adamw();
285        let optimizer = optimizer.load_record(state_optim_before_copy);
286        let state_optim_after = optimizer.to_record();
287
288        assert_eq!(state_optim_before.len(), state_optim_after.len());
289    }
290    #[test]
291    fn test_adamw_optimizer_with_amsgrad_50_steps() {
292        let device = Default::default();
293        let mut linear = given_linear_layer(
294            TensorData::from([
295                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
296                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
297                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
298                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
299                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
300                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
301            ]),
302            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
303        );
304
305        let mut optimizer = AdamWConfig::new()
306            .with_epsilon(1e-8)
307            .with_beta_1(0.9)
308            .with_beta_2(0.999)
309            .with_amsgrad(true)
310            .with_weight_decay(0.5)
311            .init();
312
313        for i in 1..=50 {
314            let x = Tensor::<TestAutodiffBackend, 2>::ones([2, 6], &device)
315                .mul_scalar(i as f32 * 0.1)
316                .require_grad();
317
318            let grads = linear.forward(x).backward();
319            let grads = GradientsParams::from_grads(grads, &linear);
320            linear = optimizer.step(LEARNING_RATE, linear, grads);
321        }
322
323        let state_updated = linear.into_record();
324        let weight_updated = state_updated.weight.to_data();
325        let bias_updated = state_updated.bias.unwrap().to_data();
326
327        let weights_expected = TensorData::from([
328            [
329                -0.7822558283805847,
330                -0.42578864097595215,
331                -0.21805696189403534,
332                -0.28366872668266296,
333                -0.46587175130844116,
334                -0.4805040955543518,
335            ],
336            [
337                -0.4722539782524109,
338                -0.5471276640892029,
339                -0.8181359767913818,
340                -0.33425918221473694,
341                -0.3805687427520752,
342                -0.7601516842842102,
343            ],
344            [
345                -0.5475167632102966,
346                -0.5057991743087769,
347                -0.763265073299408,
348                -0.3393959403038025,
349                -0.7490996718406677,
350                -0.28911691904067993,
351            ],
352            [
353                -0.7646660208702087,
354                -0.7050473093986511,
355                -0.8218720555305481,
356                -0.7647438049316406,
357                -0.5919585227966309,
358                -0.40617525577545166,
359            ],
360            [
361                -0.27588561177253723,
362                -0.7025567889213562,
363                -0.24343004822731018,
364                -0.6672990918159485,
365                -0.23728127777576447,
366                -0.556389570236206,
367            ],
368            [
369                -0.5451040267944336,
370                -0.5420684814453125,
371                -0.4348171353340149,
372                -0.3832150399684906,
373                -0.5099242925643921,
374                -0.23440153896808624,
375            ],
376        ]);
377        let bias_expected = TensorData::from([
378            -0.7473056316375732,
379            -0.3745720386505127,
380            -0.5188710689544678,
381            -0.35184532403945923,
382            -0.33705732226371765,
383            -0.4332566559314728,
384        ]);
385
386        type FT = FloatElem<TestAutodiffBackend>;
387        let tolerance = Tolerance::absolute(1e-5);
388        weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
389        bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
390    }
391    #[test]
392    fn test_adamw_optimizer_with_numbers() {
393        let linear = given_linear_layer(
394            TensorData::from([
395                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
396                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
397                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
398                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
399                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
400                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
401            ]),
402            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
403        );
404        let device = Default::default();
405        let x_1 = Tensor::<TestAutodiffBackend, 2>::from_floats(
406            [
407                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
408                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
409            ],
410            &device,
411        )
412        .require_grad();
413        let x_2 = 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            &device,
419        )
420        .require_grad();
421
422        let mut optimizer = AdamWConfig::new()
423            .with_epsilon(1e-8)
424            .with_beta_1(0.9)
425            .with_beta_2(0.999)
426            .with_weight_decay(0.5)
427            .init();
428
429        let grads = linear.forward(x_1).backward();
430        let grads = GradientsParams::from_grads(grads, &linear);
431        let linear = optimizer.step(LEARNING_RATE, linear, grads);
432
433        let grads = linear.forward(x_2).backward();
434        let grads = GradientsParams::from_grads(grads, &linear);
435        let linear = optimizer.step(LEARNING_RATE, linear, grads);
436
437        let state_updated = linear.into_record();
438        let weights_expected = TensorData::from([
439            [-0.337295, 0.117827, 0.380358, 0.296868, 0.065232, 0.046534],
440            [
441                0.057032, -0.036518, -0.382951, 0.232516, 0.173738, -0.309182,
442            ],
443            [
444                -0.038703, 0.016052, -0.313155, 0.225982, -0.295039, 0.289981,
445            ],
446            [
447                -0.314920, -0.237394, -0.387704, -0.315067, -0.095153, 0.141081,
448            ],
449            [
450                0.306815, -0.234226, 0.348083, -0.191115, 0.356002, -0.049993,
451            ],
452            [-0.035634, -0.030083, 0.104636, 0.170244, 0.009196, 0.359580],
453        ]);
454        let bias_expected = TensorData::from([
455            -0.406555, 0.067568, -0.115982, 0.096477, 0.115287, -0.007080,
456        ]);
457
458        let (weight_updated, bias_updated) = (
459            state_updated.weight.to_data(),
460            state_updated.bias.unwrap().to_data(),
461        );
462
463        let tolerance = Tolerance::absolute(1e-2);
464        bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
465        weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
466    }
467
468    #[test]
469    fn test_adamw_optimizer_with_numbers_cautious() {
470        let linear = given_linear_layer(
471            TensorData::from([
472                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
473                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
474                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
475                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
476                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
477                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
478            ]),
479            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
480        );
481        let device = Default::default();
482        let x_1 = Tensor::<TestAutodiffBackend, 2>::from_floats(
483            [
484                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
485                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
486            ],
487            &device,
488        )
489        .require_grad();
490        let x_2 = Tensor::<TestAutodiffBackend, 2>::from_floats(
491            [
492                [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
493                [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, -0.9085],
494            ],
495            &device,
496        )
497        .require_grad();
498
499        let mut optimizer = AdamWConfig::new()
500            .with_cautious_weight_decay(true)
501            .with_epsilon(1e-8)
502            .with_beta_1(0.9)
503            .with_beta_2(0.999)
504            .with_weight_decay(0.5)
505            .init();
506
507        let grads = linear.forward(x_1).backward();
508        let grads = GradientsParams::from_grads(grads, &linear);
509        let linear = optimizer.step(LEARNING_RATE, linear, grads);
510
511        let grads = linear.forward(x_2).backward();
512        let grads = GradientsParams::from_grads(grads, &linear);
513        let linear = optimizer.step(LEARNING_RATE, linear, grads);
514
515        let state_updated = linear.into_record();
516        let weights_expected = TensorData::from([
517            [-0.337295, 0.117827, 0.380358, 0.296868, 0.065232, 0.046534],
518            [
519                0.057032, -0.036518, -0.382951, 0.232516, 0.173738, -0.309182,
520            ],
521            [
522                -0.038703, 0.016052, -0.313155, 0.225982, -0.295039, 0.289981,
523            ],
524            [
525                -0.314920, -0.237394, -0.387704, -0.315067, -0.095153, 0.141081,
526            ],
527            [
528                0.306815, -0.234226, 0.348083, -0.191115, 0.356002, -0.049993,
529            ],
530            [
531                -0.035634, -0.030083, 0.104636, 0.170244, 0.009196, 0.37061332,
532            ],
533        ]);
534        let bias_expected = TensorData::from([
535            -0.406555, 0.067568, -0.115982, 0.096477, 0.115287, -0.007080,
536        ]);
537
538        let (weight_updated, bias_updated) = (
539            state_updated.weight.to_data(),
540            state_updated.bias.unwrap().to_data(),
541        );
542
543        let tolerance = Tolerance::absolute(1e-2);
544        bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
545        weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
546    }
547
548    #[test]
549    fn test_adam_optimizer_no_nan() {
550        let linear = given_linear_layer(
551            TensorData::from([
552                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
553                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
554                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
555                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
556                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
557                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
558            ]),
559            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
560        );
561
562        let x = Tensor::<TestAutodiffBackend, 2>::from_floats(
563            [
564                [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
565                [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, 0.9085],
566            ],
567            &Default::default(),
568        )
569        .require_grad();
570
571        let mut optimizer = AdamWConfig::new()
572            .with_epsilon(1e-8)
573            .with_beta_1(0.9)
574            .with_beta_2(0.999)
575            .with_weight_decay(0.5)
576            .init();
577
578        let grads = linear.forward(x.clone()).backward();
579        let grads = GradientsParams::from_grads(grads, &linear);
580        let linear = optimizer.step(LEARNING_RATE, linear, grads);
581
582        let grads = linear.forward(x).backward();
583        let grads = GradientsParams::from_grads(grads, &linear);
584        let linear = optimizer.step(LEARNING_RATE, linear, grads);
585
586        let state_updated = linear.into_record();
587        assert!(!state_updated.weight.to_data().as_slice::<f32>().unwrap()[0].is_nan());
588    }
589
590    fn given_linear_layer(weight: TensorData, bias: TensorData) -> Linear<TestAutodiffBackend> {
591        let device = Default::default();
592        let record = LinearRecord {
593            weight: Param::from_data(weight, &device),
594            bias: Some(Param::from_data(bias, &device)),
595        };
596
597        LinearConfig::new(6, 6).init(&device).load_record(record)
598    }
599
600    fn create_adamw() -> OptimizerAdaptor<AdamW, Linear<TestAutodiffBackend>, TestAutodiffBackend> {
601        let config = AdamWConfig::new();
602        AdamW {
603            momentum: AdaptiveMomentumW {
604                beta_1: config.beta_1,
605                beta_2: config.beta_2,
606                epsilon: config.epsilon,
607                amsgrad: config.amsgrad,
608            },
609            weight_decay: config.weight_decay,
610            cautious_weight_decay: false,
611        }
612        .into()
613    }
614}