feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+313
View File
@@ -0,0 +1,313 @@
// 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
}