1use crate::error::{OptimError, Result};
15use crate::optimizers::Optimizer;
16use scirs2_core::ndarray::{Ix1, ScalarOperand};
17use scirs2_core::ndarray_ext::{Array1, ArrayView1};
18use scirs2_core::numeric::Float;
19use serde::{Deserialize, Serialize};
20use std::fmt::Debug;
21
22#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct Ranger<T: Float + ScalarOperand> {
37 learning_rate: T,
39 beta1: T,
40 beta2: T,
41 epsilon: T,
42 weight_decay: T,
43
44 lookahead_k: usize,
46 lookahead_alpha: T,
47
48 momentum: Option<Array1<T>>,
50 velocity: Option<Array1<T>>,
51
52 slow_weights: Option<Array1<T>>,
54
55 step_count: usize,
57 slow_update_count: usize,
58}
59
60impl<T: Float + ScalarOperand> Default for Ranger<T> {
61 fn default() -> Self {
62 Self::new(
63 T::from(0.001).expect("Ranger: default learning_rate (0.001) must fit in T"),
64 T::from(0.9).expect("Ranger: default beta1 (0.9) must fit in T"),
65 T::from(0.999).expect("Ranger: default beta2 (0.999) must fit in T"),
66 T::from(1e-8).expect("Ranger: default epsilon (1e-8) must fit in T"),
67 T::zero(),
68 5,
69 T::from(0.5).expect("Ranger: default lookahead_alpha (0.5) must fit in T"),
70 )
71 .expect("Ranger: default hyperparameters always satisfy validation")
72 }
73}
74
75impl<T: Float + ScalarOperand> Ranger<T> {
76 pub fn new(
102 learning_rate: T,
103 beta1: T,
104 beta2: T,
105 epsilon: T,
106 weight_decay: T,
107 lookahead_k: usize,
108 lookahead_alpha: T,
109 ) -> Result<Self> {
110 let lr_f64 = crate::optimizers::scalar_to_f64(learning_rate)?;
112 let beta1_f64 = crate::optimizers::scalar_to_f64(beta1)?;
113 let beta2_f64 = crate::optimizers::scalar_to_f64(beta2)?;
114 let eps_f64 = crate::optimizers::scalar_to_f64(epsilon)?;
115 let wd_f64 = crate::optimizers::scalar_to_f64(weight_decay)?;
116 let lookahead_alpha_f64 = crate::optimizers::scalar_to_f64(lookahead_alpha)?;
117
118 if lr_f64 <= 0.0 {
119 return Err(OptimError::InvalidParameter(format!(
120 "learning_rate must be positive, got {lr_f64}"
121 )));
122 }
123 if beta1_f64 <= 0.0 || beta1_f64 >= 1.0 {
124 return Err(OptimError::InvalidParameter(format!(
125 "beta1 must be in (0, 1), got {beta1_f64}"
126 )));
127 }
128 if beta2_f64 <= 0.0 || beta2_f64 >= 1.0 {
129 return Err(OptimError::InvalidParameter(format!(
130 "beta2 must be in (0, 1), got {beta2_f64}"
131 )));
132 }
133 if eps_f64 <= 0.0 {
134 return Err(OptimError::InvalidParameter(format!(
135 "epsilon must be positive, got {eps_f64}"
136 )));
137 }
138 if wd_f64 < 0.0 {
139 return Err(OptimError::InvalidParameter(format!(
140 "weight_decay must be non-negative, got {wd_f64}"
141 )));
142 }
143 if lookahead_k == 0 {
144 return Err(OptimError::InvalidParameter(
145 "lookahead_k must be positive".to_string(),
146 ));
147 }
148 if lookahead_alpha_f64 <= 0.0 || lookahead_alpha_f64 > 1.0 {
149 return Err(OptimError::InvalidParameter(format!(
150 "lookahead_alpha must be in (0, 1], got {lookahead_alpha_f64}"
151 )));
152 }
153
154 Ok(Self {
155 learning_rate,
156 beta1,
157 beta2,
158 epsilon,
159 weight_decay,
160 lookahead_k,
161 lookahead_alpha,
162 momentum: None,
163 velocity: None,
164 slow_weights: None,
165 step_count: 0,
166 slow_update_count: 0,
167 })
168 }
169
170 pub fn step<'a, P, G>(&mut self, params: P, grads: G) -> Result<Array1<T>>
186 where
187 P: Into<ArrayView1<'a, T>>,
188 G: Into<ArrayView1<'a, T>>,
189 T: 'a,
190 {
191 self.step_view(params.into(), grads.into())
192 }
193
194 pub fn step_view(&mut self, params: ArrayView1<T>, grads: ArrayView1<T>) -> Result<Array1<T>> {
198 let n = params.len();
199
200 if grads.len() != n {
201 return Err(OptimError::DimensionMismatch(format!(
202 "Expected gradient size {}, got {}",
203 n,
204 grads.len()
205 )));
206 }
207
208 if self.slow_weights.is_none() {
210 self.slow_weights = Some(params.to_owned());
211 }
212
213 self.step_count += 1;
214 let t: T = crate::optimizers::cast_scalar(self.step_count)?;
215
216 let momentum = self.momentum.get_or_insert_with(|| Array1::zeros(n));
217 let velocity = self.velocity.get_or_insert_with(|| Array1::zeros(n));
218
219 let one = T::one();
220 let two: T = crate::optimizers::cast_scalar(2)?;
221
222 let effective_grads = if self.weight_decay > T::zero() {
224 grads.to_owned() + &(params.to_owned() * self.weight_decay)
225 } else {
226 grads.to_owned()
227 };
228
229 for i in 0..n {
231 momentum[i] = self.beta1 * momentum[i] + (one - self.beta1) * effective_grads[i];
232 }
233
234 for i in 0..n {
236 let grad_sq = effective_grads[i] * effective_grads[i];
237 velocity[i] = self.beta2 * velocity[i] + (one - self.beta2) * grad_sq;
238 }
239
240 let bias_correction1 = one - self.beta1.powf(t);
242 let bias_correction2 = one - self.beta2.powf(t);
243
244 let rho_inf = two / (one - self.beta2) - one;
246 let rho_t = rho_inf - two * t * self.beta2.powf(t) / bias_correction2;
247
248 let mut updated_params = params.to_owned();
250
251 if crate::optimizers::scalar_to_f64(rho_t)? > 4.0 {
252 let four: T = crate::optimizers::cast_scalar(4)?;
254 let rect_term = ((rho_t - four) * (rho_t - two) * rho_inf
255 / ((rho_inf - four) * (rho_inf - two) * rho_t))
256 .sqrt();
257
258 for i in 0..n {
259 let m_hat = momentum[i] / bias_correction1;
260 let v_hat = velocity[i] / bias_correction2;
261 let step_size = self.learning_rate * rect_term / (v_hat.sqrt() + self.epsilon);
262 updated_params[i] = updated_params[i] - step_size * m_hat;
263 }
264 } else {
265 for i in 0..n {
267 let m_hat = momentum[i] / bias_correction1;
268 updated_params[i] = updated_params[i] - self.learning_rate * m_hat;
269 }
270 }
271
272 if self.step_count.is_multiple_of(self.lookahead_k) {
274 let slow = self.slow_weights.get_or_insert_with(|| params.to_owned());
275 for i in 0..n {
276 slow[i] = slow[i] + self.lookahead_alpha * (updated_params[i] - slow[i]);
277 }
278 self.slow_update_count += 1;
279
280 Ok(slow.clone())
283 } else {
284 Ok(updated_params)
286 }
287 }
288
289 pub fn step_count(&self) -> usize {
291 self.step_count
292 }
293
294 pub fn slow_update_count(&self) -> usize {
296 self.slow_update_count
297 }
298
299 pub fn reset(&mut self) {
301 self.momentum = None;
302 self.velocity = None;
303 self.slow_weights = None;
304 self.step_count = 0;
305 self.slow_update_count = 0;
306 }
307
308 pub fn slow_weights(&self) -> Option<&Array1<T>> {
310 self.slow_weights.as_ref()
311 }
312
313 pub fn is_rectified(&self) -> bool {
315 if self.step_count == 0 {
316 return false;
317 }
318 let t = T::from(self.step_count)
319 .expect("Ranger: step_count must be representable in T (f32/f64)");
320 let one = T::one();
321 let two = T::from(2).expect("Ranger: integer literal 2 must be representable in T");
322 let bias_correction2 = one - self.beta2.powf(t);
323 let rho_inf = two / (one - self.beta2) - one;
324 let rho_t = rho_inf - two * t * self.beta2.powf(t) / bias_correction2;
325 rho_t
326 .to_f64()
327 .expect("Ranger: T (f32/f64) always converts to f64")
328 > 4.0
329 }
330}
331
332impl<T> Optimizer<T, Ix1> for Ranger<T>
333where
334 T: Float + ScalarOperand + Debug + Send + Sync,
335{
336 fn step(&mut self, params: &Array1<T>, gradients: &Array1<T>) -> Result<Array1<T>> {
337 self.step_view(params.view(), gradients.view())
338 }
339
340 fn get_learning_rate(&self) -> T {
341 self.learning_rate
342 }
343
344 fn set_learning_rate(&mut self, learning_rate: T) {
345 self.learning_rate = learning_rate;
346 }
347}
348
349#[cfg(test)]
350mod tests {
351 use super::*;
352 use scirs2_core::ndarray_ext::array;
353
354 #[test]
355 fn test_ranger_creation() {
356 let optimizer = Ranger::<f32>::default();
357 assert_eq!(optimizer.step_count(), 0);
358 assert_eq!(optimizer.slow_update_count(), 0);
359 }
360
361 #[test]
362 fn test_ranger_custom_creation() {
363 let optimizer = Ranger::<f32>::new(0.002, 0.95, 0.9999, 1e-7, 0.01, 6, 0.6)
364 .expect("Ranger::<f32>::new succeeds in test_ranger_custom_creation");
365 assert_eq!(optimizer.step_count(), 0);
366 }
367
368 #[test]
369 fn test_ranger_single_step() {
370 let mut optimizer = Ranger::<f32>::default();
371 let params = array![1.0, 2.0, 3.0];
372 let grads = array![0.1, 0.2, 0.3];
373
374 let updated_params = optimizer
375 .step(params.view(), grads.view())
376 .expect("step succeeds in test_ranger_single_step");
377 assert_eq!(updated_params.len(), 3);
378 assert_eq!(optimizer.step_count(), 1);
379
380 for i in 0..3 {
381 assert!(updated_params[i] < params[i]);
382 }
383 }
384
385 #[test]
386 fn test_ranger_slow_updates() {
387 let mut optimizer = Ranger::<f32>::new(0.001, 0.9, 0.999, 1e-8, 0.0, 3, 0.5)
388 .expect("Ranger::<f32>::new succeeds in test_ranger_slow_updates");
389 let mut params = array![1.0, 2.0, 3.0];
390
391 for _ in 0..3 {
392 let grads = array![0.1, 0.2, 0.3];
393 params = optimizer
394 .step(params.view(), grads.view())
395 .expect("step succeeds in test_ranger_slow_updates");
396 }
397 assert_eq!(optimizer.slow_update_count(), 1);
398 }
399
400 #[test]
401 fn test_ranger_convergence() {
402 let mut optimizer = Ranger::<f64>::new(
405 0.1, 0.9, 0.999, 1e-8, 0.0, 5, 0.5, )
413 .expect("Ranger::new succeeds in test_ranger_convergence");
414 let mut params = array![5.0];
415
416 for _ in 0..500 {
418 let grads = params.mapv(|x| 2.0 * x);
419 params = optimizer
420 .step(params.view(), grads.view())
421 .expect("step succeeds in test_ranger_convergence");
422 }
423
424 assert!(
425 params[0].abs() < 0.1,
426 "Failed to converge, got {}",
427 params[0]
428 );
429 }
430
431 #[test]
432 fn test_ranger_reset() {
433 let mut optimizer = Ranger::<f32>::default();
434 let params = array![1.0, 2.0, 3.0];
435 let grads = array![0.1, 0.2, 0.3];
436
437 for _ in 0..10 {
438 optimizer
439 .step(params.view(), grads.view())
440 .expect("step succeeds in test_ranger_reset");
441 }
442
443 optimizer.reset();
444 assert_eq!(optimizer.step_count(), 0);
445 assert_eq!(optimizer.slow_update_count(), 0);
446 assert!(optimizer.slow_weights().is_none());
447 }
448
449 #[test]
450 fn test_ranger_rectification() {
451 let mut optimizer = Ranger::<f32>::default();
452 let params = array![1.0];
453 let grads = array![0.1];
454
455 assert!(!optimizer.is_rectified());
457
458 for _ in 0..10 {
460 optimizer
461 .step(params.view(), grads.view())
462 .expect("step succeeds in test_ranger_rectification");
463 }
464 assert!(optimizer.is_rectified());
465 }
466
467 #[test]
469 fn test_ranger_optimizer_trait() {
470 let mut optimizer = Ranger::<f64>::default();
471 let params = array![1.0f64, 2.0, 3.0];
472 let grads = array![0.1f64, 0.2, 0.3];
473
474 let updated =
475 Optimizer::<f64, scirs2_core::ndarray::Ix1>::step(&mut optimizer, ¶ms, &grads)
476 .expect("trait step failed");
477 assert_eq!(updated.len(), 3);
478
479 let again = optimizer.step(¶ms, &grads).expect("ref step failed");
481 assert_eq!(again.len(), 3);
482 }
483}