Skip to content

Restore physical time in backward Event callbacks - #779

Draft
Afloat16 wants to merge 1 commit into
patrick-kidger:mainfrom
Afloat16:fix/backward-event-physical-time
Draft

Afloat16 wants to merge 1 commit into
patrick-kidger:mainfrom
Afloat16:fix/backward-event-physical-time

Conversation

@Afloat16

@Afloat16 Afloat16 commented Oct 2, 2026 •

Copy link
Copy Markdown

Validation status: the full-test prerequisite in CONTRIBUTING remains unmet. The bounded CPU recovery attempt is closed with incomplete coverage; this draft is not ready for review.

Default collection on this head: 1,021 nodes. The evidence accounts for 669 passed (617 in the new attempt plus 52 preserved checks on unchanged final source), 2 interrupted, and 350 unrun. No completed failing-node outcome was recorded. Interrupted nodes were not counted as passes; coordinator-held files did not start tests.

A single unchanged pristine test_adjoint.py::test_against control on e950aa2 also exited 137 before completing, after 382.093 s, with a native maximum RSS of 3,222,860 KiB. No cancellation signal was sent. The kill cause is unproven; host-wide OOM counters do not identify a per-process cause. This control establishes neither baseline equivalence nor a full-suite pass. The scoped correctness results below remain separate from this unresolved prerequisite.

Problem

Backward solves normalise time internally, but Event condition functions currently receive that normalised time. For example, integrating y' = t, y(2) = 2 backwards from 2 to 0 with the condition t - 0.9 misses the event entirely and returns at 0. A boolean condition such as t < 1 can instead terminate immediately at the wrong initial time.

Backward integration is a supported diffeqsolve operation. Event callbacks need to observe the physical problem at initialisation, after each step, and during event root finding.

Change

Wrap each condition with an Equinox callable module that restores physical time and the original terms, t0, t1, dt0, SaveAt, and step-size controller. Restoring terms together with t is important: a condition that calls solver.func(terms, t, y, args) would otherwise transform time twice. Explicit PyTree fields keep traced dependencies visible to custom differentiation, including ImplicitAdjoint.

The regression cases cover continuous/boolean conditions, root finding, forward controls, callback problem data, condition PyTrees, mixed-direction vmap, initial events, equal initial/final times, callable Equinox conditions, condition/term parameter gradients, and legacy discrete/steady-state callbacks.

Validation

On pristine upstream e950aa2cf1974c02816c1542f20beabf7cca496b, the 23 new regression cases produce 15 failed, 8 passed. The corrected code produces 23 passed.

Local CPU validation, Python 3.12.14 / JAX 0.11.2:

  • pytest test/test_backward_event.py -q: 23 passed
  • pytest test/test_event.py -q: 42 passed, on pristine and corrected source
  • pytest test/test_adjoint.py -q -m 'not slow': 10 passed, 1 deselected
  • Project-wide Ruff lint and format checks pass; changed-source/new-test Pyright has no errors
  • Project-wide Pyright has the same pre-existing unused type-ignore error in _misc.py:175 on pristine and corrected source

An independent SciPy solve_ivp check and the exact solution y(t) = t**2 / 2 both give event time 0.9, event state 0.405, and sensitivities dT/dthreshold = 1, dY/dthreshold = 0.9. Corrected Diffrax matches these values and centred finite differences. Pristine Diffrax misses the event and returns time/state/sensitivities approximately zero.

The slow adjoint comparison and full project test suite have not completed locally. An earlier multi-file process was killed before completion. Accelerator execution has not been tested, and no runtime or training-quality improvement is claimed.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant