// 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..