// reptile.mbt — Reptile (Nichol et al. 2018) meta-update (v0.96.0).
//
// Reptile is the simplest first-order meta-learning algorithm: after
// adapting the meta parameters on a task's training set via SGD, the
// meta parameters are interpolated toward the adapted parameters:
//
// θ_adapted = θ - inner_lr · ∇_θ L_train(θ)
// θ_new = (1 - ε) · θ + ε · θ_adapted
//
// No second-order gradient through the inner adaptation. No
// query-set evaluation needed at the meta-step. Just n_tasks different
// adapted parameters, averaged via the interpolation.
//
// Reptile is competitive with FOMAML on simple task distributions and
// even faster (no outer gradient computation). The interpolation
// step `ε` is the meta-step size; typically 0.1 - 0.5.
//
// Scope of v0.96.0:
// - reptile_meta_step: one Reptile meta-step — sample n_tasks
// tasks, adapt on each task's support set via n_inner SGD steps
// starting from the meta model, then interpolate the meta model
// toward the AVERAGE adapted parameters.
//
// Reference: Nichol et al. 2018 "On First-Order Meta-Learning
// Algorithms".
///|
/// One Reptile meta-step. Samples a meta-batch, adapts the meta
/// model on each task's support set via n_inner SGD steps, then
/// updates the meta model via:
/// θ <- (1 - ε) · θ + ε · θ_adapted
/// where θ_adapted is averaged across tasks.
///
/// Returns (updated_TaskFamily, mean_inner_loss).
pub fn reptile_meta_step(
family : TaskFamily,
n_tasks : Int,
n_support : Int,
n_inner : Int,
inner_lr : Float,
meta_step_size : Float,
seed : UInt64,
) -> (TaskFamily, Float) {
let tasks = sample_meta_batch(family, n_tasks, n_support, n_support, seed)
// Aggregated adapted parameters (averaged across tasks).
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_w1 : Array[Array[Float]] = Array::make(
h_dim, Array::make(in_dim, 0.0F),
)
let agg_b1 : Array[Float] = Array::make(h_dim, 0.0F)
let agg_w2 : Array[Array[Float]] = Array::make(
out_dim, Array::make(h_dim, 0.0F),
)
let agg_b2 : Array[Float] = Array::make(out_dim, 0.0F)
let mut total_inner_loss = 0.0F
for i in 0..