// ode_adjoint.mbt — Adjoint sensitivity method for Neural ODEs
// (v0.91.0).
//
// The adjoint method (Chen et al. 2018) computes dL/dθ without
// backpropagating through the ODE solver — instead it solves another
// ODE backward in time. Given:
//   - Forward trajectory x(t) for t in [t0, t1] (from v0.90.0)
//   - Loss gradient w.r.t. final state: a(t1) = dL/dx(t1)
//   - Jacobian of the vector field: A(t) = ∂f/∂x(x(t), t)
//
// The adjoint a(t) = dL/dx(t) satisfies:
//   da/dt = -A(t)ᵀ · a(t)
// with initial condition a(t1) = dL/dx(t1).
//
// The parameter gradient accumulates as we integrate:
//   dL/dθ = -∫[t1..t0] a(t)ᵀ · ∂f(x(t), t; θ)/∂θ dt
//
// For our MLP ODEFunc (state → tanh(w1·state + b1) → w2·hidden + b2):
//   A = ∂f/∂x = w2 · diag(1 - hidden²) · w1
//   ∂f/∂w1 = (1 - hidden²) · a · stateᵀ (per hidden neuron)
//   ∂f/∂w2 = a · hiddenᵀ
//   ∂f/∂b1 = (1 - hidden²) · a
//   ∂f/∂b2 = a
//
// Scope of v0.91.0:
//   - ode_adjoint_gradients: BPTT-free parameter gradients via the
//     adjoint method (Euler backward integration of the adjoint ODE).
//   - Returns (d_w1, d_b1, d_w2, d_b2) matching the ODEFunc shape.
//
// Reference: Chen et al. 2018 "Neural Ordinary Differential Equations"
// Section 4 (adjoint sensitivity method).

///|
/// Adjoint sensitivity: compute parameter gradients dL/dθ for the
/// ODEFunc via the adjoint method (BPTT-free). The caller provides:
//   - `func`: the ODEFunc to differentiate
///   - `trajectory`: the forward state trajectory flat
///     `[n_steps+1 × state_dim]` from `ode_euler_solve` or
///     `ode_heun_solve`
///   - `t0, t1, n_steps`: same as the forward solver
///   - `dL_dx_t1`: gradient of the loss w.r.t. the final state (length
///     state_dim)
///
/// Returns (d_w1, d_b1, d_w2, d_b2) — gradients matching ODEFunc's
/// weight shapes. Implementation uses Euler integration of the
/// adjoint ODE backward from t1 to t0.
pub fn ode_adjoint_gradients(
  func : ODEFunc,
  trajectory : Array[Float],
  t0 : Float,
  t1 : Float,
  n_steps : Int,
  dL_dx_t1 : Array[Float],
) -> (Array[Array[Float]], Array[Float], Array[Array[Float]], Array[Float]) {
  let state_dim = func.state_dim
  let hidden_dim = func.hidden_dim
  let h = (t1 - t0) / Float::from_int(n_steps)
  // Initialize adjoint at t1 with dL/dx(t1).
  let a : Array[Float] = Array::make(state_dim, 0.0F)
  for i in 0..