Files
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

476 lines
17 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 integrate
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import (
"math"
"testing"
)
// decay returns f for y' = −y, the reference every scheme must nail.
func decay(t float64, y *core.Array) (*core.Array, error) {
return core.MulF(y, -1), nil
}
// TestIntegrateODEExponential checks the adaptive solver against the
// analytic exponential decay, forward and backward in time.
func TestIntegrateODEExponential(t *testing.T) {
y0 := mustFloats(t, []float64{1})
end, err := IntegrateODE(decay, 0, 1, y0, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODE: %v", err)
}
if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-9 {
t.Fatalf("y(1) = %.14g, want %.14g", end.FloatAt(0), math.Exp(-1))
}
// Backward integration from t=1 to t=0 must return the start.
start, err := IntegrateODE(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}),
ODEOptions{RelTol: 1e-10, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODE backward: %v", err)
}
if math.Abs(start.FloatAt(0)-1) > 1e-8 {
t.Fatalf("backward y(0) = %.14g, want 1", start.FloatAt(0))
}
}
// TestIntegrateODEOscillator checks a two-dimensional linear system
// against the analytic phase rotation.
func TestIntegrateODEOscillator(t *testing.T) {
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
y0 := mustFloats(t, []float64{1, 0})
end, err := IntegrateODE(f, 0, math.Pi/2, y0, ODEOptions{RelTol: 1e-11, AbsTol: 1e-13})
if err != nil {
t.Fatalf("IntegrateODE: %v", err)
}
// y1 = cos t, y2 = −sin t, so a quarter period lands on (0, −1).
if math.Abs(end.FloatAt(0)) > 1e-7 || math.Abs(end.FloatAt(1)+1) > 1e-7 {
t.Fatalf("quarter period = (%.10g, %.10g), want (0, -1)",
end.FloatAt(0), end.FloatAt(1))
}
// A full period returns to the start.
full, err := IntegrateODE(f, 0, 2*math.Pi, y0, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODE full: %v", err)
}
for i := range 2 {
if math.Abs(full.FloatAt(i)-y0.FloatAt(i)) > 1e-6 {
t.Fatalf("full period[%d] = %v, want %v", i, full.FloatAt(i), y0.FloatAt(i))
}
}
}
// TestIntegrateRk4 checks the fixed-step scheme's fourth-order
// convergence: halving h must shrink the error by roughly sixteen.
func TestIntegrateRk4(t *testing.T) {
errAt := func(steps int) float64 {
end, err := IntegrateRK4(decay, 0, 1, mustFloats(t, []float64{1}), steps)
if err != nil {
t.Fatalf("IntegrateRK4(%d): %v", steps, err)
}
return math.Abs(end.FloatAt(0) - math.Exp(-1))
}
e10, e20 := errAt(10), errAt(20)
if e10 < 1e-13 {
t.Skipf("error already at round-off (%v)", e10)
}
ratio := e10 / e20
if ratio < 12 || ratio > 20 {
t.Fatalf("error ratio over a halved step = %.2g, want ≈ 16 for a fourth-order scheme", ratio)
}
}
// TestIntegrateBackwardEulerStiff demonstrates the reason an implicit
// scheme exists: on y' = −1000y the implicit Euler stays bounded and
// matches its closed-form damping at a step where explicit schemes
// blow up.
func TestIntegrateBackwardEulerStiff(t *testing.T) {
const lambda = 1000.0
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.MulF(y, -lambda), nil
}
// Ten steps of h = 0.01: h·λ = 10, far outside RK4's stability
// region but perfectly damped for the implicit scheme.
y0 := mustFloats(t, []float64{1})
end, err := IntegrateBackwardEuler(f, 0, 0.1, y0, 10, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateBackwardEuler: %v", err)
}
// Closed form of one implicit Euler step: y_{n+1} = y_n/(1+hλ),
// so y_10 = (1/11)^10.
damped := math.Pow(1/(1+0.01*lambda), 10)
if math.Abs(end.FloatAt(0)-damped) > 1e-9*math.Abs(damped) {
t.Fatalf("stiff result = %.12g, want %.12g", end.FloatAt(0), damped)
}
if end.FloatAt(0) <= 0 {
t.Fatalf("the implicit scheme must stay positive on decay, got %v", end.FloatAt(0))
}
}
// TestIntegrateBackwardEulerAccuracy checks that with a sane step the
// implicit scheme tracks the analytic decay as well.
func TestIntegrateBackwardEulerAccuracy(t *testing.T) {
end, err := IntegrateBackwardEuler(decay, 0, 1, mustFloats(t, []float64{1}), 1000, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateBackwardEuler: %v", err)
}
// First order: the global error is O(h), about h/2 for decay.
if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-3 {
t.Fatalf("y(1) = %.14g, want %.14g ± 1e-3", end.FloatAt(0), math.Exp(-1))
}
}
// TestIntegrateBackwardEulerVectorState pins the Newton solve on a
// two-dimensional stiff system: both modes decay at their own rate,
// and constant steps give each component the closed form
// y(t1) = y0·(1+h·λ)^{−steps} an implicit Euler step has on y' = −λy.
func TestIntegrateBackwardEulerVectorState(t *testing.T) {
const slow = 1.0
const fast = 1000.0
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{-fast * y.FloatAt(0), -slow * y.FloatAt(1)}, 2)
}
y0 := mustFloats(t, []float64{1, 1})
end, err := IntegrateBackwardEuler(f, 0, 1, y0, 10, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateBackwardEuler: %v", err)
}
wantFast := math.Pow(1/(1+0.1*fast), 10)
if math.Abs(end.FloatAt(0)-wantFast) > 1e-9*math.Abs(wantFast) {
t.Fatalf("fast mode = %.14g, want %.14g", end.FloatAt(0), wantFast)
}
wantSlow := math.Pow(1/(1+0.1*slow), 10)
if math.Abs(end.FloatAt(1)-wantSlow) > 1e-9*math.Abs(wantSlow) {
t.Fatalf("slow mode = %.14g, want %.14g", end.FloatAt(1), wantSlow)
}
}
// TestODEErrors pins the error contracts shared by the solvers.
func TestODEErrors(t *testing.T) {
wrongShape := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{1, 1}, 2)
}
y0 := mustFloats(t, []float64{1})
if _, err := IntegrateODE(wrongShape, 0, 1, y0, ODEOptions{}); err == nil {
t.Fatal("expected an error when f returns the wrong shape")
}
if _, err := IntegrateODE(wrongShape, 0, 1, y0, ODEOptions{}); err == nil {
t.Fatal("expected an error when f returns the wrong shape in RK4 as well")
}
if _, err := IntegrateRK4(decay, 0, 1, y0, 0); err == nil {
t.Fatal("expected an error for zero steps")
}
matrixState := mustFloats(t, []float64{1, 1}, 1, 2)
if _, err := IntegrateODE(decay, 0, 1, matrixState, ODEOptions{}); err == nil {
t.Fatal("expected an error for a rank-2 state")
}
empty := mustFloats(t, nil)
if _, err := IntegrateODE(decay, 0, 1, empty, ODEOptions{}); err == nil {
t.Fatal("expected an error for an empty state")
}
// A tight budget on a slow decay must report the budget, not lie.
if _, err := IntegrateODE(decay, 0, 1, y0, ODEOptions{MaxSteps: 3}); err == nil {
t.Fatal("expected an error for an exhausted step budget")
}
}
// TestIntegrateODEPathExponential checks the sampled trajectory of
// y' = −y against the analytic decay at every sample point.
func TestIntegrateODEPathExponential(t *testing.T) {
const n = 5
times, states, err := IntegrateODEPath(decay, 0, 2, mustFloats(t, []float64{1}),
n, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODEPath: %v", err)
}
if len(times) != n || len(states) != n {
t.Fatalf("lengths (%d, %d), want (%d, %d)", len(times), len(states), n, n)
}
for i := range n {
want := 0.5 * float64(i)
if math.Abs(times[i]-want) > 1e-12 {
t.Fatalf("times[%d] = %.14g, want %.14g", i, times[i], want)
}
got := states[i].FloatAt(0)
if math.Abs(got-math.Exp(-want)) > 1e-9 {
t.Fatalf("y(%.1f) = %.14g, want %.14g", want, got, math.Exp(-want))
}
}
}
// TestIntegrateODEPathBackward samples a backwards integration: the
// times descend and the analytic law still holds per sample.
func TestIntegrateODEPathBackward(t *testing.T) {
times, states, err := IntegrateODEPath(decay, 2, 0, mustFloats(t, []float64{math.Exp(-2)}),
3, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODEPath backward: %v", err)
}
for i, want := range []float64{2, 1, 0} {
if math.Abs(times[i]-want) > 1e-12 {
t.Fatalf("times[%d] = %.14g, want %g", i, times[i], want)
}
if math.Abs(states[i].FloatAt(0)-math.Exp(-want)) > 1e-9 {
t.Fatalf("y(%g) = %.14g, want %.14g", want, states[i].FloatAt(0), math.Exp(-want))
}
}
}
// TestIntegrateODEPathStatesIndependent checks the returned states do
// not alias one another or the initial vector: a later integration
// must never rewrite an earlier sample.
func TestIntegrateODEPathStatesIndependent(t *testing.T) {
y0 := mustFloats(t, []float64{1, 0})
f := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
_, states, err := IntegrateODEPath(f, 0, 1, y0, 4, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODEPath: %v", err)
}
states[3].RawFloats()[0] = 99 // must not leak into y0 or the other samples
if y0.FloatAt(0) != 1 {
t.Fatal("mutating a sample changed the caller's initial state")
}
if states[2].FloatAt(0) == 99 {
t.Fatal("samples alias one another")
}
}
// TestIntegrateODEPathErrors pins the path-specific error contract.
func TestIntegrateODEPathErrors(t *testing.T) {
y0 := mustFloats(t, []float64{1})
if _, _, err := IntegrateODEPath(decay, 0, 1, y0, 1, ODEOptions{}); err == nil {
t.Fatal("expected an error for a single sample")
}
// An f that fails past the midpoint must fail the whole path.
boom := func(t float64, y *core.Array) (*core.Array, error) {
if t > 0.6 {
return nil, base.Errf("detector tripped")
}
return core.MulF(y, -1), nil
}
if _, _, err := IntegrateODEPath(boom, 0, 1, y0, 5, ODEOptions{}); err == nil {
t.Fatal("expected the operator error to propagate")
}
if _, _, err := IntegrateODEPath(decay, 0, 1, mustFloats(t, nil), 3, ODEOptions{}); err == nil {
t.Fatal("expected an error for an empty state")
}
}
// TestIntegrateODEEventsProjectile drops a projectile and watches its
// height: the crossing of zero is the flight time 2v₀/g, analytically
// known, with the state's velocity the exact mirror of the launch.
func TestIntegrateODEEventsProjectile(t *testing.T) {
const g = 9.81
const v0 = 10
f := func(now float64, y *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{y.FloatAt(1), -g}), nil
}
height := func(now float64, y *core.Array) (float64, error) {
return y.FloatAt(0), nil
}
hits, final, err := IntegrateODEEvents(f, 0, 5, mustFloats(t, []float64{0, v0}),
[]ODEWatch{{Function: height, Direction: -1}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14})
if err != nil {
t.Fatalf("IntegrateODEEvents: %v", err)
}
if len(hits) != 1 {
t.Fatalf("hits = %d, want 1", len(hits))
}
wantT := 2 * v0 / g
if math.Abs(hits[0].Time-wantT) > 1e-9 {
t.Fatalf("impact at %.14g, want %.14g", hits[0].Time, wantT)
}
if math.Abs(hits[0].State.FloatAt(1)+v0) > 1e-7 {
t.Fatalf("impact speed %.10g, want %.10g", hits[0].State.FloatAt(1), float64(-v0))
}
if hits[0].Rising {
t.Fatal("height crossing on the way down must not be rising")
}
if final.Len() != 2 {
t.Fatalf("final state shape %v", final.Shape())
}
}
// TestIntegrateODEEventsOscillator watches the oscillator's position
// over five periods: cos crosses zero once per half period, and the
// falling-only filter keeps every second one, at π/2 + 2πk.
func TestIntegrateODEEventsOscillator(t *testing.T) {
f := func(now float64, y *core.Array) (*core.Array, error) {
return mustFloats(t, []float64{y.FloatAt(1), -y.FloatAt(0)}), nil
}
position := func(now float64, y *core.Array) (float64, error) {
return y.FloatAt(0), nil
}
hits, _, err := IntegrateODEEvents(f, 0, 5*2*math.Pi, mustFloats(t, []float64{1, 0}),
[]ODEWatch{{Function: position, Direction: -1}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14})
if err != nil {
t.Fatalf("IntegrateODEEvents: %v", err)
}
if len(hits) != 5 {
t.Fatalf("hits = %d, want 5 falling crossings", len(hits))
}
for k, hit := range hits {
want := math.Pi/2 + float64(2*k)*math.Pi
if math.Abs(hit.Time-want) > 1e-8 {
t.Fatalf("hit %d at %.12g, want %.12g", k, hit.Time, want)
}
}
// Without the filter every half period fires: ten crossings.
hits, _, err = IntegrateODEEvents(f, 0, 5*2*math.Pi, mustFloats(t, []float64{1, 0}),
[]ODEWatch{{Function: position}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14})
if err != nil {
t.Fatalf("IntegrateODEEvents unfiltered: %v", err)
}
if len(hits) != 10 {
t.Fatalf("unfiltered hits = %d, want 10", len(hits))
}
for i := 1; i < len(hits); i++ {
if hits[i].Time <= hits[i-1].Time {
t.Fatal("hits are not sorted by time")
}
}
}
// TestIntegrateODEStepsErrors pins the error contract of the step
// recorder.
func TestIntegrateODEStepsErrors(t *testing.T) {
y0 := mustFloats(t, []float64{1})
if _, _, err := IntegrateODESteps(decay, 0, 1, y0, ODEOptions{MaxSteps: 2}); err == nil {
t.Fatal("expected an error for an exhausted step budget")
}
if _, _, err := IntegrateODESteps(decay, 0, 1, mustFloats(t, nil), ODEOptions{}); err == nil {
t.Fatal("expected an error for an empty state")
}
}
// TestIntegrateODEStepsForward records the accepted steps of a decay:
// the trace starts at the initial state, ends exactly at t1, and every
// node sits on the analytic curve at solver tolerance.
func TestIntegrateODEStepsForward(t *testing.T) {
times, states, err := IntegrateODESteps(decay, 0, 1, mustFloats(t, []float64{1}),
ODEOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODESteps: %v", err)
}
if len(times) != len(states) || len(times) < 3 {
t.Fatalf("trace has %d times and %d states, want matching lengths of at least 3",
len(times), len(states))
}
if times[0] != 0 || times[len(times)-1] != 1 {
t.Fatalf("trace spans [%g, %g], want [0, 1]", times[0], times[len(times)-1])
}
for i := 1; i < len(times); i++ {
if times[i] <= times[i-1] {
t.Fatalf("times not strictly increasing at %d: %g after %g", i, times[i], times[i-1])
}
}
if states[0].FloatAt(0) != 1 {
t.Fatalf("first state = %v, want the initial state 1", states[0].FloatAt(0))
}
for i := range times {
if math.Abs(states[i].FloatAt(0)-math.Exp(-times[i])) > 1e-6 {
t.Fatalf("y(%g) = %.14g, want %.14g", times[i], states[i].FloatAt(0), math.Exp(-times[i]))
}
}
}
// TestIntegrateODEStepsDegenerate covers the zero-span trace and a
// backward pass with descending times.
func TestIntegrateODEStepsDegenerate(t *testing.T) {
times, states, err := IntegrateODESteps(decay, 1, 1, mustFloats(t, []float64{1}), ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODESteps zero span: %v", err)
}
if len(times) != 1 || times[0] != 1 || states[0].FloatAt(0) != 1 {
t.Fatalf("zero-span trace = %v, want the single initial node", times)
}
times, states, err = IntegrateODESteps(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}),
ODEOptions{RelTol: 1e-8, AbsTol: 1e-12})
if err != nil {
t.Fatalf("IntegrateODESteps backward: %v", err)
}
if times[0] != 1 || times[len(times)-1] != 0 {
t.Fatalf("backward trace spans [%g, %g], want [1, 0]", times[0], times[len(times)-1])
}
for i := 1; i < len(times); i++ {
if times[i] >= times[i-1] {
t.Fatalf("backward times not strictly decreasing at %d", i)
}
}
}
// TestIntegrateODEEventsErrors pins the validation and error paths.
func TestIntegrateODEEventsErrors(t *testing.T) {
f := decay
if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}),
nil, ODEOptions{}); err == nil {
t.Fatal("no watches: want an error")
}
if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}),
[]ODEWatch{{}}, ODEOptions{}); err == nil {
t.Fatal("empty watch: want an error")
}
boom := func(t float64, y *core.Array) (float64, error) {
return 0, base.Errf("watch failed")
}
if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}),
[]ODEWatch{{Function: boom}}, ODEOptions{}); err == nil {
t.Fatal("watch error: want an error")
}
// An f that fails on the refinement path must surface.
broken := func(t float64, y *core.Array) (*core.Array, error) {
if t > 0.5 {
return nil, base.Errf("integrand failed")
}
return core.MulF(y, -1), nil
}
cross := func(t float64, y *core.Array) (float64, error) {
return 1 - t, nil
}
if _, _, err := IntegrateODEEvents(broken, 0, 2, mustFloats(t, []float64{1}),
[]ODEWatch{{Function: cross}}, ODEOptions{}); err == nil {
t.Fatal("refinement error: want an error")
}
}
// TestODEArrivalGuard pins the ulp-level arrival rule: an endpoint
// missed by rounding terminates the loop, a genuine remaining distance
// does not, and a plain integration across magnitudes still lands.
func TestODEArrivalGuard(t *testing.T) {
if !odeArrived(1, 1) {
t.Fatal("equal times must count as arrived")
}
if !odeArrived(1, 1+4*2.220446049250313e-16) {
t.Fatal("a four-ulp miss must count as arrived")
}
if odeArrived(1, 1.1) {
t.Fatal("a genuine remaining distance must not count as arrived")
}
if odeArrived(0, -1e-20) {
t.Fatal("a tiny but representable distance near zero must not count as arrived")
}
// An integration whose endpoints differ well beyond Sterbenz still
// terminates and reports the endpoint state.
f := func(tt float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(0) * 0}, 1)
}
y0, _ := core.FromFloats([]float64{3}, 1)
got, err := IntegrateODE(f, 1e9, 1e9+0.5, y0, ODEOptions{})
if err != nil {
t.Fatalf("IntegrateODE across magnitudes: %v", err)
}
if got.FloatAt(0) != 3 {
t.Fatalf("constant state changed: %g", got.FloatAt(0))
}
}