1use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand};
37use scirs2_core::numeric::Float;
38use std::fmt::Debug;
39
40use crate::error::{OptimError, Result};
41use crate::optimizers::Optimizer;
42
43#[derive(Debug, Clone, Copy, PartialEq)]
48pub enum MAMLVariant {
49 SecondOrder,
57 FirstOrder,
61 Reptile,
65}
66
67#[derive(Debug, Clone)]
81pub struct TaskBatch<A: Float + ScalarOperand + Debug> {
82 pub initial_params: Array<A, IxDyn>,
84 pub inner_gradients: Vec<Array<A, IxDyn>>,
86 pub final_loss_grad: Array<A, IxDyn>,
88}
89
90pub type InnerAdaptResult<A, D> = Result<(Array<A, D>, Vec<Array<A, D>>)>;
114
115#[derive(Debug, Clone)]
116pub struct MAML<A: Float + ScalarOperand + Debug> {
117 meta_lr: A,
119 inner_lr: A,
121 inner_steps: usize,
123 variant: MAMLVariant,
125 weight_decay: A,
127 meta_params: Option<Array<A, IxDyn>>,
130 step_count: usize,
132}
133
134impl<A: Float + ScalarOperand + Debug> MAML<A> {
135 pub fn new(meta_lr: A) -> Self {
147 let default_inner =
148 A::from(0.01).expect("MAML: failed to convert default inner_lr constant");
149 Self {
150 meta_lr,
151 inner_lr: default_inner,
152 inner_steps: 5,
153 variant: MAMLVariant::FirstOrder,
154 weight_decay: A::zero(),
155 meta_params: None,
156 step_count: 0,
157 }
158 }
159
160 pub fn with_inner_lr(mut self, alpha: A) -> Self {
162 self.inner_lr = alpha;
163 self
164 }
165
166 pub fn with_inner_steps(mut self, k: usize) -> Self {
171 self.inner_steps = if k == 0 { 1 } else { k };
172 self
173 }
174
175 pub fn with_variant(mut self, v: MAMLVariant) -> Self {
178 self.variant = v;
179 self
180 }
181
182 pub fn with_weight_decay(mut self, wd: A) -> Self {
184 self.weight_decay = wd;
185 self
186 }
187
188 pub fn get_meta_lr(&self) -> A {
190 self.meta_lr
191 }
192
193 pub fn get_inner_lr(&self) -> A {
195 self.inner_lr
196 }
197
198 pub fn get_inner_steps(&self) -> usize {
200 self.inner_steps
201 }
202
203 pub fn get_variant(&self) -> MAMLVariant {
205 self.variant
206 }
207
208 pub fn get_weight_decay(&self) -> A {
210 self.weight_decay
211 }
212
213 pub fn get_step_count(&self) -> usize {
215 self.step_count
216 }
217
218 pub fn meta_params(&self) -> Option<&Array<A, IxDyn>> {
221 self.meta_params.as_ref()
222 }
223
224 pub fn reset(&mut self) {
226 self.meta_params = None;
227 self.step_count = 0;
228 }
229
230 pub fn inner_adapt<D: Dimension>(
233 &self,
234 params: &Array<A, D>,
235 gradients: &Array<A, D>,
236 ) -> Result<Array<A, D>> {
237 if params.shape() != gradients.shape() {
238 return Err(OptimError::InvalidParameter(format!(
239 "inner_adapt: params shape {:?} does not match gradients shape {:?}",
240 params.shape(),
241 gradients.shape()
242 )));
243 }
244 Ok(params - &(gradients * self.inner_lr))
245 }
246
247 pub fn inner_adapt_multi_step<D, F>(
255 &self,
256 params: &Array<A, D>,
257 mut loss_grad_fn: F,
258 ) -> InnerAdaptResult<A, D>
259 where
260 D: Dimension,
261 F: FnMut(&Array<A, D>) -> Array<A, D>,
262 {
263 let mut current = params.to_owned();
264 let mut trajectory: Vec<Array<A, D>> = Vec::with_capacity(self.inner_steps);
265 for _ in 0..self.inner_steps {
266 let grad = loss_grad_fn(¤t);
267 if grad.shape() != current.shape() {
268 return Err(OptimError::InvalidParameter(format!(
269 "inner_adapt_multi_step: gradient shape {:?} does not match parameter shape {:?}",
270 grad.shape(),
271 current.shape()
272 )));
273 }
274 current = ¤t - &(&grad * self.inner_lr);
275 trajectory.push(grad);
276 }
277 Ok((current, trajectory))
278 }
279
280 fn task_meta_gradient(&self, task: &TaskBatch<A>) -> Result<Array<A, IxDyn>> {
285 if task.inner_gradients.is_empty() {
286 return Err(OptimError::InvalidParameter(
287 "MAML::task_meta_gradient: inner_gradients must not be empty".to_string(),
288 ));
289 }
290 let theta_shape = task.initial_params.shape();
291 if task.final_loss_grad.shape() != theta_shape {
292 return Err(OptimError::InvalidParameter(format!(
293 "MAML::task_meta_gradient: final_loss_grad shape {:?} does not match initial_params shape {:?}",
294 task.final_loss_grad.shape(),
295 theta_shape
296 )));
297 }
298 for (idx, g) in task.inner_gradients.iter().enumerate() {
299 if g.shape() != theta_shape {
300 return Err(OptimError::InvalidParameter(format!(
301 "MAML::task_meta_gradient: inner_gradients[{}] shape {:?} does not match initial_params shape {:?}",
302 idx,
303 g.shape(),
304 theta_shape
305 )));
306 }
307 }
308
309 match self.variant {
310 MAMLVariant::FirstOrder => Ok(task.final_loss_grad.clone()),
311 MAMLVariant::SecondOrder => {
312 let k = task.inner_gradients.len();
323 if k < 2 {
324 return Ok(task.final_loss_grad.clone());
325 }
326 let g_first = &task.inner_gradients[0];
327 let g_last = &task.inner_gradients[k - 1];
328 let steps_minus_one: A = crate::optimizers::cast_scalar(k - 1)?;
329 let denom = self.inner_lr * steps_minus_one;
330 if denom.abs() <= A::epsilon() {
331 return Ok(task.final_loss_grad.clone());
332 }
333 let hessian_approx = (g_last - g_first) / denom;
334 let correction = &(&hessian_approx * self.inner_lr) * &task.final_loss_grad;
337 Ok(&task.final_loss_grad - &correction)
338 }
339 MAMLVariant::Reptile => {
340 let mut sum_grads = Array::<A, IxDyn>::zeros(task.initial_params.raw_dim());
343 for g in &task.inner_gradients {
344 sum_grads = &sum_grads + g;
345 }
346 let adapted = &task.initial_params - &(&sum_grads * self.inner_lr);
347 let alpha = self.inner_lr;
352 if alpha.abs() <= A::epsilon() {
353 return Err(OptimError::InvalidConfig(
354 "MAML(Reptile): inner_lr must be non-zero".to_string(),
355 ));
356 }
357 Ok((&task.initial_params - &adapted) / alpha)
358 }
359 }
360 }
361
362 pub fn meta_step<D: Dimension>(
372 &mut self,
373 task_batches: &[TaskBatch<A>],
374 ) -> Result<Array<A, IxDyn>> {
375 if task_batches.is_empty() {
376 return Err(OptimError::InvalidParameter(
377 "MAML::meta_step: task_batches must not be empty".to_string(),
378 ));
379 }
380 let ref_shape = task_batches[0].initial_params.shape().to_vec();
381 for (idx, t) in task_batches.iter().enumerate().skip(1) {
382 if t.initial_params.shape() != ref_shape.as_slice() {
383 return Err(OptimError::InvalidParameter(format!(
384 "MAML::meta_step: task_batches[{}].initial_params shape {:?} differs from task_batches[0] shape {:?}",
385 idx,
386 t.initial_params.shape(),
387 ref_shape
388 )));
389 }
390 }
391
392 match self.meta_params.as_ref() {
395 None => {
396 self.meta_params = Some(task_batches[0].initial_params.clone());
397 }
398 Some(stored) => {
399 if stored.shape() != ref_shape.as_slice() {
400 return Err(OptimError::InvalidParameter(format!(
401 "MAML::meta_step: stored meta_params shape {:?} differs from task initial_params shape {:?}",
402 stored.shape(),
403 ref_shape
404 )));
405 }
406 }
407 }
408
409 let mut accumulator = Array::<A, IxDyn>::zeros(IxDyn(&ref_shape));
411 for task in task_batches {
412 let g = self.task_meta_gradient(task)?;
413 accumulator = &accumulator + &g;
414 }
415 let n: A = crate::optimizers::cast_scalar(task_batches.len())?;
416 let mean_meta_grad = &accumulator / n;
417
418 let mut updated = self
420 .meta_params
421 .as_ref()
422 .expect("MAML: meta_params must be initialised at this point")
423 .clone();
424 if self.weight_decay != A::zero() {
425 let decay = self.meta_lr * self.weight_decay;
426 updated = &updated - &(&updated * decay);
427 }
428 updated = &updated - &(&mean_meta_grad * self.meta_lr);
429 self.meta_params = Some(updated.clone());
430 self.step_count += 1;
431
432 let _ = std::marker::PhantomData::<D>;
435 Ok(updated)
436 }
437}
438
439impl<A, D> Optimizer<A, D> for MAML<A>
440where
441 A: Float + ScalarOperand + Debug,
442 D: Dimension,
443{
444 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
448 if params.shape() != gradients.shape() {
449 return Err(OptimError::InvalidParameter(format!(
450 "MAML::step: params shape {:?} does not match gradients shape {:?}",
451 params.shape(),
452 gradients.shape()
453 )));
454 }
455 let mut updated = params - &(gradients * self.meta_lr);
456 if self.weight_decay != A::zero() {
457 let decay = self.meta_lr * self.weight_decay;
458 updated = &updated - &(params * decay);
459 }
460 self.meta_params = Some(updated.to_owned().into_dyn());
463 self.step_count += 1;
464 Ok(updated)
465 }
466
467 fn get_learning_rate(&self) -> A {
468 self.meta_lr
469 }
470
471 fn set_learning_rate(&mut self, learning_rate: A) {
472 self.meta_lr = learning_rate;
473 }
474}
475
476#[cfg(test)]
477mod tests {
478 use super::*;
479 use scirs2_core::ndarray::{Array1, IxDyn};
480
481 fn quadratic_grad(theta: &Array1<f64>, target: &Array1<f64>) -> Array1<f64> {
484 theta - target
485 }
486
487 fn quadratic_loss(theta: &Array1<f64>, target: &Array1<f64>) -> f64 {
488 theta
489 .iter()
490 .zip(target.iter())
491 .map(|(t, tg)| 0.5 * (t - tg).powi(2))
492 .sum()
493 }
494
495 fn make_quadratic_task(
498 theta: &Array1<f64>,
499 target: &Array1<f64>,
500 inner_lr: f64,
501 inner_steps: usize,
502 ) -> TaskBatch<f64> {
503 let mut current = theta.clone();
504 let mut grads: Vec<Array<f64, IxDyn>> = Vec::with_capacity(inner_steps);
505 for _ in 0..inner_steps {
506 let g = quadratic_grad(¤t, target);
507 grads.push(g.clone().into_dyn());
508 current = ¤t - &(&g * inner_lr);
509 }
510 let final_grad = quadratic_grad(¤t, target);
511 TaskBatch {
512 initial_params: theta.clone().into_dyn(),
513 inner_gradients: grads,
514 final_loss_grad: final_grad.into_dyn(),
515 }
516 }
517
518 #[test]
519 fn test_default_config_values() {
520 let opt: MAML<f64> = MAML::new(0.1);
521 assert!((opt.get_meta_lr() - 0.1).abs() < 1e-12);
522 assert!((opt.get_inner_lr() - 0.01).abs() < 1e-12);
523 assert_eq!(opt.get_inner_steps(), 5);
524 assert_eq!(opt.get_variant(), MAMLVariant::FirstOrder);
525 assert!((opt.get_weight_decay() - 0.0).abs() < 1e-12);
526 assert_eq!(opt.get_step_count(), 0);
527 assert!(opt.meta_params().is_none());
528 }
529
530 #[test]
531 fn test_builder_pattern() {
532 let opt: MAML<f64> = MAML::new(0.1)
533 .with_inner_lr(0.05)
534 .with_inner_steps(7)
535 .with_variant(MAMLVariant::SecondOrder)
536 .with_weight_decay(1e-4);
537 assert!((opt.get_inner_lr() - 0.05).abs() < 1e-12);
538 assert_eq!(opt.get_inner_steps(), 7);
539 assert_eq!(opt.get_variant(), MAMLVariant::SecondOrder);
540 assert!((opt.get_weight_decay() - 1e-4).abs() < 1e-12);
541
542 let clamped: MAML<f64> = MAML::new(0.1).with_inner_steps(0);
544 assert_eq!(clamped.get_inner_steps(), 1);
545 }
546
547 #[test]
548 fn test_inner_adapt_basic() {
549 let opt: MAML<f64> = MAML::new(0.1).with_inner_lr(0.1);
552 let theta = Array1::from_vec(vec![1.0, -2.0, 0.5]);
553 let grad = theta.clone();
554 let adapted = opt.inner_adapt(&theta, &grad).expect("inner_adapt failed");
555 let target = Array1::from_vec(vec![0.0, 0.0, 0.0]);
556 let loss_before = quadratic_loss(&theta, &target);
557 let loss_after = quadratic_loss(&adapted, &target);
558 assert!(
559 loss_after < loss_before,
560 "inner_adapt should decrease loss: before={loss_before}, after={loss_after}"
561 );
562 for (a, t) in adapted.iter().zip(theta.iter()) {
564 assert!((a - 0.9 * t).abs() < 1e-12);
565 }
566 }
567
568 #[test]
569 fn test_inner_adapt_multi_step_returns_trajectory() {
570 let opt: MAML<f64> = MAML::new(0.1).with_inner_lr(0.05).with_inner_steps(4);
571 let theta = Array1::from_vec(vec![1.0, -1.0, 2.0]);
572 let target = Array1::from_vec(vec![0.0, 0.0, 0.0]);
573 let target_clone = target.clone();
574 let (adapted, traj) = opt
575 .inner_adapt_multi_step(&theta, move |t| quadratic_grad(t, &target_clone))
576 .expect("inner_adapt_multi_step failed");
577 assert_eq!(traj.len(), opt.get_inner_steps());
578 let mut manual = theta.clone();
580 for _ in 0..opt.get_inner_steps() {
581 let g = quadratic_grad(&manual, &target);
582 manual = &manual - &(&g * opt.get_inner_lr());
583 }
584 for (a, m) in adapted.iter().zip(manual.iter()) {
585 assert!(
586 (a - m).abs() < 1e-12,
587 "adapted={a}, manual={m} differ beyond tolerance"
588 );
589 }
590 }
591
592 #[test]
593 fn test_meta_step_first_order_decreases_meta_loss() {
594 let targets = [
597 Array1::from_vec(vec![1.0, 0.0]),
598 Array1::from_vec(vec![-1.0, 0.0]),
599 Array1::from_vec(vec![0.0, 1.0]),
600 Array1::from_vec(vec![0.0, -1.0]),
601 ];
602 let mut opt: MAML<f64> = MAML::new(0.1)
603 .with_inner_lr(0.05)
604 .with_inner_steps(3)
605 .with_variant(MAMLVariant::FirstOrder);
606 let mut theta = Array1::from_vec(vec![3.0, -3.0]).into_dyn();
607
608 let initial_loss =
610 mean_adapted_loss(&theta, &targets, opt.get_inner_lr(), opt.get_inner_steps());
611
612 for _ in 0..50 {
613 let batches: Vec<TaskBatch<f64>> = targets
614 .iter()
615 .map(|tgt| {
616 let theta_1d = theta
617 .clone()
618 .into_dimensionality::<scirs2_core::ndarray::Ix1>()
619 .expect("test: theta should be 1-D");
620 make_quadratic_task(&theta_1d, tgt, opt.get_inner_lr(), opt.get_inner_steps())
621 })
622 .collect();
623 theta = opt
624 .meta_step::<scirs2_core::ndarray::Ix1>(&batches)
625 .expect("meta_step failed");
626 }
627
628 let final_loss =
629 mean_adapted_loss(&theta, &targets, opt.get_inner_lr(), opt.get_inner_steps());
630 assert!(
631 final_loss < initial_loss,
632 "FOMAML meta-loss did not decrease: initial={initial_loss}, final={final_loss}"
633 );
634 }
635
636 #[test]
637 fn test_meta_step_second_order_decreases_meta_loss() {
638 let targets = [
639 Array1::from_vec(vec![1.0, 0.0]),
640 Array1::from_vec(vec![-1.0, 0.0]),
641 Array1::from_vec(vec![0.0, 1.0]),
642 Array1::from_vec(vec![0.0, -1.0]),
643 ];
644 let mut opt: MAML<f64> = MAML::new(0.1)
645 .with_inner_lr(0.05)
646 .with_inner_steps(3)
647 .with_variant(MAMLVariant::SecondOrder);
648 let mut theta = Array1::from_vec(vec![3.0, -3.0]).into_dyn();
649
650 let initial_loss =
651 mean_adapted_loss(&theta, &targets, opt.get_inner_lr(), opt.get_inner_steps());
652
653 for _ in 0..50 {
654 let batches: Vec<TaskBatch<f64>> = targets
655 .iter()
656 .map(|tgt| {
657 let theta_1d = theta
658 .clone()
659 .into_dimensionality::<scirs2_core::ndarray::Ix1>()
660 .expect("test: theta should be 1-D");
661 make_quadratic_task(&theta_1d, tgt, opt.get_inner_lr(), opt.get_inner_steps())
662 })
663 .collect();
664 theta = opt
665 .meta_step::<scirs2_core::ndarray::Ix1>(&batches)
666 .expect("meta_step failed");
667 }
668
669 let final_loss =
670 mean_adapted_loss(&theta, &targets, opt.get_inner_lr(), opt.get_inner_steps());
671 assert!(
672 final_loss < initial_loss,
673 "SecondOrder MAML meta-loss did not decrease: initial={initial_loss}, final={final_loss}"
674 );
675 }
676
677 #[test]
678 fn test_reptile_variant_uses_difference_form() {
679 let mut opt: MAML<f64> = MAML::new(0.1)
680 .with_inner_lr(0.05)
681 .with_inner_steps(4)
682 .with_variant(MAMLVariant::Reptile);
683 let theta = Array1::from_vec(vec![2.0, -2.0]);
684 let target = Array1::from_vec(vec![0.5, -0.5]);
685 let task = make_quadratic_task(&theta, &target, opt.get_inner_lr(), opt.get_inner_steps());
686
687 let mut adapted = theta.clone();
690 for _ in 0..opt.get_inner_steps() {
691 let g = quadratic_grad(&adapted, &target);
692 adapted = &adapted - &(&g * opt.get_inner_lr());
693 }
694 let expected_meta_grad = (&theta - &adapted) / opt.get_inner_lr();
695
696 let updated = opt
698 .meta_step::<scirs2_core::ndarray::Ix1>(std::slice::from_ref(&task))
699 .expect("meta_step failed");
700 let expected_updated = &theta - &(&expected_meta_grad * opt.get_meta_lr());
701 for (u, e) in updated.iter().zip(expected_updated.iter()) {
702 assert!(
703 (u - e).abs() < 1e-10,
704 "Reptile meta-update mismatch: got {u}, expected {e}"
705 );
706 }
707
708 let direction = &theta - &adapted;
710 for (e, d) in expected_meta_grad.iter().zip(direction.iter()) {
711 assert!(
712 e.signum() == d.signum() || d.abs() < 1e-12,
713 "Reptile direction does not match (initial - adapted)"
714 );
715 }
716 }
717
718 #[test]
719 fn test_fomaml_cheaper_than_secondorder() {
720 let theta = Array1::from_vec(vec![1.0, -1.0, 0.5, -0.5]);
725 let target = Array1::from_vec(vec![0.0, 0.0, 0.0, 0.0]);
726 let task = make_quadratic_task(&theta, &target, 0.05, 5);
727
728 let mut fomaml: MAML<f64> = MAML::new(0.1)
729 .with_inner_lr(0.05)
730 .with_inner_steps(5)
731 .with_variant(MAMLVariant::FirstOrder);
732 let mut secondorder: MAML<f64> = MAML::new(0.1)
733 .with_inner_lr(0.05)
734 .with_inner_steps(5)
735 .with_variant(MAMLVariant::SecondOrder);
736
737 let a = fomaml
738 .meta_step::<scirs2_core::ndarray::Ix1>(std::slice::from_ref(&task))
739 .expect("FOMAML meta_step failed");
740 let b = secondorder
741 .meta_step::<scirs2_core::ndarray::Ix1>(std::slice::from_ref(&task))
742 .expect("SecondOrder meta_step failed");
743
744 for v in a.iter().chain(b.iter()) {
745 assert!(v.is_finite(), "Meta-update produced non-finite value: {v}");
746 }
747 assert_eq!(fomaml.get_step_count(), 1);
748 assert_eq!(secondorder.get_step_count(), 1);
749 }
750
751 #[test]
752 fn test_variant_switching_preserves_meta_params() {
753 let theta = Array1::from_vec(vec![1.0, 2.0, 3.0]);
754 let target = Array1::from_vec(vec![0.0, 0.0, 0.0]);
755 let task = make_quadratic_task(&theta, &target, 0.05, 3);
756
757 let mut opt: MAML<f64> = MAML::new(0.05)
758 .with_inner_lr(0.05)
759 .with_inner_steps(3)
760 .with_variant(MAMLVariant::FirstOrder);
761 let updated = opt
762 .meta_step::<scirs2_core::ndarray::Ix1>(std::slice::from_ref(&task))
763 .expect("meta_step failed");
764
765 opt = opt.with_variant(MAMLVariant::SecondOrder);
767 assert_eq!(opt.get_variant(), MAMLVariant::SecondOrder);
768
769 let stored = opt
771 .meta_params()
772 .expect("meta_params should be retained across variant switch");
773 for (s, u) in stored.iter().zip(updated.iter()) {
774 assert!((s - u).abs() < 1e-12);
775 }
776 }
777
778 #[test]
779 fn test_optimizer_trait_step_is_meta_lr_descent() {
780 let mut opt: MAML<f64> = MAML::new(0.05);
781 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
782 let grads = Array1::from_vec(vec![0.1, -0.2, 0.3]);
783 let updated = opt.step(¶ms, &grads).expect("step failed");
784 for ((p, g), u) in params.iter().zip(grads.iter()).zip(updated.iter()) {
785 let expected = p - 0.05 * g;
786 assert!(
787 (u - expected).abs() < 1e-12,
788 "Optimizer::step produced {u}, expected {expected}"
789 );
790 }
791 assert_eq!(opt.get_step_count(), 1);
792 assert!(opt.meta_params().is_some());
793 assert!(
794 (Optimizer::<f64, scirs2_core::ndarray::Ix1>::get_learning_rate(&opt) - 0.05).abs()
795 < 1e-12
796 );
797 }
798
799 #[test]
800 fn test_zero_gradients_no_change() {
801 let theta = Array1::from_vec(vec![1.5, -0.5, 2.0]);
804 let zero = Array1::from_vec(vec![0.0, 0.0, 0.0]);
805
806 for variant in [
807 MAMLVariant::FirstOrder,
808 MAMLVariant::SecondOrder,
809 MAMLVariant::Reptile,
810 ] {
811 let mut opt: MAML<f64> = MAML::new(0.1)
812 .with_inner_lr(0.05)
813 .with_inner_steps(3)
814 .with_variant(variant);
815 let task = TaskBatch {
816 initial_params: theta.clone().into_dyn(),
817 inner_gradients: vec![zero.clone().into_dyn(); opt.get_inner_steps()],
818 final_loss_grad: zero.clone().into_dyn(),
819 };
820 let updated = opt
821 .meta_step::<scirs2_core::ndarray::Ix1>(std::slice::from_ref(&task))
822 .expect("meta_step failed");
823 for (t, u) in theta.iter().zip(updated.iter()) {
824 assert!(
825 (t - u).abs() < 1e-12,
826 "Variant {:?}: theta should not move with zero gradients (got {} vs {})",
827 variant,
828 t,
829 u
830 );
831 }
832 }
833 }
834
835 #[test]
836 fn test_dimension_mismatch_errors() {
837 let mut opt: MAML<f64> = MAML::new(0.1).with_inner_lr(0.05).with_inner_steps(2);
838 let theta = Array1::from_vec(vec![1.0, 2.0]);
839 let theta_wrong = Array1::from_vec(vec![1.0, 2.0, 3.0]);
840
841 let bad_task = TaskBatch {
843 initial_params: theta.clone().into_dyn(),
844 inner_gradients: vec![Array1::from_vec(vec![0.1, 0.1]).into_dyn()],
845 final_loss_grad: Array1::from_vec(vec![0.1, 0.1, 0.1]).into_dyn(),
846 };
847 let err = opt
848 .meta_step::<scirs2_core::ndarray::Ix1>(std::slice::from_ref(&bad_task))
849 .expect_err("expected InvalidParameter error");
850 assert!(matches!(err, OptimError::InvalidParameter(_)));
851
852 let bad_task2 = TaskBatch {
854 initial_params: theta.clone().into_dyn(),
855 inner_gradients: vec![Array1::from_vec(vec![0.1, 0.1, 0.1]).into_dyn()],
856 final_loss_grad: Array1::from_vec(vec![0.1, 0.1]).into_dyn(),
857 };
858 let err2 = opt
859 .meta_step::<scirs2_core::ndarray::Ix1>(std::slice::from_ref(&bad_task2))
860 .expect_err("expected InvalidParameter error");
861 assert!(matches!(err2, OptimError::InvalidParameter(_)));
862
863 let good_a = make_quadratic_task(
865 &theta,
866 &Array1::from_vec(vec![0.0, 0.0]),
867 opt.get_inner_lr(),
868 opt.get_inner_steps(),
869 );
870 let good_b = make_quadratic_task(
871 &theta_wrong,
872 &Array1::from_vec(vec![0.0, 0.0, 0.0]),
873 opt.get_inner_lr(),
874 opt.get_inner_steps(),
875 );
876 let err3 = opt
877 .meta_step::<scirs2_core::ndarray::Ix1>(&[good_a, good_b])
878 .expect_err("expected InvalidParameter error");
879 assert!(matches!(err3, OptimError::InvalidParameter(_)));
880
881 let mut opt2: MAML<f64> = MAML::new(0.1);
883 let p = Array1::from_vec(vec![1.0, 2.0]);
884 let g = Array1::from_vec(vec![1.0, 2.0, 3.0]);
885 let err4 = opt2.step(&p, &g).expect_err("expected InvalidParameter");
886 assert!(matches!(err4, OptimError::InvalidParameter(_)));
887 }
888
889 #[test]
890 fn test_weight_decay_shrinks_params() {
891 let mut opt: MAML<f64> = MAML::new(0.1)
894 .with_inner_lr(0.05)
895 .with_inner_steps(2)
896 .with_weight_decay(0.5)
897 .with_variant(MAMLVariant::FirstOrder);
898 let theta = Array1::from_vec(vec![1.0, -2.0, 4.0]);
899 let zero = Array1::from_vec(vec![0.0, 0.0, 0.0]);
900 let task = TaskBatch {
901 initial_params: theta.clone().into_dyn(),
902 inner_gradients: vec![zero.clone().into_dyn(); 2],
903 final_loss_grad: zero.into_dyn(),
904 };
905 let updated = opt
906 .meta_step::<scirs2_core::ndarray::Ix1>(std::slice::from_ref(&task))
907 .expect("meta_step failed");
908 let factor = 1.0 - opt.get_meta_lr() * opt.get_weight_decay();
909 for (t, u) in theta.iter().zip(updated.iter()) {
910 assert!(
911 (u - t * factor).abs() < 1e-12,
912 "Weight decay mismatch: got {u}, expected {}",
913 t * factor
914 );
915 assert!(
916 u.abs() < t.abs(),
917 "Weight decay should shrink magnitude: |{u}| !< |{t}|"
918 );
919 }
920
921 let mut opt2: MAML<f64> = MAML::new(0.1).with_weight_decay(0.5);
923 let params = Array1::from_vec(vec![2.0, -4.0]);
924 let zeros = Array1::from_vec(vec![0.0, 0.0]);
925 let step_updated = opt2.step(¶ms, &zeros).expect("step failed");
926 for (p, u) in params.iter().zip(step_updated.iter()) {
927 assert!(
928 (u - p * (1.0 - 0.1 * 0.5)).abs() < 1e-12,
929 "Optimizer::step weight decay mismatch"
930 );
931 }
932 }
933
934 #[test]
935 fn test_reset_clears_meta_params() {
936 let mut opt: MAML<f64> = MAML::new(0.1);
937 let params = Array1::from_vec(vec![1.0, 2.0]);
938 let grads = Array1::from_vec(vec![0.1, 0.2]);
939 let _ = opt.step(¶ms, &grads).expect("step failed");
940 assert!(opt.meta_params().is_some());
941 assert_eq!(opt.get_step_count(), 1);
942
943 opt.reset();
944 assert!(opt.meta_params().is_none());
945 assert_eq!(opt.get_step_count(), 0);
946 }
947
948 fn mean_adapted_loss(
951 theta: &Array<f64, IxDyn>,
952 targets: &[Array1<f64>],
953 inner_lr: f64,
954 inner_steps: usize,
955 ) -> f64 {
956 let theta_1d = theta
957 .clone()
958 .into_dimensionality::<scirs2_core::ndarray::Ix1>()
959 .expect("test: theta should be 1-D");
960 let mut total = 0.0;
961 for tgt in targets {
962 let mut current = theta_1d.clone();
963 for _ in 0..inner_steps {
964 let g = quadratic_grad(¤t, tgt);
965 current = ¤t - &(&g * inner_lr);
966 }
967 total += quadratic_loss(¤t, tgt);
968 }
969 total / (targets.len() as f64)
970 }
971}