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#[derive(Config, Debug)]
16pub struct AdamWConfig {
17 #[config(default = 0.9)]
19 beta_1: f32,
20 #[config(default = 0.999)]
22 beta_2: f32,
23 #[config(default = 1e-5)]
25 epsilon: f32,
26 #[config(default = 1e-4)]
28 weight_decay: f32,
29
30 #[config(default = false)]
34 cautious_weight_decay: bool,
35
36 #[config(default = false)]
38 amsgrad: bool,
39 grad_clipping: Option<GradientClippingConfig>,
41}
42
43#[derive(Clone)]
52pub struct AdamW {
53 momentum: AdaptiveMomentumW,
54 weight_decay: f32,
55 cautious_weight_decay: bool,
56}
57
58#[derive(Record, Clone, new)]
60pub struct AdamWState<B: Backend, const D: usize> {
61 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 fn step<const D: usize>(
70 &self,
71 lr: LearningRate,
73 tensor: Tensor<B, D>,
75 grad: Tensor<B, D>,
77 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 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 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 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 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 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 state.moment_1 = state
177 .moment_1
178 .mul_scalar(self.beta_1)
179 .add(grad.clone().mul_scalar(factor_1));
180
181 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 state.time += 1;
197
198 state
199 } else {
200 let moment_1 = grad.clone().mul_scalar(factor_1);
202
203 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 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}