Files
tensor/integrate/odecolloc.go
T

549 lines
18 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package integrate
import (
"math"
"slices"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Boundary value problems by collocation, the mesh-based sibling of
// the shooting method in IntegrateBoundary. Instead of marching a
// single trajectory and tuning its free start, the solver lays a mesh
// over [t0, t1], represents the solution by a cubic on every mesh
// interval, drives the whole discrete system to zero by a damped
// Newton iteration over a numerically assembled Jacobian, and then
// halves the intervals whose residual estimate is past tolerance and
// solves again, until every interval sits inside the tolerance or the
// node budget runs out.
//
// The scheme is the classic three-point Lobatto IIIA collocation, not
// the Kierzenka-Shampine variant: the collocation polynomial on each
// interval satisfies the ODE at both endpoints and the midpoint,
// which makes the nodal values fourth-order accurate in the interval
// width. The unknowns follow the shape scipy's solve_bvp solves for:
// the state y at every mesh node and the slope s = f(t, y) at every
// mesh node. Between the nodes the solution is the piecewise cubic
// Hermite through (t, y, s), which is exactly the collocation
// polynomial, so that triple is the solver's S-slope representation
// of the continuous answer.
// CollocationOptions tunes SolveBoundaryCollocation. RelTol ≤ 0 means
// 1e-6, AbsTol ≤ 0 means 1e-9, InitialNodes ≤ 0 means 10, MaxNodes ≤ 0
// means 256 and MaxIterations ≤ 0 means 40.
type CollocationOptions struct {
// RelTol and AbsTol scale the mesh-refinement estimate: an
// interval whose root-mean-square of residual over
// AbsTol + RelTol·|slope| stays above 1 is halved. The same pair
// floors the Newton convergence, an order of magnitude below it.
RelTol float64
AbsTol float64
// InitialNodes is the interval count of the uniform starting
// mesh.
InitialNodes int
// MaxNodes bounds the refined mesh. The Newton matrix is factored
// by the library's dense LU, so the cap also bounds the per-round
// cost; a two-component system lives comfortably at 256, well
// inside memory.
MaxNodes int
// MaxIterations bounds the damped Newton rounds on each mesh.
MaxIterations int
}
// CollocationSolution carries the solved problem: Mesh holds the node
// times, Values[k] the state at Mesh[k] and Slopes[k] the derivative
// y' = f(t, y) there. The piecewise cubic Hermite through
// (Mesh, Values, Slopes) is the collocation solution itself, so
// interpolating from that data between the nodes is exact to the
// solver's tolerance.
type CollocationSolution struct {
Mesh []float64
Values []*core.Array
Slopes []*core.Array
}
// SolveBoundaryCollocation solves the two-point boundary value
// problem y' = f(t, y) on [t0, t1] by three-point Lobatto IIIA
// collocation on an adaptively refined mesh, returning the mesh, the
// nodal states and the nodal slopes. The boundary conditions follow
// the BoundaryConditions contract of IntegrateBoundary: Start lists
// the components prescribed at t0 with values read from y0, End the
// components prescribed at t1 with EndValues, and exactly n
// conditions must be given in total, because the collocation system
// is square. The initial guess interpolates linearly between the
// prescribed endpoint states and reads its slopes from f.
//
// Refusal is part of the contract: inconsistent boundary conditions
// (fewer or more than n conditions, out-of-range or repeated
// indices, EndValues of the wrong length), a non-positive interval,
// a starting mesh past MaxNodes, refinement that would grow past
// MaxNodes, a singular Newton matrix or an iteration that cannot
// converge are errors, never silent answers.
func SolveBoundaryCollocation(f func(t float64, y *core.Array) (*core.Array, error),
t0, t1 float64, y0 *core.Array, bc BoundaryConditions,
opts CollocationOptions) (*CollocationSolution, error) {
const name = "SolveBoundaryCollocation"
if opts.RelTol <= 0 {
opts.RelTol = 1e-6
}
if opts.AbsTol <= 0 {
opts.AbsTol = 1e-9
}
if opts.InitialNodes <= 0 {
opts.InitialNodes = 10
}
if opts.MaxNodes <= 0 {
opts.MaxNodes = 256
}
if opts.MaxIterations <= 0 {
opts.MaxIterations = 40
}
if y0.NDim() != 1 {
return nil, base.Errf("%s: the state must be a vector, got shape %s", name, base.ShapeText(y0.Shape()))
}
if y0.Len() == 0 {
return nil, base.Errf("%s: the state must not be empty", name)
}
if y0.Dtype() == core.Complex {
return nil, base.Errf("%s: complex states are not supported", name)
}
if !(t1 > t0) {
return nil, base.Errf("%s: the interval must have positive length, got [%g, %g]", name, t0, t1)
}
n := y0.Len()
seed := make([]float64, n)
for i := range n {
seed[i] = y0.FloatAt(i)
if math.IsNaN(seed[i]) || math.IsInf(seed[i], 0) {
return nil, base.Errf("%s: the state holds the non-finite value %g at %d", name, seed[i], i)
}
}
if len(bc.End) == 0 {
return nil, base.Errf("%s: End must prescribe at least one component at t1", name)
}
if len(bc.Start)+len(bc.End) != n {
return nil, base.Errf("%s: %d conditions at t0 and %d at t1 for a state of length %d, want %d in total",
name, len(bc.Start), len(bc.End), n, n)
}
if len(bc.EndValues) != len(bc.End) {
return nil, base.Errf("%s: EndValues has length %d, want %d to match End",
name, len(bc.EndValues), len(bc.End))
}
inStart := make(map[int]bool, len(bc.Start))
for _, j := range bc.Start {
if j < 0 || j >= n {
return nil, base.Errf("%s: Start index %d out of range for a state of length %d", name, j, n)
}
if inStart[j] {
return nil, base.Errf("%s: Start prescribes component %d twice", name, j)
}
inStart[j] = true
}
inEnd := make(map[int]bool, len(bc.End))
for _, j := range bc.End {
if j < 0 || j >= n {
return nil, base.Errf("%s: End index %d out of range for a state of length %d", name, j, n)
}
if inEnd[j] {
return nil, base.Errf("%s: End prescribes component %d twice", name, j)
}
inEnd[j] = true
}
endState := cloneDenseSlice(seed)
for q, j := range bc.End {
endState[j] = bc.EndValues[q]
}
if opts.InitialNodes < 1 {
return nil, base.Errf("%s: InitialNodes must be ≥ 1, got %d", name, opts.InitialNodes)
}
if opts.InitialNodes+1 > opts.MaxNodes {
return nil, base.Errf("%s: the starting mesh of %d intervals already exceeds MaxNodes=%d",
name, opts.InitialNodes, opts.MaxNodes)
}
// The starting mesh and guess: uniform in t, linear between the
// prescribed endpoint states, slopes read from f.
mesh := make([]float64, opts.InitialNodes+1)
for k := range mesh {
mesh[k] = t0 + (t1-t0)*float64(k)/float64(opts.InitialNodes)
}
stride := 2 * n
z := make([]float64, stride*(opts.InitialNodes+1))
for k := range opts.InitialNodes + 1 {
theta := (mesh[k] - t0) / (t1 - t0)
for i := range n {
z[k*stride+i] = seed[i] + theta*(endState[i]-seed[i])
}
}
for k := range opts.InitialNodes + 1 {
sv, err := odeEval(name, f, mesh[k], z[k*stride:k*stride+n], n, nil)
if err != nil {
return nil, err
}
copy(z[k*stride+n:(k+1)*stride], sv)
}
// collocResidual writes the discrete system for the unknown vector
// zz into dst: one slope-definition block per node, one
// collocation block per interval (the midpoint value eliminated
// through the cubic Hermite it belongs to), then the boundary
// rows. The row count equals the unknown count exactly.
collocResidual := func(dst, zz, msh []float64) error {
nodes := len(msh)
for k := range nodes {
fv, err := odeEval(name, f, msh[k], zz[k*stride:k*stride+n], n, nil)
if err != nil {
return err
}
for i := range n {
dst[k*n+i] = zz[k*stride+n+i] - fv[i]
}
}
baseRow := nodes * n
ym := make([]float64, n)
for i := range nodes - 1 {
h := msh[i+1] - msh[i]
for j := range n {
ym[j] = (zz[i*stride+j]+zz[(i+1)*stride+j])/2 + h*(zz[i*stride+n+j]-zz[(i+1)*stride+n+j])/8
}
fm, err := odeEval(name, f, msh[i]+h/2, ym, n, nil)
if err != nil {
return err
}
for j := range n {
dst[baseRow+i*n+j] = zz[(i+1)*stride+j] - zz[i*stride+j] -
h*(zz[i*stride+n+j]+4*fm[j]+zz[(i+1)*stride+n+j])/6
}
}
last := len(dst) - n
for p, j := range bc.Start {
dst[last+p] = zz[j] - seed[j]
}
for q, j := range bc.End {
dst[last+len(bc.Start)+q] = zz[(nodes-1)*stride+j] - bc.EndValues[q]
}
return nil
}
// collocJac assembles the Newton matrix by central differences:
// the slope rows differentiate y − f(t, y) over their node's y
// block (their slope columns are the exact −I), the collocation
// rows differentiate the interval map over the four blocks
// (yᵢ, sᵢ, yᵢ₊₁, sᵢ₊₁), and the boundary rows enter exactly.
collocJac := func(zz, msh []float64) ([][]float64, error) {
nodes := len(msh)
size := stride * nodes
jac := make([][]float64, size)
for i := range jac {
jac[i] = make([]float64, size)
}
for k := range nodes {
// The slope rows are s − f(y): as a function of the node's
// y block the residual is −f(y), and the slope columns
// carry the +I.
block := func(y []float64) ([]float64, error) {
fv, err := odeEval(name, f, msh[k], y, n, nil)
if err != nil {
return nil, err
}
out := make([]float64, n)
for i := range n {
out[i] = -fv[i]
}
return out, nil
}
if err := collocNumJac(name, block, zz[k*stride:k*stride+n], k*n, k*stride, n, n, jac); err != nil {
return nil, err
}
for i := range n {
// The slope rows are s − f(y), so the slope columns
// carry +I.
jac[k*n+i][k*stride+n+i] = 1
}
}
baseRow := nodes * n
w := make([]float64, 4*n)
ym := make([]float64, n)
for i := range nodes - 1 {
h := msh[i+1] - msh[i]
copy(w[0:n], zz[i*stride:i*stride+n])
copy(w[n:2*n], zz[i*stride+n:(i+1)*stride])
copy(w[2*n:3*n], zz[(i+1)*stride:(i+1)*stride+n])
copy(w[3*n:4*n], zz[(i+1)*stride+n:(i+2)*stride])
interval := func(x []float64) ([]float64, error) {
for j := range n {
ym[j] = (x[j]+x[2*n+j])/2 + h*(x[n+j]-x[3*n+j])/8
}
fm, err := odeEval(name, f, msh[i]+h/2, ym, n, nil)
if err != nil {
return nil, err
}
out := make([]float64, n)
for j := range n {
out[j] = x[2*n+j] - x[j] - h*(x[n+j]+4*fm[j]+x[3*n+j])/6
}
return out, nil
}
if err := collocNumJac(name, interval, w, baseRow+i*n, i*stride, n, 4*n, jac); err != nil {
return nil, err
}
}
last := size - n
for p, j := range bc.Start {
jac[last+p][j] = 1
}
for q, j := range bc.End {
jac[last+len(bc.Start)+q][(nodes-1)*stride+j] = 1
}
return jac, nil
}
// newtonSolve drives the damped Newton on the fixed mesh: the
// Jacobian is frozen from the seed and rebuilt twice when
// convergence drags, as odeNewton does, and each round backtracks
// along the step until the residual infinity norm actually falls.
newtonSolve := func(zz, msh []float64) error {
size := stride * len(msh)
r := make([]float64, size)
trialZ := make([]float64, size)
trialR := make([]float64, size)
col := make([]float64, size)
if err := collocResidual(r, zz, msh); err != nil {
return err
}
for i := range r {
if math.IsNaN(r[i]) || math.IsInf(r[i], 0) {
return base.Errf("%s: the residual returned the non-finite value %g at row %d", name, r[i], i)
}
}
scale := normInfOfStep(zz)
// The algebraic floor sits three orders below the mesh
// tolerance: the refinement estimator reads the true ODE
// residual of the interpolant, and that reading must not be
// dominated by the residual the Newton iteration left.
limit := math.Max(0.001*(opts.AbsTol+opts.RelTol*scale), 8*base.EpsF*(scale+1))
worst := normInfOfStep(r)
var work [][]float64
var perm []int
for iteration := 0; iteration < opts.MaxIterations; iteration++ {
if worst <= limit {
return nil
}
if iteration == 0 || iteration == 4 || iteration == 10 {
jac, err := collocJac(zz, msh)
if err != nil {
return err
}
// Factor a working copy: base.Factor consumes its
// argument in place, and the pristine matrix is not
// needed again before the next rebuild.
work = make([][]float64, size)
for i := range jac {
work[i] = cloneDenseSlice(jac[i])
}
perm, _ = base.Factor(work)
if err := base.CheckSingular(name, work); err != nil {
return base.Errf("%s: %w, singular collocation Newton matrix", name, errNewtonStalled)
}
}
for i := range col {
col[i] = -r[i]
}
base.PermuteColumn(col, perm)
base.SolveColumn(work, col)
accepted := false
factor := 1.0
for range 40 {
for i := range zz {
trialZ[i] = zz[i] + factor*col[i]
}
if err := collocResidual(trialR, trialZ, msh); err == nil {
finite := true
cand := 0.0
for i := range trialR {
v := trialR[i]
if math.IsNaN(v) || math.IsInf(v, 0) {
finite = false
break
}
cand = math.Max(cand, math.Abs(v))
}
if finite && cand <= (1-1e-4*factor)*worst {
copy(zz, trialZ)
// The residual travels with the accepted point,
// so the next round solves against the state
// this round left behind.
copy(r, trialR)
worst = cand
accepted = true
break
}
}
factor /= 2
}
if !accepted {
return base.Errf("%s: %w: the residual cannot be reduced below %g by damping",
name, errNewtonStalled, worst)
}
}
return base.Errf("%s: %w after %d rounds, residual %g", name, errNewtonStalled, opts.MaxIterations, worst)
}
// estimateRefinement returns the intervals whose root-mean-square
// scaled residual is past 1, measured from the collocation
// solution's own cubic Hermite: the ODE is evaluated at three
// interior quadrature points of every interval (the nodes one half
// plus or minus half the square root of three sevenths and the
// midpoint, weights 49/180, 16/45, 49/180) and the mismatch with
// the Hermite derivative is quadrature-weighted (the endpoints
// contribute exactly nothing, the Hermite slope is the ODE slope
// there by construction).
estimateRefinement := func(zz, msh []float64) ([]int, error) {
nodes := len(msh)
theta := [3]float64{0.5 * (1 - math.Sqrt(3.0/7)), 0.5, 0.5 * (1 + math.Sqrt(3.0/7))}
weight := [3]float64{49.0 / 180, 16.0 / 45, 49.0 / 180}
var bad []int
val := make([]float64, n)
for i := range nodes - 1 {
h := msh[i+1] - msh[i]
sum := 0.0
for pt := range 3 {
th := theta[pt]
th2 := th * th
th3 := th2 * th
h00 := 2*th3 - 3*th2 + 1
h10 := th3 - 2*th2 + th
h01 := -2*th3 + 3*th2
h11 := th3 - th2
hd00 := 6*th2 - 6*th
hd10 := 3*th2 - 4*th + 1
hd01 := -6*th2 + 6*th
hd11 := 3*th2 - 2*th
for j := range n {
val[j] = h00*zz[i*stride+j] + h*h10*zz[i*stride+n+j] +
h01*zz[(i+1)*stride+j] + h*h11*zz[(i+1)*stride+n+j]
}
fv, err := odeEval(name, f, msh[i]+th*h, val, n, nil)
if err != nil {
return nil, err
}
for j := range n {
der := hd00*zz[i*stride+j]/h + hd10*zz[i*stride+n+j] +
hd01*zz[(i+1)*stride+j]/h + hd11*zz[(i+1)*stride+n+j]
slope := math.Max(math.Abs(zz[i*stride+n+j]), math.Abs(zz[(i+1)*stride+n+j]))
sc := opts.AbsTol + opts.RelTol*math.Max(slope, math.Abs(fv[j]))
ratio := (der - fv[j]) / sc
sum += weight[pt] * ratio * ratio
}
}
if math.Sqrt(sum) > 1 {
bad = append(bad, i)
}
}
return bad, nil
}
// refineMesh inserts the midpoint of every interval in bad, with
// the new node's state and slope taken from the collocation
// solution's own cubic Hermite at the midpoint.
refineMesh := func(zz, msh []float64, bad []int) ([]float64, []float64, error) {
nodes := len(msh)
if nodes+len(bad) > opts.MaxNodes {
return nil, nil, base.Errf("%s: refining %d intervals would grow the mesh to %d nodes past MaxNodes=%d",
name, len(bad), nodes+len(bad), opts.MaxNodes)
}
badSet := make(map[int]bool, len(bad))
for _, i := range bad {
badSet[i] = true
}
newMesh := make([]float64, 0, nodes+len(bad))
newZ := make([]float64, 0, len(zz)+2*n*len(bad))
push := func(t float64, y, s []float64) {
newMesh = append(newMesh, t)
newZ = append(newZ, y...)
newZ = append(newZ, s...)
}
for i := range nodes - 1 {
push(msh[i], zz[i*stride:i*stride+n], zz[i*stride+n:(i+1)*stride])
if !badSet[i] {
continue
}
h := msh[i+1] - msh[i]
mid := (msh[i] + msh[i+1]) / 2
ym := make([]float64, n)
sm := make([]float64, n)
for j := range n {
ym[j] = (zz[i*stride+j]+zz[(i+1)*stride+j])/2 + h*(zz[i*stride+n+j]-zz[(i+1)*stride+n+j])/8
sm[j] = 1.5*(zz[(i+1)*stride+j]-zz[i*stride+j])/h - (zz[i*stride+n+j]+zz[(i+1)*stride+n+j])/4
}
push(mid, ym, sm)
}
push(msh[nodes-1], zz[(nodes-1)*stride:(nodes-1)*stride+n], zz[(nodes-1)*stride+n:nodes*stride])
return newMesh, newZ, nil
}
for pass := 0; ; pass++ {
if pass >= 100 {
return nil, base.Errf("%s: refinement did not settle within 100 passes", name)
}
if err := newtonSolve(z, mesh); err != nil {
return nil, err
}
bad, err := estimateRefinement(z, mesh)
if err != nil {
return nil, err
}
if len(bad) == 0 {
break
}
mesh, z, err = refineMesh(z, mesh, bad)
if err != nil {
return nil, err
}
}
solution := &CollocationSolution{
Mesh: slices.Clone(mesh),
Values: make([]*core.Array, len(mesh)),
Slopes: make([]*core.Array, len(mesh)),
}
for k := range mesh {
solution.Values[k] = arrayFromVector(z[k*stride : k*stride+n])
solution.Slopes[k] = arrayFromVector(z[k*stride+n : (k+1)*stride])
}
return solution, nil
}
// collocNumJac fills jac[row0+i][col0+c] with the central-difference
// derivative of g's i-th output against x's c-th entry, one column
// per entry of x. g receives its own perturbed copy of x and returns
// a fresh output slice, so nothing aliases.
func collocNumJac(name string, g func(x []float64) ([]float64, error), x []float64,
row0, col0, rows, cols int, jac [][]float64) error {
xp := make([]float64, cols)
xm := make([]float64, cols)
for c := range cols {
eps := math.Sqrt(base.EpsF) * math.Max(1, math.Abs(x[c]))
copy(xp, x)
copy(xm, x)
xp[c] += eps
xm[c] -= eps
rp, e1 := g(xp)
if e1 != nil {
return base.Errf("%s: %w", name, e1)
}
rm, e2 := g(xm)
if e2 != nil {
return base.Errf("%s: %w", name, e2)
}
for i := range rows {
jac[row0+i][col0+c] = (rp[i] - rm[i]) / (2 * eps)
}
}
return nil
}