Repository navigation
Exact implicit-function sensitivities for continuous callback event times - #570
Draft
ChrisRackauckas-Claude wants to merge 4 commits into
Draft
ChrisRackauckas-Claude wants to merge 4 commits into
ChrisRackauckas-Claude wants to merge 4 commits into
Conversation
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
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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, andu(t; θ)the dense-output interpolant of the step containing the event, as a function of everything it depends on (u0,p,t0,dt, …). DefineG(t, θ) = g(u(t; θ), t, θ). The event time τ(θ) solvesG(τ, θ) = 0. WithG_t ≠ 0(a transversal crossing), the implicit function theorem givesImplementation (
_implicit_root):ignore_derivatives.ForwardDiff.Dual{EventTimeTag}(τ₀, 1)givesG(τ₀, θ) + G_t ε. The tagged dual always interpolates the state, even when τ₀ equals a stored step endpoint, soGis evaluated at fixed time τ₀.τ̂ = τ₀ - (G - ignore_derivatives(G)) / G_t. The numerator is exactly+0.0in the primal, soτ̂ ≡ τ₀. Enzyme differentiatesGwith respect to θ at fixed τ₀, sodτ̂/dθ = -G_θ/G_t. Because the numerator is zero, the sensitivity ofG_titself drops out.State at the event.
u(τ) = I(τ; θ)has sensitivity∂I/∂θ|_τ + u′(τ)·dτ/dθ.integrator.u = integrator(τ).Tnumerically, master leavesuandtalone. The storedu_Tcarries∂I/∂θ|_T + u′(T)·dT/dθ. So under Enzyme the time change setsu ← u_T - (I(T) - I(τ))andt ← τ. The difference of equal values is+0.0, and the sensitivity becomes∂I/∂θ|_T + u′·dτ/dθ. Later steps start from τ and carrydτ/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 attf. The adaptive kernel already assignstspan[2]directly.Other branches:
find_callback_time, which skips root finding, applies the same correction.NoRootFindkeeps the step end.G_tzero or non-finite, orGnon-finite, raises an error ("tangential or singular crossing") instead of returning a silent gradient.Supporting pieces:
get_conditionmethod forForwardDiff.Dual{EventTimeTag}times.bθsevaluation for that tag only, becauseSimpleDiffEq.bθsrequires θ to share the coefficients' element type.Regression tests (
test/enzyme/gradients.jl)Scalar
gpu_find_rootsensitivities, exact derivative 1 atrtol = 10eps(T), bothLeftRootFindandRightRootFind:(t-p) + 1e4(t-p)^3(t-p) + (t-p)^3atp = 1e9expm1(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:u′ = uevent;p = 0.5, 0.75) and off them (0.74, 0.76), with reset and state-doubling affects;x = [a, b, t0, tf, c, r, dt]foru′ = a,u(t0) = b, reset tor, conditionu + 10t - c. Hereτ = (c - b + a·t0)/(a + 10)andL = r + a(tf - τ), so∇L = [1 - 10c/121, 1/11, -1/11, 1, -1/11, 1, 0]. Run forc = 8.14, 8.25, 8.36, where 8.25 puts τ exactly on a step endpoint, and for bothadaptive = falseandtrue;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)
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…]: thet0sensitivity was lost and a spuriousdtsensitivity appeared.Passing after:
GROUP=Enzyme julia --project -e 'using Pkg; Pkg.test()'on 561cad4The reviewer's
single-event-endpoint-regression.jlpasses 6/6. Theirensemble-endpoint-probe.jlgivesdp = -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 = falseandtrue:u + 10t - c;u - c;t - c;t0 = 0.05,tf = 1.3;All pass, with value error ≤ 2.7e-15 and gradient error ≤ 6.4e-14.
Nonlinear
u′ = a·uwith conditionu - cand an event at τ = 0.74, 0.75 (aligned) and 0.76, against the analyticL = r·exp(a(tf - τ)):dt = 0.25: relative gradient error 1e-6 to 8e-8, matching Tsit5's discretization error (value error 7e-9);dt = 0.01: 4e-14;The aligned case behaves like its neighbours.
Other checks run locally (Julia 1.12.4, CPU)
GROUP=CPUandGROUP=JLArrays(includesGPU callback event detection | 4 4):Testing DiffEqGPU tests passed.GROUP=QA: passes 27/27 withJET static analysis | 75 75when 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.forward-root-probes.jlgave 72/72 non-AD root parity with master's function in round 2. The ITP code is unchanged since then.--check, typos andgit diff --checkare clean on all changed files.Not verified
GPUTsit5(Vern7/9, Rosenbrock). Their interpolants are generic arithmetic, but no Enzyme test covers them. Vector callbacks do not go through this path.LLVM error: Failed to materialize symbols) instead of the message. Examples areone(t) - pand((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
x - x = +0.0for finitex, andu - (+0.0) = ubitwise. 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.tspan[2]whilet_end = t0 + n·dtcarries a different sensitivity (dtort0differentiated). It gives the derivative of the interpolated solution attspan[2], consistent with the non-aligned case.SimpleDiffEq.bθs's polynomials. The alternative is relaxing that function'sθ::Tsignature upstream.Fixes #533
Please ignore this draft until reviewed by @ChrisRackauckas.
Risk assessment
🤖 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
EnsembleGPUKernelintegrators (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 bywithin_autodiff()or by dispatch on the newEventTimeTagdual, 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.561cad4, three checks fail:CUDA Tests (Julia 1),buildkite/.../julia-oneapi-run-tests-on-julia-v1andbuildkite/.../julia-oneapi-run-tests-on-julia-v1-dot-10. All three are pre-existing.4f726d5fails with the identical error: theMethodError: Cannot convert Int64 to EnzymeAllocationRecordintest/enzyme/cuda_records.jl:9, which aborts the Enzyme group.bacef85and145135dalso fail CUDA, and805d50a's run was cancelled. A side effect is that the newgradients.jltests never run on CUDA in CI. The CUDA testsets that run before that point pass on the PR, as they do on master.DiffEqBase.get_conditionand callsSimpleDiffEq.bθs, both non-public, but master already uses both in the same way. The newignore_derivativesis 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.conditionfunctions are now called withForwardDiff.Dualuandt. A condition that annotatest::Float64or writes into a fixed-type buffer will now fail under AD.cuda_records.jlbreaks beforegradients.jlruns; 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 aDualt; the Tsit5bθspolynomials 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)