// 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.. Array[Float] {
let n = final_pred.length()
let grad : Array[Float] = Array::make(n, 0.0F)
if n <= 0 {
return grad
}
let scale = 2.0F / Float::from_int(n)
for i in 0..