476 lines
17 KiB
Go
476 lines
17 KiB
Go
// 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))
|
|||
|
|
}
|
|||
|
|
}
|