// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "slices" ode "sourcedock.dev/petrbalvin/tensor/integrate" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // Adjoint sensitivities of an initial value problem. Fitting the // parameters of a differential equation to data asks for dL/dθ when // the trajectory y(t; θ) is produced by an ODE solve, and // differentiating through every solver step is neither necessary nor // cheap. The adjoint method runs the dynamics once forward, then // integrates the adjoint state λ(t) = ∂L/∂y(t) backward along // λ' = −(∂f/∂y)ᵀλ from the loss gradient at the endpoint, carrying a // per-parameter accumulator β' = −(∂f/∂θ)ᵀλ alongside; both // Jacobian-vector products come from one automatic-differentiation // backward pass per evaluation. The cost is one more ODE solve // whatever the parameter count, which is what makes whole-trajectory // fitting tractable. // AdjointODE differentiates the solution of y' = f(t, y) at t1 with // respect to the initial state and to the parameters the given // function closes over. lossGrad is ∂L/∂y(t1), the seed the loss // itself contributes; the return values are ∂L/∂y0 and, parallel to // params, ∂L/∂θk, each shaped like its parameter. Backward-in-time // problems (t1 < t0) work. // // The forward trajectory is recorded at the adaptive solver's accepted // steps and handed to the backward pass through cubic Hermite // interpolation, fourth-order accurate like the Dormand-Prince pair // that produced it; the augmented adjoint system is integrated by the // same adaptive solver in reverse. The parameters' own accumulated // gradients are left untouched: every evaluation runs a reverse pass // that computes into a local map and commits nothing, and the answers // travel out as return values instead. // // A nil function, a non-vector start or loss seed, a parameter that // does not require grad, and an f that ignores its state or returns a // wrong shape are errors, never silent zeros. func AdjointODE(f func(t float64, y *Tensor) (*Tensor, error), params []*Tensor, t0, t1 float64, y0 *core.Array, lossGrad *core.Array, opts ode.ODEOptions) (*core.Array, []*core.Array, error) { const name = "AdjointODE" if f == nil { return nil, nil, errf("%s: f must not be nil", name) } if y0 == nil || y0.NDim() != 1 || y0.Len() == 0 { return nil, nil, errf("%s: the state must be a non-empty vector", name) } if y0.Dtype() == core.Complex { return nil, nil, errf("%s: complex states are not supported", name) } dim := y0.Len() if lossGrad == nil || lossGrad.NDim() != 1 || lossGrad.Len() != dim { return nil, nil, errf("%s: lossGrad must be a vector of length %d", name, dim) } sizes := make([]int, len(params)) total := 0 for k, p := range params { if p == nil || !p.RequiresGrad() { return nil, nil, errf("%s: parameter %d does not require grad", name, k) } if p.Data().NDim() == 0 || p.Data().Len() == 0 { return nil, nil, errf("%s: parameter %d must be a non-empty tensor of rank at least 1", name, k) } sizes[k] = p.Data().Len() total += sizes[k] } if t0 == t1 { // No dynamics: the endpoint is the initial state. seed := flatFloats(lossGrad) gradY0, err := core.FromFloats(seed, dim) if err != nil { return nil, nil, errf("%s: %w", name, err) } return gradY0, zeroBlocks(params, sizes), nil } // Forward pass with the dynamics evaluated for value only. forward := func(t float64, ya *core.Array) (*core.Array, error) { out, err := f(t, FromArray(ya, false)) if err != nil { return nil, errf("%s: %w", name, err) } return out.Data(), nil } times, states, err := ode.IntegrateODESteps(forward, t0, t1, y0, opts) if err != nil { return nil, nil, errf("%s: %w", name, err) } // Cubic Hermite interpolation wants the derivative of the dynamics // at every recorded node; one detached pass supplies them all. slopes := make([][]float64, len(times)) for i, node := range states { out, err := f(times[i], FromArray(node, false)) if err != nil { return nil, nil, errf("%s: %w", name, err) } if out.Data().NDim() != 1 || out.Data().Len() != dim { return nil, nil, errf("%s: f returned shape %s, want a vector of length %d", name, prettyShape(out.Data().Shape()), dim) } slopes[i] = flatFloats(out.Data()) } trace := newODETrace(times, states, slopes, dim) vjp := func(t float64, y []float64, lam []float64) ([]float64, []float64, error) { data, err := core.FromFloats(y, dim) if err != nil { return nil, nil, errf("%s: %w", name, err) } yLeaf := FromArray(data, true) out, err := f(t, yLeaf) if err != nil { return nil, nil, errf("%s: %w", name, err) } if out.Data().NDim() != 1 || out.Data().Len() != dim { return nil, nil, errf("%s: f returned shape %s, want a vector of length %d", name, prettyShape(out.Data().Shape()), dim) } seedData, err := core.FromFloats(lam, dim) if err != nil { return nil, nil, errf("%s: %w", name, err) } weighted, err := out.Mul(FromArray(seedData, false)) if err != nil { return nil, nil, errf("%s: %w", name, err) } summed, err := weighted.Sum() if err != nil { return nil, nil, errf("%s: %w", name, err) } // The Jacobian-vector products come from the reverse pass // alone: reverseGrads computes every reached tensor's gradient // into a local map and commits nothing, so no leaf the dynamics // close over, in params or out, is written at all, and the // callers' own accumulated gradients survive untouched. grads, err := summed.reverseGrads() if err != nil { return nil, nil, errf("%s: %w", name, err) } g := grads[yLeaf] if g == nil || g.Len() != dim { return nil, nil, errf("%s: f did not yield a state gradient of length %d", name, dim) } gY := flatFloats(g) gTheta := make([]float64, 0, total) for k, p := range params { now := grads[p] if now == nil { return nil, nil, errf("%s: parameter %d is disconnected from the dynamics, no gradient flowed", name, k) } nowFloats := flatFloats(now) gTheta = append(gTheta, nowFloats...) } return gY, gTheta, nil } rhs := func(t float64, z *core.Array) (*core.Array, error) { lam := make([]float64, dim) if !z.Strided() && z.Dtype() == core.Float { copy(lam, z.RawFloats()[:dim]) } else { for i := range lam { lam[i] = z.FloatAt(i) } } yArr, err := trace.at(t) if err != nil { return nil, err } gY, gTheta, err := vjp(t, yArr, lam) if err != nil { return nil, err } out := make([]float64, dim+total) for i := range gY { out[i] = -gY[i] } for i := range gTheta { out[dim+i] = -gTheta[i] } return core.FromFloats(out, dim+total) } zStart := make([]float64, dim+total) if !lossGrad.Strided() && lossGrad.Dtype() == core.Float { copy(zStart, lossGrad.RawFloats()[:dim]) } else { for i := range dim { zStart[i] = lossGrad.FloatAt(i) } } zEnd, err := ode.IntegrateODE(rhs, t1, t0, mustFromFloats(zStart, dim+total), opts) if err != nil { return nil, nil, errf("%s: %w", name, err) } seed := make([]float64, dim) for i := range seed { seed[i] = zEnd.FloatAt(i) } gradY0, err := core.FromFloats(seed, dim) if err != nil { return nil, nil, errf("%s: %w", name, err) } blocks := make([]*core.Array, len(params)) offset := dim for k, p := range params { block := make([]float64, sizes[k]) if !zEnd.Strided() && zEnd.Dtype() == core.Float { copy(block, zEnd.RawFloats()[offset:offset+sizes[k]]) } else { for j := range block { block[j] = zEnd.FloatAt(offset + j) } } shaped, err := core.FromFloats(block, p.Data().Shape()...) if err != nil { return nil, nil, errf("%s: %w", name, err) } blocks[k] = shaped offset += sizes[k] } return gradY0, blocks, nil } // zeroBlocks builds zero arrays shaped like each parameter, the // degenerate answer of a span with no dynamics. func zeroBlocks(params []*Tensor, sizes []int) []*core.Array { blocks := make([]*core.Array, len(params)) for k, p := range params { zeros := make([]float64, sizes[k]) blocks[k], _ = core.FromFloats(zeros, p.Data().Shape()...) } return blocks } // mustFromFloats wraps a plain construction that cannot fail: the // length always matches the single-dimension shape. func mustFromFloats(vals []float64, n int) *core.Array { a, _ := core.FromFloats(vals, n) return a } // odeTrace is a recorded forward trajectory, kept ascending in time, // with the dynamics' derivative at every node so the interpolation is // cubic Hermite. type odeTrace struct { times []float64 states [][]float64 slopes [][]float64 dim int // buf is the interpolation buffer at reuses: the solver calls at // once per right-hand-side evaluation, so the array is borrowed // rather than allocated each time. Callers must consume the answer // before the next call. buf []float64 } // newODETrace flattens and time-orders the recorded nodes. func newODETrace(times []float64, states []*core.Array, slopes [][]float64, dim int) *odeTrace { tr := &odeTrace{times: append([]float64{}, times...), slopes: slopes, dim: dim} tr.states = make([][]float64, len(states)) for i, s := range states { tr.states[i] = flatFloats(s) } if len(times) > 1 && times[1] < times[0] { // A backward-in-time pass records descending nodes; flip to // ascending so the search below has one convention. for i, j := 0, len(tr.times)-1; i < j; i, j = i+1, j-1 { tr.times[i], tr.times[j] = tr.times[j], tr.times[i] tr.states[i], tr.states[j] = tr.states[j], tr.states[i] tr.slopes[i], tr.slopes[j] = tr.slopes[j], tr.slopes[i] } } return tr } // at interpolates the state at t by cubic Hermite over the enclosing // recorded interval, fourth-order accurate in the step size and never // leaving the recorded span. A trace with fewer than two nodes offers // no interval to interpolate over, so the accessor refuses instead of // indexing out of range. The returned slice belongs to the trace and // stays valid only until the next call. func (tr *odeTrace) at(t float64) ([]float64, error) { n := len(tr.times) if n < 2 { return nil, errf("AdjointODE: the solver recorded %d trajectory nodes, the adjoint interpolation needs at least two", n) } idx, _ := slices.BinarySearch(tr.times, t) i := min(max(idx-1, 0), n-2) h := tr.times[i+1] - tr.times[i] s := (t - tr.times[i]) / h h00 := s*s*(2*s-3) + 1 h10 := s * (s - 1) * (s - 1) h01 := s * s * (3 - 2*s) h11 := s * s * (s - 1) if len(tr.buf) != tr.dim { tr.buf = make([]float64, tr.dim) } out := tr.buf si, si1 := tr.states[i], tr.states[i+1] li, li1 := tr.slopes[i], tr.slopes[i+1] for j := range out { out[j] = h00*si[j] + h01*si1[j] + h*(h10*li[j]+h11*li1[j]) } return out, nil }