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
|
||
}
|