350 lines
12 KiB
Go
350 lines
12 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package grad
|
||
|
||
import (
|
||
"math"
|
||
"testing"
|
||
|
||
ode "sourcedock.dev/petrbalvin/tensor/integrate"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
)
|
||
|
||
// decayWith builds f for y' = −θy with θ a scalar parameter leaf.
|
||
func decayWith(theta *Tensor) func(t float64, y *Tensor) (*Tensor, error) {
|
||
return func(t float64, y *Tensor) (*Tensor, error) {
|
||
rate, err := theta.Neg()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return y.Mul(rate)
|
||
}
|
||
}
|
||
|
||
// scalarVector returns a length-1 real array.
|
||
func scalarVector(t *testing.T, v float64) *core.Array {
|
||
t.Helper()
|
||
a, err := core.FromFloats([]float64{v}, 1)
|
||
if err != nil {
|
||
t.Fatalf("FromFloats: %v", err)
|
||
}
|
||
return a
|
||
}
|
||
|
||
// TestAdjointODEDecay differentiates y' = −θy with L = y(1): the exact
|
||
// sensitivities are dL/dy0 = e^{−θ} and dL/dθ = −e^{−θ}, and a run
|
||
// with no parameters at all must still answer the initial-state one.
|
||
func TestAdjointODEDecay(t *testing.T) {
|
||
const want = 0.4965853037914095 // e^{−0.7}
|
||
theta, err := FromFloat64s([]float64{0.7}, true, 1)
|
||
if err != nil {
|
||
t.Fatalf("FromFloat64s: %v", err)
|
||
}
|
||
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0, 1,
|
||
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||
if err != nil {
|
||
t.Fatalf("AdjointODE: %v", err)
|
||
}
|
||
if math.Abs(gradY0.FloatAt(0)-want) > 1e-6 {
|
||
t.Fatalf("dL/dy0 = %.14g, want %.14g", gradY0.FloatAt(0), want)
|
||
}
|
||
if len(blocks) != 1 || math.Abs(blocks[0].FloatAt(0)+want) > 1e-6 {
|
||
t.Fatalf("dL/dθ = %v, want −%.14g", blocks[0].FloatAt(0), want)
|
||
}
|
||
// The parameter's own accumulated gradient must be untouched.
|
||
if theta.Grad() != nil {
|
||
t.Fatal("AdjointODE must leave the parameters' gradients untouched")
|
||
}
|
||
solo, soloBlocks, err := AdjointODE(decayWith(theta), nil, 0, 1,
|
||
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||
if err != nil {
|
||
t.Fatalf("AdjointODE without parameters: %v", err)
|
||
}
|
||
if len(soloBlocks) != 0 || math.Abs(solo.FloatAt(0)-want) > 1e-6 {
|
||
t.Fatalf("parameter-free run: dL/dy0 = %.14g with %d blocks", solo.FloatAt(0), len(soloBlocks))
|
||
}
|
||
}
|
||
|
||
// TestAdjointODEOscillator differentiates the harmonic oscillator
|
||
// y” = −ω²y with L = y(1) and y0 = (1, 0): dL/dω = −sin(ω),
|
||
// dL/dy0 = (cos ω, sin ω/ω).
|
||
func TestAdjointODEOscillator(t *testing.T) {
|
||
omega, err := FromFloat64s([]float64{1.3}, true, 1)
|
||
if err != nil {
|
||
t.Fatalf("FromFloat64s: %v", err)
|
||
}
|
||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||
u, err := y.Slice(0, 0, 1)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
v, err := y.Slice(0, 1, 2)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
sq, err := omega.Mul(omega)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
acc, err := u.Mul(sq)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
drag, err := acc.Scale(-1)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return v.Concat(drag, 0)
|
||
}
|
||
y0, err := core.FromFloats([]float64{1, 0}, 2)
|
||
if err != nil {
|
||
t.Fatalf("FromFloats: %v", err)
|
||
}
|
||
seed, err := core.FromFloats([]float64{1, 0}, 2)
|
||
if err != nil {
|
||
t.Fatalf("FromFloats: %v", err)
|
||
}
|
||
gradY0, blocks, err := AdjointODE(f, []*Tensor{omega}, 0, 1, y0, seed,
|
||
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||
if err != nil {
|
||
t.Fatalf("AdjointODE: %v", err)
|
||
}
|
||
if math.Abs(gradY0.FloatAt(0)-math.Cos(1.3)) > 1e-6 {
|
||
t.Fatalf("dL/du0 = %.14g, want cos(1.3) = %.14g", gradY0.FloatAt(0), math.Cos(1.3))
|
||
}
|
||
if math.Abs(gradY0.FloatAt(1)-math.Sin(1.3)/1.3) > 1e-6 {
|
||
t.Fatalf("dL/dv0 = %.14g, want sin(1.3)/1.3 = %.14g",
|
||
gradY0.FloatAt(1), math.Sin(1.3)/1.3)
|
||
}
|
||
if math.Abs(blocks[0].FloatAt(0)+math.Sin(1.3)) > 1e-6 {
|
||
t.Fatalf("dL/dω = %.14g, want −sin(1.3) = %.14g", blocks[0].FloatAt(0), -math.Sin(1.3))
|
||
}
|
||
}
|
||
|
||
// TestAdjointODEBackwardTime runs the forward pass itself backwards
|
||
// (t1 < t0): y(t) = e^{−θ(t−1)} from y(1) = 1 has y(0) = e^θ and both
|
||
// sensitivities equal e^θ.
|
||
func TestAdjointODEBackwardTime(t *testing.T) {
|
||
const want = 1.6487212707001282 // e^{0.5}
|
||
theta, _ := FromFloat64s([]float64{0.5}, true, 1)
|
||
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 1, 0,
|
||
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||
if err != nil {
|
||
t.Fatalf("AdjointODE: %v", err)
|
||
}
|
||
if math.Abs(gradY0.FloatAt(0)-want) > 1e-6 {
|
||
t.Fatalf("dL/dy0 = %.14g, want %.14g", gradY0.FloatAt(0), want)
|
||
}
|
||
if math.Abs(blocks[0].FloatAt(0)-want) > 1e-6 {
|
||
t.Fatalf("dL/dθ = %.14g, want %.14g", blocks[0].FloatAt(0), want)
|
||
}
|
||
}
|
||
|
||
// TestAdjointODEFiniteDifference checks a nonlinear two-parameter
|
||
// system against central differences of the forward solve itself,
|
||
// perturbing the parameter leaves around the adjoint run.
|
||
func TestAdjointODEFiniteDifference(t *testing.T) {
|
||
th1, _ := FromFloat64s([]float64{0.8}, true, 1)
|
||
th2, _ := FromFloat64s([]float64{1.1}, true, 1)
|
||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||
y1, err := y.Slice(0, 0, 1)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
y2, err := y.Slice(0, 1, 2)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
drag, err := y1.Mul(th1)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
drive, err := y2.Mul(th2)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
r1, err := drive.Sub(drag)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
pool, err := y1.Mul(y2)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
r2, err := pool.Scale(-1)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return r1.Concat(r2, 0)
|
||
}
|
||
y0, _ := core.FromFloats([]float64{1, 0.5}, 2)
|
||
seed, _ := core.FromFloats([]float64{1, 2}, 2)
|
||
gradY0, blocks, err := AdjointODE(f, []*Tensor{th1, th2}, 0, 1, y0, seed,
|
||
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-13})
|
||
if err != nil {
|
||
t.Fatalf("AdjointODE: %v", err)
|
||
}
|
||
// Central differences on the loss y1(1) + 2·y2(1), with the leaf
|
||
// data swapped out for the perturbed values.
|
||
forwardLoss := func() float64 {
|
||
end, err := ode.IntegrateODE(func(t float64, ya *core.Array) (*core.Array, error) {
|
||
out, err := f(t, FromArray(ya, false))
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return out.Data(), nil
|
||
}, 0, 1, y0, ode.ODEOptions{RelTol: 1e-11, AbsTol: 1e-15})
|
||
if err != nil {
|
||
t.Fatalf("forward solve: %v", err)
|
||
}
|
||
return end.FloatAt(0) + 2*end.FloatAt(1)
|
||
}
|
||
perturb := func(p *Tensor, eps float64) {
|
||
swapped, _ := core.FromFloats([]float64{p.Data().FloatAt(0) + eps}, 1)
|
||
p.ReplaceWith(swapped)
|
||
}
|
||
restore := func(p *Tensor, v float64) {
|
||
orig, _ := core.FromFloats([]float64{v}, 1)
|
||
p.ReplaceWith(orig)
|
||
}
|
||
const eps = 1e-5
|
||
for k, p := range []*Tensor{th1, th2} {
|
||
orig := p.Data().FloatAt(0)
|
||
perturb(p, eps)
|
||
up := forwardLoss()
|
||
perturb(p, -2*eps)
|
||
down := forwardLoss()
|
||
restore(p, orig)
|
||
fd := (up - down) / (2 * eps)
|
||
got := blocks[k].FloatAt(0)
|
||
if math.Abs(got-fd) > 1e-3*math.Max(1, math.Abs(fd)) {
|
||
t.Fatalf("dL/dθ%d: adjoint %.8g, finite difference %.8g", k+1, got, fd)
|
||
}
|
||
}
|
||
if gradY0.Len() != 2 {
|
||
t.Fatalf("dL/dy0 has length %d, want 2", gradY0.Len())
|
||
}
|
||
}
|
||
|
||
// TestAdjointODEDegenerateSpan returns the loss seed unchanged and
|
||
// zero parameter gradients when the span carries no dynamics.
|
||
func TestAdjointODEDegenerateSpan(t *testing.T) {
|
||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0.5, 0.5,
|
||
scalarVector(t, 1), scalarVector(t, 2), ode.ODEOptions{})
|
||
if err != nil {
|
||
t.Fatalf("AdjointODE: %v", err)
|
||
}
|
||
if gradY0.FloatAt(0) != 2 {
|
||
t.Fatalf("dL/dy0 = %.14g, want the seed 2", gradY0.FloatAt(0))
|
||
}
|
||
if blocks[0].FloatAt(0) != 0 {
|
||
t.Fatalf("dL/dθ = %.14g, want 0", blocks[0].FloatAt(0))
|
||
}
|
||
}
|
||
|
||
// TestAdjointODEErrors pins the validation contract.
|
||
func TestAdjointODEErrors(t *testing.T) {
|
||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||
y0 := scalarVector(t, 1)
|
||
seed := scalarVector(t, 1)
|
||
rank2, _ := core.FromFloats([]float64{1, 1}, 1, 2)
|
||
if _, _, err := AdjointODE(nil, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error for a nil function")
|
||
}
|
||
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, nil, seed, ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error for a nil state")
|
||
}
|
||
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, rank2, seed, ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error for a rank-2 state")
|
||
}
|
||
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, y0, nil, ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error for a nil loss seed")
|
||
}
|
||
badSeed, _ := core.FromFloats([]float64{1, 1}, 2)
|
||
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, y0, badSeed, ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error for a loss seed of the wrong length")
|
||
}
|
||
frozen := FromArray(scalarVector(t, 0.7), false)
|
||
if _, _, err := AdjointODE(decayWith(frozen), []*Tensor{frozen}, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error for a parameter that does not require grad")
|
||
}
|
||
wrongShape := func(t float64, y *Tensor) (*Tensor, error) {
|
||
three, _ := core.FromFloats([]float64{1, 2, 3}, 3)
|
||
return FromArray(three, false), nil
|
||
}
|
||
if _, _, err := AdjointODE(wrongShape, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error for a wrong-shaped derivative")
|
||
}
|
||
ignoresState := func(t float64, y *Tensor) (*Tensor, error) {
|
||
one, _ := core.FromFloats([]float64{1}, 1)
|
||
return FromArray(one, false), nil
|
||
}
|
||
if _, _, err := AdjointODE(ignoresState, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error when f ignores its state")
|
||
}
|
||
}
|
||
|
||
// TestAdjointODERestoresGradientsOnError pins the error-path contract:
|
||
// when f consumes the parameter but ignores the state, no state
|
||
// gradient can flow and the run errors, and the caller's own
|
||
// accumulated parameter gradient must come back untouched instead of
|
||
// being polluted by the aborted pass: the reverse sweep is not run at
|
||
// all, so nothing writes the leaf gradients on the way out.
|
||
func TestAdjointODERestoresGradientsOnError(t *testing.T) {
|
||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||
preset, _ := core.FromFloats([]float64{3}, 1)
|
||
theta.SetGrad(preset)
|
||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||
return theta.Scale(2)
|
||
}
|
||
if _, _, err := AdjointODE(f, []*Tensor{theta}, 0, 1, scalarVector(t, 1),
|
||
scalarVector(t, 1), ode.ODEOptions{}); err == nil {
|
||
t.Fatal("expected an error when f ignores its state")
|
||
}
|
||
if theta.Grad() == nil || theta.Grad().FloatAt(0) != 3 {
|
||
t.Fatalf("the parameter's gradient was not restored: %v, want 3", theta.Grad())
|
||
}
|
||
}
|
||
|
||
// TestAdjointODEDisconnectedParameter pins that a parameter the
|
||
// dynamics never touch is an error, not a silent zero gradient block.
|
||
func TestAdjointODEDisconnectedParameter(t *testing.T) {
|
||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||
return y.Scale(2)
|
||
}
|
||
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, scalarVector(t, 1),
|
||
scalarVector(t, 1), ode.ODEOptions{})
|
||
if err == nil {
|
||
t.Fatal("expected an error for a parameter disconnected from the dynamics")
|
||
}
|
||
if blocks != nil {
|
||
t.Fatalf("an errored run returned blocks: %v", blocks)
|
||
}
|
||
}
|
||
|
||
// TestAdjointODESuccessKeepsPresetGradient pins the same restore on
|
||
// the success path: a preset accumulated gradient survives the run.
|
||
func TestAdjointODESuccessKeepsPresetGradient(t *testing.T) {
|
||
theta, _ := FromFloat64s([]float64{0.5}, true, 1)
|
||
preset, _ := core.FromFloats([]float64{7}, 1)
|
||
theta.SetGrad(preset)
|
||
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0, 1,
|
||
scalarVector(t, 1), scalarVector(t, 2), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||
if err != nil {
|
||
t.Fatalf("AdjointODE: %v", err)
|
||
}
|
||
const want = 0.6065306597126334 // e^{−0.5}
|
||
if math.Abs(gradY0.FloatAt(0)-2*want) > 1e-6 {
|
||
t.Fatalf("dL/dy0 = %.14g, want 2e^{−0.5}", gradY0.FloatAt(0))
|
||
}
|
||
if theta.Grad() == nil || theta.Grad().FloatAt(0) != 7 {
|
||
t.Fatalf("the preset gradient did not survive a successful run: %v", theta.Grad())
|
||
}
|
||
if blocks[0].FloatAt(0) >= 0 {
|
||
t.Fatalf("dL/dθ = %v, want negative", blocks[0].FloatAt(0))
|
||
}
|
||
}
|