Skip to main content

optirs_core/optimizers/
meta_sgd.rs

1// Meta-SGD optimizer with per-parameter learnable learning rates
2//
3// Meta-SGD extends MAML by learning not only the model initialization but also
4// per-parameter learning rates. This allows the model to adapt more effectively
5// to new tasks by using different learning rates for different parameters.
6//
7// Reference: Li, Z., Zhou, F., Chen, F., & Li, H. (2017).
8// "Meta-SGD: Learning to Learn Quickly for Few-Shot Learning"
9
10use 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/// Meta-SGD optimizer with per-parameter learnable learning rates
18///
19/// Implements the Meta-SGD algorithm which learns per-parameter learning rates
20/// alongside the model parameters. Each parameter gets its own adaptive learning
21/// rate that is updated based on the meta-gradient.
22///
23/// # Algorithm
24///
25/// For each step:
26/// 1. Initialize per-parameter learning rates alpha_i to base_lr (if first step)
27/// 2. Compute parameter update: delta_i = alpha_i * grad_i
28/// 3. Update parameters: theta_i = theta_i - delta_i
29/// 4. Update per-parameter LRs: alpha_i = alpha_i - alpha_lr * grad_i * delta_i
30/// 5. Clamp alpha_i to [1e-8, 10.0]
31///
32/// The per-parameter learning rates evolve over time, allowing the optimizer to
33/// automatically discover the best learning rate for each parameter dimension.
34///
35/// # Examples
36///
37/// ```
38/// use scirs2_core::ndarray::Array1;
39/// use optirs_core::optimizers::{MetaSGD, Optimizer};
40///
41/// let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
42/// let gradients = Array1::from_vec(vec![0.1, -0.2, 0.3]);
43///
44/// let mut optimizer = MetaSGD::new(0.01);
45/// let new_params = optimizer.step(&params, &gradients).expect("step failed");
46/// ```
47#[derive(Debug, Clone)]
48pub struct MetaSGD<A: Float + ScalarOperand + Debug> {
49    /// Base learning rate (used to initialize per-parameter LRs)
50    base_lr: A,
51    /// Learning rate for updating per-parameter learning rates (meta-learning rate)
52    alpha_lr: A,
53    /// Number of inner adaptation steps
54    inner_steps: usize,
55    /// Per-parameter learnable learning rates
56    per_param_lr: Option<Array<A, IxDyn>>,
57    /// Count of steps taken
58    step_count: usize,
59}
60
61impl<A: Float + ScalarOperand + Debug> MetaSGD<A> {
62    /// Creates a new Meta-SGD optimizer with the given base learning rate
63    ///
64    /// Defaults:
65    /// - alpha_lr: 0.001
66    /// - inner_steps: 5
67    ///
68    /// # Arguments
69    ///
70    /// * `base_lr` - Base learning rate for initializing per-parameter LRs
71    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    /// Sets the meta-learning rate for updating per-parameter learning rates
82    ///
83    /// # Arguments
84    ///
85    /// * `lr` - Learning rate for the per-parameter LR updates
86    pub fn with_alpha_lr(mut self, lr: A) -> Self {
87        self.alpha_lr = lr;
88        self
89    }
90
91    /// Sets the number of inner adaptation steps
92    ///
93    /// # Arguments
94    ///
95    /// * `n` - Number of inner steps (must be >= 1)
96    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    /// Returns the base learning rate
102    pub fn get_base_lr(&self) -> A {
103        self.base_lr
104    }
105
106    /// Returns the meta-learning rate (alpha_lr)
107    pub fn get_alpha_lr(&self) -> A {
108        self.alpha_lr
109    }
110
111    /// Returns the number of inner adaptation steps
112    pub fn get_inner_steps(&self) -> usize {
113        self.inner_steps
114    }
115
116    /// Returns the number of steps taken so far
117    pub fn get_step_count(&self) -> usize {
118        self.step_count
119    }
120
121    /// Returns a reference to the current per-parameter learning rates, if initialized
122    pub fn get_per_param_lr(&self) -> Option<&Array<A, IxDyn>> {
123        self.per_param_lr.as_ref()
124    }
125
126    /// Resets the per-parameter learning rates (they will be re-initialized on next step)
127    pub fn reset_per_param_lr(&mut self) {
128        self.per_param_lr = None;
129    }
130
131    /// Performs the inner adaptation loop, recomputing the gradient at every step
132    ///
133    /// Meta-SGD's inner loop is defined on *fresh* gradients: after each inner update
134    /// the loss gradient must be re-evaluated at the newly adapted parameters. The
135    /// plain [`Optimizer::step`] entry point cannot do that — it only receives a single
136    /// pre-computed gradient array, so it necessarily reuses the same gradient for
137    /// every inner step (a first-order approximation). Use this method whenever you
138    /// can evaluate gradients on demand.
139    ///
140    /// # Arguments
141    ///
142    /// * `params` - Current parameter values
143    /// * `grad_fn` - Closure returning the loss gradient at the parameters it is given
144    ///
145    /// # Errors
146    ///
147    /// Propagates any error returned by `grad_fn`, and fails if the closure returns a
148    /// gradient whose shape does not match the parameters.
149    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(&params_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            // delta = per_param_lr * grad(theta_k)
192            let delta = &per_param_lr * &gradient_dyn;
193            cumulative_delta = &cumulative_delta + &delta;
194
195            let adapted_dyn = adapted.into_dyn() - &delta;
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        // Meta-gradient uses the gradient evaluated at the final adapted parameters.
207        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    /// Initializes (or re-initializes on shape change) the per-parameter learning rates
219    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    /// Clamp learning rate values to the valid range [min_val, max_val]
233    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    /// Performs one Meta-SGD outer step from a single pre-computed gradient
252    ///
253    /// # Inner-loop caveat
254    ///
255    /// The [`Optimizer`] trait supplies one gradient array per call, so when
256    /// `inner_steps > 1` this method necessarily **reuses the same gradient** for
257    /// every inner adaptation step instead of re-evaluating the loss at each inner
258    /// iterate. That is a deliberate first-order approximation of Meta-SGD, valid
259    /// when the inner steps are small, and it is what the trait contract allows.
260    ///
261    /// Use [`MetaSGD::step_with_closure`] to get the exact algorithm with gradients
262    /// recomputed at every inner step.
263    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        // Step 1: Initialize (or reset on shape change) per-parameter learning rates
283        self.ensure_per_param_lr(&params_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        // Step 2-3: Apply inner adaptation steps using per-parameter learning rates
294        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            // delta = per_param_lr * gradients
299            let delta = &per_param_lr * &gradients_dyn;
300            // Accumulate total parameter change for meta-gradient
301            cumulative_delta = &cumulative_delta + &delta;
302            // Update adapted params
303            adapted_params = &adapted_params - &delta;
304        }
305
306        // Step 4: Update per-parameter learning rates using meta-gradient
307        // The meta-gradient for alpha is: grad * cumulative_delta
308        // This encourages learning rates that reduce the loss
309        let meta_gradient = &gradients_dyn * &cumulative_delta;
310        let mut updated_lr = &per_param_lr - &(&meta_gradient * self.alpha_lr);
311
312        // Step 5: Clamp per-parameter learning rates
313        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        // Convert back to original dimension
319        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        // Reset per-param LRs so they re-initialize with new base_lr
334        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(&params, &gradients).expect("step failed");
371
372        // With inner_steps=1, base_lr=0.1:
373        // delta = per_param_lr * gradients = [0.1*0.5, 0.1*(-0.5), 0.1*0.0] = [0.05, -0.05, 0.0]
374        // new_params = params - delta = [0.95, 2.05, 3.0]
375        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        // Per-param LR should be initialized now
381        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        // First step initializes per-param LRs
394        let _ = optimizer.step(&params, &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        // The parameter with larger gradient (dim 0) should have its LR adjusted more
402        // than the parameter with smaller gradient (dim 1)
403        // meta_gradient = grad * delta = grad * (lr * grad) = lr * grad^2
404        // For dim 0: meta_grad = 0.1 * 1.0^2 = 0.1
405        //   new_lr = 0.1 - 0.01 * 0.1 = 0.099
406        // For dim 1: meta_grad = 0.1 * 0.001^2 = 0.0000001
407        //   new_lr = 0.1 - 0.01 * 0.0000001 ≈ 0.1
408        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        // Optimize f(x) = x^2, gradient = 2x
419        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 = &params * 2.0;
427            params = optimizer.step(&params, &gradients).expect("step failed");
428        }
429
430        // After many steps, params should be close to zero
431        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        // Use very large alpha_lr to force per-param LRs to be clamped
442        let mut optimizer = MetaSGD::new(0.1_f64)
443            .with_alpha_lr(100.0) // Extremely large meta-LR
444            .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        // Run a step - the large alpha_lr should cause LRs to hit clamp bounds
450        let _ = optimizer.step(&params, &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        // All LR values should be within [1e-8, 10.0]
457        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(&params, &gradients).expect("step failed");
473
474        // With zero gradients, params should not change
475        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(&params, &gradients).expect("step failed");
490        assert!(optimizer.get_per_param_lr().is_some());
491
492        // Setting learning rate should reset per-param LRs
493        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(&params, &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(&params, &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    /// Regression test for the reused-gradient inner loop.
534    ///
535    /// `Optimizer::step` can only reuse the single gradient it is handed, so multiple
536    /// inner steps move linearly. `step_with_closure` re-evaluates the gradient at each
537    /// inner iterate, which for f(x) = x^2 must produce a strictly different (and
538    /// smaller in magnitude) total displacement.
539    #[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(&params, &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(&params, |p: &Array1<f64>| Ok(p.mapv(|x| 2.0 * x)))
550            .expect("closure step failed");
551
552        // Reusing g = 2.0 three times: 1 - 3 * 0.1 * 2 = 0.4
553        assert!((reused_out[0] - 0.4).abs() < 1e-12, "got {}", reused_out[0]);
554
555        // Recomputing: x <- x * (1 - 0.2) each step => 0.8^3 = 0.512
556        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    /// The closure API must surface gradient-shape errors instead of panicking.
567    #[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(&params, |_p: &Array1<f64>| Ok(Array1::zeros(3)));
573        assert!(result.is_err());
574    }
575}