Files
tensor/grad/adjoint_test.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

350 lines
12 KiB
Go
Raw 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 (
"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))
}
}