optirs_core/optimizers/
reptile.rs1use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
13use scirs2_core::numeric::Float;
14use std::fmt::Debug;
15
16use crate::error::{OptimError, Result};
17use crate::optimizers::Optimizer;
18
19#[derive(Debug, Clone)]
48pub struct ReptileOptimizer<A: Float + ScalarOperand + Debug> {
49 learning_rate: A,
51 inner_lr: A,
53 inner_steps: usize,
55 epsilon: A,
57 step_count: usize,
59}
60
61impl<A: Float + ScalarOperand + Debug> ReptileOptimizer<A> {
62 pub fn new(lr: A) -> Self {
73 Self {
74 learning_rate: lr,
75 inner_lr: lr,
76 inner_steps: 5,
77 epsilon: lr,
78 step_count: 0,
79 }
80 }
81
82 pub fn with_inner_steps(mut self, n: usize) -> Self {
90 self.inner_steps = if n == 0 { 1 } else { n };
91 self
92 }
93
94 pub fn with_epsilon(mut self, e: A) -> Self {
103 self.epsilon = e;
104 self
105 }
106
107 pub fn with_inner_lr(mut self, lr: A) -> Self {
115 self.inner_lr = lr;
116 self
117 }
118
119 pub fn get_inner_steps(&self) -> usize {
121 self.inner_steps
122 }
123
124 pub fn get_epsilon(&self) -> A {
126 self.epsilon
127 }
128
129 pub fn get_inner_lr(&self) -> A {
131 self.inner_lr
132 }
133
134 pub fn get_step_count(&self) -> usize {
136 self.step_count
137 }
138}
139
140impl<A, D> Optimizer<A, D> for ReptileOptimizer<A>
141where
142 A: Float + ScalarOperand + Debug,
143 D: Dimension,
144{
145 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
146 let params_dyn = params.to_owned().into_dyn();
148 let gradients_dyn = gradients.to_owned().into_dyn();
149
150 let theta_original = params_dyn.clone();
152
153 let mut theta_adapted = params_dyn;
158 for _ in 0..self.inner_steps {
159 theta_adapted = &theta_adapted - &(&gradients_dyn * self.inner_lr);
160 }
161
162 let meta_direction = &theta_adapted - &theta_original;
164
165 let updated_params = &theta_original + &(&meta_direction * self.epsilon);
167
168 self.step_count += 1;
169
170 updated_params.into_dimensionality::<D>().map_err(|e| {
172 OptimError::DimensionMismatch(format!(
173 "Reptile: failed to convert back to original dimensionality: {e}"
174 ))
175 })
176 }
177
178 fn get_learning_rate(&self) -> A {
179 self.learning_rate
180 }
181
182 fn set_learning_rate(&mut self, learning_rate: A) {
183 self.learning_rate = learning_rate;
184 self.epsilon = learning_rate;
185 }
186}
187
188#[cfg(test)]
189mod tests {
190 use super::*;
191 use scirs2_core::ndarray::Array1;
192
193 #[test]
194 fn test_reptile_basic_creation() {
195 let optimizer: ReptileOptimizer<f64> = ReptileOptimizer::new(0.01);
196 assert!(
197 (Optimizer::<f64, scirs2_core::ndarray::Ix1>::get_learning_rate(&optimizer) - 0.01)
198 .abs()
199 < 1e-10
200 );
201 assert_eq!(optimizer.get_inner_steps(), 5);
202 assert!((optimizer.get_epsilon() - 0.01).abs() < 1e-10);
203 assert!((optimizer.get_inner_lr() - 0.01).abs() < 1e-10);
204 assert_eq!(optimizer.get_step_count(), 0);
205 }
206
207 #[test]
208 fn test_reptile_builder_pattern() {
209 let optimizer: ReptileOptimizer<f64> = ReptileOptimizer::new(0.01)
210 .with_inner_steps(10)
211 .with_epsilon(0.05)
212 .with_inner_lr(0.001);
213
214 assert_eq!(optimizer.get_inner_steps(), 10);
215 assert!((optimizer.get_epsilon() - 0.05).abs() < 1e-10);
216 assert!((optimizer.get_inner_lr() - 0.001).abs() < 1e-10);
217 }
218
219 #[test]
220 fn test_reptile_step_works() {
221 let mut optimizer = ReptileOptimizer::new(0.1_f64)
222 .with_inner_steps(1)
223 .with_epsilon(1.0)
224 .with_inner_lr(0.1);
225
226 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
227 let gradients = Array1::from_vec(vec![0.5, -0.5, 0.0]);
228
229 let new_params = optimizer.step(¶ms, &gradients).expect("step failed");
230
231 assert!((new_params[0] - 0.95).abs() < 1e-10);
236 assert!((new_params[1] - 2.05).abs() < 1e-10);
237 assert!((new_params[2] - 3.0).abs() < 1e-10);
238 assert_eq!(optimizer.get_step_count(), 1);
239 }
240
241 #[test]
242 fn test_reptile_convergence_toward_minimum() {
243 let mut optimizer = ReptileOptimizer::new(0.1_f64)
246 .with_inner_steps(3)
247 .with_epsilon(0.5)
248 .with_inner_lr(0.1);
249
250 let mut params = Array1::from_vec(vec![5.0, -3.0, 2.0]);
251
252 for _ in 0..100 {
253 let gradients = ¶ms * 2.0; params = optimizer.step(¶ms, &gradients).expect("step failed");
255 }
256
257 for &val in params.iter() {
259 assert!(
260 val.abs() < 0.1,
261 "Parameter {val} did not converge to near zero"
262 );
263 }
264 }
265
266 #[test]
267 fn test_reptile_multiple_steps_decrement_count() {
268 let mut optimizer = ReptileOptimizer::new(0.01_f64);
269 let params = Array1::from_vec(vec![1.0, 2.0]);
270 let gradients = Array1::from_vec(vec![0.1, 0.2]);
271
272 for i in 0..5 {
273 let _new_params = optimizer.step(¶ms, &gradients).expect("step failed");
274 assert_eq!(optimizer.get_step_count(), i + 1);
275 }
276 assert_eq!(optimizer.get_step_count(), 5);
277 }
278
279 #[test]
280 fn test_reptile_zero_gradient() {
281 let mut optimizer = ReptileOptimizer::new(0.1_f64).with_inner_steps(5);
282
283 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
284 let gradients = Array1::from_vec(vec![0.0, 0.0, 0.0]);
285
286 let new_params = optimizer.step(¶ms, &gradients).expect("step failed");
287
288 for (p, np) in params.iter().zip(new_params.iter()) {
290 assert!(
291 (*p - *np).abs() < 1e-12,
292 "Params changed with zero gradient"
293 );
294 }
295 }
296
297 #[test]
298 fn test_reptile_inner_steps_zero_clamps_to_one() {
299 let optimizer: ReptileOptimizer<f64> = ReptileOptimizer::new(0.01).with_inner_steps(0);
300 assert_eq!(optimizer.get_inner_steps(), 1);
301 }
302
303 #[test]
304 fn test_reptile_set_learning_rate() {
305 let mut optimizer: ReptileOptimizer<f64> = ReptileOptimizer::new(0.01);
306 Optimizer::<f64, scirs2_core::ndarray::Ix1>::set_learning_rate(&mut optimizer, 0.05);
307 assert!(
308 (Optimizer::<f64, scirs2_core::ndarray::Ix1>::get_learning_rate(&optimizer) - 0.05)
309 .abs()
310 < 1e-10
311 );
312 assert!((optimizer.get_epsilon() - 0.05).abs() < 1e-10);
313 }
314
315 #[test]
316 fn test_reptile_multiple_inner_steps_effect() {
317 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
319 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
320
321 let mut opt_1step = ReptileOptimizer::new(0.1_f64)
322 .with_inner_steps(1)
323 .with_epsilon(1.0)
324 .with_inner_lr(0.1);
325
326 let mut opt_5steps = ReptileOptimizer::new(0.1_f64)
327 .with_inner_steps(5)
328 .with_epsilon(1.0)
329 .with_inner_lr(0.1);
330
331 let result_1 = opt_1step.step(¶ms, &gradients).expect("step failed");
332 let result_5 = opt_5steps.step(¶ms, &gradients).expect("step failed");
333
334 let diff_1: f64 = params
336 .iter()
337 .zip(result_1.iter())
338 .map(|(a, b)| (*a - *b).powi(2))
339 .sum();
340 let diff_5: f64 = params
341 .iter()
342 .zip(result_5.iter())
343 .map(|(a, b)| (*a - *b).powi(2))
344 .sum();
345
346 assert!(
347 diff_5 > diff_1,
348 "More inner steps should cause larger displacement: diff_5={diff_5}, diff_1={diff_1}"
349 );
350 }
351}