Files
tensor/grad/adjoint.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

314 lines
11 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
}