Skip to main content

sklears_multioutput/regularization/
meta_learning.rs

1//! Meta-Learning for Multi-Task Learning
2//!
3//! This method learns meta-parameters that can quickly adapt to new tasks.
4//! It uses a model-agnostic meta-learning (MAML) approach adapted for multi-task scenarios.
5#![allow(non_snake_case)] // Standard ML notation: X for feature matrices, K for kernels
6
7// Use SciRS2-Core for arrays and random number generation (SciRS2 Policy)
8use scirs2_core::ndarray::{Array1, Array2, ArrayView2, Axis};
9use scirs2_core::random::thread_rng;
10use scirs2_core::random::RandNormal;
11use sklears_core::{
12    error::{Result as SklResult, SklearsError},
13    traits::{Estimator, Fit, Predict, Untrained},
14    types::Float,
15};
16use std::collections::HashMap;
17
18/// Meta-Learning for Multi-Task Learning
19///
20/// This method learns meta-parameters that can quickly adapt to new tasks.
21/// It uses a model-agnostic meta-learning (MAML) approach adapted for multi-task scenarios.
22///
23/// # Examples
24///
25/// ```
26/// use sklears_multioutput::regularization::MetaLearningMultiTask;
27/// use sklears_core::traits::{Predict, Fit};
28/// // Use SciRS2-Core for arrays and random number generation (SciRS2 Policy)
29/// use scirs2_core::ndarray::array;
30/// use std::collections::HashMap;
31///
32/// let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [4.0, 4.0]];
33/// let mut y_tasks = HashMap::new();
34/// y_tasks.insert("task1".to_string(), array![[1.0], [2.0], [1.5], [2.5]]);
35/// y_tasks.insert("task2".to_string(), array![[0.5], [1.0], [0.8], [1.2]]);
36///
37/// let meta_learning = MetaLearningMultiTask::new()
38///     .meta_learning_rate(0.01)
39///     .inner_learning_rate(0.1)
40///     .n_inner_steps(5)
41///     .max_iter(1000);
42/// ```
43#[derive(Debug, Clone)]
44pub struct MetaLearningMultiTask<S = Untrained> {
45    pub(crate) state: S,
46    /// Meta-learning rate for updating meta-parameters
47    pub(crate) meta_learning_rate: Float,
48    /// Inner learning rate for task-specific adaptation
49    pub(crate) inner_learning_rate: Float,
50    /// Number of inner gradient steps per task
51    pub(crate) n_inner_steps: usize,
52    /// Maximum meta-iterations
53    pub(crate) max_iter: usize,
54    /// Convergence tolerance
55    pub(crate) tolerance: Float,
56    /// Task configurations
57    pub(crate) task_outputs: HashMap<String, usize>,
58    /// Include intercept term
59    pub(crate) fit_intercept: bool,
60    /// Random state for reproducible meta-learning
61    pub(crate) random_state: Option<u64>,
62}
63
64/// Trained state for MetaLearningMultiTask
65#[derive(Debug, Clone)]
66pub struct MetaLearningMultiTaskTrained {
67    /// Meta-parameters (initialization for new tasks)
68    pub(crate) meta_parameters: Array2<Float>,
69    /// Meta-intercepts
70    pub(crate) meta_intercepts: Array1<Float>,
71    /// Task-specific adapted parameters
72    pub(crate) task_parameters: HashMap<String, Array2<Float>>,
73    /// Task-specific adapted intercepts
74    pub(crate) task_intercepts: HashMap<String, Array1<Float>>,
75    /// Number of input features
76    pub(crate) n_features: usize,
77    #[allow(dead_code)]
78    /// Task configurations
79    pub(crate) task_outputs: HashMap<String, usize>,
80    /// Training parameters
81    pub(crate) meta_learning_rate: Float,
82    pub(crate) inner_learning_rate: Float,
83    pub(crate) n_inner_steps: usize,
84    /// Training iterations performed
85    pub(crate) n_iter: usize,
86}
87
88impl MetaLearningMultiTask<Untrained> {
89    /// Create a new MetaLearningMultiTask instance
90    pub fn new() -> Self {
91        Self {
92            state: Untrained,
93            meta_learning_rate: 0.01,
94            inner_learning_rate: 0.1,
95            n_inner_steps: 5,
96            max_iter: 1000,
97            tolerance: 1e-4,
98            task_outputs: HashMap::new(),
99            fit_intercept: true,
100            random_state: None,
101        }
102    }
103
104    /// Set meta-learning rate
105    pub fn meta_learning_rate(mut self, lr: Float) -> Self {
106        self.meta_learning_rate = lr;
107        self
108    }
109
110    /// Set inner learning rate
111    pub fn inner_learning_rate(mut self, lr: Float) -> Self {
112        self.inner_learning_rate = lr;
113        self
114    }
115
116    /// Set number of inner gradient steps
117    pub fn n_inner_steps(mut self, steps: usize) -> Self {
118        self.n_inner_steps = steps;
119        self
120    }
121
122    /// Set maximum iterations
123    pub fn max_iter(mut self, max_iter: usize) -> Self {
124        self.max_iter = max_iter;
125        self
126    }
127
128    /// Set tolerance
129    pub fn tolerance(mut self, tolerance: Float) -> Self {
130        self.tolerance = tolerance;
131        self
132    }
133
134    /// Set random state
135    pub fn random_state(mut self, seed: u64) -> Self {
136        self.random_state = Some(seed);
137        self
138    }
139
140    /// Set task outputs
141    pub fn task_outputs(mut self, outputs: &[(&str, usize)]) -> Self {
142        self.task_outputs = outputs
143            .iter()
144            .map(|(name, size)| (name.to_string(), *size))
145            .collect();
146        self
147    }
148}
149
150impl Default for MetaLearningMultiTask<Untrained> {
151    fn default() -> Self {
152        Self::new()
153    }
154}
155
156impl Estimator for MetaLearningMultiTask<Untrained> {
157    type Config = ();
158    type Error = SklearsError;
159    type Float = Float;
160
161    fn config(&self) -> &Self::Config {
162        &()
163    }
164}
165
166impl Fit<ArrayView2<'_, Float>, HashMap<String, Array2<Float>>>
167    for MetaLearningMultiTask<Untrained>
168{
169    type Fitted = MetaLearningMultiTask<MetaLearningMultiTaskTrained>;
170
171    fn fit(
172        self,
173        X: &ArrayView2<'_, Float>,
174        y: &HashMap<String, Array2<Float>>,
175    ) -> SklResult<Self::Fitted> {
176        let x = X.to_owned();
177        let (n_samples, n_features) = x.dim();
178
179        if n_samples == 0 || n_features == 0 {
180            return Err(SklearsError::InvalidInput("Empty input data".to_string()));
181        }
182
183        // Initialize meta-parameters
184        let mut rng_gen = thread_rng();
185
186        // Use first task to determine output size for meta-parameters
187        let first_task_outputs = y.values().next().expect("operation should succeed").ncols();
188        let mut meta_parameters = Array2::<Float>::zeros((n_features, first_task_outputs));
189        let normal_dist = RandNormal::new(0.0, 0.1).expect("operation should succeed");
190        for i in 0..n_features {
191            for j in 0..first_task_outputs {
192                meta_parameters[[i, j]] = rng_gen.sample(normal_dist);
193            }
194        }
195        let mut meta_intercepts = Array1::<Float>::zeros(first_task_outputs);
196
197        let _task_names: Vec<String> = y.keys().cloned().collect();
198        let mut task_parameters: HashMap<String, Array2<Float>> = HashMap::new();
199        let mut task_intercepts: HashMap<String, Array1<Float>> = HashMap::new();
200
201        // Meta-learning loop
202        let mut prev_loss = Float::INFINITY;
203        let mut n_iter = 0;
204
205        for iteration in 0..self.max_iter {
206            let mut total_meta_loss = 0.0;
207            let mut meta_grad_sum: Array2<Float> = Array2::<Float>::zeros(meta_parameters.dim());
208            let mut meta_intercept_grad_sum: Array1<Float> =
209                Array1::<Float>::zeros(meta_intercepts.len());
210
211            // For each task, perform inner loop adaptation
212            for (task_name, y_task) in y {
213                // Initialize task parameters from meta-parameters
214                let mut task_params = meta_parameters.clone();
215                let mut task_intercept = meta_intercepts.clone();
216
217                // Inner loop: adapt to specific task
218                for _inner_step in 0..self.n_inner_steps {
219                    // Compute predictions
220                    let predictions = x.dot(&task_params);
221                    let predictions_with_intercept = &predictions + &task_intercept;
222
223                    // Compute residuals
224                    let residuals = &predictions_with_intercept - y_task;
225
226                    // Compute gradients
227                    let grad_params = x.t().dot(&residuals) / (n_samples as Float);
228                    let grad_intercept = residuals.sum_axis(Axis(0)) / (n_samples as Float);
229
230                    // Update task-specific parameters
231                    task_params -= &(&grad_params * self.inner_learning_rate);
232                    task_intercept -= &(&grad_intercept * self.inner_learning_rate);
233                }
234
235                // Compute final loss for this task
236                let final_predictions = x.dot(&task_params);
237                let final_predictions_with_intercept = &final_predictions + &task_intercept;
238                let final_residuals = &final_predictions_with_intercept - y_task;
239                let task_loss = final_residuals.mapv(|x| x * x).sum();
240                total_meta_loss += task_loss;
241
242                // Compute meta-gradients (how changes in meta-parameters affect final loss)
243                let meta_grad_params = x.t().dot(&final_residuals) / (n_samples as Float);
244                let meta_grad_intercept = final_residuals.sum_axis(Axis(0)) / (n_samples as Float);
245
246                meta_grad_sum = meta_grad_sum + meta_grad_params;
247                meta_intercept_grad_sum = meta_intercept_grad_sum + meta_grad_intercept;
248
249                // Store adapted parameters
250                task_parameters.insert(task_name.clone(), task_params);
251                task_intercepts.insert(task_name.clone(), task_intercept);
252            }
253
254            // Update meta-parameters
255            let n_tasks = y.len() as Float;
256            meta_parameters -= &(&(meta_grad_sum / n_tasks) * self.meta_learning_rate);
257            meta_intercepts -= &(&(meta_intercept_grad_sum / n_tasks) * self.meta_learning_rate);
258
259            // Check convergence
260            if (prev_loss - total_meta_loss).abs() < self.tolerance {
261                n_iter = iteration + 1;
262                break;
263            }
264            prev_loss = total_meta_loss;
265            n_iter = iteration + 1;
266        }
267
268        Ok(MetaLearningMultiTask {
269            state: MetaLearningMultiTaskTrained {
270                meta_parameters,
271                meta_intercepts,
272                task_parameters,
273                task_intercepts,
274                n_features,
275                task_outputs: self.task_outputs.clone(),
276                meta_learning_rate: self.meta_learning_rate,
277                inner_learning_rate: self.inner_learning_rate,
278                n_inner_steps: self.n_inner_steps,
279                n_iter,
280            },
281            meta_learning_rate: self.meta_learning_rate,
282            inner_learning_rate: self.inner_learning_rate,
283            n_inner_steps: self.n_inner_steps,
284            max_iter: self.max_iter,
285            tolerance: self.tolerance,
286            task_outputs: self.task_outputs,
287            fit_intercept: self.fit_intercept,
288            random_state: self.random_state,
289        })
290    }
291}
292
293impl Predict<ArrayView2<'_, Float>, HashMap<String, Array2<Float>>>
294    for MetaLearningMultiTask<MetaLearningMultiTaskTrained>
295{
296    fn predict(&self, X: &ArrayView2<'_, Float>) -> SklResult<HashMap<String, Array2<Float>>> {
297        let x = X.to_owned();
298        let (_n_samples, n_features) = x.dim();
299
300        if n_features != self.state.n_features {
301            return Err(SklearsError::InvalidInput(
302                "Number of features doesn't match training data".to_string(),
303            ));
304        }
305
306        let mut predictions = HashMap::new();
307
308        for (task_name, coef) in &self.state.task_parameters {
309            let task_predictions = x.dot(coef);
310            let intercept = &self.state.task_intercepts[task_name];
311            let final_predictions = &task_predictions + intercept;
312            predictions.insert(task_name.clone(), final_predictions);
313        }
314
315        Ok(predictions)
316    }
317}
318
319impl MetaLearningMultiTask<MetaLearningMultiTaskTrained> {
320    /// Adapt meta-parameters to a new task with few examples
321    pub fn adapt_to_new_task(
322        &self,
323        X: &ArrayView2<Float>,
324        y: &Array2<Float>,
325        n_adaptation_steps: usize,
326    ) -> SklResult<(Array2<Float>, Array1<Float>)> {
327        let x = X.to_owned();
328        let (n_samples, n_features) = x.dim();
329
330        if n_features != self.state.n_features {
331            return Err(SklearsError::InvalidInput(
332                "Number of features doesn't match training data".to_string(),
333            ));
334        }
335
336        // Start with meta-parameters
337        let mut adapted_params = self.state.meta_parameters.clone();
338        let mut adapted_intercept = self.state.meta_intercepts.clone();
339
340        // Perform adaptation steps
341        for _step in 0..n_adaptation_steps {
342            // Compute predictions
343            let predictions = x.dot(&adapted_params);
344            let predictions_with_intercept = &predictions + &adapted_intercept;
345
346            // Compute residuals
347            let residuals = &predictions_with_intercept - y;
348
349            // Compute gradients
350            let grad_params = x.t().dot(&residuals) / (n_samples as Float);
351            let grad_intercept = residuals.sum_axis(Axis(0)) / (n_samples as Float);
352
353            // Update parameters
354            adapted_params -= &(&grad_params * self.state.inner_learning_rate);
355            adapted_intercept -= &(&grad_intercept * self.state.inner_learning_rate);
356        }
357
358        Ok((adapted_params, adapted_intercept))
359    }
360
361    /// Get meta-parameters for initialization of new tasks
362    pub fn get_meta_parameters(&self) -> (&Array2<Float>, &Array1<Float>) {
363        (&self.state.meta_parameters, &self.state.meta_intercepts)
364    }
365}
366
367impl MetaLearningMultiTaskTrained {
368    /// Get meta-parameters
369    pub fn meta_parameters(&self) -> &Array2<Float> {
370        &self.meta_parameters
371    }
372
373    /// Get meta-intercepts
374    pub fn meta_intercepts(&self) -> &Array1<Float> {
375        &self.meta_intercepts
376    }
377
378    /// Get task-specific parameters
379    pub fn task_parameters(&self, task_name: &str) -> Option<&Array2<Float>> {
380        self.task_parameters.get(task_name)
381    }
382
383    /// Get task-specific intercepts
384    pub fn task_intercepts(&self, task_name: &str) -> Option<&Array1<Float>> {
385        self.task_intercepts.get(task_name)
386    }
387
388    /// Get number of iterations performed
389    pub fn n_iter(&self) -> usize {
390        self.n_iter
391    }
392
393    /// Get meta-learning parameters
394    pub fn meta_learning_config(&self) -> (Float, Float, usize) {
395        (
396            self.meta_learning_rate,
397            self.inner_learning_rate,
398            self.n_inner_steps,
399        )
400    }
401}