Files

209 lines
6.8 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 (
"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
}