// Copyright (c) 2026 Petr Balvín (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 }