209 lines
6.8 KiB
Go
209 lines
6.8 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"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"cmp"
|
|||
|
|
"math"
|
|||
|
|
"slices"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// Event detection along an ODE trajectory. Impact times, resonance
|
|||
|
|
// crossings and threshold passages are all the same question: when
|
|||
|
|
// does a scalar watch function g(t, y) cross zero on the way to t1?
|
|||
|
|
// The integrator checks every accepted step for a sign change of each
|
|||
|
|
// watch and, when one appears, narrows the crossing by bisection that
|
|||
|
|
// re-integrates the step's interval, so the event time is as accurate
|
|||
|
|
// as the integrator itself and needs no dense-output machinery.
|
|||
|
|
//
|
|||
|
|
// Watches are only compared across the boundaries of accepted steps:
|
|||
|
|
// a watch that touches zero and returns to its sign inside one step
|
|||
|
|
// goes unnoticed, the same blind spot every step-based detector has.
|
|||
|
|
// Sign changes are searched strictly after t0, so a watch sitting on
|
|||
|
|
// zero at the initial state does not fire until it leaves and returns.
|
|||
|
|
|
|||
|
|
// ODEWatch reports a scalar quantity to watch along the trajectory.
|
|||
|
|
// Direction filters the crossings: +1 records only rising crossings
|
|||
|
|
// (watch going from negative to non-negative), −1 only falling ones,
|
|||
|
|
// 0 both.
|
|||
|
|
type ODEWatch struct {
|
|||
|
|
Function func(t float64, y *core.Array) (float64, error)
|
|||
|
|
Direction int
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// ODEEventHit records one zero crossing: the time it happened, the
|
|||
|
|
// state there, and which watch fired.
|
|||
|
|
type ODEEventHit struct {
|
|||
|
|
Time float64
|
|||
|
|
State *core.Array
|
|||
|
|
Watch int
|
|||
|
|
Rising bool
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// IntegrateODEEvents integrates y' = f(t, y) from t0 to t1 exactly
|
|||
|
|
// like IntegrateODE and, alongside the final state, returns every
|
|||
|
|
// zero crossing of the watches, sorted by time. The watch functions
|
|||
|
|
// must tolerate being called at any time inside [t0, t1]; they are
|
|||
|
|
// also called at the step boundaries the integrator accepts, and the
|
|||
|
|
// first accepted step is seeded with the watch value at its start
|
|||
|
|
// state, so a crossing inside it is detected like any other. A watch
|
|||
|
|
// error, a non-finite watch value included, aborts the integration.
|
|||
|
|
func IntegrateODEEvents(f func(t float64, y *core.Array) (*core.Array, error),
|
|||
|
|
t0, t1 float64, y0 *core.Array, watches []ODEWatch, opts ODEOptions) ([]ODEEventHit, *core.Array, error) {
|
|||
|
|
if len(watches) == 0 {
|
|||
|
|
return nil, nil, base.Errf("IntegrateODEEvents: at least one watch is needed")
|
|||
|
|
}
|
|||
|
|
for i := range watches {
|
|||
|
|
if watches[i].Function == nil {
|
|||
|
|
return nil, nil, base.Errf("IntegrateODEEvents: watch %d has no function", i)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// gPrev[i] carries the watch value at the last accepted boundary;
|
|||
|
|
// havePrev becomes true after the first evaluation.
|
|||
|
|
gPrev := make([]float64, len(watches))
|
|||
|
|
havePrev := false
|
|||
|
|
var hits []ODEEventHit
|
|||
|
|
|
|||
|
|
watch := func(tPrev, tNow float64, yPrev, yNow []float64) (bool, error) {
|
|||
|
|
for i := range watches {
|
|||
|
|
if !havePrev {
|
|||
|
|
// First accepted step: the watch value at the step's
|
|||
|
|
// start state seeds the comparison, so a crossing
|
|||
|
|
// inside the very first step is seen like any other
|
|||
|
|
// one instead of hiding behind a missing gPrev.
|
|||
|
|
g0, err := watches[i].Function(tPrev, wrapVector(yPrev))
|
|||
|
|
if err != nil {
|
|||
|
|
return false, base.Errf("IntegrateODEEvents: watch %d: %w", i, err)
|
|||
|
|
}
|
|||
|
|
if math.IsNaN(g0) || math.IsInf(g0, 0) {
|
|||
|
|
return false, base.Errf("IntegrateODEEvents: watch %d returned the non-finite value %g at t=%g", i, g0, tPrev)
|
|||
|
|
}
|
|||
|
|
gPrev[i] = g0
|
|||
|
|
}
|
|||
|
|
g, err := watches[i].Function(tNow, wrapVector(yNow))
|
|||
|
|
if err != nil {
|
|||
|
|
return false, base.Errf("IntegrateODEEvents: watch %d: %w", i, err)
|
|||
|
|
}
|
|||
|
|
// A non-finite watch value compares false against every
|
|||
|
|
// sign test and would masquerade as a crossing (or swallow
|
|||
|
|
// one), so it is refused like quad's non-finite integrand.
|
|||
|
|
if math.IsNaN(g) || math.IsInf(g, 0) {
|
|||
|
|
return false, base.Errf("IntegrateODEEvents: watch %d returned the non-finite value %g at t=%g", i, g, tNow)
|
|||
|
|
}
|
|||
|
|
gp := gPrev[i]
|
|||
|
|
if gp == 0 {
|
|||
|
|
gPrev[i] = g
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
if g == 0 {
|
|||
|
|
// The watch landed exactly on zero at the accepted
|
|||
|
|
// boundary: the crossing is here, at (tNow, yNow), no
|
|||
|
|
// bisection needed. The leaving interval starts from
|
|||
|
|
// zero, so the next step stays silent, exactly the
|
|||
|
|
// "does not fire until it leaves and returns" the
|
|||
|
|
// documented zero-at-start rule spells out. The clamp
|
|||
|
|
// h = t1 − t makes the final boundary land here for
|
|||
|
|
// round numbers, so dropping it would lose events.
|
|||
|
|
rising := gp < 0
|
|||
|
|
if watches[i].Direction > 0 && !rising || watches[i].Direction < 0 && rising {
|
|||
|
|
gPrev[i] = g
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
state := make([]float64, len(yNow))
|
|||
|
|
copy(state, yNow)
|
|||
|
|
hits = append(hits, ODEEventHit{
|
|||
|
|
Time: tNow,
|
|||
|
|
State: arrayFromVector(state),
|
|||
|
|
Watch: i,
|
|||
|
|
Rising: rising,
|
|||
|
|
})
|
|||
|
|
gPrev[i] = g
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
if (gp < 0) == (g < 0) {
|
|||
|
|
gPrev[i] = g
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
// A crossing between tPrev and tNow: bisect on the
|
|||
|
|
// re-integrated watch from the step's start state.
|
|||
|
|
rising := gp < 0
|
|||
|
|
if watches[i].Direction != 0 {
|
|||
|
|
if watches[i].Direction > 0 && !rising {
|
|||
|
|
gPrev[i] = g
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
if watches[i].Direction < 0 && rising {
|
|||
|
|
gPrev[i] = g
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
tHit, yHit, err := refineEvent(f, tPrev, tNow, yPrev, gPrev[i], watches[i].Function, opts)
|
|||
|
|
if err != nil {
|
|||
|
|
return false, base.Errf("IntegrateODEEvents: %w", err)
|
|||
|
|
}
|
|||
|
|
hits = append(hits, ODEEventHit{
|
|||
|
|
Time: tHit,
|
|||
|
|
State: yHit,
|
|||
|
|
Watch: i,
|
|||
|
|
Rising: rising,
|
|||
|
|
})
|
|||
|
|
gPrev[i] = g
|
|||
|
|
}
|
|||
|
|
havePrev = true
|
|||
|
|
return false, nil
|
|||
|
|
}
|
|||
|
|
final, err := odeRun(f, t0, t1, y0, opts, watch)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, nil, err
|
|||
|
|
}
|
|||
|
|
slices.SortFunc(hits, func(a, b ODEEventHit) int {
|
|||
|
|
return cmp.Compare(a.Time, b.Time)
|
|||
|
|
})
|
|||
|
|
return hits, final, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// refineEvent narrows a zero crossing of g between tPrev and tNow by
|
|||
|
|
// bisection, evaluating g by re-integrating from the step's start
|
|||
|
|
// state. Both endpoint values are known to have opposite signs, which
|
|||
|
|
// bisection turns into the crossing at integrator accuracy.
|
|||
|
|
func refineEvent(f func(t float64, y *core.Array) (*core.Array, error),
|
|||
|
|
tPrev, tNow float64, yPrev []float64, gPrev float64,
|
|||
|
|
g func(t float64, y *core.Array) (float64, error), opts ODEOptions) (float64, *core.Array, error) {
|
|||
|
|
lo, hi := tPrev, tNow
|
|||
|
|
for range 100 {
|
|||
|
|
mid := (lo + hi) / 2
|
|||
|
|
if mid == lo || mid == hi {
|
|||
|
|
break
|
|||
|
|
}
|
|||
|
|
yMid, err := IntegrateODE(f, tPrev, mid, wrapVector(yPrev), opts)
|
|||
|
|
if err != nil {
|
|||
|
|
return 0, nil, err
|
|||
|
|
}
|
|||
|
|
gm, err := g(mid, yMid)
|
|||
|
|
if err != nil {
|
|||
|
|
return 0, nil, err
|
|||
|
|
}
|
|||
|
|
if gm == 0 {
|
|||
|
|
return mid, yMid, nil
|
|||
|
|
}
|
|||
|
|
if (gPrev < 0) == (gm < 0) {
|
|||
|
|
lo = mid
|
|||
|
|
gPrev = gm
|
|||
|
|
} else {
|
|||
|
|
hi = mid
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
tHit := (lo + hi) / 2
|
|||
|
|
yHit, err := IntegrateODE(f, tPrev, tHit, wrapVector(yPrev), opts)
|
|||
|
|
if err != nil {
|
|||
|
|
return 0, nil, err
|
|||
|
|
}
|
|||
|
|
return tHit, yHit, nil
|
|||
|
|
}
|