Files
tensor/integrate/symplectic2.go
T

286 lines
11 KiB
Go
Raw Permalink 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"
"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 carry a discrete state with no place in a continuous
// integrator, so they follow Int into the standing "int states cannot
// integrate" refusal. The float dtypes, float16 included, keep the
// treatment they carry today.
2026-09-03 10:00:00 +02:00
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
}