feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+313
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user