371 lines
14 KiB
Go
371 lines
14 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package integrate
|
||
|
||
import (
|
||
"math"
|
||
"testing"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
)
|
||
|
||
// TestYoshidaWeightsIdentity pins the triple-jump weights: the
|
||
// composition V(w1·h)·V(w0·h)·V(w1·h) is fourth order exactly when
|
||
// w0 + 2·w1 = 1, with w1 positive and w0 negative.
|
||
func TestYoshidaWeightsIdentity(t *testing.T) {
|
||
if math.Abs(yoshidaW0+2*yoshidaW1-1) > 1e-15 {
|
||
t.Fatalf("w0 + 2·w1 = %.17g, want 1", yoshidaW0+2*yoshidaW1)
|
||
}
|
||
if yoshidaW1 <= 0 || yoshidaW0 >= 0 {
|
||
t.Fatalf("weights %g and %g, want w1 positive and w0 negative", yoshidaW0, yoshidaW1)
|
||
}
|
||
c := math.Cbrt(2)
|
||
if math.Abs(yoshidaW1-1/(2-c)) > 1e-15 || math.Abs(yoshidaW0+c/(2-c)) > 1e-15 {
|
||
t.Fatalf("weights %.17g and %.17g are not the triple-jump choice", yoshidaW0, yoshidaW1)
|
||
}
|
||
}
|
||
|
||
// harmonicEnergy returns the energy of a one-dimensional harmonic
|
||
// state.
|
||
func harmonicEnergy(q, p []float64, omega float64) float64 {
|
||
return 0.5 * (p[0]*p[0] + omega*omega*q[0]*q[0])
|
||
}
|
||
|
||
// yoshidaEnergyBand integrates the harmonic oscillator over periods
|
||
// and returns the largest energy deviation from the initial one.
|
||
func yoshidaEnergyBand(t *testing.T, perPeriod, periods int, verlet bool) float64 {
|
||
t.Helper()
|
||
const omega = 1.0
|
||
accel := func(q *core.Array) (*core.Array, error) {
|
||
return core.MulF(q, -omega*omega), nil
|
||
}
|
||
q0 := mustFloats(t, []float64{1})
|
||
p0 := mustFloats(t, []float64{0})
|
||
steps := perPeriod * periods
|
||
var positions, momenta []*core.Array
|
||
var err error
|
||
if verlet {
|
||
positions, momenta, err = IntegrateVerlet(accel, 0, float64(periods)*2*math.Pi, q0, p0, steps)
|
||
} else {
|
||
positions, momenta, err = IntegrateYoshida4(accel, 0, float64(periods)*2*math.Pi, q0, p0, steps)
|
||
}
|
||
if err != nil {
|
||
t.Fatalf("integrator: %v", err)
|
||
}
|
||
e0 := harmonicEnergy([]float64{positions[0].FloatAt(0)}, []float64{momenta[0].FloatAt(0)}, omega)
|
||
worst := 0.0
|
||
q := make([]float64, 1)
|
||
p := make([]float64, 1)
|
||
for s := range positions {
|
||
q[0] = positions[s].FloatAt(0)
|
||
p[0] = momenta[s].FloatAt(0)
|
||
if d := math.Abs(harmonicEnergy(q, p, omega) - e0); d > worst {
|
||
worst = d
|
||
}
|
||
}
|
||
return worst
|
||
}
|
||
|
||
// TestYoshidaHarmonicFourthOrder pins the order: the harmonic
|
||
// oscillator's energy-band error scales as h⁴ for Yoshida 4 (ratios
|
||
// near 16 on halved steps) where Verlet's scales as h² (ratios near
|
||
// 4), measured on the same problem.
|
||
func TestYoshidaHarmonicFourthOrder(t *testing.T) {
|
||
for _, tc := range []struct {
|
||
name string
|
||
verlet bool
|
||
low float64
|
||
high float64
|
||
}{{"Verlet", true, 3, 5.5}, {"Yoshida4", false, 12, 20}} {
|
||
previous := 0.0
|
||
for _, perPeriod := range []int{20, 40, 80} {
|
||
band := yoshidaEnergyBand(t, perPeriod, 4, tc.verlet)
|
||
t.Logf("%s steps/period %d: energy band %.3g", tc.name, perPeriod, band)
|
||
if previous > 0 {
|
||
if r := previous / band; r < tc.low || r > tc.high {
|
||
t.Fatalf("%s: error ratio %.2f at %d steps/period, want the band [%g, %g]",
|
||
tc.name, r, perPeriod, tc.low, tc.high)
|
||
}
|
||
}
|
||
previous = band
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestYoshidaKeplerEnergyBand pins the long-time fidelity: the
|
||
// Kepler two-body problem on an eccentric orbit keeps its energy in
|
||
// a narrow band over twenty periods, and the orbit closes.
|
||
func TestYoshidaKeplerEnergyBand(t *testing.T) {
|
||
const eccentricity = 0.5
|
||
accel := func(q *core.Array) (*core.Array, error) {
|
||
x, y := q.FloatAt(0), q.FloatAt(1)
|
||
r3 := math.Pow(x*x+y*y, 1.5)
|
||
return core.FromFloats([]float64{-x / r3, -y / r3}, 2)
|
||
}
|
||
q0 := mustFloats(t, []float64{1 + eccentricity, 0})
|
||
p0 := mustFloats(t, []float64{0, math.Sqrt((1 - eccentricity) / (1 + eccentricity))})
|
||
periods := 20
|
||
stepsPerPeriod := 200
|
||
positions, momenta, err := IntegrateYoshida4(accel, 0, float64(periods)*2*math.Pi, q0, p0, periods*stepsPerPeriod)
|
||
if err != nil {
|
||
t.Fatalf("IntegrateYoshida4: %v", err)
|
||
}
|
||
energy := func(s int) float64 {
|
||
vx := momenta[s].FloatAt(0)
|
||
vy := momenta[s].FloatAt(1)
|
||
x := positions[s].FloatAt(0)
|
||
y := positions[s].FloatAt(1)
|
||
return 0.5*(vx*vx+vy*vy) - 1/math.Hypot(x, y)
|
||
}
|
||
e0 := energy(0)
|
||
worst := 0.0
|
||
for s := range positions {
|
||
if d := math.Abs(energy(s) - e0); d > worst {
|
||
worst = d
|
||
}
|
||
}
|
||
t.Logf("Kepler e=0.5 over %d periods: energy band %.3g (E0 = %.6g)", periods, worst, e0)
|
||
if worst > 1e-4*math.Abs(e0) {
|
||
t.Fatalf("energy drifted by %.3g over %d periods", worst, periods)
|
||
}
|
||
// The closing error is the accumulated per-period phase error,
|
||
// fourth order in h, not an energy drift: at 200 steps per period
|
||
// it stays a few thousandths of the orbit radius.
|
||
last := len(positions) - 1
|
||
if math.Hypot(positions[last].FloatAt(0)-(1+eccentricity), positions[last].FloatAt(1)) > 5e-3 {
|
||
t.Fatalf("the orbit did not close: q = (%.8g, %.8g)",
|
||
positions[last].FloatAt(0), positions[last].FloatAt(1))
|
||
}
|
||
}
|
||
|
||
func TestYoshidaErrors(t *testing.T) {
|
||
accel := func(q *core.Array) (*core.Array, error) {
|
||
return core.MulF(q, -1), nil
|
||
}
|
||
q0 := mustFloats(t, []float64{1})
|
||
p0 := mustFloats(t, []float64{0})
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, q0, p0, 0); err == nil {
|
||
t.Fatal("zero steps: want an error")
|
||
}
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, q0, mustFloats(t, []float64{0, 1}), 5); err == nil {
|
||
t.Fatal("mismatched vectors: want an error")
|
||
}
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, mustFloats(t, []float64{}), mustFloats(t, []float64{}), 5); err == nil {
|
||
t.Fatal("empty state: want an error")
|
||
}
|
||
boom := func(*core.Array) (*core.Array, error) { return nil, base.Errf("accel failed") }
|
||
if _, _, err := IntegrateYoshida4(boom, 0, 1, q0, p0, 5); err == nil || !stringsContains(err, "accel failed") {
|
||
t.Fatal("a failing accel: want the error propagated")
|
||
}
|
||
wrong := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{1, 1}), nil }
|
||
if _, _, err := IntegrateYoshida4(wrong, 0, 1, q0, p0, 5); err == nil {
|
||
t.Fatal("wrong accel shape: want an error")
|
||
}
|
||
// A non-finite acceleration is refused instead of publishing NaNs.
|
||
nan := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{math.NaN()}), nil }
|
||
if _, _, err := IntegrateYoshida4(nan, 0, 1, q0, p0, 5); err == nil || !stringsContains(err, "non-finite") {
|
||
t.Fatal("a NaN accel: want an error")
|
||
}
|
||
// Int and complex states are refused, float32 states integrate.
|
||
pi0, err := core.FromInts([]int64{1}, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, pi0, p0, 5); err == nil {
|
||
t.Fatal("an int state: want an error")
|
||
}
|
||
q32, err := core.FromFloat32s([]float32{1}, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
p32, err := core.FromFloat32s([]float32{0}, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, q32, p32, 5); err != nil {
|
||
t.Fatalf("a float32 state: %v", err)
|
||
}
|
||
}
|
||
|
||
// midpointQP is the gradient of H(q, p) = q·p: the flow is
|
||
// q' = q, p' = −p, whose midpoint solution is the Cayley transform
|
||
// q_n = q0·((1+h/2)/(1−h/2))ⁿ, p_n = p0·((1−h/2)/(1+h/2))ⁿ, and
|
||
// H = q·p is conserved exactly.
|
||
func midpointQP(z *core.Array) (*core.Array, error) {
|
||
return core.FromFloats([]float64{z.FloatAt(1), z.FloatAt(0)}, 2)
|
||
}
|
||
|
||
// TestMidpointCayleyExact pins the analytic solution: for the
|
||
// nonseparable H = q·p the midpoint iterates must land on the Cayley
|
||
// transform and conserve H to rounding.
|
||
func TestMidpointCayleyExact(t *testing.T) {
|
||
const (
|
||
q0 = 1.2
|
||
p0 = 0.7
|
||
steps = 50
|
||
t1 = 1.0
|
||
)
|
||
h := t1 / steps
|
||
positions, momenta, err := IntegrateMidpoint(midpointQP, 0, t1,
|
||
mustFloats(t, []float64{q0}), mustFloats(t, []float64{p0}), steps, MidpointOptions{})
|
||
if err != nil {
|
||
t.Fatalf("IntegrateMidpoint: %v", err)
|
||
}
|
||
up := (1 + h/2) / (1 - h/2)
|
||
down := (1 - h/2) / (1 + h/2)
|
||
for s := range positions {
|
||
q := positions[s].FloatAt(0)
|
||
p := momenta[s].FloatAt(0)
|
||
wantQ := q0 * math.Pow(up, float64(s))
|
||
wantP := p0 * math.Pow(down, float64(s))
|
||
if math.Abs(q-wantQ) > 1e-12 || math.Abs(p-wantP) > 1e-12 {
|
||
t.Fatalf("step %d: (q, p) = (%.14g, %.14g), want (%.14g, %.14g)", s, q, p, wantQ, wantP)
|
||
}
|
||
if math.Abs(q*p-q0*p0) > 1e-12 {
|
||
t.Fatalf("step %d: H = %.17g, want %.17g", s, q*p, q0*p0)
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestMidpointFixedPointExact pins the equilibrium property: an
|
||
// equilibrium of the flow is a fixed point of the midpoint map,
|
||
// exactly, with no roundoff drift.
|
||
func TestMidpointFixedPointExact(t *testing.T) {
|
||
positions, momenta, err := IntegrateMidpoint(midpointQP, 0, 0.5,
|
||
mustFloats(t, []float64{0}), mustFloats(t, []float64{0}), 3, MidpointOptions{})
|
||
if err != nil {
|
||
t.Fatalf("IntegrateMidpoint: %v", err)
|
||
}
|
||
for s := range positions {
|
||
if positions[s].FloatAt(0) != 0 || momenta[s].FloatAt(0) != 0 {
|
||
t.Fatalf("step %d: the fixed point moved to (%g, %g)",
|
||
s, positions[s].FloatAt(0), momenta[s].FloatAt(0))
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestMidpointNonseparableBand pins the energy band on a genuinely
|
||
// nonseparable two-degree system: H = ½(q₁²+1)(p₁²+1) +
|
||
// ½(q₂²+1)(p₂²+1) cannot be written as T(p) + V(q), and the midpoint
|
||
// rule must keep H inside a bounded band over the whole run.
|
||
func TestMidpointNonseparableBand(t *testing.T) {
|
||
gradH := func(z *core.Array) (*core.Array, error) {
|
||
q1, q2 := z.FloatAt(0), z.FloatAt(1)
|
||
p1, p2 := z.FloatAt(2), z.FloatAt(3)
|
||
return core.FromFloats([]float64{
|
||
q1 * (p1*p1 + 1), q2 * (p2*p2 + 1),
|
||
p1 * (q1*q1 + 1), p2 * (q2*q2 + 1),
|
||
}, 4)
|
||
}
|
||
hamiltonian := func(q1, q2, p1, p2 float64) float64 {
|
||
return 0.5*(q1*q1+1)*(p1*p1+1) + 0.5*(q2*q2+1)*(p2*p2+1)
|
||
}
|
||
q0 := mustFloats(t, []float64{0.5, -0.4})
|
||
p0 := mustFloats(t, []float64{0.3, 0.8})
|
||
h0 := hamiltonian(0.5, -0.4, 0.3, 0.8)
|
||
positions, momenta, err := IntegrateMidpoint(gradH, 0, 10, q0, p0, 1000, MidpointOptions{})
|
||
if err != nil {
|
||
t.Fatalf("IntegrateMidpoint: %v", err)
|
||
}
|
||
worst := 0.0
|
||
for s := range positions {
|
||
h := hamiltonian(positions[s].FloatAt(0), positions[s].FloatAt(1),
|
||
momenta[s].FloatAt(0), momenta[s].FloatAt(1))
|
||
if d := math.Abs(h - h0); d > worst {
|
||
worst = d
|
||
}
|
||
}
|
||
t.Logf("nonseparable two-degree band over t = 10 at h = 0.01: %.3g (H0 = %.6g)", worst, h0)
|
||
if worst > 1e-3*h0 {
|
||
t.Fatalf("energy drifted by %.3g (H0 = %.6g)", worst, h0)
|
||
}
|
||
}
|
||
|
||
func TestMidpointErrors(t *testing.T) {
|
||
q0 := mustFloats(t, []float64{1})
|
||
p0 := mustFloats(t, []float64{0})
|
||
if _, _, err := IntegrateMidpoint(midpointQP, 0, 1, q0, p0, 0, MidpointOptions{}); err == nil {
|
||
t.Fatal("zero steps: want an error")
|
||
}
|
||
if _, _, err := IntegrateMidpoint(midpointQP, 0, 1, q0, mustFloats(t, []float64{0, 1}), 5, MidpointOptions{}); err == nil {
|
||
t.Fatal("mismatched vectors: want an error")
|
||
}
|
||
if _, _, err := IntegrateMidpoint(midpointQP, 0, 1, mustFloats(t, []float64{}), mustFloats(t, []float64{}), 5, MidpointOptions{}); err == nil {
|
||
t.Fatal("empty state: want an error")
|
||
}
|
||
pi0, err := core.FromInts([]int64{1}, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if _, _, err := IntegrateMidpoint(midpointQP, 0, 1, pi0, p0, 5, MidpointOptions{}); err == nil {
|
||
t.Fatal("an int state: want an error")
|
||
}
|
||
boom := func(*core.Array) (*core.Array, error) { return nil, base.Errf("gradH failed") }
|
||
if _, _, err := IntegrateMidpoint(boom, 0, 1, q0, p0, 5, MidpointOptions{}); err == nil || !stringsContains(err, "gradH failed") {
|
||
t.Fatal("a failing gradient: want the error propagated")
|
||
}
|
||
wrong := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{1}), nil }
|
||
if _, _, err := IntegrateMidpoint(wrong, 0, 1, q0, p0, 5, MidpointOptions{}); err == nil {
|
||
t.Fatal("wrong gradient shape: want an error")
|
||
}
|
||
nan := func(*core.Array) (*core.Array, error) { return mustFloats(t, []float64{math.NaN(), 1}), nil }
|
||
if _, _, err := IntegrateMidpoint(nan, 0, 1, q0, p0, 5, MidpointOptions{}); err == nil || !stringsContains(err, "non-finite") {
|
||
t.Fatal("a NaN gradient: want an error")
|
||
}
|
||
}
|
||
|
||
// TestVerletFamilyComplexRefusal pins the dtype contract across the
|
||
// whole family: complex states are refused everywhere.
|
||
func TestVerletFamilyComplexRefusal(t *testing.T) {
|
||
accel := func(q *core.Array) (*core.Array, error) { return core.MulF(q, -1), nil }
|
||
gradH := func(z *core.Array) (*core.Array, error) { return core.FromFloats([]float64{0, 0}, 2) }
|
||
qi, err := core.FromFloat32s([]float32{1}, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
pi, err := core.FromFloat32s([]float32{0}, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, qi, pi, 5); err != nil {
|
||
t.Fatalf("a float32 Yoshida state: %v", err)
|
||
}
|
||
if _, _, err := IntegrateMidpoint(gradH, 0, 1, qi, pi, 5, MidpointOptions{}); err != nil {
|
||
t.Fatalf("a float32 midpoint state: %v", err)
|
||
}
|
||
// A complex state is refused on both integrators.
|
||
complexArr := core.New(core.Complex, 1)
|
||
zeros := mustFloats(t, []float64{0})
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, complexArr, zeros, 5); err == nil || !stringsContains(err, "complex") {
|
||
t.Fatal("a complex Yoshida state: want an error")
|
||
}
|
||
if _, _, err := IntegrateMidpoint(gradH, 0, 1, complexArr, zeros, 5, MidpointOptions{}); err == nil || !stringsContains(err, "complex") {
|
||
t.Fatal("a complex midpoint state: want an error")
|
||
}
|
||
}
|
||
|
||
// TestVerletFamilyNonFiniteRefusal pins the up-front refusal of
|
||
// non-finite initial states on both new integrators.
|
||
func TestVerletFamilyNonFiniteRefusal(t *testing.T) {
|
||
accel := func(q *core.Array) (*core.Array, error) { return core.MulF(q, -1), nil }
|
||
gradH := func(z *core.Array) (*core.Array, error) { return core.FromFloats([]float64{0, 0}, 2) }
|
||
nanQ := mustFloats(t, []float64{math.NaN()})
|
||
p0 := mustFloats(t, []float64{0})
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, nanQ, p0, 5); err == nil || !stringsContains(err, "non-finite") {
|
||
t.Fatal("a NaN position: want an error")
|
||
}
|
||
if _, _, err := IntegrateMidpoint(gradH, 0, 1, nanQ, p0, 5, MidpointOptions{}); err == nil || !stringsContains(err, "non-finite") {
|
||
t.Fatal("a NaN position: want an error")
|
||
}
|
||
q0 := mustFloats(t, []float64{1})
|
||
nanP := mustFloats(t, []float64{math.Inf(-1)})
|
||
if _, _, err := IntegrateYoshida4(accel, 0, 1, q0, nanP, 5); err == nil || !stringsContains(err, "non-finite") {
|
||
t.Fatal("an infinite momentum: want an error")
|
||
}
|
||
if _, _, err := IntegrateMidpoint(gradH, 0, 1, q0, nanP, 5, MidpointOptions{}); err == nil || !stringsContains(err, "non-finite") {
|
||
t.Fatal("an infinite momentum: want an error")
|
||
}
|
||
}
|