140 lines
5.2 KiB
Go
140 lines
5.2 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package integrate
|
||
|
|
|
||
|
|
import (
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/optim"
|
||
|
|
)
|
||
|
|
|
||
|
|
import "math"
|
||
|
|
|
||
|
|
// Two-point boundary value problems by shooting. The differential
|
||
|
|
// equation is integrated as an initial value problem whose free
|
||
|
|
// starting components are chosen so the trajectory lands on the
|
||
|
|
// prescribed end values: the mismatch at t1 is a function of the free
|
||
|
|
// starting components alone, and FindRootSystem drives that mismatch
|
||
|
|
// to zero. Linear problems give a mismatch linear in the unknowns,
|
||
|
|
// which the damped Newton settles in one step; nonlinear ones cost a
|
||
|
|
// handful of trajectory integrations more.
|
||
|
|
|
||
|
|
// BoundaryConditions fixes the state of a boundary value problem at
|
||
|
|
// the two ends of the interval. Start lists the components prescribed
|
||
|
|
// at t0, whose values are read from the initial state; End lists the
|
||
|
|
// components prescribed at t1 with their values in EndValues, parallel
|
||
|
|
// to End. A component may carry a condition at both ends, as a
|
||
|
|
// second-order equation written as a first-order system does for its
|
||
|
|
// position; the remaining components are the shooting unknowns, seeded
|
||
|
|
// from the initial state.
|
||
|
|
type BoundaryConditions struct {
|
||
|
|
Start []int
|
||
|
|
End []int
|
||
|
|
EndValues []float64
|
||
|
|
}
|
||
|
|
|
||
|
|
// IntegrateBoundary solves the two-point boundary value problem y' =
|
||
|
|
// f(t, y) on [t0, t1] under the given conditions by shooting, and
|
||
|
|
// returns the trajectory sampled at nSamples evenly spaced times, the
|
||
|
|
// same contract as IntegrateODEPath. states[0] carries the initial
|
||
|
|
// state with the shooting unknowns replaced by the values that satisfy
|
||
|
|
// the end conditions. Backward integration (t1 < t0) works.
|
||
|
|
//
|
||
|
|
// The count of conditions decides solvability: len(Start) + len(End)
|
||
|
|
// must equal the state length, with the components in Start the ones
|
||
|
|
// whose initial values are known. An index out of range, a repeated
|
||
|
|
// index within Start or within End, an EndValues of the wrong length,
|
||
|
|
// a truncated sample count or an exhausted step budget in any trial
|
||
|
|
// trajectory is an error, never a silent answer. The shooting root
|
||
|
|
// find is local: a start whose basin holds no matching trajectory, or
|
||
|
|
// a trial that blows up on the way to t1, reports the failure.
|
||
|
|
func IntegrateBoundary(f func(t float64, y *core.Array) (*core.Array, error),
|
||
|
|
t0, t1 float64, y0 *core.Array, bc BoundaryConditions, nSamples int,
|
||
|
|
opts ODEOptions) ([]float64, []*core.Array, error) {
|
||
|
|
const name = "IntegrateBoundary"
|
||
|
|
y, err := odeCheck("IntegrateBoundary", y0, &opts)
|
||
|
|
if err != nil {
|
||
|
|
return nil, nil, base.Errf("%s: %w", name, err)
|
||
|
|
}
|
||
|
|
n := len(y)
|
||
|
|
if nSamples < 2 {
|
||
|
|
return nil, nil, base.Errf("%s: nSamples must be ≥ 2, got %d", name, nSamples)
|
||
|
|
}
|
||
|
|
if len(bc.End) == 0 {
|
||
|
|
return nil, nil, base.Errf("%s: End must prescribe at least one component at t1", name)
|
||
|
|
}
|
||
|
|
if len(bc.Start)+len(bc.End) != n {
|
||
|
|
return nil, 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, 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, nil, base.Errf("%s: Start index %d out of range for a state of length %d", name, j, n)
|
||
|
|
}
|
||
|
|
if inStart[j] {
|
||
|
|
return nil, 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, nil, base.Errf("%s: End index %d out of range for a state of length %d", name, j, n)
|
||
|
|
}
|
||
|
|
if inEnd[j] {
|
||
|
|
return nil, nil, base.Errf("%s: End prescribes component %d twice", name, j)
|
||
|
|
}
|
||
|
|
inEnd[j] = true
|
||
|
|
}
|
||
|
|
free := make([]int, 0, n-len(bc.Start))
|
||
|
|
for j := range n {
|
||
|
|
if !inStart[j] {
|
||
|
|
free = append(free, j)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
// The end-point mismatch as a function of the free starting
|
||
|
|
// components. Each evaluation is one full trajectory, so the
|
||
|
|
// integration tolerance bounds the noise the root find must see
|
||
|
|
// and its target sits just above that noise.
|
||
|
|
residual := func(x *core.Array) (*core.Array, error) {
|
||
|
|
u0 := cloneDenseSlice(y)
|
||
|
|
for k, j := range free {
|
||
|
|
u0[j] = x.FloatAt(k)
|
||
|
|
}
|
||
|
|
u1, err := IntegrateODE(f, t0, t1, wrapVector(u0), opts)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
r := make([]float64, len(bc.End))
|
||
|
|
for k, j := range bc.End {
|
||
|
|
r[k] = u1.FloatAt(j) - bc.EndValues[k]
|
||
|
|
}
|
||
|
|
return wrapVector(r), nil
|
||
|
|
}
|
||
|
|
guess := make([]float64, len(free))
|
||
|
|
for k, j := range free {
|
||
|
|
guess[k] = y[j]
|
||
|
|
}
|
||
|
|
tol := math.Max(10*opts.RelTol, 1e-12)
|
||
|
|
solution, _, err := optim.FindRootSystem(residual, wrapVector(guess),
|
||
|
|
optim.RootSystemOptions{Tolerance: tol})
|
||
|
|
if err != nil {
|
||
|
|
return nil, nil, base.Errf("%s: %w", name, err)
|
||
|
|
}
|
||
|
|
u0 := cloneDenseSlice(y)
|
||
|
|
for k, j := range free {
|
||
|
|
u0[j] = solution.FloatAt(k)
|
||
|
|
}
|
||
|
|
times, states, err := IntegrateODEPath(f, t0, t1, wrapVector(u0), nSamples, opts)
|
||
|
|
if err != nil {
|
||
|
|
return nil, nil, base.Errf("%s: %w", name, err)
|
||
|
|
}
|
||
|
|
return times, states, nil
|
||
|
|
}
|