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