// ode_func.mbt — ODEFunc: dynamics field for Neural ODEs (v0.89.0).
//
// The Neural ODE (Chen et al. 2018) replaces discrete-depth ResNet
// blocks with a continuous-depth ODE:
//
// dx/dt = f(x, t; θ) (vector field, parametrized by θ)
// x(t) = x(0) + ∫_0^t f(x(s), s; θ) ds
//
// This file ships the `ODEFunc` primitive — a small MLP that maps
// (state, time) → state derivative. The MLP is intentionally shallow
// (1 hidden layer with tanh) to keep the per-step cost low, since the
// ODE solver will call `f` many times per forward pass.
//
// Reference: Chen et al. 2018 "Neural Ordinary Differential Equations".
//
// Scope of v0.89.0:
// - ODEFunc struct + constructor (MLP w1, b1, w2, b2)
// - ode_func_eval: forward — given (state, t), returns f(x, t)
///|
/// ODEFunc: a small MLP parameterizing the vector field
/// dx/dt = f(x, t; θ). Two-layer MLP: state → Linear → tanh → Linear →
/// state_derivative. The `t` input is broadcast-added to the hidden
/// layer (time-conditioning).
pub struct ODEFunc {
state_dim : Int
hidden_dim : Int
// Linear1: (hidden_dim × state_dim) + bias of length hidden_dim
w1 : Array[Array[Float]]
b1 : Array[Float]
// Linear2: (state_dim × hidden_dim) + bias of length state_dim
w2 : Array[Array[Float]]
b2 : Array[Float]
}
///|
/// Build fresh ODEFunc. Weights init via xavier_normal scaled by
/// sqrt(2/state_dim) for the input layer and sqrt(2/hidden_dim) for
/// the output layer. Zero biases.
pub fn ODEFunc::new(
state_dim : Int,
hidden_dim : Int,
seed : UInt64,
) -> ODEFunc {
let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let std1 = sqrtf(2.0F / Float::from_int(state_dim))
let w1 = xavier_normal(hidden_dim, state_dim, std1, rng1)
let b1 : Array[Float] = Array::make(hidden_dim, 0.0F)
let rng2 = Xoshiro::from_state(seed + 4UL, seed + 5UL, seed + 6UL, seed + 7UL)
let std2 = sqrtf(2.0F / Float::from_int(hidden_dim))
let w2 = xavier_normal(state_dim, hidden_dim, std2, rng2)
let b2 : Array[Float] = Array::make(state_dim, 0.0F)
{ state_dim, hidden_dim, w1, b1, w2, b2 }
}
///|
/// Evaluate the vector field at (state, t). Returns the state
/// derivative f(state, t). Time `t` is broadcast-added to the hidden
/// layer (single scalar addition — broadcast across hidden_dim).
pub fn ode_func_eval(
func : ODEFunc,
state : Array[Float],
t : Float,
) -> Array[Float] {
// hidden = tanh(w1 · state + b1 + t)
let hidden : Array[Float] = Array::make(func.hidden_dim, 0.0F)
for i in 0..