// maml.mbt — MAML (Model-Agnostic Meta-Learning) inner/outer loop
// meta-update (v0.94.0).
//
// MAML (Finn et al. 2017) learns an initialization θ such that one or
// a few gradient steps on a new task's support set yields good
// performance on the query set:
//
// For each meta-batch:
// for each task T_i:
// # Inner loop: adapt on T_i's support set
// θ'_i = θ - α · ∇_θ L_sup(θ; T_i) (one or more steps)
// # Outer loop: meta-update θ using query losses at θ'_i
// θ <- θ - β · ∇_θ Σ_i L_query(θ'_i; T_i)
//
// The full second-order MAML propagates gradients through the inner
// adaptation (chain rule through the inner SGD step). For v0.94.0 we
// use the FOMAML approximation — treat θ'_i as independent of θ for
// the outer gradient — which is the canonical practical variant used
// in most MAML implementations on small models.
//
// Scope of v0.94.0:
// - mlp_model_mse_loss: per-batch MSE loss on a (model, inputs,
// targets) tuple
// - mlp_model_mse_grad_per_sample: per-sample analytic gradients of
// MSE w.r.t. model weights (returned as gradient tuples)
// - mlp_model_sgd_step: one SGD step on the model weights
// - maml_meta_step: one MAML meta-step (inner loop on support +
// outer FOMAML update on query)
//
// Reference: Finn et al. 2017 "Model-Agnostic Meta-Learning for Fast
// Adaptation of Deep Networks".
///|
/// Compute per-batch MSE loss for an MLPModel. `inputs` flat
/// `[batch × input_dim]`, `targets` flat `[batch × output_dim]`.
/// Returns scalar (1/N) · Σ (pred - target)² averaged over the batch.
pub fn mlp_model_mse_loss(
model : MLPModel,
inputs : Array[Float],
targets : Array[Float],
batch : Int,
) -> Float {
if batch <= 0 {
return 0.0F
}
let in_dim = model.input_dim
let out_dim = model.output_dim
let mut sum_sq = 0.0F
for n in 0.. MLPGrad {
let in_dim = model.input_dim
let out_dim = model.output_dim
let h_dim = model.hidden_dim
let d_w1 : Array[Array[Float]] = Array::make(
h_dim, Array::make(in_dim, 0.0F),
)
let d_b1 : Array[Float] = Array::make(h_dim, 0.0F)
let d_w2 : Array[Array[Float]] = Array::make(
out_dim, Array::make(h_dim, 0.0F),
)
let d_b2 : Array[Float] = Array::make(out_dim, 0.0F)
if batch <= 0 {
return { d_w1, d_b1, d_w2, d_b2 }
}
let scale = 2.0F / Float::from_int(batch)
for n in 0.. MLPModel {
let (new_w1, new_b1) = sgd_update_arrays(
flatten_2d(model.w1), model.b1, flatten_2d(grad.d_w1), grad.d_b1,
lr,
)
let (new_w2, new_b2) = sgd_update_arrays(
flatten_2d(model.w2), model.b2, flatten_2d(grad.d_w2), grad.d_b2,
lr,
)
{
..model,
w1: unflatten_2d(new_w1, model.hidden_dim, model.input_dim),
b1: new_b1,
w2: unflatten_2d(new_w2, model.output_dim, model.hidden_dim),
b2: new_b2,
}
}
///|
/// One MAML meta-step. Samples a meta-batch of tasks, runs `n_inner`
/// SGD steps on each task's support set starting from the meta model,
// / then computes the FOMAML meta-gradient (treats the adapted θ'_i
/// as independent of θ). Applies the meta-update to the family's meta
/// model. Returns the updated TaskFamily + the mean meta-loss.
///
/// FOMAML approximation: instead of propagating gradients through the
/// inner adaptation (full second-order MAML), we compute the gradient
/// of L_query at the adapted point and apply it directly to θ. This
/// is the canonical practical MAML variant for small models; full
/// second-order requires either autograd or finite-difference Hessian
/// (deferred).
pub fn maml_meta_step(
family : TaskFamily,
n_tasks : Int,
n_support : Int,
n_query : Int,
n_inner : Int,
inner_lr : Float,
outer_lr : Float,
seed : UInt64,
) -> (TaskFamily, Float) {
let tasks = sample_meta_batch(family, n_tasks, n_support, n_query, seed)
// Aggregated outer gradient (FOMAML — sum of adapted-model query
// gradients, treated as if they were gradients w.r.t. θ).
let meta = family.meta_model
let in_dim = meta.input_dim
let out_dim = meta.output_dim
let h_dim = meta.hidden_dim
let agg_d_w1 : Array[Array[Float]] = Array::make(
h_dim, Array::make(in_dim, 0.0F),
)
let agg_d_b1 : Array[Float] = Array::make(h_dim, 0.0F)
let agg_d_w2 : Array[Array[Float]] = Array::make(
out_dim, Array::make(h_dim, 0.0F),
)
let agg_d_b2 : Array[Float] = Array::make(out_dim, 0.0F)
let mut total_loss = 0.0F
for i in 0..