Files
tensor/integrate/odecolloc.go
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

549 lines
18 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 (
"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
}