Skip to content

Exact implicit-function sensitivities for continuous callback event times - #570

Draft
ChrisRackauckas-Claude wants to merge 4 commits into
SciML:masterfrom
ChrisRackauckas-Claude:fix-533-enzyme-event-time-grad
Draft

ChrisRackauckas-Claude wants to merge 4 commits into
SciML:masterfrom
ChrisRackauckas-Claude:fix-533-enzyme-event-time-grad

Conversation

@ChrisRackauckas-Claude

@ChrisRackauckas-Claude ChrisRackauckas-Claude commented Oct 2, 2026 •

Copy link
Copy Markdown
Member

Reverse-mode Enzyme gradients through parameter-dependent continuous-callback event times were wrong. Enzyme differentiated the floating-point ITP bracketing iterations in gpu_find_root, which do not carry the derivative of the mathematical root. Endpoint shortcuts also dropped the event time's sensitivity when an event landed exactly on a step end. Under Enzyme only (within_autodiff()), this PR attaches the implicit-function-theorem sensitivity to the event time and carries it through the state and the final save. Outside Enzyme, every changed branch returns exactly what master returns: the ITP loop is untouched, so roots and iteration counts are identical.

Sensitivity formula

Let g(u, t, θ) be the condition, and u(t; θ) the dense-output interpolant of the step containing the event, as a function of everything it depends on (u0, p, t0, dt, …). Define G(t, θ) = g(u(t; θ), t, θ). The event time τ(θ) solves G(τ, θ) = 0. With G_t ≠ 0 (a transversal crossing), the implicit function theorem gives

dτ/dθ = -G_θ(τ, θ) / G_t(τ, θ)
G_t = ∂g/∂u · u′(τ) + ∂g/∂t
G_θ = ∂g/∂u · ∂u(τ; θ)/∂θ |_(τ fixed) + ∂g/∂θ

Implementation (_implicit_root):

  • τ₀ is the converged root, detached with ignore_derivatives.
  • One evaluation at a ForwardDiff.Dual{EventTimeTag}(τ₀, 1) gives G(τ₀, θ) + G_t ε. The tagged dual always interpolates the state, even when τ₀ equals a stored step endpoint, so G is evaluated at fixed time τ₀.
  • It returns τ̂ = τ₀ - (G - ignore_derivatives(G)) / G_t. The numerator is exactly +0.0 in the primal, so τ̂ ≡ τ₀. Enzyme differentiates G with respect to θ at fixed τ₀, so dτ̂/dθ = -G_θ/G_t. Because the numerator is zero, the sensitivity of G_t itself drops out.

State at the event. u(τ) = I(τ; θ) has sensitivity ∂I/∂θ|_τ + u′(τ)·dτ/dθ.

  • Off the grid, this comes from the existing integrator.u = integrator(τ).
  • When τ equals the step end T numerically, master leaves u and t alone. The stored u_T carries ∂I/∂θ|_T + u′(T)·dT/dθ. So under Enzyme the time change sets u ← u_T - (I(T) - I(τ)) and t ← τ. The difference of equal values is +0.0, and the sensitivity becomes ∂I/∂θ|_T + u′·dτ/dθ. Later steps start from τ and carry dτ/dθ.

Final save. When the fixed-step grid ends exactly at tspan[2], u(tf) = u_end - (I(t_end) - I(tf)) under Enzyme. This is the same construction and gives the interpolated solution's sensitivity at tf. The adaptive kernel already assigns tspan[2] directly.

Other branches:

  • The exact-zero-at-step-end branch of find_callback_time, which skips root finding, applies the same correction.
  • NoRootFind keeps the step end.
  • A zero sign at the start of a step reports no event, so it has no root.
  • G_t zero or non-finite, or G non-finite, raises an error ("tangential or singular crossing") instead of returning a silent gradient.

Supporting pieces:

  • A get_condition method for ForwardDiff.Dual{EventTimeTag} times.
  • A local Tsit5 bθs evaluation for that tag only, because SimpleDiffEq.bθs requires θ to share the coefficients' element type.

Regression tests (test/enzyme/gradients.jl)

  • Scalar gpu_find_root sensitivities, exact derivative 1 at rtol = 10eps(T), both LeftRootFind and RightRootFind:

    • Float32 (t-p) + 1e4(t-p)^3
    • Float64, same condition
    • Float64 (t-p) + (t-p)^3 at p = 1e9
    • Float32 expm1(25000(t-p))

    Each case also checks that the AD primal === the non-AD root.

  • Degenerate crossings raise the error: zero slope at the root, a slope overflowing Float32, and a time derivative that cancels.

  • Full ensemble solves (GPUTsit5, EnsembleGPUKernel(CPU())), all with analytic expectations:

    • linear and curved conditions in Float32 and Float64;
    • a nonlinear u′ = u event;
    • time-only events exactly on step endpoints (p = 0.5, 0.75) and off them (0.74, 0.76), with reset and state-doubling affects;
    • all problem fields: x = [a, b, t0, tf, c, r, dt] for u′ = a, u(t0) = b, reset to r, condition u + 10t - c. Here τ = (c - b + a·t0)/(a + 10) and L = r + a(tf - τ), so ∇L = [1 - 10c/121, 1/11, -1/11, 1, -1/11, 1, 0]. Run for c = 8.14, 8.25, 8.36, where 8.25 puts τ exactly on a step endpoint, and for both adaptive = false and true;
    • six state events per solve: condition u - c, c = 0.15, reset to 0, several events within one original step, ∇L = [1, 1, -1, 1, -6, 6, 0].

Failing before: the same test file against unmodified master 4f726d5 (outer testset so every failure prints)

Continuous event root sensitivities: Float32 curved 0.99989045f0, Float64 p = 1e9 0.0, Float32 expm1 1.0022331f0 (≈ 1; Left and Right)
Event root sensitivities reject degenerate crossings: 3 fail (no error)
Parameter-dependent continuous event time gradients (Float32): Float32[-1.0000005, -0.4999997, -1.0000005] ≈ -1 (and curved -1.0002105, -0.9997793, ...)
Parameter-dependent continuous event time gradients (Float64): [-1.000000000000007, -0.6899945810473186, -1.0000000000000053] ≈ -1
Event time gradients at step endpoints: Evaluated 0.0 ≈ ∓1 at p = 0.5 and 0.75 (both affects)
Event time gradients for all problem fields (adaptive = false), c = 8.25:
   Evaluated: [0.24999999999999994, 0.0, 0.0, 0.0, 0.0, 1.0, 0.9999999999999998] ≈ [0.318…, 0.0909…, -0.0909…, 1.0, -0.0909…, 1.0, 0.0]
gradients.jl | 54 pass  17 fail  71 total

The previous PR head 2eabab5 gives 70 pass / 1 fail. The failure is the reviewer's aligned case, [0.318…, 0.0909…, 0.0, 0.99…, -0.0909…, 1.0, 0.2727…]: the t0 sensitivity was lost and a spurious dt sensitivity appeared.

Passing after: GROUP=Enzyme julia --project -e 'using Pkg; Pkg.test()' on 561cad4

Continuous event root sensitivities (8 testsets)               |    2      2  each
Event root sensitivities reject degenerate crossings           |    3      3
Parameter-dependent continuous event time gradients (Float32)  |    4      4
Parameter-dependent continuous event time gradients (Float64)  |    4      4
Nonlinear continuous event time gradients                      |    2      2
Event time gradients at step endpoints (reset_zero!)           |    8      8
Event time gradients at step endpoints (double_state!)         |    8      8
Event time gradients for all problem fields (adaptive = false) |    8      8
Event time gradients for all problem fields (adaptive = true)  |    8      8
     Testing DiffEqGPU tests passed

The reviewer's single-event-endpoint-regression.jl passes 6/6. Their ensemble-endpoint-probe.jl gives dp = -1.0, -0.999999999999988, -0.9999999999999787.

Additional probes (scratch, not committed)

  • 28 cases, compared against a closed-form ForwardDiff gradient over all seven inputs, for adaptive = false and true:

    • mixed u + 10t - c;
    • state-only u - c;
    • time-only t - c;
    • repeated state events (1, 3 and 6 events, several per original step);
    • shifted t0 = 0.05, tf = 1.3;
    • event times aligned and not aligned with the grid.

    All pass, with value error ≤ 2.7e-15 and gradient error ≤ 6.4e-14.

  • Nonlinear u′ = a·u with condition u - c and an event at τ = 0.74, 0.75 (aligned) and 0.76, against the analytic L = r·exp(a(tf - τ)):

    • fixed dt = 0.25: relative gradient error 1e-6 to 8e-8, matching Tsit5's discretization error (value error 7e-9);
    • fixed dt = 0.01: 4e-14;
    • adaptive: ≤ 2e-11.

    The aligned case behaves like its neighbours.

Other checks run locally (Julia 1.12.4, CPU)

  • GROUP=CPU and GROUP=JLArrays (includes GPU callback event detection | 4 4): Testing DiffEqGPU tests passed.
  • GROUP=QA: passes 27/27 with JET static analysis | 75 75 when run alone. A run concurrent with five other Julia jobs hit Aqua's 60 s persistent-task precompile timeout, as the reviewer also saw once.
  • The reviewer's forward-root-probes.jl gave 72/72 non-AD root parity with master's function in round 2. The ITP code is unchanged since then.
  • Runic --check, typos and git diff --check are clean on all changed files.

Not verified

  • GPU/CUDA is unverified. That covers forward callback kernels and the CUDA Enzyme path, including whether the error path and the dual-time interpolation compile under Enzyme on device. No GPU on this machine.
  • Julia LTS and pre-release.
  • Event gradients for integrators other than GPUTsit5 (Vern7/9, Rosenbrock). Their interpolants are generic arithmetic, but no Enzyme test covers them. Vector callbacks do not go through this path.
  • Known Enzyme limitation: a degenerate condition whose time derivative is a compile-time constant hits an Enzyme JIT failure (LLVM error: Failed to materialize symbols) instead of the message. Examples are one(t) - p and ((t - p) * 1f38) * 10f0. It still errors rather than returning a gradient. Kernel conditions read integrator state, so their slopes are runtime values; the runtime versions raise the message and are tested. I could not reduce this to a standalone Enzyme reproducer.

Points a reviewer should push back on

  • Primal preservation relies on exact cancellation: x - x = +0.0 for finite x, and u - (+0.0) = u bitwise. The cancelling terms are the same function evaluated at equal values. A non-finite condition value at the root is rejected. A non-finite interpolated state at an endpoint event would turn into NaN under Enzyme, but the solution is already non-finite in that case.
  • Final-save correction without events: it also applies when no event occurred, if the fixed-step grid ends exactly at tspan[2] while t_end = t0 + n·dt carries a different sensitivity (dt or t0 differentiated). It gives the derivative of the interpolated solution at tspan[2], consistent with the non-aligned case.
  • Under Enzyme, tangential or singular crossings now throw. Forward solves never throw.
  • Extra work only under Enzyme: one dual condition evaluation per located event, plus two interpolations per endpoint-aligned event or grid end.
  • Duplicated polynomials: the local Tsit5 coefficient evaluation duplicates SimpleDiffEq.bθs's polynomials. The alternative is relaxing that function's θ::T signature upstream.

Fixes #533

Please ignore this draft until reviewed by @ChrisRackauckas.

Risk assessment

  • Risk: medium
  • Blast radius: the callback root-finding path in the EnsembleGPUKernel integrators, under Enzyme reverse mode only (event-time and endpoint sensitivities). The forward (non-AD) root finding is byte-identical to master, with the same roots and iteration counts. No public API change.
  • Evidence: 28 analytic cases (state-only, time-only, mixed and multi-event conditions; aligned and non-aligned; fixed-step and adaptive) match the implicit-function gradients to <= 6.4e-14. The new test file gives 54/17 on master and passes on this head. GROUP=Enzyme, CPU, JLArrays and QA pass on CPU.
  • Not verified: GPU/CUDA (no GPU on the build machine), Julia LTS and pre-release, non-Tsit5 kernel integrators under Enzyme, and the documented Enzyme JIT error for a compile-time-constant degenerate slope.
  • Independent review: Codex CLI (gpt-6-astra) rated it MERGE, medium risk, after four rounds (round 1 found an inaccurate finite-difference slope; round 2 the endpoint-aligned event; round 3 state-dependent endpoint conditions; each was fixed with regressions).
  • Merge: needs human review: it touches shared event paths, and CUDA is unverified.

🤖 Generated with Claude Code 2.1.287 (model: claude-opus-5-5), transcript /home/crackauc/sandbox/goals/issue-backlog/jobs/fix-DiffEqGPU.jl-533/log.txt on amdci2.julia.csail.mit.edu

Risk assessment

  • Risk: medium
  • Blast radius: Continuous-callback event handling in the EnsembleGPUKernel integrators (gpu_find_root, find_callback_time, change_t_via_interpolation!, the fixed-step kernel's final save, and the Tsit5 interpolant). Every changed branch is guarded by within_autodiff() or by dispatch on the new EventTimeTag dual, so in normal solves it returns the same result as master. Only Enzyme reverse-mode users of event callbacks see different behavior: their gradients change, and degenerate crossings now throw. No public API is added or removed.
  • Evidence: On head 561cad4, three checks fail: CUDA Tests (Julia 1), buildkite/.../julia-oneapi-run-tests-on-julia-v1 and buildkite/.../julia-oneapi-run-tests-on-julia-v1-dot-10. All three are pre-existing.
    • CUDA: master 4f726d5 fails with the identical error: the MethodError: Cannot convert Int64 to EnzymeAllocationRecord in test/enzyme/cuda_records.jl:9, which aborts the Enzyme group. bacef85 and 145135d also fail CUDA, and 805d50a's run was cancelled. A side effect is that the new gradients.jl tests never run on CUDA in CI. The CUDA testsets that run before that point pass on the PR, as they do on master.
    • oneAPI: both jobs fail on master builds 1776, 1775, 1756 and 1747. The Buildkite logs need an auth token, so I could not compare their failure text.
    • Everything else passes: CPU (1/lts/pre/x86), Enzyme (1/lts), JLArrays, OpenCL, QA, Downgrade, Documentation, Runic, typos, AMDGPU v1/v1.10 and Metal v1/v1.10.
    • No regressions. No tests are weakened; the PR only adds 160 test lines.
    • Dependency internals: it extends DiffEqBase.get_condition and calls SimpleDiffEq.bθs, both non-public, but master already uses both in the same way. The new ignore_derivatives is exported by EnzymeCore 0.8.22, and the Downgrade job (compat floor 0.8.21) loads the package, which suggests it is also defined at the floor.
    • Not disclosed in the description: under Enzyme, user condition functions are now called with ForwardDiff.Dual u and t. A condition that annotates t::Float64 or writes into a fixed-type buffer will now fail under AD.
  • Independent review: Devin CLI (Mac) / fusion-claude-opus-5-5-high-sidekick-swe-2-medium rated it medium, medium-high confidence: the change is well-tested on CPU and does not alter normal solves, but it is numerically subtle work on a shared event path. Its CUDA/oneAPI Enzyme path is never exercised because of a pre-existing CI breakage, and it adds new throwing behavior under AD.
  • Merge: needs human review (the PR is a draft that explicitly awaits @ChrisRackauckas; CUDA Enzyme event gradients are unverified because cuda_records.jl breaks before gradients.jl runs; oneAPI logs were not compared; Enzyme users now get an error for tangential or singular crossings, including a step-end exact zero with zero slope; the description does not say user conditions must accept a Dual t; the Tsit5 bθs polynomials are copied locally; there are no GitHub review comments, only the offline Codex review the description cites).

🤖 Risk assessment posted by an AI agent (fleet master) — harness: Devin CLI (Mac) / fusion-claude-opus-5-5-high-sidekick-swe-2-medium; dispatched by Devin CLI 3000.11.3 (Mac) fleet-master caretaker, model Fusion (claude-opus-5-5 medium + swe-2 medium)
Conversation: local transcript ~/.local/share/devin/cli/transcripts/speckle-whale.json (Chris's Mac)

ChrisRackauckas and others added 2 commits October 2, 2026 19:17
Reverse-mode Enzyme gradients through parameter-dependent event times were
wrong because they differentiated the ITP bracketing iterations, which do not
carry the derivative of the mathematical root. gpu_find_root now returns the
converged root with the implicit-function-theorem tangent -f_θ/f_t attached
through a correction term that is identically zero in the primal, so the
forward root and its rootfind side are unchanged.

Fixes SciML#533

Co-Authored-By: Chris Rackauckas <accounts@chrisrackauckas.com>
Co-Authored-By: Claude <noreply@anthropic.com>
Agent-Harness: Claude Code 2.1.287
Agent-Model: claude-opus-5-5
Agent-Session: local session, transcript at /home/crackauc/sandbox/goals/issue-backlog/jobs/fix-DiffEqGPU.jl-533/log.txt on amdci2.julia.csail.mit.edu
The implicit-function correction now takes the condition's time derivative
from a ForwardDiff dual at the converged root instead of a finite difference,
which was inaccurate for curved Float32 conditions and large time origins.
The correction runs only under Enzyme (`within_autodiff()`), so the ordinary
forward root-finding path is identical to master. Tangential or singular
crossings (zero or non-finite time derivative) raise an error instead of
returning a zero gradient. Event-time duals use a private tag so they always
interpolate the state, and the Tsit5 interpolant gets a local coefficient
evaluation for them because SimpleDiffEq.bθs requires matching element types.

Co-Authored-By: Chris Rackauckas <accounts@chrisrackauckas.com>
Co-Authored-By: Claude <noreply@anthropic.com>
Agent-Harness: Claude Code 2.1.287
Agent-Model: claude-opus-5-5
Agent-Session: local session, transcript at /home/crackauc/sandbox/goals/issue-backlog/jobs/fix-DiffEqGPU.jl-533/log.txt on amdci2.julia.csail.mit.edu
@ChrisRackauckas-Claude ChrisRackauckas-Claude changed the title Attach implicit-function sensitivities to continuous callback roots Exact implicit-function sensitivities for continuous callback event times Oct 3, 2026
ChrisRackauckas and others added 2 commits October 3, 2026 04:11
An event whose condition is exactly zero at the step end bypassed root finding
and returned `top_t` with no sensitivity, and both the time change and the
final save skip interpolation when the times are equal, so the gradient was
silently zero (e.g. `u′ = 1`, event at `t = p = 0.75`, `dt = 0.25`). Under
Enzyme only, that branch now applies the implicit-function correction, and the
equal-time paths add `u - (I(t_step) - I(t_event))`, which is exactly `u` in the
primal but carries the event time's sensitivity relative to the step end. The
non-AD forward path is unchanged.

The slope of the implicit correction is no longer wrapped in
`ignore_derivatives`: the correction's numerator is exactly zero, so the
slope's own sensitivity drops out, and Enzyme cannot apply
`ignore_derivatives` to the constant zero slope of a time-independent
condition.

Co-Authored-By: Chris Rackauckas <accounts@chrisrackauckas.com>
Co-Authored-By: Claude <noreply@anthropic.com>
Agent-Harness: Claude Code 2.1.287
Agent-Model: claude-opus-5-5
Agent-Session: local session, transcript at /home/crackauc/sandbox/goals/issue-backlog/jobs/fix-DiffEqGPU.jl-533/log.txt on amdci2.julia.csail.mit.edu
The numerator of the implicit-function correction evaluated the condition
through `get_condition`'s endpoint shortcut, so for an event exactly on a
step endpoint it used the stored step-end state, whose sensitivity follows
the moving endpoint (`u′(T)·∂T/∂θ`). That dropped the initial-time
sensitivity and introduced a spurious step-size sensitivity for
state-dependent conditions. The value and the time partial now both come
from the single dual evaluation at the detached root, which always
interpolates, so ∂g/∂θ holds the evaluation time fixed. A non-finite
condition value at the root is reported as degenerate.

Co-Authored-By: Chris Rackauckas <accounts@chrisrackauckas.com>
Co-Authored-By: Claude <noreply@anthropic.com>
Agent-Harness: Claude Code 2.1.287
Agent-Model: claude-opus-5-5
Agent-Session: local session, transcript at /home/crackauc/sandbox/goals/issue-backlog/jobs/fix-DiffEqGPU.jl-533/log.txt on amdci2.julia.csail.mit.edu

This branch has not been deployed

No deployments
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.

Enzyme gradients through parameter-dependent continuous event times are incorrect

2 participants