1use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand};
11use scirs2_core::numeric::Float;
12use std::fmt::Debug;
13
14use crate::error::{OptimError, Result};
15use crate::optimizers::Optimizer;
16
17#[derive(Debug, Clone)]
48pub struct MetaSGD<A: Float + ScalarOperand + Debug> {
49 base_lr: A,
51 alpha_lr: A,
53 inner_steps: usize,
55 per_param_lr: Option<Array<A, IxDyn>>,
57 step_count: usize,
59}
60
61impl<A: Float + ScalarOperand + Debug> MetaSGD<A> {
62 pub fn new(base_lr: A) -> Self {
72 Self {
73 base_lr,
74 alpha_lr: A::from(0.001).expect("MetaSGD: failed to convert alpha_lr constant"),
75 inner_steps: 5,
76 per_param_lr: None,
77 step_count: 0,
78 }
79 }
80
81 pub fn with_alpha_lr(mut self, lr: A) -> Self {
87 self.alpha_lr = lr;
88 self
89 }
90
91 pub fn with_inner_steps(mut self, n: usize) -> Self {
97 self.inner_steps = if n == 0 { 1 } else { n };
98 self
99 }
100
101 pub fn get_base_lr(&self) -> A {
103 self.base_lr
104 }
105
106 pub fn get_alpha_lr(&self) -> A {
108 self.alpha_lr
109 }
110
111 pub fn get_inner_steps(&self) -> usize {
113 self.inner_steps
114 }
115
116 pub fn get_step_count(&self) -> usize {
118 self.step_count
119 }
120
121 pub fn get_per_param_lr(&self) -> Option<&Array<A, IxDyn>> {
123 self.per_param_lr.as_ref()
124 }
125
126 pub fn reset_per_param_lr(&mut self) {
128 self.per_param_lr = None;
129 }
130
131 pub fn step_with_closure<D, F>(
150 &mut self,
151 params: &Array<A, D>,
152 mut grad_fn: F,
153 ) -> Result<Array<A, D>>
154 where
155 D: Dimension,
156 F: FnMut(&Array<A, D>) -> Result<Array<A, D>>,
157 {
158 let min_lr = A::from(1e-8).ok_or_else(|| {
159 OptimError::InvalidConfig("MetaSGD: failed to convert min_lr constant".to_string())
160 })?;
161 let max_lr = A::from(10.0).ok_or_else(|| {
162 OptimError::InvalidConfig("MetaSGD: failed to convert max_lr constant".to_string())
163 })?;
164
165 let params_dyn = params.to_owned().into_dyn();
166 self.ensure_per_param_lr(¶ms_dyn);
167
168 let per_param_lr = self
169 .per_param_lr
170 .as_ref()
171 .ok_or_else(|| {
172 OptimError::InvalidConfig("MetaSGD: per_param_lr not initialized".to_string())
173 })?
174 .clone();
175
176 let mut adapted = params.to_owned();
177 let mut cumulative_delta = Array::<A, IxDyn>::zeros(params_dyn.raw_dim());
178 let mut last_gradient: Option<Array<A, IxDyn>> = None;
179
180 for _ in 0..self.inner_steps {
181 let gradient = grad_fn(&adapted)?;
182 if gradient.shape() != adapted.shape() {
183 return Err(OptimError::DimensionMismatch(format!(
184 "MetaSGD: gradient shape {:?} does not match parameter shape {:?}",
185 gradient.shape(),
186 adapted.shape()
187 )));
188 }
189 let gradient_dyn = gradient.into_dyn();
190
191 let delta = &per_param_lr * &gradient_dyn;
193 cumulative_delta = &cumulative_delta + δ
194
195 let adapted_dyn = adapted.into_dyn() - δ
196 adapted = adapted_dyn.into_dimensionality::<D>().map_err(|e| {
197 OptimError::DimensionMismatch(format!(
198 "MetaSGD: failed to restore parameter dimensionality: {}",
199 e
200 ))
201 })?;
202
203 last_gradient = Some(gradient_dyn);
204 }
205
206 if let Some(final_gradient) = last_gradient {
208 let meta_gradient = &final_gradient * &cumulative_delta;
209 let mut updated_lr = &per_param_lr - &(&meta_gradient * self.alpha_lr);
210 Self::clamp_lr_array(&mut updated_lr, min_lr, max_lr);
211 self.per_param_lr = Some(updated_lr);
212 }
213
214 self.step_count += 1;
215 Ok(adapted)
216 }
217
218 fn ensure_per_param_lr(&mut self, params_dyn: &Array<A, IxDyn>) {
220 let needs_init = match self.per_param_lr.as_ref() {
221 Some(lr) => lr.raw_dim() != params_dyn.raw_dim(),
222 None => true,
223 };
224 if needs_init {
225 self.per_param_lr = Some(Array::<A, IxDyn>::from_elem(
226 params_dyn.raw_dim(),
227 self.base_lr,
228 ));
229 }
230 }
231
232 fn clamp_lr_array(lr_array: &mut Array<A, IxDyn>, min_val: A, max_val: A) {
234 lr_array.mapv_inplace(|v| {
235 if v < min_val {
236 min_val
237 } else if v > max_val {
238 max_val
239 } else {
240 v
241 }
242 });
243 }
244}
245
246impl<A, D> Optimizer<A, D> for MetaSGD<A>
247where
248 A: Float + ScalarOperand + Debug,
249 D: Dimension,
250{
251 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
264 if params.shape() != gradients.shape() {
265 return Err(OptimError::DimensionMismatch(format!(
266 "MetaSGD: gradient shape {:?} does not match parameter shape {:?}",
267 gradients.shape(),
268 params.shape()
269 )));
270 }
271
272 let params_dyn = params.to_owned().into_dyn();
273 let gradients_dyn = gradients.to_owned().into_dyn();
274
275 let min_lr = A::from(1e-8).ok_or_else(|| {
276 OptimError::InvalidConfig("MetaSGD: failed to convert min_lr constant".to_string())
277 })?;
278 let max_lr = A::from(10.0).ok_or_else(|| {
279 OptimError::InvalidConfig("MetaSGD: failed to convert max_lr constant".to_string())
280 })?;
281
282 self.ensure_per_param_lr(¶ms_dyn);
284
285 let per_param_lr = self
286 .per_param_lr
287 .as_ref()
288 .ok_or_else(|| {
289 OptimError::InvalidConfig("MetaSGD: per_param_lr not initialized".to_string())
290 })?
291 .clone();
292
293 let mut adapted_params = params_dyn.clone();
295 let mut cumulative_delta = Array::<A, IxDyn>::zeros(params_dyn.raw_dim());
296
297 for _ in 0..self.inner_steps {
298 let delta = &per_param_lr * &gradients_dyn;
300 cumulative_delta = &cumulative_delta + δ
302 adapted_params = &adapted_params - δ
304 }
305
306 let meta_gradient = &gradients_dyn * &cumulative_delta;
310 let mut updated_lr = &per_param_lr - &(&meta_gradient * self.alpha_lr);
311
312 Self::clamp_lr_array(&mut updated_lr, min_lr, max_lr);
314
315 self.per_param_lr = Some(updated_lr);
316 self.step_count += 1;
317
318 adapted_params.into_dimensionality::<D>().map_err(|e| {
320 OptimError::DimensionMismatch(format!(
321 "MetaSGD: failed to convert back to original dimensionality: {}",
322 e
323 ))
324 })
325 }
326
327 fn get_learning_rate(&self) -> A {
328 self.base_lr
329 }
330
331 fn set_learning_rate(&mut self, learning_rate: A) {
332 self.base_lr = learning_rate;
333 self.per_param_lr = None;
335 }
336}
337
338#[cfg(test)]
339mod tests {
340 use super::*;
341 use scirs2_core::ndarray::Array1;
342
343 #[test]
344 fn test_meta_sgd_basic_creation() {
345 let optimizer: MetaSGD<f64> = MetaSGD::new(0.01);
346 assert!((optimizer.get_base_lr() - 0.01).abs() < 1e-10);
347 assert!((optimizer.get_alpha_lr() - 0.001).abs() < 1e-10);
348 assert_eq!(optimizer.get_inner_steps(), 5);
349 assert_eq!(optimizer.get_step_count(), 0);
350 assert!(optimizer.get_per_param_lr().is_none());
351 }
352
353 #[test]
354 fn test_meta_sgd_builder_pattern() {
355 let optimizer: MetaSGD<f64> = MetaSGD::new(0.01)
356 .with_alpha_lr(0.0001)
357 .with_inner_steps(10);
358
359 assert!((optimizer.get_alpha_lr() - 0.0001).abs() < 1e-10);
360 assert_eq!(optimizer.get_inner_steps(), 10);
361 }
362
363 #[test]
364 fn test_meta_sgd_step_works() {
365 let mut optimizer = MetaSGD::new(0.1_f64).with_inner_steps(1);
366
367 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
368 let gradients = Array1::from_vec(vec![0.5, -0.5, 0.0]);
369
370 let new_params = optimizer.step(¶ms, &gradients).expect("step failed");
371
372 assert!((new_params[0] - 0.95).abs() < 1e-10);
376 assert!((new_params[1] - 2.05).abs() < 1e-10);
377 assert!((new_params[2] - 3.0).abs() < 1e-10);
378 assert_eq!(optimizer.get_step_count(), 1);
379
380 assert!(optimizer.get_per_param_lr().is_some());
382 }
383
384 #[test]
385 fn test_meta_sgd_per_param_lr_adaptation() {
386 let mut optimizer = MetaSGD::new(0.1_f64)
387 .with_alpha_lr(0.01)
388 .with_inner_steps(1);
389
390 let params = Array1::from_vec(vec![1.0, 2.0]);
391 let gradients = Array1::from_vec(vec![1.0, 0.001]);
392
393 let _ = optimizer.step(¶ms, &gradients).expect("step failed");
395
396 let lr_after_first = optimizer
397 .get_per_param_lr()
398 .expect("per_param_lr should exist")
399 .clone();
400
401 let lr_diff_0 = (lr_after_first[0] - 0.1_f64).abs();
409 let lr_diff_1 = (lr_after_first[1] - 0.1_f64).abs();
410 assert!(
411 lr_diff_0 > lr_diff_1,
412 "Larger gradient dimension should have more LR change: diff_0={lr_diff_0}, diff_1={lr_diff_1}"
413 );
414 }
415
416 #[test]
417 fn test_meta_sgd_convergence_toward_minimum() {
418 let mut optimizer = MetaSGD::new(0.05_f64)
420 .with_alpha_lr(0.0001)
421 .with_inner_steps(1);
422
423 let mut params = Array1::from_vec(vec![5.0, -3.0, 2.0]);
424
425 for _ in 0..200 {
426 let gradients = ¶ms * 2.0;
427 params = optimizer.step(¶ms, &gradients).expect("step failed");
428 }
429
430 for &val in params.iter() {
432 assert!(
433 val.abs() < 0.5,
434 "Parameter {val} did not converge to near zero"
435 );
436 }
437 }
438
439 #[test]
440 fn test_meta_sgd_lr_clamping() {
441 let mut optimizer = MetaSGD::new(0.1_f64)
443 .with_alpha_lr(100.0) .with_inner_steps(1);
445
446 let params = Array1::from_vec(vec![1.0, 2.0]);
447 let gradients = Array1::from_vec(vec![1.0, -1.0]);
448
449 let _ = optimizer.step(¶ms, &gradients).expect("step failed");
451
452 let per_param_lr = optimizer
453 .get_per_param_lr()
454 .expect("per_param_lr should exist");
455
456 for &lr in per_param_lr.iter() {
458 assert!(
459 (1e-8..=10.0).contains(&lr),
460 "Per-param LR {lr} is out of clamped range [1e-8, 10.0]"
461 );
462 }
463 }
464
465 #[test]
466 fn test_meta_sgd_zero_gradient() {
467 let mut optimizer = MetaSGD::new(0.1_f64).with_inner_steps(3);
468
469 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
470 let gradients = Array1::from_vec(vec![0.0, 0.0, 0.0]);
471
472 let new_params = optimizer.step(¶ms, &gradients).expect("step failed");
473
474 for (p, np) in params.iter().zip(new_params.iter()) {
476 assert!(
477 (*p - *np).abs() < 1e-12,
478 "Params changed with zero gradient"
479 );
480 }
481 }
482
483 #[test]
484 fn test_meta_sgd_set_learning_rate_resets_per_param() {
485 let mut optimizer = MetaSGD::new(0.1_f64);
486 let params = Array1::from_vec(vec![1.0, 2.0]);
487 let gradients = Array1::from_vec(vec![0.1, 0.2]);
488
489 let _ = optimizer.step(¶ms, &gradients).expect("step failed");
490 assert!(optimizer.get_per_param_lr().is_some());
491
492 Optimizer::<f64, scirs2_core::ndarray::Ix1>::set_learning_rate(&mut optimizer, 0.05);
494 assert!(optimizer.get_per_param_lr().is_none());
495 assert!(
496 (Optimizer::<f64, scirs2_core::ndarray::Ix1>::get_learning_rate(&optimizer) - 0.05)
497 .abs()
498 < 1e-10
499 );
500 }
501
502 #[test]
503 fn test_meta_sgd_inner_steps_zero_clamps_to_one() {
504 let optimizer: MetaSGD<f64> = MetaSGD::new(0.01).with_inner_steps(0);
505 assert_eq!(optimizer.get_inner_steps(), 1);
506 }
507
508 #[test]
509 fn test_meta_sgd_multiple_steps_count() {
510 let mut optimizer = MetaSGD::new(0.01_f64);
511 let params = Array1::from_vec(vec![1.0, 2.0]);
512 let gradients = Array1::from_vec(vec![0.1, 0.2]);
513
514 for i in 0..5 {
515 let _ = optimizer.step(¶ms, &gradients).expect("step failed");
516 assert_eq!(optimizer.get_step_count(), i + 1);
517 }
518 }
519
520 #[test]
521 fn test_meta_sgd_reset_per_param_lr() {
522 let mut optimizer = MetaSGD::new(0.1_f64);
523 let params = Array1::from_vec(vec![1.0]);
524 let gradients = Array1::from_vec(vec![0.1]);
525
526 let _ = optimizer.step(¶ms, &gradients).expect("step failed");
527 assert!(optimizer.get_per_param_lr().is_some());
528
529 optimizer.reset_per_param_lr();
530 assert!(optimizer.get_per_param_lr().is_none());
531 }
532
533 #[test]
540 fn test_meta_sgd_step_with_closure_recomputes_gradients() {
541 let params = Array1::from_vec(vec![1.0f64]);
542
543 let mut reused = MetaSGD::new(0.1_f64).with_alpha_lr(0.0).with_inner_steps(3);
544 let gradients = params.mapv(|x| 2.0 * x);
545 let reused_out = reused.step(¶ms, &gradients).expect("step failed");
546
547 let mut recomputed = MetaSGD::new(0.1_f64).with_alpha_lr(0.0).with_inner_steps(3);
548 let recomputed_out = recomputed
549 .step_with_closure(¶ms, |p: &Array1<f64>| Ok(p.mapv(|x| 2.0 * x)))
550 .expect("closure step failed");
551
552 assert!((reused_out[0] - 0.4).abs() < 1e-12, "got {}", reused_out[0]);
554
555 assert!(
557 (recomputed_out[0] - 0.512).abs() < 1e-12,
558 "got {}",
559 recomputed_out[0]
560 );
561
562 assert!((reused_out[0] - recomputed_out[0]).abs() > 1e-3);
563 assert_eq!(recomputed.get_step_count(), 1);
564 }
565
566 #[test]
568 fn test_meta_sgd_step_with_closure_shape_mismatch() {
569 let mut optimizer = MetaSGD::new(0.1_f64).with_inner_steps(1);
570 let params = Array1::from_vec(vec![1.0f64, 2.0]);
571
572 let result = optimizer.step_with_closure(¶ms, |_p: &Array1<f64>| Ok(Array1::zeros(3)));
573 assert!(result.is_err());
574 }
575}