// neural_ode_agent.mbt — NeuralODEAgent: full forward + adjoint
// training step + SGD (v0.92.0).
//
// Wires the v0.89.0 ODEFunc primitive + v0.90.0 ODESolver (Euler or
// Heun) + v0.91.0 adjoint sensitivity into a single training step.
// The agent fits a continuous-depth transformation from initial state
// x(0) to a target final state — useful as a building block for
// normalizing flows, residual-style continuous-time models, and
// time-series prediction.
//
// Scope of v0.92.0:
//   - NeuralODEAgent struct + constructor (wraps ODEFunc + solver +
//     integration window)
//   - neural_ode_forward: full forward — integrate from x0 over
//     [t0, t1] with n_steps; returns final_state + trajectory
//   - neural_ode_train_step: one SGD step — forward + MSE loss +
//     adjoint backward + SGD on ODEFunc weights. Returns
//     (updated_agent, loss).
//
// Reference: Chen et al. 2018 "Neural Ordinary Differential Equations".

///|
/// Choice of ODE solver for the NeuralODEAgent.
pub enum NeuralODESolver {
  Euler
  Heun
}

///|
/// NeuralODEAgent: ODEFunc + integration window + solver choice.
/// Stores the time interval [t0, t1] and step count; ODEFunc owns the
/// parameters.
pub struct NeuralODEAgent {
  func : ODEFunc
  t0 : Float
  t1 : Float
  n_steps : Int
  solver : NeuralODESolver
}

///|
/// Build a fresh NeuralODEAgent with the given ODEFunc and
/// integration window. Default solver: Heun (2nd-order, better
/// accuracy per step).
pub fn NeuralODEAgent::new(
  func : ODEFunc,
  t0 : Float,
  t1 : Float,
  n_steps : Int,
) -> NeuralODEAgent {
  { func, t0, t1, n_steps, solver: Heun }
}

///|
/// Set the solver (Euler or Heun). Returns an updated agent.
pub fn neural_ode_set_solver(
  agent : NeuralODEAgent,
  solver : NeuralODESolver,
) -> NeuralODEAgent {
  { ..agent, solver }
}

///|
/// Full forward: integrate ODEFunc from x0 over [t0, t1] with
/// n_steps using the agent's chosen solver. Returns the final state
/// and the full trajectory (flat [n_steps+1 × state_dim]).
pub fn neural_ode_forward(
  agent : NeuralODEAgent,
  x0 : Array[Float],
) -> (Array[Float], Array[Float]) {
  match agent.solver {
    Euler => ode_euler_solve(
      agent.func, x0, agent.t0, agent.t1, agent.n_steps,
    )
    Heun => ode_heun_solve(
      agent.func, x0, agent.t0, agent.t1, agent.n_steps,
    )
  }
}

///|
/// One full training step: forward + MSE loss + adjoint backward +
/// SGD step on ODEFunc weights. Returns (updated_agent, loss).
pub fn neural_ode_train_step(
  agent : NeuralODEAgent,
  x0 : Array[Float],
  target : Array[Float],
  lr : Float,
) -> (NeuralODEAgent, Float) {
  // 1. Forward.
  let (final_state, trajectory) = neural_ode_forward(agent, x0)
  // 2. Loss (MSE).
  let state_dim = agent.func.state_dim
  let mut sum_sq = 0.0F
  for i in 0..