1use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand, Zip};
6use scirs2_core::numeric::Float;
7use std::fmt::Debug;
8
9use crate::error::{OptimError, Result};
10use crate::optimizers::Optimizer;
11
12#[derive(Debug, Clone)]
56pub struct RAdam<A: Float + ScalarOperand + Debug> {
57 learning_rate: A,
59 beta1: A,
61 beta2: A,
63 epsilon: A,
65 weight_decay: A,
67 m: Option<Vec<Array<A, IxDyn>>>,
69 v: Option<Vec<Array<A, IxDyn>>>,
71 t: Vec<usize>,
73 rho_inf: A,
75}
76
77impl<A: Float + ScalarOperand + Debug + Send + Sync> RAdam<A> {
78 pub fn new(learning_rate: A) -> Self {
84 let beta2 = A::from(0.999).expect("RAdam: default beta2 (0.999) must fit in A");
85 Self {
86 learning_rate,
87 beta1: A::from(0.9).expect("RAdam: default beta1 (0.9) must fit in A"),
88 beta2,
89 epsilon: A::from(1e-8).expect("RAdam: default epsilon (1e-8) must fit in A"),
90 weight_decay: A::zero(),
91 m: None,
92 v: None,
93 t: Vec::new(),
94 rho_inf: A::from(2.0).expect("RAdam: integer literal 2.0 must fit in A")
95 / (A::one() - beta2)
96 - A::one(),
97 }
98 }
99
100 pub fn new_with_config(
110 learning_rate: A,
111 beta1: A,
112 beta2: A,
113 epsilon: A,
114 weight_decay: A,
115 ) -> Self {
116 Self {
117 learning_rate,
118 beta1,
119 beta2,
120 epsilon,
121 weight_decay,
122 m: None,
123 v: None,
124 t: Vec::new(),
125 rho_inf: A::from(2.0).expect("RAdam: integer literal 2.0 must fit in A")
126 / (A::one() - beta2)
127 - A::one(),
128 }
129 }
130
131 pub fn set_beta1(&mut self, beta1: A) -> &mut Self {
133 self.beta1 = beta1;
134 self
135 }
136
137 pub fn get_beta1(&self) -> A {
139 self.beta1
140 }
141
142 pub fn set_beta2(&mut self, beta2: A) -> &mut Self {
144 self.beta2 = beta2;
145 self.rho_inf = A::from(2.0).expect("RAdam: integer literal 2.0 must fit in A")
147 / (A::one() - beta2)
148 - A::one();
149 self
150 }
151
152 pub fn get_beta2(&self) -> A {
154 self.beta2
155 }
156
157 pub fn set_epsilon(&mut self, epsilon: A) -> &mut Self {
159 self.epsilon = epsilon;
160 self
161 }
162
163 pub fn get_epsilon(&self) -> A {
165 self.epsilon
166 }
167
168 pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
170 self.weight_decay = weight_decay;
171 self
172 }
173
174 pub fn get_weight_decay(&self) -> A {
176 self.weight_decay
177 }
178
179 pub fn learning_rate(&self) -> A {
181 self.learning_rate
182 }
183
184 pub fn set_lr(&mut self, lr: A) {
186 self.learning_rate = lr;
187 }
188
189 pub fn reset(&mut self) {
191 self.m = None;
192 self.v = None;
193 self.t.clear();
194 }
195
196 pub fn timestep(&self, index: usize) -> usize {
200 self.t.get(index).copied().unwrap_or(0)
201 }
202
203 pub fn rho_inf(&self) -> A {
205 self.rho_inf
206 }
207
208 pub fn rho_t(&self, t: usize) -> Option<A> {
212 if t == 0 {
213 return None;
214 }
215 let exp = i32::try_from(t).ok()?;
216 let two = A::one() + A::one();
217 let t_f = A::from(t)?;
218 let beta2_t = self.beta2.powi(exp);
219 Some(self.rho_inf - two * t_f * beta2_t / (A::one() - beta2_t))
220 }
221
222 pub fn rectification_term(&self, t: usize) -> Option<A> {
229 let rho_t = self.rho_t(t)?;
230 let two = A::one() + A::one();
231 let four = two + two;
232 if rho_t <= four {
233 return None;
234 }
235 let rho_inf = self.rho_inf;
236 let numerator = (rho_t - four) * (rho_t - two) * rho_inf;
237 let denominator = (rho_inf - four) * (rho_inf - two) * rho_t;
238 if denominator <= A::zero() {
239 return None;
240 }
241 Some((numerator / denominator).sqrt())
242 }
243
244 fn advance_state(&mut self, index: usize, dim: &IxDyn) -> Result<usize> {
246 let m = self.m.get_or_insert_with(Vec::new);
247 let v = self.v.get_or_insert_with(Vec::new);
248 while m.len() <= index {
249 m.push(Array::zeros(dim.clone()));
250 }
251 while v.len() <= index {
252 v.push(Array::zeros(dim.clone()));
253 }
254 while self.t.len() <= index {
255 self.t.push(0);
256 }
257
258 if m[index].raw_dim() != *dim || v[index].raw_dim() != *dim {
260 m[index] = Array::zeros(dim.clone());
261 v[index] = Array::zeros(dim.clone());
262 self.t[index] = 0;
263 }
264
265 let next = self.t[index].checked_add(1).ok_or_else(|| {
266 OptimError::InvalidConfig(
267 "Timestep counter overflow - too many optimization steps".to_string(),
268 )
269 })?;
270 self.t[index] = next;
271 Ok(next)
272 }
273
274 pub fn step_inplace_indexed<D: Dimension>(
276 &mut self,
277 index: usize,
278 params: &mut Array<A, D>,
279 gradients: &Array<A, D>,
280 ) -> Result<()> {
281 if params.shape() != gradients.shape() {
282 return Err(OptimError::DimensionMismatch(format!(
283 "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
284 params.shape(),
285 gradients.shape()
286 )));
287 }
288
289 let dim = params.raw_dim().into_dyn();
290 let t = self.advance_state(index, &dim)?;
291 let exp = i32::try_from(t).map_err(|_| {
292 OptimError::InvalidConfig(
293 "Timestep too large for bias correction calculation".to_string(),
294 )
295 })?;
296
297 let beta1 = self.beta1;
298 let beta2 = self.beta2;
299 let lr = self.learning_rate;
300 let eps = self.epsilon;
301 let weight_decay = self.weight_decay;
302 let use_weight_decay = weight_decay > A::zero();
303 let one = A::one();
304 let bias_correction1 = one - beta1.powi(exp);
305 let bias_correction2 = one - beta2.powi(exp);
306
307 let rect = self.rectification_term(t);
310
311 let m = self
312 .m
313 .as_mut()
314 .ok_or_else(|| OptimError::InvalidConfig("RAdam state not initialized".to_string()))?;
315 let v = self
316 .v
317 .as_mut()
318 .ok_or_else(|| OptimError::InvalidConfig("RAdam state not initialized".to_string()))?;
319
320 let mut params_view = params.view_mut().into_dyn();
321 let gradients_view = gradients.view().into_dyn();
322
323 Zip::from(&mut params_view)
324 .and(&gradients_view)
325 .and(&mut m[index])
326 .and(&mut v[index])
327 .for_each(|p, &g, m_i, v_i| {
328 let grad = if use_weight_decay {
329 g + weight_decay * *p
330 } else {
331 g
332 };
333 *m_i = *m_i * beta1 + grad * (one - beta1);
334 *v_i = *v_i * beta2 + grad * grad * (one - beta2);
335 let m_hat = *m_i / bias_correction1;
336 match rect {
337 Some(r_t) => {
338 let v_hat = *v_i / bias_correction2;
339 *p = *p - lr * r_t * m_hat / (v_hat.sqrt() + eps);
340 }
341 None => {
342 *p = *p - lr * m_hat;
343 }
344 }
345 });
346
347 Ok(())
348 }
349
350 pub fn step_inplace<D: Dimension>(
352 &mut self,
353 params: &mut Array<A, D>,
354 gradients: &Array<A, D>,
355 ) -> Result<()> {
356 self.step_inplace_indexed(0, params, gradients)
357 }
358
359 pub fn step_indexed<D: Dimension>(
363 &mut self,
364 index: usize,
365 params: &Array<A, D>,
366 gradients: &Array<A, D>,
367 ) -> Result<Array<A, D>> {
368 let mut updated = params.to_owned();
369 self.step_inplace_indexed(index, &mut updated, gradients)?;
370 Ok(updated)
371 }
372}
373
374impl<A, D> Optimizer<A, D> for RAdam<A>
375where
376 A: Float + ScalarOperand + Debug + Send + Sync + std::convert::From<f64>,
377 D: Dimension,
378{
379 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
380 self.step_indexed(0, params, gradients)
381 }
382
383 fn step_list(
384 &mut self,
385 params_list: &[&Array<A, D>],
386 gradients_list: &[&Array<A, D>],
387 ) -> Result<Vec<Array<A, D>>> {
388 if params_list.len() != gradients_list.len() {
389 return Err(OptimError::InvalidConfig(format!(
390 "Number of parameter arrays ({}) does not match number of gradient arrays ({})",
391 params_list.len(),
392 gradients_list.len()
393 )));
394 }
395
396 let mut results = Vec::with_capacity(params_list.len());
397 for (index, (params, grads)) in params_list.iter().zip(gradients_list.iter()).enumerate() {
398 results.push(self.step_indexed(index, params, grads)?);
399 }
400 Ok(results)
401 }
402
403 fn get_learning_rate(&self) -> A {
404 self.learning_rate
405 }
406
407 fn set_learning_rate(&mut self, learning_rate: A) {
408 self.learning_rate = learning_rate;
409 }
410}
411
412#[cfg(test)]
413mod tests {
414 use super::*;
415 use scirs2_core::ndarray::Array1;
416
417 #[test]
418 fn test_radam_step() {
419 let params = Array1::zeros(3);
421 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
422
423 let mut optimizer = RAdam::new(0.01);
425
426 let new_params = optimizer
428 .step(¶ms, &gradients)
429 .expect("optimizer.step succeeds in test_radam_step");
430
431 assert!(new_params.iter().all(|&x| x != 0.0));
433
434 for i in 1..3 {
437 assert!(new_params[i].abs() > new_params[i - 1].abs());
438 }
439 }
440
441 #[test]
442 fn test_radam_multiple_steps() {
443 let mut params = Array1::zeros(3);
445 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
446
447 let mut optimizer = RAdam::new(0.01);
449
450 for _ in 0..100 {
452 params = optimizer
453 .step(¶ms, &gradients)
454 .expect("optimizer.step succeeds in test_radam_multiple_steps");
455 }
456
457 for i in 1..3 {
460 assert!(params[i].abs() > params[i - 1].abs());
461 }
462 }
463
464 #[test]
465 fn test_radam_weight_decay() {
466 let params = Array1::from_vec(vec![0.1, 0.2, 0.3]);
468 let gradients = Array1::from_vec(vec![0.01, 0.01, 0.01]);
469
470 let mut optimizer = RAdam::new_with_config(
472 0.01, 0.9, 0.999, 1e-8, 0.1, );
474
475 let new_params = optimizer
477 .step(¶ms, &gradients)
478 .expect("optimizer.step succeeds in test_radam_weight_decay");
479
480 for i in 0..3 {
482 assert!(new_params[i].abs() < params[i].abs());
483 }
484 }
485
486 #[test]
505 fn test_radam_reset() {
506 let params = Array1::zeros(3);
508 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
509
510 let mut optimizer = RAdam::new(0.01);
512
513 optimizer
515 .step(¶ms, &gradients)
516 .expect("optimizer.step succeeds in test_radam_reset");
517 assert_eq!(optimizer.timestep(0), 1);
518 assert!(optimizer.m.is_some());
519 assert!(optimizer.v.is_some());
520
521 optimizer.reset();
523 assert_eq!(optimizer.timestep(0), 0);
524 assert!(optimizer.m.is_none());
525 assert!(optimizer.v.is_none());
526 }
527}