sklears_multioutput/regularization/
meta_learning.rs1#![allow(non_snake_case)] use 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#[derive(Debug, Clone)]
44pub struct MetaLearningMultiTask<S = Untrained> {
45 pub(crate) state: S,
46 pub(crate) meta_learning_rate: Float,
48 pub(crate) inner_learning_rate: Float,
50 pub(crate) n_inner_steps: usize,
52 pub(crate) max_iter: usize,
54 pub(crate) tolerance: Float,
56 pub(crate) task_outputs: HashMap<String, usize>,
58 pub(crate) fit_intercept: bool,
60 pub(crate) random_state: Option<u64>,
62}
63
64#[derive(Debug, Clone)]
66pub struct MetaLearningMultiTaskTrained {
67 pub(crate) meta_parameters: Array2<Float>,
69 pub(crate) meta_intercepts: Array1<Float>,
71 pub(crate) task_parameters: HashMap<String, Array2<Float>>,
73 pub(crate) task_intercepts: HashMap<String, Array1<Float>>,
75 pub(crate) n_features: usize,
77 #[allow(dead_code)]
78 pub(crate) task_outputs: HashMap<String, usize>,
80 pub(crate) meta_learning_rate: Float,
82 pub(crate) inner_learning_rate: Float,
83 pub(crate) n_inner_steps: usize,
84 pub(crate) n_iter: usize,
86}
87
88impl MetaLearningMultiTask<Untrained> {
89 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 pub fn meta_learning_rate(mut self, lr: Float) -> Self {
106 self.meta_learning_rate = lr;
107 self
108 }
109
110 pub fn inner_learning_rate(mut self, lr: Float) -> Self {
112 self.inner_learning_rate = lr;
113 self
114 }
115
116 pub fn n_inner_steps(mut self, steps: usize) -> Self {
118 self.n_inner_steps = steps;
119 self
120 }
121
122 pub fn max_iter(mut self, max_iter: usize) -> Self {
124 self.max_iter = max_iter;
125 self
126 }
127
128 pub fn tolerance(mut self, tolerance: Float) -> Self {
130 self.tolerance = tolerance;
131 self
132 }
133
134 pub fn random_state(mut self, seed: u64) -> Self {
136 self.random_state = Some(seed);
137 self
138 }
139
140 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 let mut rng_gen = thread_rng();
185
186 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 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 (task_name, y_task) in y {
213 let mut task_params = meta_parameters.clone();
215 let mut task_intercept = meta_intercepts.clone();
216
217 for _inner_step in 0..self.n_inner_steps {
219 let predictions = x.dot(&task_params);
221 let predictions_with_intercept = &predictions + &task_intercept;
222
223 let residuals = &predictions_with_intercept - y_task;
225
226 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 task_params -= &(&grad_params * self.inner_learning_rate);
232 task_intercept -= &(&grad_intercept * self.inner_learning_rate);
233 }
234
235 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 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 task_parameters.insert(task_name.clone(), task_params);
251 task_intercepts.insert(task_name.clone(), task_intercept);
252 }
253
254 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 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 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 let mut adapted_params = self.state.meta_parameters.clone();
338 let mut adapted_intercept = self.state.meta_intercepts.clone();
339
340 for _step in 0..n_adaptation_steps {
342 let predictions = x.dot(&adapted_params);
344 let predictions_with_intercept = &predictions + &adapted_intercept;
345
346 let residuals = &predictions_with_intercept - y;
348
349 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 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 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 pub fn meta_parameters(&self) -> &Array2<Float> {
370 &self.meta_parameters
371 }
372
373 pub fn meta_intercepts(&self) -> &Array1<Float> {
375 &self.meta_intercepts
376 }
377
378 pub fn task_parameters(&self, task_name: &str) -> Option<&Array2<Float>> {
380 self.task_parameters.get(task_name)
381 }
382
383 pub fn task_intercepts(&self, task_name: &str) -> Option<&Array1<Float>> {
385 self.task_intercepts.get(task_name)
386 }
387
388 pub fn n_iter(&self) -> usize {
390 self.n_iter
391 }
392
393 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}