Skip to main content

optirs_core/optimizers/
ranger.rs

1// OptiRS - Ranger Optimizer
2// RAdam + Lookahead combination for improved convergence and stability
3// Reference: "Ranger - a synergistic optimizer" by Less Wright (2019)
4//
5// Ranger combines:
6// 1. RAdam (Rectified Adam) - Adaptive learning rate with variance rectification
7// 2. Lookahead - Slow and fast weight updates for stability
8//
9// This combination provides:
10// - Fast convergence from RAdam
11// - Stability and reduced variance from Lookahead
12// - Better generalization than either optimizer alone
13
14use 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/// Ranger optimizer configuration
23///
24/// Ranger combines RAdam (Rectified Adam) with Lookahead mechanism.
25/// This standalone implementation integrates both algorithms efficiently.
26///
27/// # Key Features
28/// - Fast convergence from RAdam's variance rectification
29/// - Stability from Lookahead's slow weight trajectory
30/// - Reduced sensitivity to hyperparameter choices
31/// - Better generalization than Adam or RAdam alone
32///
33/// # Type Parameters
34/// - `T`: Floating-point type (f32 or f64)
35#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct Ranger<T: Float + ScalarOperand> {
37    // RAdam parameters
38    learning_rate: T,
39    beta1: T,
40    beta2: T,
41    epsilon: T,
42    weight_decay: T,
43
44    // Lookahead parameters
45    lookahead_k: usize,
46    lookahead_alpha: T,
47
48    // RAdam state
49    momentum: Option<Array1<T>>,
50    velocity: Option<Array1<T>>,
51
52    // Lookahead state
53    slow_weights: Option<Array1<T>>,
54
55    // Step counters
56    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    /// Create a new Ranger optimizer
77    ///
78    /// # Arguments
79    /// - `learning_rate`: Learning rate for RAdam (typically 0.001)
80    /// - `beta1`: First moment decay rate (typically 0.9)
81    /// - `beta2`: Second moment decay rate (typically 0.999)
82    /// - `epsilon`: Small constant for numerical stability (typically 1e-8)
83    /// - `weight_decay`: L2 regularization coefficient (typically 0.0)
84    /// - `lookahead_k`: Number of fast updates per slow update (typically 5-6)
85    /// - `lookahead_alpha`: Interpolation factor for slow weights (typically 0.5)
86    ///
87    /// # Example
88    /// ```
89    /// use optirs_core::optimizers::Ranger;
90    ///
91    /// let optimizer = Ranger::<f32>::new(
92    ///     0.001,  // learning_rate
93    ///     0.9,    // beta1
94    ///     0.999,  // beta2
95    ///     1e-8,   // epsilon
96    ///     0.0,    // weight_decay
97    ///     5,      // lookahead_k
98    ///     0.5     // lookahead_alpha
99    /// ).expect("Ranger::new succeeds for finite, in-range default hyperparameters");
100    /// ```
101    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        // Validate parameters
111        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    /// Perform a single optimization step
171    ///
172    /// Combines RAdam (fast weights) with Lookahead (slow weights)
173    ///
174    /// # Example
175    /// ```
176    /// use optirs_core::optimizers::Ranger;
177    /// use scirs2_core::ndarray_ext::array;
178    ///
179    /// let mut optimizer = Ranger::<f32>::default();
180    /// let params = array![1.0, 2.0, 3.0];
181    /// let grads = array![0.1, 0.2, 0.3];
182    ///
183    /// let updated_params = optimizer.step(params.view(), grads.view()).expect("optimizer.step succeeds");
184    /// ```
185    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    /// Perform a single optimization step on borrowed views
195    ///
196    /// This is the concrete implementation behind [`Ranger::step`].
197    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        // Initialize state on first step
209        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        // Apply weight decay if configured
223        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        // RAdam: Update biased first moment
230        for i in 0..n {
231            momentum[i] = self.beta1 * momentum[i] + (one - self.beta1) * effective_grads[i];
232        }
233
234        // RAdam: Update biased second moment
235        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        // RAdam: Compute bias correction
241        let bias_correction1 = one - self.beta1.powf(t);
242        let bias_correction2 = one - self.beta2.powf(t);
243
244        // RAdam: Compute SMA (Simple Moving Average) length
245        let rho_inf = two / (one - self.beta2) - one;
246        let rho_t = rho_inf - two * t * self.beta2.powf(t) / bias_correction2;
247
248        // RAdam: Apply variance rectification
249        let mut updated_params = params.to_owned();
250
251        if crate::optimizers::scalar_to_f64(rho_t)? > 4.0 {
252            // Use adaptive learning rate with variance rectification
253            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            // Use simple momentum update during warmup
266            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        // Lookahead: Update slow weights every k steps
273        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            // Synchronize fast weights with slow weights
281            // This is the key to Lookahead: we return the slow weights after update
282            Ok(slow.clone())
283        } else {
284            // Between slow updates, return fast weights
285            Ok(updated_params)
286        }
287    }
288
289    /// Get the number of optimization steps performed
290    pub fn step_count(&self) -> usize {
291        self.step_count
292    }
293
294    /// Get the number of slow weight updates performed
295    pub fn slow_update_count(&self) -> usize {
296        self.slow_update_count
297    }
298
299    /// Reset the optimizer state
300    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    /// Get the slow weights (Lookahead trajectory)
309    pub fn slow_weights(&self) -> Option<&Array1<T>> {
310        self.slow_weights.as_ref()
311    }
312
313    /// Check if variance rectification is active
314    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        // Use higher learning rate for this simple convex problem
403        // Default 0.001 is tuned for neural networks
404        let mut optimizer = Ranger::<f64>::new(
405            0.1,   // learning_rate: higher for simple problem
406            0.9,   // beta1
407            0.999, // beta2
408            1e-8,  // epsilon
409            0.0,   // weight_decay
410            5,     // lookahead_k
411            0.5,   // lookahead_alpha
412        )
413        .expect("Ranger::new succeeds in test_ranger_convergence");
414        let mut params = array![5.0];
415
416        // Ranger combines RAdam (adaptive LR) with Lookahead (slow updates)
417        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        // Initially not rectified
456        assert!(!optimizer.is_rectified());
457
458        // After several steps, should be rectified
459        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    /// Ranger must be usable through the generic `Optimizer` trait.
468    #[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, &params, &grads)
476                .expect("trait step failed");
477        assert_eq!(updated.len(), 3);
478
479        // The generic inherent `step` also accepts plain references.
480        let again = optimizer.step(&params, &grads).expect("ref step failed");
481        assert_eq!(again.len(), 3);
482    }
483}