Skip to main content

optirs_core/optimizers/
maml.rs

1// Model-Agnostic Meta-Learning (MAML) optimizer
2//
3// Implements MAML (Finn et al., 2017): "Model-Agnostic Meta-Learning for Fast
4// Adaptation of Deep Networks". MAML learns an initialization `theta` such that
5// a small number of gradient-descent steps on a new task quickly reach high
6// performance.
7//
8// Algorithm (canonical, K inner steps, batch of tasks T_i):
9//
10//   Inner (per task T_i):
11//       theta_i^{(0)} = theta
12//       for k = 0..K-1:
13//           g_k       = grad_theta L_{T_i}(theta_i^{(k)})
14//           theta_i^{(k+1)} = theta_i^{(k)} - alpha * g_k
15//
16//   Outer (meta-update):
17//       theta <- theta - beta * (1/N) * sum_i grad_theta L_{T_i}(theta_i^{(K)})
18//
19// The outer gradient passes *through* the inner updates, producing a second-
20// order term involving the Hessian of L:
21//
22//       grad_theta L_{T_i}(theta_i^{(K)}) = J_i^T * final_loss_grad_i,
23//       where J_i = prod_{k=0}^{K-1} (I - alpha * H(theta_i^{(k)})).
24//
25// FOMAML (First-Order MAML, Finn et al., 2017) drops the Hessian term, using
26// J_i ~= I, i.e. meta_grad_i := final_loss_grad_i. Reptile (Nichol et al.,
27// 2018) approximates the meta-gradient as (theta - theta_i^{(K)}) / alpha,
28// which empirically yields a closely related update direction.
29//
30// References:
31//   Finn, C., Abbeel, P., Levine, S. (2017). "Model-Agnostic Meta-Learning for
32//     Fast Adaptation of Deep Networks", ICML.
33//   Nichol, A., Achiam, J., Schulman, J. (2018). "On First-Order Meta-Learning
34//     Algorithms", arXiv:1803.02999.
35
36use 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/// Variant of MAML to use.
44///
45/// The variants share the same inner-loop adaptation procedure but differ in
46/// how the outer (meta-) gradient is computed.
47#[derive(Debug, Clone, Copy, PartialEq)]
48pub enum MAMLVariant {
49    /// Full MAML with a second-order term. Because OptiRS does not maintain a
50    /// computation graph, the Hessian-vector product is approximated by a
51    /// finite difference of inner gradients along the adaptation trajectory.
52    /// This yields a meta-gradient of the form
53    /// `final_loss_grad - alpha * (g_K - g_0) / K (elementwise) * final_loss_grad`,
54    /// which captures the leading second-order correction without ever
55    /// instantiating the Hessian.
56    SecondOrder,
57    /// First-Order MAML (FOMAML): drops the Hessian term entirely. The
58    /// meta-gradient is simply `final_loss_grad`. Cheaper than SecondOrder
59    /// and usually competitive in practice.
60    FirstOrder,
61    /// Reptile-style update: the meta-gradient is
62    /// `(initial_params - adapted_params) / alpha`. No second-order math is
63    /// required and the resulting direction is empirically similar to FOMAML.
64    Reptile,
65}
66
67/// Per-task data used by the meta-update.
68///
69/// For each task `T_i` the caller supplies:
70/// * `initial_params` – the meta-parameters `theta` at the start of the inner
71///   loop (`theta_i^{(0)}`).
72/// * `inner_gradients` – the gradients `g_0, g_1, ..., g_{K-1}` evaluated at
73///   each inner iterate `theta_i^{(0)}, ..., theta_i^{(K-1)}`. The length of
74///   this vector must be at least 1 and matches the number of inner steps.
75/// * `final_loss_grad` – the gradient of the meta-loss evaluated at the
76///   *adapted* parameters `theta_i^{(K)}` (sometimes denoted `∂L/∂theta'`).
77///
78/// All arrays must share the same shape; mismatches surface as
79/// `OptimError::InvalidParameter`.
80#[derive(Debug, Clone)]
81pub struct TaskBatch<A: Float + ScalarOperand + Debug> {
82    /// Parameters at the start of the inner loop (`theta_i^{(0)}`).
83    pub initial_params: Array<A, IxDyn>,
84    /// Sequence of gradients along the inner trajectory.
85    pub inner_gradients: Vec<Array<A, IxDyn>>,
86    /// Gradient of the meta-loss at the adapted parameters.
87    pub final_loss_grad: Array<A, IxDyn>,
88}
89
90/// MAML optimizer.
91///
92/// Maintains the meta-parameters `theta` and implements inner-loop adaptation
93/// plus outer-loop meta-updates. The struct also implements [`Optimizer`] so
94/// it can be used as a drop-in plain-SGD optimizer (with learning rate equal
95/// to the meta-learning rate `beta`).
96///
97/// # Examples
98///
99/// ```
100/// use scirs2_core::ndarray::Array1;
101/// use optirs_core::optimizers::{MAML, MAMLVariant, Optimizer};
102///
103/// // Use MAML as a regular optimizer (SGD with meta_lr as the step size).
104/// let params = Array1::from_vec(vec![1.0_f64, 2.0, 3.0]);
105/// let grads = Array1::from_vec(vec![0.1, 0.2, 0.3]);
106/// let mut opt = MAML::new(0.05).with_variant(MAMLVariant::FirstOrder);
107/// let updated = opt.step(&params, &grads).expect("step failed");
108/// assert!((updated[0] - (1.0 - 0.05 * 0.1)).abs() < 1e-12);
109/// ```
110/// Result of a multi-step inner adaptation: the final adapted parameters
111/// together with the per-step gradient trajectory (see
112/// [`MAML::inner_adapt_multi_step`]).
113pub 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    /// Outer / meta learning rate `beta`.
118    meta_lr: A,
119    /// Inner learning rate `alpha`.
120    inner_lr: A,
121    /// Number of inner-loop adaptation steps `K`.
122    inner_steps: usize,
123    /// Which MAML variant to use for the outer gradient.
124    variant: MAMLVariant,
125    /// L2 weight decay applied to the meta-parameters during the outer step.
126    weight_decay: A,
127    /// Current meta-parameters `theta`, lazily initialised on the first call
128    /// to [`MAML::meta_step`] (or to the `Optimizer::step` shortcut).
129    meta_params: Option<Array<A, IxDyn>>,
130    /// Number of outer (meta) steps applied so far.
131    step_count: usize,
132}
133
134impl<A: Float + ScalarOperand + Debug> MAML<A> {
135    /// Creates a new MAML optimizer with the given meta-learning rate `beta`.
136    ///
137    /// Defaults:
138    /// * `inner_lr` (`alpha`): `0.01`
139    /// * `inner_steps` (`K`): `5`
140    /// * `variant`: [`MAMLVariant::FirstOrder`]
141    /// * `weight_decay`: `0`
142    ///
143    /// # Arguments
144    ///
145    /// * `meta_lr` – outer learning rate `beta` used for the meta-update.
146    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    /// Sets the inner-loop learning rate `alpha`.
161    pub fn with_inner_lr(mut self, alpha: A) -> Self {
162        self.inner_lr = alpha;
163        self
164    }
165
166    /// Sets the number of inner-loop adaptation steps `K`.
167    ///
168    /// Zero is interpreted as one step (the inner loop must take at least one
169    /// step to produce a meta-gradient).
170    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    /// Selects the MAML variant ([`MAMLVariant::SecondOrder`],
176    /// [`MAMLVariant::FirstOrder`] or [`MAMLVariant::Reptile`]).
177    pub fn with_variant(mut self, v: MAMLVariant) -> Self {
178        self.variant = v;
179        self
180    }
181
182    /// Sets the L2 weight-decay coefficient applied during the outer update.
183    pub fn with_weight_decay(mut self, wd: A) -> Self {
184        self.weight_decay = wd;
185        self
186    }
187
188    /// Returns the meta-learning rate `beta`.
189    pub fn get_meta_lr(&self) -> A {
190        self.meta_lr
191    }
192
193    /// Returns the inner-loop learning rate `alpha`.
194    pub fn get_inner_lr(&self) -> A {
195        self.inner_lr
196    }
197
198    /// Returns the number of inner-loop steps `K`.
199    pub fn get_inner_steps(&self) -> usize {
200        self.inner_steps
201    }
202
203    /// Returns the active MAML variant.
204    pub fn get_variant(&self) -> MAMLVariant {
205        self.variant
206    }
207
208    /// Returns the configured weight-decay coefficient.
209    pub fn get_weight_decay(&self) -> A {
210        self.weight_decay
211    }
212
213    /// Returns the number of outer meta-steps applied so far.
214    pub fn get_step_count(&self) -> usize {
215        self.step_count
216    }
217
218    /// Returns a reference to the current meta-parameters, if they have been
219    /// initialised by a prior call to [`MAML::meta_step`] or [`Optimizer::step`].
220    pub fn meta_params(&self) -> Option<&Array<A, IxDyn>> {
221        self.meta_params.as_ref()
222    }
223
224    /// Clears the stored meta-parameters and resets the step counter.
225    pub fn reset(&mut self) {
226        self.meta_params = None;
227        self.step_count = 0;
228    }
229
230    /// Performs a single inner-loop adaptation step:
231    /// `params' = params - inner_lr * gradients`.
232    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    /// Performs `inner_steps` adaptation steps starting from `params`, using
248    /// `loss_grad_fn` to compute the gradient at each iterate.
249    ///
250    /// Returns the final adapted parameters together with the trajectory of
251    /// gradients evaluated at iterates `theta^{(0)}, theta^{(1)}, ...,
252    /// theta^{(K-1)}`. The returned vector therefore has length
253    /// `inner_steps`, suitable for direct use in a [`TaskBatch`].
254    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(&current);
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 = &current - &(&grad * self.inner_lr);
275            trajectory.push(grad);
276        }
277        Ok((current, trajectory))
278    }
279
280    /// Computes the meta-gradient contribution for a single task batch.
281    ///
282    /// The returned array always has dynamic dimensionality and is suitable
283    /// for accumulation across the task batch.
284    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                // Hessian-free finite-difference approximation along the inner
313                // trajectory. With K inner gradients g_0, ..., g_{K-1} we use
314                //
315                //     H * v  ~=  ((g_{K-1} - g_0) / (alpha * (K - 1))) (elementwise) * v
316                //
317                // which is the leading term of a first-order Taylor expansion
318                // of the gradient field along the inner step direction. For
319                // K = 1 there is no trajectory information to extract a
320                // Hessian estimate from, so the SecondOrder variant falls
321                // back to FOMAML.
322                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                // meta_grad = (I - alpha * H)^T * final_loss_grad
335                //          = final_loss_grad - alpha * H (elementwise) * final_loss_grad
336                let correction = &(&hessian_approx * self.inner_lr) * &task.final_loss_grad;
337                Ok(&task.final_loss_grad - &correction)
338            }
339            MAMLVariant::Reptile => {
340                // Reconstruct the adapted parameters from the trajectory:
341                //   theta_K = theta_0 - alpha * sum_k g_k
342                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                // meta_grad = (theta_0 - theta_K) / alpha = sum_k g_k. We
348                // intentionally compute via the difference form below to keep
349                // the semantics of "Reptile uses (initial - adapted) / alpha"
350                // explicit in the code.
351                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    /// Performs one outer (meta-) update step over a batch of tasks.
363    ///
364    /// The meta-gradient is averaged over all task batches, optionally
365    /// augmented with L2 weight decay, and applied as
366    /// `theta <- theta - meta_lr * mean_meta_grad`. Returns the updated
367    /// meta-parameters.
368    ///
369    /// All task batches must share the same parameter shape and, when the
370    /// optimizer already holds meta-parameters, must also match that shape.
371    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        // Initialise (or validate) the stored meta-parameters from the first
393        // task's initial parameters.
394        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        // Average meta-gradient across tasks.
410        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        // Outer update with optional decoupled weight decay (AdamW-style).
419        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        // Sanity check: the trait signature uses generic `D`; we want to
433        // ensure dimensionality still matches the caller-provided arrays.
434        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    /// Drop-in plain-SGD step using the meta-learning rate. Useful when
445    /// embedding MAML inside a standard supervised training loop that has not
446    /// (yet) been refactored to use [`MAML::meta_step`].
447    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        // Cache the parameters as meta_params so subsequent meta_step calls
461        // operate on a consistent state.
462        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    /// Helper: quadratic task `L_i(theta) = 0.5 * (theta - target_i)^2`.
482    /// Gradient is `theta - target_i`.
483    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    /// Helper: build a TaskBatch from a quadratic task by running the inner
496    /// loop manually.
497    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(&current, target);
507            grads.push(g.clone().into_dyn());
508            current = &current - &(&g * inner_lr);
509        }
510        let final_grad = quadratic_grad(&current, 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        // Zero inner steps must clamp up to 1.
543        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        // Quadratic loss with target = 0: gradient = theta. A single inner
550        // step with alpha=0.1 must reduce |theta|, hence the loss.
551        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        // The exact arithmetic: adapted = theta - 0.1 * theta = 0.9 * theta.
563        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        // Adapted params must equal what a manual loop would produce.
579        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        // Distribution of quadratic tasks with random-ish targets. The meta
595        // optimum lies at the mean of the task targets.
596        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        // Initial mean adapted loss.
609        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        // The Reptile meta-gradient is `(theta_0 - theta_K) / alpha`. Build
688        // the same quantity by hand from the task.
689        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        // Apply one meta-step and compare against theta - meta_lr * expected.
697        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        // Direction check: meta-gradient sign must match (initial - adapted).
709        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        // Verifies both variants run end-to-end on the same task batch and
721        // produce finite outputs. The "cheaper" property is structural —
722        // FOMAML executes a single vector copy whereas SecondOrder also
723        // computes a finite-difference Hessian-vector product.
724        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        // Switch variant via the builder-style setter on a mutable reference.
766        opt = opt.with_variant(MAMLVariant::SecondOrder);
767        assert_eq!(opt.get_variant(), MAMLVariant::SecondOrder);
768
769        // Meta-params survived the variant change.
770        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(&params, &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        // When every per-task gradient is zero the meta-gradient is zero and
802        // the meta-parameters must not move.
803        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        // final_loss_grad shape mismatch
842        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        // inner_gradients shape mismatch
853        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        // Inconsistent shapes across task batches.
864        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        // Optimizer::step shape mismatch.
882        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        // Apply meta_step with zero meta-gradient but non-zero weight decay:
892        // updated = theta - meta_lr * wd * theta = (1 - meta_lr * wd) * theta.
893        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        // Also verify via the Optimizer::step path.
922        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(&params, &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(&params, &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    /// Mean loss across tasks after applying `inner_steps` inner-loop steps
949    /// from `theta` toward each task's target. Used by the convergence tests.
950    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(&current, tgt);
965                current = &current - &(&g * inner_lr);
966            }
967            total += quadratic_loss(&current, tgt);
968        }
969        total / (targets.len() as f64)
970    }
971}