Files
tensor/integrate/odecolloc_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

274 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 integrate
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// collocCubicSystem is y” = 6t written as u' = v, v' = 6t, whose
// exact solution with u(0) = 0, u(1) = 1 is u = t³, v = 3t². The
// three-point Lobatto IIIA collocation reproduces a cubic exactly, so
// the discrete solve is the exact answer, not an approximation.
func collocCubicSystem(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), 6 * t}, 2)
}
// collocHermite evaluates the solution's piecewise cubic Hermite
// through (mesh, values, slopes) at time tau, the documented
// continuous representation of the collocation answer.
func collocHermite(t *testing.T, sol *CollocationSolution, tau float64, component int) float64 {
t.Helper()
if tau < sol.Mesh[0] || tau > sol.Mesh[len(sol.Mesh)-1] {
t.Fatalf("time %g outside the mesh", tau)
}
lo, hi := 0, len(sol.Mesh)-1
for hi-lo > 1 {
mid := (lo + hi) / 2
if sol.Mesh[mid] <= tau {
lo = mid
} else {
hi = mid
}
}
h := sol.Mesh[lo+1] - sol.Mesh[lo]
th := (tau - sol.Mesh[lo]) / h
th2 := th * th
th3 := th2 * th
y0 := sol.Values[lo].FloatAt(component)
y1 := sol.Values[lo+1].FloatAt(component)
s0 := sol.Slopes[lo].FloatAt(component)
s1 := sol.Slopes[lo+1].FloatAt(component)
return (2*th3-3*th2+1)*y0 + h*(th3-2*th2+th)*s0 +
(-2*th3+3*th2)*y1 + h*(th3-th2)*s1
}
// TestSolveBoundaryCollocationCubicExact solves the linear problem
// u” = 6t with u(0) = 0, u(1) = 1 on a uniform mesh: the cubic
// collocation reproduces t³ exactly, the Newton residual drops to
// machine precision in one step, and no refinement is needed, so the
// mesh keeps its initial size.
func TestSolveBoundaryCollocationCubicExact(t *testing.T) {
sol, err := SolveBoundaryCollocation(collocCubicSystem, 0, 1,
mustFloats(t, []float64{0, 0}),
BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}},
CollocationOptions{RelTol: 1e-7, InitialNodes: 8, MaxNodes: 64})
if err != nil {
t.Fatalf("SolveBoundaryCollocation: %v", err)
}
if len(sol.Mesh) != 9 {
t.Fatalf("mesh grew to %d nodes on a cubic-exact problem, want the initial 9", len(sol.Mesh))
}
for k, tk := range sol.Mesh {
if math.Abs(sol.Values[k].FloatAt(0)-tk*tk*tk) > 1e-12 {
t.Fatalf("u(%.6g) = %.14g, want %.14g", tk, sol.Values[k].FloatAt(0), tk*tk*tk)
}
if math.Abs(sol.Slopes[k].FloatAt(0)-3*tk*tk) > 1e-12 {
t.Fatalf("u'(%.6g) = %.14g, want %.14g", tk, sol.Slopes[k].FloatAt(0), 3*tk*tk)
}
if math.Abs(sol.Values[k].FloatAt(1)-3*tk*tk) > 1e-12 {
t.Fatalf("v(%.6g) = %.14g, want %.14g", tk, sol.Values[k].FloatAt(1), 3*tk*tk)
}
}
// A one-component state carries exactly one condition: an
// endpoint-only prescription solves y' = y backward from t1.
sol1, err := SolveBoundaryCollocation(
func(t float64, y *core.Array) (*core.Array, error) { return core.MulF(y, 1), nil },
0, 1, mustFloats(t, []float64{1}),
BoundaryConditions{End: []int{0}, EndValues: []float64{math.E}},
CollocationOptions{})
if err != nil {
t.Fatalf("one-component solve: %v", err)
}
for k, tk := range sol1.Mesh {
if math.Abs(sol1.Values[k].FloatAt(0)-math.Exp(tk)) > 1e-6 {
t.Fatalf("y(%.6g) = %.14g, want %.14g", tk, sol1.Values[k].FloatAt(0), math.Exp(tk))
}
}
}
// TestSolveBoundaryCollocationBratuMatchesShooting solves Bratu's
// equation u” + e^u = 0 with u(0) = u(1) = 0, λ = 1, and requires
// the collocation answer to agree with the shooting method's answer
// through the existing IntegrateBoundary.
func TestSolveBoundaryCollocationBratuMatchesShooting(t *testing.T) {
bratu := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -math.Exp(y.FloatAt(0))}, 2)
}
sol, err := SolveBoundaryCollocation(bratu, 0, 1,
mustFloats(t, []float64{0, 0.4}),
BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{0}},
CollocationOptions{RelTol: 1e-7, AbsTol: 1e-10, InitialNodes: 10, MaxNodes: 600, MaxIterations: 60})
if err != nil {
t.Fatalf("collocation: %v", err)
}
if len(sol.Mesh) <= 10 {
t.Fatalf("the mesh never grew past the initial 10 intervals (refinement instrument): %d nodes", len(sol.Mesh))
}
times, states, err := IntegrateBoundary(bratu, 0, 1,
mustFloats(t, []float64{0, 0.4}),
BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{0}},
9, ODEOptions{RelTol: 1e-10, AbsTol: 1e-13})
if err != nil {
t.Fatalf("shooting: %v", err)
}
worstU, worstS := 0.0, 0.0
for i := 1; i < len(times)-1; i++ {
if d := math.Abs(collocHermite(t, sol, times[i], 0) - states[i].FloatAt(0)); d > worstU {
worstU = d
}
if d := math.Abs(collocHermite(t, sol, times[i], 1) - states[i].FloatAt(1)); d > worstS {
worstS = d
}
}
t.Logf("Bratu λ=1: worst state difference %.3g, worst slope difference %.3g", worstU, worstS)
if worstU > 1e-5 {
t.Fatalf("collocation and shooting disagree on u by %.3g", worstU)
}
if worstS > 1e-4 {
t.Fatalf("collocation and shooting disagree on u' by %.3g", worstS)
}
}
// TestSolveBoundaryCollocationRefinementLayer pins the refinement
// loop with a linear boundary-layer problem u” = −100·u' scaled as
// u' = v, v' = −100v, whose solution 1 − e^(−100t) needs intervals
// clustered near t = 0. A loose tolerance must leave the initial
// mesh alone; a tight one must refine it and land on the solution.
func TestSolveBoundaryCollocationRefinementLayer(t *testing.T) {
layer := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -100 * y.FloatAt(1)}, 2)
}
exact := func(t float64) float64 { return 1 - math.Exp(-100*t) }
bc := BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}}
y0 := mustFloats(t, []float64{0, 0})
loose, err := SolveBoundaryCollocation(layer, 0, 1, y0, bc,
CollocationOptions{AbsTol: 1e6, RelTol: 1, InitialNodes: 10, MaxNodes: 4000})
if err != nil {
t.Fatalf("loose solve: %v", err)
}
if len(loose.Mesh) != 11 {
t.Fatalf("a loose tolerance still refined to %d nodes", len(loose.Mesh))
}
tight, err := SolveBoundaryCollocation(layer, 0, 1, y0, bc,
CollocationOptions{RelTol: 1e-6, AbsTol: 1e-9, InitialNodes: 10, MaxNodes: 4000, MaxIterations: 60})
if err != nil {
t.Fatalf("tight solve: %v", err)
}
t.Logf("layer problem: mesh refined from 11 to %d nodes", len(tight.Mesh))
if len(tight.Mesh) <= 11 {
t.Fatal("the tight solve never refined the mesh (refinement instrument)")
}
worst := 0.0
for k, tk := range tight.Mesh {
if d := math.Abs(tight.Values[k].FloatAt(0) - exact(tk)); d > worst {
worst = d
}
}
t.Logf("layer problem: worst nodal error %.3g", worst)
if worst > 1e-3 {
t.Fatalf("refined layer error %.3g too large", worst)
}
for k := range tight.Mesh {
if k > 0 && tight.Mesh[k] <= tight.Mesh[k-1] {
t.Fatalf("the mesh is not increasing at %d", k)
}
}
}
// TestSolveBoundaryCollocationErrors pins the refusal contract.
func TestSolveBoundaryCollocationErrors(t *testing.T) {
const name = "SolveBoundaryCollocation"
good := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2)
}
y0 := mustFloats(t, []float64{0, 0.5})
bc := BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}}
// Inconsistent boundary conditions: wrong counts, repeated and
// out-of-range indices, wrong EndValues length.
if _, err := SolveBoundaryCollocation(good, 0, 1, y0,
BoundaryConditions{End: []int{0}, EndValues: []float64{1}}, CollocationOptions{}); err == nil || !stringsContains(err, "in total") {
t.Fatalf("%s: one condition for a two-component state: %v", name, err)
}
if _, err := SolveBoundaryCollocation(good, 0, 1, y0,
BoundaryConditions{Start: []int{0, 1}, End: []int{0, 1}, EndValues: []float64{1, 2}},
CollocationOptions{}); err == nil || !stringsContains(err, "in total") {
t.Fatalf("%s: four conditions for a two-component state: %v", name, err)
}
if _, err := SolveBoundaryCollocation(good, 0, 1, y0,
BoundaryConditions{Start: []int{0}, End: []int{0}}, CollocationOptions{}); err == nil || !stringsContains(err, "EndValues") {
t.Fatalf("%s: EndValues mismatch: %v", name, err)
}
if _, err := SolveBoundaryCollocation(good, 0, 1, y0,
BoundaryConditions{Start: []int{0}, End: []int{2}, EndValues: []float64{1}},
CollocationOptions{}); err == nil || !stringsContains(err, "out of range") {
t.Fatalf("%s: End index out of range: %v", name, err)
}
// The duplicate-index cases need the total count to be right
// first, so they run on a three-component state.
good3 := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0), y.FloatAt(2)}, 3)
}
if _, err := SolveBoundaryCollocation(good3, 0, 1, mustFloats(t, []float64{0, 0.5, 1}),
BoundaryConditions{Start: []int{0, 0}, End: []int{2}, EndValues: []float64{1}},
CollocationOptions{}); err == nil || !stringsContains(err, "twice") {
t.Fatalf("%s: repeated Start index: %v", name, err)
}
if _, err := SolveBoundaryCollocation(good3, 0, 1, mustFloats(t, []float64{0, 0.5, 1}),
BoundaryConditions{Start: []int{0}, End: []int{2, 2}, EndValues: []float64{1, 2}},
CollocationOptions{}); err == nil || !stringsContains(err, "twice") {
t.Fatalf("%s: repeated End index: %v", name, err)
}
if _, err := SolveBoundaryCollocation(good, 0, 1, y0,
BoundaryConditions{}, CollocationOptions{}); err == nil || !stringsContains(err, "End must prescribe") {
t.Fatalf("%s: nothing prescribed at t1: %v", name, err)
}
// State and interval gates.
if _, err := SolveBoundaryCollocation(good, 0, 1, mustFloats(t, []float64{1, 0, 2}),
bc, CollocationOptions{}); err == nil || !stringsContains(err, "in total") {
t.Fatalf("%s: three components with two conditions: %v", name, err)
}
if _, err := SolveBoundaryCollocation(good, 0, 1, mustFloats(t, nil), bc, CollocationOptions{}); err == nil {
t.Fatalf("%s: empty state accepted", name)
}
if _, err := SolveBoundaryCollocation(good, 0, 1, mustFloats(t, []float64{1, 0, 2}, 1, 3), bc, CollocationOptions{}); err == nil || !stringsContains(err, "vector") {
t.Fatalf("%s: a rank-2 state accepted", name)
}
if _, err := SolveBoundaryCollocation(good, 0, 0, y0, bc, CollocationOptions{}); err == nil || !stringsContains(err, "positive length") {
t.Fatalf("%s: an empty interval accepted", name)
}
if _, err := SolveBoundaryCollocation(good, 1, 0, y0, bc, CollocationOptions{}); err == nil || !stringsContains(err, "positive length") {
t.Fatalf("%s: a backward interval accepted", name)
}
if _, err := SolveBoundaryCollocation(good, 0, 1, mustFloats(t, []float64{math.NaN(), 0}),
bc, CollocationOptions{}); err == nil || !stringsContains(err, "non-finite") {
t.Fatalf("%s: a NaN state accepted", name)
}
// Mesh gates.
if _, err := SolveBoundaryCollocation(good, 0, 1, y0, bc,
CollocationOptions{InitialNodes: 300, MaxNodes: 200}); err == nil || !stringsContains(err, "MaxNodes") {
t.Fatalf("%s: a starting mesh past MaxNodes accepted", name)
}
// A right-hand side of the wrong shape surfaces with its name.
bad := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{1, 2, 3}, 3)
}
if _, err := SolveBoundaryCollocation(bad, 0, 1, y0, bc, CollocationOptions{}); err == nil || !stringsContains(err, "want a vector") {
t.Fatalf("%s: a wrong-shaped f accepted", name)
}
// Refinement past MaxNodes is refused with its name: the layer
// problem demands far more than 24 intervals at this tolerance.
layer := func(t float64, y *core.Array) (*core.Array, error) {
return core.FromFloats([]float64{y.FloatAt(1), -100 * y.FloatAt(1)}, 2)
}
if _, err := SolveBoundaryCollocation(layer, 0, 1, mustFloats(t, []float64{0, 0}),
BoundaryConditions{Start: []int{0}, End: []int{0}, EndValues: []float64{1}},
CollocationOptions{RelTol: 1e-4, AbsTol: 1e-6, InitialNodes: 8, MaxNodes: 24}); err == nil || !stringsContains(err, "MaxNodes") {
t.Fatalf("%s: refinement past MaxNodes accepted", name)
}
}