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