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#[derive(Config, Debug)]
19pub struct AdanConfig {
20 #[config(default = 0.98)]
22 beta_1: f32,
23 #[config(default = 0.92)]
25 beta_2: f32,
26 #[config(default = 0.99)]
28 beta_3: f32,
29 #[config(default = 1e-8)]
31 epsilon: f32,
32 #[config(default = 0.0)]
34 weight_decay: f32,
35 #[config(default = false)]
37 no_prox: bool,
38 grad_clipping: Option<GradientClippingConfig>,
40}
41
42#[derive(Clone)]
49pub struct Adan {
50 momentum: AdaptiveNesterovMomentum,
51 weight_decay: f32,
52 no_prox: bool,
53}
54
55#[derive(Record, Clone, new)]
57pub struct AdanState<B: Backend, const D: usize> {
58 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 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 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#[derive(Record, Clone, new)]
132pub struct AdaptiveNesterovMomentumState<B: Backend, const D: usize> {
133 pub time: usize,
135 pub exp_avg: Tensor<B, D>,
137 pub exp_avg_sq: Tensor<B, D>,
139 pub exp_avg_diff: Tensor<B, D>,
141 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}