Skip to main content

optirs_core/optimizers/
rmsprop.rs

1// RMSprop optimizer implementation
2
3use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand, Zip};
4use scirs2_core::numeric::Float;
5use std::fmt::Debug;
6
7use crate::error::{OptimError, Result};
8use crate::optimizers::Optimizer;
9
10/// RMSprop optimizer
11///
12/// Implements the RMSprop optimization algorithm as proposed by Geoffrey Hinton
13/// in his Coursera course "Neural Networks for Machine Learning".
14///
15/// Formula:
16/// v_t = rho * v_{t-1} + (1 - rho) * g_t^2
17/// param_t = param_{t-1} - learning_rate * g_t / (sqrt(v_t) + epsilon)
18///
19/// # Examples
20///
21/// ```
22/// use scirs2_core::ndarray::Array1;
23/// use optirs_core::optimizers::{RMSprop, Optimizer};
24///
25/// // Initialize parameters and gradients
26/// let params = Array1::zeros(5);
27/// let gradients = Array1::from_vec(vec![0.1, 0.2, -0.3, 0.0, 0.5]);
28///
29/// // Create an RMSprop optimizer with learning rate 0.001
30/// let mut optimizer = RMSprop::new(0.001);
31///
32/// // Update parameters
33/// let new_params = optimizer.step(&params, &gradients).expect("optimizer.step succeeds");
34/// ```
35#[derive(Debug, Clone)]
36pub struct RMSprop<A: Float + ScalarOperand + Debug> {
37    /// Learning rate
38    learning_rate: A,
39    /// Decay rate for the moving average of squared gradients
40    rho: A,
41    /// Small constant for numerical stability
42    epsilon: A,
43    /// Weight decay factor (L2 regularization)
44    weight_decay: A,
45    /// Moving average of squared gradients, one slot per parameter-tensor index
46    v: Option<Vec<Array<A, IxDyn>>>,
47}
48
49impl<A: Float + ScalarOperand + Debug + Send + Sync> RMSprop<A> {
50    /// Creates a new RMSprop optimizer with the given learning rate and default settings
51    ///
52    /// # Arguments
53    ///
54    /// * `learning_rate` - The learning rate for parameter updates
55    pub fn new(learning_rate: A) -> Self {
56        Self {
57            learning_rate,
58            rho: A::from(0.9).expect("RMSprop: default rho (0.9) must fit in A"),
59            epsilon: A::from(1e-8).expect("RMSprop: default epsilon (1e-8) must fit in A"),
60            weight_decay: A::zero(),
61            v: None,
62        }
63    }
64
65    /// Creates a new RMSprop optimizer with the full configuration
66    ///
67    /// # Arguments
68    ///
69    /// * `learning_rate` - The learning rate for parameter updates
70    /// * `rho` - Decay rate for the moving average of squared gradients (default: 0.9)
71    /// * `epsilon` - Small constant for numerical stability (default: 1e-8)
72    /// * `weight_decay` - Weight decay factor for L2 regularization (default: 0.0)
73    pub fn new_with_config(learning_rate: A, rho: A, epsilon: A, weight_decay: A) -> Self {
74        Self {
75            learning_rate,
76            rho,
77            epsilon,
78            weight_decay,
79            v: None,
80        }
81    }
82
83    /// Sets the rho parameter
84    pub fn set_rho(&mut self, rho: A) -> &mut Self {
85        self.rho = rho;
86        self
87    }
88
89    /// Gets the rho parameter
90    pub fn get_rho(&self) -> A {
91        self.rho
92    }
93
94    /// Sets the epsilon parameter
95    pub fn set_epsilon(&mut self, epsilon: A) -> &mut Self {
96        self.epsilon = epsilon;
97        self
98    }
99
100    /// Gets the epsilon parameter
101    pub fn get_epsilon(&self) -> A {
102        self.epsilon
103    }
104
105    /// Sets the weight decay parameter
106    pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
107        self.weight_decay = weight_decay;
108        self
109    }
110
111    /// Gets the weight decay parameter
112    pub fn get_weight_decay(&self) -> A {
113        self.weight_decay
114    }
115
116    /// Resets the internal state of the optimizer
117    pub fn reset(&mut self) {
118        self.v = None;
119    }
120
121    /// Ensures a state slot exists for `index` and matches `dim`
122    fn ensure_state(&mut self, index: usize, dim: &IxDyn) {
123        let v = self.v.get_or_insert_with(Vec::new);
124        while v.len() <= index {
125            v.push(Array::zeros(dim.clone()));
126        }
127        if v[index].raw_dim() != *dim {
128            v[index] = Array::zeros(dim.clone());
129        }
130    }
131
132    /// Performs an RMSprop update for the parameter tensor at `index`
133    ///
134    /// Each `index` owns an independent moving average, so several parameter tensors
135    /// can be optimized by a single `RMSprop` instance without interference.
136    pub fn step_indexed<D: Dimension>(
137        &mut self,
138        index: usize,
139        params: &Array<A, D>,
140        gradients: &Array<A, D>,
141    ) -> Result<Array<A, D>> {
142        if params.shape() != gradients.shape() {
143            return Err(OptimError::DimensionMismatch(format!(
144                "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
145                params.shape(),
146                gradients.shape()
147            )));
148        }
149
150        let dim = params.raw_dim().into_dyn();
151        self.ensure_state(index, &dim);
152
153        let lr = self.learning_rate;
154        let rho = self.rho;
155        let eps = self.epsilon;
156        let weight_decay = self.weight_decay;
157        let use_weight_decay = weight_decay > A::zero();
158        let one = A::one();
159
160        let v = self.v.as_mut().ok_or_else(|| {
161            OptimError::InvalidConfig("RMSprop state not initialized".to_string())
162        })?;
163
164        let mut updated = params.to_owned();
165        let mut params_view = updated.view_mut().into_dyn();
166        let gradients_view = gradients.view().into_dyn();
167
168        Zip::from(&mut params_view)
169            .and(&gradients_view)
170            .and(&mut v[index])
171            .for_each(|p, &g, v_i| {
172                let grad = if use_weight_decay {
173                    g + weight_decay * *p
174                } else {
175                    g
176                };
177                *v_i = *v_i * rho + grad * grad * (one - rho);
178                *p = *p - lr * grad / (v_i.sqrt() + eps);
179            });
180        drop(params_view);
181
182        Ok(updated)
183    }
184}
185
186impl<A, D> Optimizer<A, D> for RMSprop<A>
187where
188    A: Float + ScalarOperand + Debug + Send + Sync,
189    D: Dimension,
190{
191    fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
192        self.step_indexed(0, params, gradients)
193    }
194
195    fn step_list(
196        &mut self,
197        params_list: &[&Array<A, D>],
198        gradients_list: &[&Array<A, D>],
199    ) -> Result<Vec<Array<A, D>>> {
200        if params_list.len() != gradients_list.len() {
201            return Err(OptimError::InvalidConfig(format!(
202                "Number of parameter arrays ({}) does not match number of gradient arrays ({})",
203                params_list.len(),
204                gradients_list.len()
205            )));
206        }
207
208        let mut results = Vec::with_capacity(params_list.len());
209        for (index, (params, grads)) in params_list.iter().zip(gradients_list.iter()).enumerate() {
210            results.push(self.step_indexed(index, params, grads)?);
211        }
212        Ok(results)
213    }
214
215    fn get_learning_rate(&self) -> A {
216        self.learning_rate
217    }
218
219    fn set_learning_rate(&mut self, learning_rate: A) {
220        self.learning_rate = learning_rate;
221    }
222}