285 lines
11 KiB
Go
285 lines
11 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package integrate
|
||
|
||
import (
|
||
"math"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
"sourcedock.dev/petrbalvin/tensor/optim"
|
||
)
|
||
|
||
// Higher-order symplectic integration, extending the velocity-Verlet
|
||
// leapfrog of symplectic.go in the two directions it cannot go by
|
||
// itself: to fourth order while staying explicit (Yoshida's
|
||
// composition of three leapfrog sub-steps), and to non-separable
|
||
// Hamiltonians at all (the implicit midpoint rule, whose implicit
|
||
// stage is a root find per step). Both keep the leapfrog's reason for
|
||
// being: the exact conservation of a shadow Hamiltonian, so the true
|
||
// energy oscillates in a bounded band instead of drifting, over
|
||
// arbitrarily long runs, at a fixed step by design.
|
||
|
||
// yoshidaW1 and yoshidaW0 are Yoshida's triple-jump weights: the
|
||
// composition V(w1·h)·V(w0·h)·V(w1·h) of three velocity-Verlet
|
||
// sub-steps is fourth order exactly when w0 + 2·w1 = 1, and the
|
||
// classic choice w1 = 1/(2 − ∛2), w0 = −∛2/(2 − ∛2) satisfies that
|
||
// identity with w1 positive and w0 negative (the middle sub-step runs
|
||
// backwards in time, which is what buys the order).
|
||
var (
|
||
yoshidaW1 = 1 / (2 - math.Cbrt(2))
|
||
yoshidaW0 = -math.Cbrt(2) * yoshidaW1
|
||
)
|
||
|
||
// IntegrateYoshida4 integrates a separable Hamiltonian system with
|
||
// unit masses over an even time grid by Yoshida's fourth-order
|
||
// composition of the kick-drift-kick leapfrog: three velocity-Verlet
|
||
// sub-steps of widths w1·h, w0·h and w1·h per step, with the weights
|
||
// above. The contract is IntegrateVerlet's: accel returns the
|
||
// acceleration −∂V/∂q at a position, q0 and p0 are the initial
|
||
// position and momentum (velocity), steps fixes the number of equal
|
||
// steps h = (t1−t0)/steps, positions[s] and momenta[s] sit at
|
||
// t0 + s·h, and the state may be float64 or float32. The step size
|
||
// stays fixed by design. The error contract is IntegrateVerlet's
|
||
// too: a mismatched pair, a complex or int state, an empty state or
|
||
// a non-finite acceleration is an error, never a silently corrupted
|
||
// trajectory.
|
||
func IntegrateYoshida4(accel func(q *core.Array) (*core.Array, error),
|
||
t0, t1 float64, q0, p0 *core.Array, steps int) (positions, momenta []*core.Array, err error) {
|
||
const name = "IntegrateYoshida4"
|
||
if steps <= 0 {
|
||
return nil, nil, base.Errf("%s: steps must be ≥ 1, got %d", name, steps)
|
||
}
|
||
q, p, n, verr := verletPrelude(name, q0, p0)
|
||
if verr != nil {
|
||
return nil, nil, verr
|
||
}
|
||
eval := verletForce(name, accel, n)
|
||
a := make([]float64, n)
|
||
if err := eval(q, a); err != nil {
|
||
return nil, nil, err
|
||
}
|
||
h := (t1 - t0) / float64(steps)
|
||
positions = make([]*core.Array, steps+1)
|
||
momenta = make([]*core.Array, steps+1)
|
||
positions[0] = arrayFromVector(q)
|
||
momenta[0] = arrayFromVector(p)
|
||
weights := [3]float64{yoshidaW1, yoshidaW0, yoshidaW1}
|
||
for s := 1; s <= steps; s++ {
|
||
// Three kick-drift-kick sub-steps. The acceleration that ends
|
||
// one sub-step is exactly the one the next sub-step's first
|
||
// kick needs, so the whole step costs three force evaluations.
|
||
for _, w := range weights {
|
||
tau := w * h
|
||
for i := range n {
|
||
p[i] += tau / 2 * a[i]
|
||
q[i] += tau * p[i]
|
||
}
|
||
if err := eval(q, a); err != nil {
|
||
return nil, nil, err
|
||
}
|
||
for i := range n {
|
||
p[i] += tau / 2 * a[i]
|
||
}
|
||
}
|
||
positions[s] = arrayFromVector(q)
|
||
momenta[s] = arrayFromVector(p)
|
||
}
|
||
return positions, momenta, nil
|
||
}
|
||
|
||
// verletPrelude is the validation the Verlet family shares: the
|
||
// dtype, shape and finiteness gates of IntegrateVerlet and the
|
||
// element-wise read of the initial state (RawFloats backs float64
|
||
// payloads only, so a float32 state would come through as nil). It
|
||
// returns the working position and momentum and the state length.
|
||
func verletPrelude(name string, q0, p0 *core.Array) (q, p []float64, n int, err error) {
|
||
if q0.Dtype() == core.Complex || p0.Dtype() == core.Complex {
|
||
return nil, nil, 0, base.Errf("%s: complex states are not supported", name)
|
||
}
|
||
// The whole integer class follows Int into the standing refusal:
|
||
// bool and the narrow widths carry a discrete state, which has no
|
||
// place in a continuous integrator, and the wording is Int's own.
|
||
if integerState(q0.Dtype()) || integerState(p0.Dtype()) {
|
||
return nil, nil, 0, base.Errf("%s: int states cannot integrate, use float or float32 states", name)
|
||
}
|
||
if q0.NDim() != 1 || p0.NDim() != 1 || q0.Len() != p0.Len() {
|
||
return nil, nil, 0, base.Errf("%s: position and momentum must be vectors of equal length, got %s and %s",
|
||
name, base.ShapeText(q0.Shape()), base.ShapeText(p0.Shape()))
|
||
}
|
||
n = q0.Len()
|
||
if n == 0 {
|
||
return nil, nil, 0, base.Errf("%s: the state must not be empty", name)
|
||
}
|
||
q = make([]float64, n)
|
||
p = make([]float64, n)
|
||
for i := range n {
|
||
q[i] = q0.FloatAt(i)
|
||
p[i] = p0.FloatAt(i)
|
||
if math.IsNaN(q[i]) || math.IsInf(q[i], 0) {
|
||
return nil, nil, 0, base.Errf("%s: q0 holds the non-finite value %g at %d", name, q[i], i)
|
||
}
|
||
if math.IsNaN(p[i]) || math.IsInf(p[i], 0) {
|
||
return nil, nil, 0, base.Errf("%s: p0 holds the non-finite value %g at %d", name, p[i], i)
|
||
}
|
||
}
|
||
return q, p, n, nil
|
||
}
|
||
|
||
// integerState reports whether dt is one of the integer-class state
|
||
// dtypes the symplectic family refuses: bool and the narrow integer
|
||
// widths follow Int into the standing "int states cannot integrate"
|
||
// refusal, exactly as the round's follow-Int rule requires. The float
|
||
// dtypes, float16 included, keep the treatment they carry today.
|
||
func integerState(dt core.Dtype) bool {
|
||
switch dt {
|
||
case core.Bool, core.Int, core.Int8, core.Uint8, core.Int16, core.Uint16, core.Int32, core.Uint32:
|
||
return true
|
||
}
|
||
return false
|
||
}
|
||
|
||
// verletForce wraps the acceleration callback the way the Verlet
|
||
// family evaluates it: shape-checked, non-finite-refused, and read
|
||
// into a plain float64 buffer. One cached view per wrapper serves
|
||
// every call: the family evaluates the force at the same position
|
||
// buffer throughout a run, so the wrapper is built once.
|
||
func verletForce(name string, accel func(q *core.Array) (*core.Array, error), n int) func(x []float64, out []float64) error {
|
||
var views odeViews
|
||
return func(x []float64, out []float64) error {
|
||
v, err := accel(views.of(x))
|
||
if err != nil {
|
||
return base.Errf("%s: %w", name, err)
|
||
}
|
||
if v.NDim() != 1 || v.Len() != n {
|
||
return base.Errf("%s: accel returned shape %s, want a vector of length %d",
|
||
name, base.ShapeText(v.Shape()), n)
|
||
}
|
||
readVector(out, v)
|
||
// A non-finite acceleration would flow through the kicks
|
||
// silently, and the published trajectory would be NaN with a
|
||
// nil error.
|
||
for i := range n {
|
||
if math.IsNaN(out[i]) || math.IsInf(out[i], 0) {
|
||
return base.Errf("%s: accel returned the non-finite value %g at coordinate %d", name, out[i], i)
|
||
}
|
||
}
|
||
return nil
|
||
}
|
||
}
|
||
|
||
// MidpointOptions tunes the per-step implicit stage of
|
||
// IntegrateMidpoint. Tolerance ≤ 0 means 1e-13 (the stage solve is a
|
||
// root find, and the energy band the rule is famous for wants it
|
||
// tight); MaxIterations ≤ 0 means 100.
|
||
type MidpointOptions struct {
|
||
Tolerance float64
|
||
MaxIterations int
|
||
}
|
||
|
||
// IntegrateMidpoint integrates the general Hamiltonian flow
|
||
// dz/dt = J·∇H(z), z = (q, p), by the implicit midpoint rule
|
||
// z_{n+1} = z_n + h·J·∇H((z_n + z_{n+1})/2) over an even time grid.
|
||
// gradH returns the gradient of H as the stacked vector
|
||
// (∂H/∂q, ∂H/∂p); H itself never has to be separable, which is the
|
||
// rule's claim over the leapfrog family. Each step's implicit stage
|
||
// is solved by the library's damped Newton root find for systems
|
||
// (optim.FindRootSystem under the given options), seeded with the
|
||
// current state, so the per-step work is a handful of gradient
|
||
// evaluations. Equilibria map to themselves exactly, and every
|
||
// quadratic invariant of the flow, H included when H is quadratic,
|
||
// is conserved to rounding. The contract otherwise mirrors
|
||
// IntegrateVerlet: q0 and p0 are the initial position and momentum,
|
||
// steps fixes the number of equal steps h = (t1−t0)/steps,
|
||
// positions[s] and momenta[s] sit at t0 + s·h, the step size stays
|
||
// fixed by design, and a mismatched pair, a complex or int state, an
|
||
// empty state, a non-finite gradient or a root find that cannot
|
||
// converge is an error, never a silent answer.
|
||
func IntegrateMidpoint(gradH func(z *core.Array) (*core.Array, error),
|
||
t0, t1 float64, q0, p0 *core.Array, steps int, opts MidpointOptions) (positions, momenta []*core.Array, err error) {
|
||
const name = "IntegrateMidpoint"
|
||
if opts.Tolerance <= 0 {
|
||
opts.Tolerance = 1e-13
|
||
}
|
||
if opts.MaxIterations <= 0 {
|
||
opts.MaxIterations = 100
|
||
}
|
||
if steps <= 0 {
|
||
return nil, nil, base.Errf("%s: steps must be ≥ 1, got %d", name, steps)
|
||
}
|
||
q, p, n, verr := verletPrelude(name, q0, p0)
|
||
if verr != nil {
|
||
return nil, nil, verr
|
||
}
|
||
m := 2 * n
|
||
z := make([]float64, m)
|
||
copy(z, q)
|
||
copy(z[n:], p)
|
||
// One cached read-only view serves every gradient call: mid is the
|
||
// run's own stable buffer, so the wrapper is built once.
|
||
var views odeViews
|
||
// The gradient evaluator with the midpoint rule's finiteness gate:
|
||
// a non-finite gradient would flow through the stage equation
|
||
// silently.
|
||
gradient := func(gz []float64, out []float64) error {
|
||
v, err := gradH(views.of(gz))
|
||
if err != nil {
|
||
return base.Errf("%s: %w", name, err)
|
||
}
|
||
if v.NDim() != 1 || v.Len() != m {
|
||
return base.Errf("%s: gradH returned shape %s, want a vector of length %d",
|
||
name, base.ShapeText(v.Shape()), m)
|
||
}
|
||
readVector(out, v)
|
||
for i := range m {
|
||
if math.IsNaN(out[i]) || math.IsInf(out[i], 0) {
|
||
return base.Errf("%s: gradH returned the non-finite value %g at coordinate %d", name, out[i], i)
|
||
}
|
||
}
|
||
return nil
|
||
}
|
||
mid := make([]float64, m)
|
||
grad := make([]float64, m)
|
||
h := (t1 - t0) / float64(steps)
|
||
// The stage residual: x is the candidate z_{n+1}, and the flow
|
||
// J·∇H flips the two halves with a sign: q' = ∂H/∂p, p' = −∂H/∂q.
|
||
// FindRootSystem copies its input before every residual call and
|
||
// reads the returned residual before the next one, so reusing mid,
|
||
// grad and out across calls is the same arithmetic the fresh
|
||
// buffers would give.
|
||
out := make([]float64, m)
|
||
stage := func(x *core.Array) (*core.Array, error) {
|
||
for i := range m {
|
||
mid[i] = (z[i] + x.FloatAt(i)) / 2
|
||
}
|
||
if err := gradient(mid, grad); err != nil {
|
||
return nil, err
|
||
}
|
||
for i := range n {
|
||
out[i] = x.FloatAt(i) - z[i] - h*grad[n+i]
|
||
out[n+i] = x.FloatAt(n+i) - z[n+i] + h*grad[i]
|
||
}
|
||
return wrapVector(out), nil
|
||
}
|
||
positions = make([]*core.Array, steps+1)
|
||
momenta = make([]*core.Array, steps+1)
|
||
positions[0] = arrayFromVector(z[:n])
|
||
momenta[0] = arrayFromVector(z[n:])
|
||
for s := 1; s <= steps; s++ {
|
||
solution, _, rerr := optim.FindRootSystem(stage, wrapVector(z), optim.RootSystemOptions{
|
||
Tolerance: opts.Tolerance,
|
||
MaxIterations: opts.MaxIterations,
|
||
})
|
||
if rerr != nil {
|
||
return nil, nil, base.Errf("%s: %w", name, rerr)
|
||
}
|
||
for i := range m {
|
||
z[i] = solution.FloatAt(i)
|
||
}
|
||
positions[s] = arrayFromVector(z[:n])
|
||
momenta[s] = arrayFromVector(z[n:])
|
||
}
|
||
return positions, momenta, nil
|
||
}
|