optirs_core/optimizers/
rmsprop.rs1use 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#[derive(Debug, Clone)]
36pub struct RMSprop<A: Float + ScalarOperand + Debug> {
37 learning_rate: A,
39 rho: A,
41 epsilon: A,
43 weight_decay: A,
45 v: Option<Vec<Array<A, IxDyn>>>,
47}
48
49impl<A: Float + ScalarOperand + Debug + Send + Sync> RMSprop<A> {
50 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 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 pub fn set_rho(&mut self, rho: A) -> &mut Self {
85 self.rho = rho;
86 self
87 }
88
89 pub fn get_rho(&self) -> A {
91 self.rho
92 }
93
94 pub fn set_epsilon(&mut self, epsilon: A) -> &mut Self {
96 self.epsilon = epsilon;
97 self
98 }
99
100 pub fn get_epsilon(&self) -> A {
102 self.epsilon
103 }
104
105 pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
107 self.weight_decay = weight_decay;
108 self
109 }
110
111 pub fn get_weight_decay(&self) -> A {
113 self.weight_decay
114 }
115
116 pub fn reset(&mut self) {
118 self.v = None;
119 }
120
121 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 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}