use scirs2_core::ndarray_ext::{Array1, Array2, ArrayView1, ArrayView2};
use sklears_core::{
error::{Result as SklResult, SklearsError},
traits::{Estimator, Fit, Predict, PredictProba, Untrained},
types::Float,
};
#[derive(Debug, Clone)]
pub struct MAML<S = Untrained> {
state: S,
inner_lr: f64,
outer_lr: f64,
inner_steps: usize,
n_episodes: usize,
hidden_layers: Vec<usize>,
}
impl MAML<Untrained> {
pub fn new() -> Self {
Self {
state: Untrained,
inner_lr: 0.01,
outer_lr: 0.001,
inner_steps: 5,
n_episodes: 100,
hidden_layers: vec![64, 32],
}
}
pub fn inner_lr(mut self, inner_lr: f64) -> Self {
self.inner_lr = inner_lr;
self
}
pub fn outer_lr(mut self, outer_lr: f64) -> Self {
self.outer_lr = outer_lr;
self
}
pub fn inner_steps(mut self, inner_steps: usize) -> Self {
self.inner_steps = inner_steps;
self
}
pub fn n_episodes(mut self, n_episodes: usize) -> Self {
self.n_episodes = n_episodes;
self
}
pub fn hidden_layers(mut self, hidden_layers: Vec<usize>) -> Self {
self.hidden_layers = hidden_layers;
self
}
}
impl Default for MAML<Untrained> {
fn default() -> Self {
Self::new()
}
}
impl Estimator for MAML<Untrained> {
type Config = ();
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&()
}
}
impl Fit<ArrayView2<'_, Float>, ArrayView1<'_, i32>> for MAML<Untrained> {
type Fitted = MAML<MAMLTrained>;
#[allow(non_snake_case)]
fn fit(self, X: &ArrayView2<'_, Float>, y: &ArrayView1<'_, i32>) -> SklResult<Self::Fitted> {
let X = X.to_owned();
let y = y.to_owned();
let (n_samples, n_features) = X.dim();
let mut classes = std::collections::HashSet::new();
for &label in y.iter() {
if label != -1 {
classes.insert(label);
}
}
let classes: Vec<i32> = classes.into_iter().collect();
let n_classes = classes.len();
let mut layer_sizes = vec![n_features];
layer_sizes.extend(&self.hidden_layers);
layer_sizes.push(n_classes);
let mut weights = Vec::new();
let mut biases = Vec::new();
for i in 0..layer_sizes.len() - 1 {
let in_size = layer_sizes[i];
let out_size = layer_sizes[i + 1];
let w = Array2::zeros((in_size, out_size));
let b = Array1::zeros(out_size);
weights.push(w);
biases.push(b);
}
Ok(MAML {
state: MAMLTrained {
meta_weights: weights,
meta_biases: biases,
classes: Array1::from(classes),
},
inner_lr: self.inner_lr,
outer_lr: self.outer_lr,
inner_steps: self.inner_steps,
n_episodes: self.n_episodes,
hidden_layers: self.hidden_layers,
})
}
}
impl Predict<ArrayView2<'_, Float>, Array1<i32>> for MAML<MAMLTrained> {
fn predict(&self, X: &ArrayView2<'_, Float>) -> SklResult<Array1<i32>> {
let n_test = X.nrows();
let n_classes = self.state.classes.len();
let mut predictions = Array1::zeros(n_test);
for i in 0..n_test {
predictions[i] = self.state.classes[i % n_classes];
}
Ok(predictions)
}
}
impl PredictProba<ArrayView2<'_, Float>, Array2<f64>> for MAML<MAMLTrained> {
fn predict_proba(&self, X: &ArrayView2<'_, Float>) -> SklResult<Array2<f64>> {
let n_test = X.nrows();
let n_classes = self.state.classes.len();
let mut probabilities = Array2::zeros((n_test, n_classes));
for i in 0..n_test {
for j in 0..n_classes {
probabilities[[i, j]] = 1.0 / n_classes as f64;
}
}
Ok(probabilities)
}
}
#[derive(Debug, Clone)]
pub struct MAMLTrained {
pub meta_weights: Vec<Array2<f64>>,
pub meta_biases: Vec<Array1<f64>>,
pub classes: Array1<i32>,
}