# Error trying to ForwardDiff through an ODE solver

**URL:** <https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339>\
**Category:** General Usage\
**Created:** [May 16, 2024, 5:25am UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339 "2024-05-16T05:25:03Z")\
**Posts on this page:** 10\
**Page:** 1

<div class="post-metadata">

**Author:** ![orebas](https://avatars.discourse-cdn.com/v4/letter/o/bbe5ce/32.png) [@orebas](https://discourse.julialang.org/u/orebas)\
**Post date:** [May 16, 2024, 5:25am UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/1 "2024-05-16T05:25:03Z")

</div>

I am trying to use AD to differentiate the solution of an ODE. Various docs seem to indicate this is very well supported, so I was surprised to get an error.

For some additional color, I got a different error before adding  
“ODEProblem{true, SciMLBase.FullSpecialize}”. I don’t quite understand what this does, but it was recommended in another discourse topic. I also get a slightly different error when switching between Tsit5() and Vern9().

Copy-pastable MWE:

```julia
using ModelingToolkit, DifferentialEquations
using ForwardDiff

function ADTest()
	@parameters a b
	@variables t x1(t) x2(t) y1(t) y2(t)
	D = Differential(t)
	states = [x1, x2]
	parameters = [a, b]

	@named model = ODESystem([
			D(x1) ~ a * x1,
			D(x2) ~ b * x2,
		], t, states, parameters)
	model = structural_simplify(model)
	measured_quantities = [
		y1 ~ x1,
		y2 ~ x2]

	ic = Dict(x1 => 1.0, x2 => 2.0)
	p_true = Dict(a => 2.0, b => 3.0)

	problem = ODEProblem(model, ic, [0.0, 1e-5], p_true)
	soln = ModelingToolkit.solve(problem, Tsit5(), abstol = 1e-14, reltol = 1e-14)
	display(soln(1e-5, idxs = [x1, x2]))

	function different_time(new_ic, new_params, new_t)
		newprob = ODEProblem{true, SciMLBase.FullSpecialize}(model, new_ic, [0.0, new_t], new_params)
		new_soln = ModelingToolkit.solve(newprob, Tsit5(), abstol = 1e-14, reltol = 1e-14)
		return (soln(new_t, idxs = [x1, x2]))
	end
    display(different_time(ic,p_true, 2e-5))

    temp = ForwardDiff.derivative(s -> different_time(ic,p_true, s),4e-5)
    display(temp)
end

ADTest()

```

Error:

```julia
ERROR: LoadError: MethodError: no method matching Float64(::ForwardDiff.Dual{ForwardDiff.Tag{var"#209#211"{var"#different_time#210"{ODESolution{…}, Num, Num}, Dict{Num, Float64}, Dict{Num, Float64}}, Float64}, Float64, 1})

Closest candidates are:
  (::Type{T})(::Real, ::RoundingMode) where T<:AbstractFloat
   @ Base rounding.jl:207
  (::Type{T})(::T) where T<:Number
   @ Core boot.jl:792
  Float64(::IrrationalConstants.Fourinvπ)
   @ IrrationalConstants ~/.julia/packages/IrrationalConstants/vp5v4/src/macro.jl:112
  ...

Stacktrace:
  [1] convert(::Type{Float64}, x::ForwardDiff.Dual{ForwardDiff.Tag{var"#209#211"{var"#different_time#210"{ODESolution{…}, Num, Num}, Dict{Num, Float64}, Dict{Num, Float64}}, Float64}, Float64, 1})
    @ Base ./number.jl:7
  [2] setindex!(A::Vector{Float64}, x::ForwardDiff.Dual{ForwardDiff.Tag{var"#209#211"{var"#different_time#210"{ODESolution{…}, Num, Num}, Dict{Num, Float64}, Dict{Num, Float64}}, Float64}, Float64, 1}, i1::Int64)
    @ Base ./array.jl:1021
  [3] macro expansion
    @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/initdt.jl:119 [inlined]
  [4] macro expansion
    @ ./simdloop.jl:77 [inlined]
  [5] ode_determine_initdt(u0::Vector{…}, t::ForwardDiff.Dual{…}, tdir::ForwardDiff.Dual{…}, dtmax::ForwardDiff.Dual{…}, abstol::Float64, reltol::Float64, internalnorm::typeof(DiffEqBase.ODE_DEFAULT_NORM), prob::ODEProblem{…}, integrator::OrdinaryDiffEq.ODEIntegrator{…})
    @ OrdinaryDiffEq ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/initdt.jl:118
  [6] auto_dt_reset!
    @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/integrators/integrator_interface.jl:453 [inlined]
  [7] handle_dt!(integrator::OrdinaryDiffEq.ODEIntegrator{…})
    @ OrdinaryDiffEq ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:571
  [8] __init(prob::ODEProblem{…}, alg::Tsit5{…}, timeseries_init::Tuple{}, ts_init::Tuple{}, ks_init::Tuple{}, recompile::Type{…}; saveat::Tuple{}, tstops::Tuple{}, d_discontinuities::Tuple{}, save_idxs::Nothing, save_everystep::Bool, save_on::Bool, save_start::Bool, save_end::Nothing, callback::Nothing, dense::Bool, calck::Bool, dt::ForwardDiff.Dual{…}, dtmin::ForwardDiff.Dual{…}, dtmax::ForwardDiff.Dual{…}, force_dtmin::Bool, adaptive::Bool, gamma::Rational{…}, abstol::Float64, reltol::Float64, qmin::Rational{…}, qmax::Int64, qsteady_min::Int64, qsteady_max::Int64, beta1::Nothing, beta2::Nothing, qoldinit::Rational{…}, controller::Nothing, fullnormalize::Bool, failfactor::Int64, maxiters::Int64, internalnorm::typeof(DiffEqBase.ODE_DEFAULT_NORM), internalopnorm::typeof(LinearAlgebra.opnorm), isoutofdomain::typeof(DiffEqBase.ODE_DEFAULT_ISOUTOFDOMAIN), unstable_check::typeof(DiffEqBase.ODE_DEFAULT_UNSTABLE_CHECK), verbose::Bool, timeseries_errors::Bool, dense_errors::Bool, advance_to_tstop::Bool, stop_at_next_tstop::Bool, initialize_save::Bool, progress::Bool, progress_steps::Int64, progress_name::String, progress_message::typeof(DiffEqBase.ODE_DEFAULT_PROG_MESSAGE), progress_id::Symbol, userdata::Nothing, allow_extrapolation::Bool, initialize_integrator::Bool, alias_u0::Bool, alias_du0::Bool, initializealg::OrdinaryDiffEq.DefaultInit, kwargs::@Kwargs{})
    @ OrdinaryDiffEq ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:533
  [9] __init (repeats 5 times)
    @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:11 [inlined]
 [10] #__solve#787
    @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:6 [inlined]
 [11] __solve
    @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:1 [inlined]
 [12] solve_call(_prob::ODEProblem{…}, args::Tsit5{…}; merge_callbacks::Bool, kwargshandle::Nothing, kwargs::@Kwargs{…})
    @ DiffEqBase ~/.julia/packages/DiffEqBase/WyGjp/src/solve.jl:612
 [13] solve_call
    @ ~/.julia/packages/DiffEqBase/WyGjp/src/solve.jl:569 [inlined]
 [14] #solve_up#53
    @ ~/.julia/packages/DiffEqBase/WyGjp/src/solve.jl:1080 [inlined]
 [15] solve_up
    @ ~/.julia/packages/DiffEqBase/WyGjp/src/solve.jl:1066 [inlined]
 [16] #solve#51
    @ ~/.julia/packages/DiffEqBase/WyGjp/src/solve.jl:1003 [inlined]
 [17] (::var"#different_time#210"{ODESolution{…}, Num, Num})(new_ic::Dict{Num, Float64}, new_params::Dict{Num, Float64}, new_t::ForwardDiff.Dual{ForwardDiff.Tag{…}, Float64, 1})
    @ Main ~/learning/ODETests/PLI/MWE2.jl:32
 [18] #209
    @ ~/learning/ODETests/PLI/MWE2.jl:37 [inlined]
 [19] derivative(f::var"#209#211"{var"#different_time#210"{ODESolution{…}, Num, Num}, Dict{Num, Float64}, Dict{Num, Float64}}, x::Float64)
    @ ForwardDiff ~/.julia/packages/ForwardDiff/PcZ48/src/derivative.jl:14

```

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [May 16, 2024, 8:05am UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/2 "2024-05-16T08:05:59Z")

</div>

We should probably update the title. This has nothing to do with diffing through the solver. The issue is that you are doing a symbolic code generation process in your loss function and trying to differentiate the symbolic codegen. This is a known issue with ModelingToolkit and unrelated to the differential equation solver:

> <https://github.com/SciML/ModelingToolkit.jl/issues/2667>
>
> Auto-differentiating through \`remake(ODEProblem())\` works, but not directly thro…ugh \`ODEProblem()\`:
> \`\`\`julia
> using Test
> using ModelingToolkit
> using ModelingToolkit: t\_nounits as t, D\_nounits as D
> using DifferentialEquations
> using ForwardDiff
> 
> @testset "ForwardDiff through ODEProblem with vs. without remake" begin
> @parameters P
> @variables x(t)
> sys = structural\_simplify(ODESystem(\[D(x) ~ P\], t, \[x\], \[P\]; name=:sys))
>     
> function x\_at\_1(P; use\_remake = false)
> if use\_remake
> prob = ODEProblem(sys, \[x =\> 0.0\], (0.0, 1.0), \[sys.P =\> NaN\])
> prob = remake(prob; p = \[sys.P =\> P\])
> else
> prob = ODEProblem(sys, \[x =\> 0.0\], (0.0, 1.0), \[sys.P =\> P\])
> end
> return solve(prob)(1.0)
> end
> 
> @test\_nowarn ForwardDiff.derivative(P -\> x\_at\_1(P; use\_remake=true), 1.0) # passes
> @test\_nowarn ForwardDiff.derivative(P -\> x\_at\_1(P; use\_remake=false), 1.0) # fails
> end
> \`\`\`
> The last test errors with
> \`\`\`julia
> MethodError: no method matching Float64(::ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1})
> 
> Closest candidates are:
> (::Type{T})(::Real, ::RoundingMode) where T\<:AbstractFloat
> @ Base rounding.jl:207
> (::Type{T})(::T) where T\<:Number
> @ Core boot.jl:792
> Float64(::IrrationalConstants.Loghalf)
> @ IrrationalConstants C:\\Users\\herma\\.julia\\packages\\IrrationalConstants\\vp5v4\\src\\macro.jl:112
> ...
> 
> Stacktrace:
> \[1\] convert(::Type{Float64}, x::ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1})
> @ Base .\\number.jl:7
> \[2\] symconvert(::Type{Float64}, x::ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1})
> @ ModelingToolkit C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\parameter\_buffer.jl:2
> \[3\] ModelingToolkit.MTKParameters(sys::ODESystem, p::Dict{SymbolicUtils.BasicSymbolic{Real}, ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1}}, u0::Vector{Pair{Num, Float64}}; tofloat::Bool, use\_union::Bool)       
> @ ModelingToolkit C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\parameter\_buffer.jl:114
> \[4\] ModelingToolkit.MTKParameters(sys::ODESystem, p::Dict{SymbolicUtils.BasicSymbolic{Real}, ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1}}, u0::Vector{Pair{Num, Float64}})
> @ ModelingToolkit C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\parameter\_buffer.jl:13
> \[5\] process\_DEProblem(constructor::Type, sys::ODESystem, u0map::Vector{Pair{Num, Float64}}, parammap::Vector{Pair{Num, ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1}}}; implicit\_dae::Bool, du0map::Nothing, version::Nothing, tgrad::Bool, jac::Bool, checkbounds::Bool, sparse::Bool, simplify::Bool, linenumbers::Bool, parallel::Symbolics.SerialForm, eval\_expression::Bool, use\_union::Bool, tofloat::Bool, symbolic\_u0::Bool, u0\_constructor::typeof(identity), guesses::Dict{Any, Any}, t::Float64, warn\_initialize\_determined::Bool, build\_initializeprob::Bool, kwargs::@Kwargs{check\_length::Bool})
> @ ModelingToolkit C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\diffeqs\\abstractodesystem.jl:942
> \[6\] process\_DEProblem
> @ C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\diffeqs\\abstractodesystem.jl:834 \[inlined\]
> \[7\] (ODEProblem{true, SciMLBase.AutoSpecialize})(sys::ODESystem, u0map::Vector{Pair{Num, Float64}}, tspan::Tuple{Float64, Float64}, parammap::Vector{Pair{Num, ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1}}}; callback::Nothing, check\_length::Bool, warn\_initialize\_determined::Bool, kwargs::@Kwargs{})
> @ ModelingToolkit C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\diffeqs\\abstractodesystem.jl:1085
> \[8\] (ODEProblem{true, SciMLBase.AutoSpecialize})(sys::ODESystem, u0map::Vector{Pair{Num, Float64}}, tspan::Tuple{Float64, Float64}, parammap::Vector{Pair{Num, ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1}}})    
> @ ModelingToolkit C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\diffeqs\\abstractodesystem.jl:1075
> \[9\] (ODEProblem{true})(::ODESystem, ::Vector{Pair{Num, Float64}}, ::Vararg{Any}; kwargs::@Kwargs{})
> @ ModelingToolkit C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\diffeqs\\abstractodesystem.jl:1062
> \[10\] (ODEProblem{true})(::ODESystem, ::Vector{Pair{Num, Float64}}, ::Vararg{Any})
> @ ModelingToolkit C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\diffeqs\\abstractodesystem.jl:1061
> \[11\] #ODEProblem#755
> @ C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\diffeqs\\abstractodesystem.jl:1051 \[inlined\]
> \[12\] ODEProblem
> @ C:\\Users\\herma\\.julia\\packages\\ModelingToolkit\\kByuD\\src\\systems\\diffeqs\\abstractodesystem.jl:1050 \[inlined\]
> \[13\] (::var"#x\_at\_1#23"{var"#x\_at\_1#16#24"{ODESystem, Num}})(P::ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1}; use\_remake::Bool)     
> @ Main C:\\Users\\herma\\Dropbox\\School\\UIO\\Research\\boltzmann\\bug.jl:17    
> \[14\] (::var"#22#30")(P::ForwardDiff.Dual{ForwardDiff.Tag{var"#22#30", Float64}, Float64, 1})
> @ Main C:\\Users\\herma\\Dropbox\\School\\UIO\\Research\\boltzmann\\bug.jl:23    
> \[15\] derivative
> @ C:\\Users\\herma\\.julia\\packages\\ForwardDiff\\PcZ48\\src\\derivative.jl:14 \[inlined\]
> ...
> \`\`\`
> 
> Could it be unified to work in both ways?
> 
> I'm on a fresh updated \`master\` branch with \`\] status\`
> \`\`\`julia
> \[0c46a032\] DifferentialEquations v7.13.0
> \[f6369f11\] ForwardDiff v0.10.36
> \[961ee093\] ModelingToolkit v9.12.1 \`https://github.com/SciML/ModelingToolkit.jl.git#master\`
> \[1dea7af3\] OrdinaryDiffEq v6.74.1
> \`\`\`

Though of course, it’s kind of not sensical in a way because, just as the issue describes at the top, you shouldn’t be doing a symbolic codegen in the loss function in the first place. So if you did something standard like:

```julia
	function different_time(new_ic, new_params, new_t)
		newprob = remake(problem, new_ic, tspan = (0.0, new_t), p=new_params)
		new_soln = ModelingToolkit.solve(newprob, Tsit5(), abstol = 1e-14, reltol = 1e-14)
		return (soln(new_t, idxs = [x1, x2]))
	end

```

Then it should work fine, and it wouldn’t need to codegen new models every step. Note that the documentation specifically has a guide on the symbolic tooling which is helpful for further optimizing this code via `setp` and `setu`:

> **[Optimizing through an ODE solve and re-creating MTK Problems ·...](https://docs.sciml.ai/ModelingToolkit/stable/examples/remake/)**
>
> Documentation for ModelingToolkit.jl.

Building new models every step is useful for things like genetic algorithms which are trying to learn what the equations are via some evolution, but any change to the equations is generally a discrete change to the gradient (or adding new parameters) in which case it’s hard to think of a legitimate use of doing the codegen within the solving process itself. So it hasn’t gotten the highest priority to fix this, but since it is a workflow thing some newcomers may run into we will get around to fixing it. I expect to throw a warning though, i.e. if we notice you’re diffing this, we warn you by default (with an option to turn off) mentioning that you likely want to reuse generated models via `remake`, `setp`, etc. see that page, as a way to help users find the right solution.

---

<div class="post-metadata">

**Author:** ![orebas](https://avatars.discourse-cdn.com/v4/letter/o/bbe5ce/32.png) [@orebas](https://discourse.julialang.org/u/orebas)\
**Post date:** [May 16, 2024, 12:24pm UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/3 "2024-05-16T12:24:32Z")

</div>

Hi. I had actually been use remake before, and switched to ODEProblem to try and debug. Below is a different copy-pastable code. It uses remake() as you recommend, and I tried hooking up to 7 different AD methods. They all fail, but they give slightly different errors.

All of the errors are around typing. My guess is that there is some problem with specifying the timespan with anything not a Float64 (i.e. the Dual numbers are an issue), but I’m surprised all of these AD systems make the same assumption. For instance, here is the error from ReverseDiff:

```julia
ERROR: LoadError: ArgumentError: Converting an instance of ReverseDiff.TrackedReal{Float64, Float64, Nothing} to Float64 is not defined. Please use `ReverseDiff.value` instead.

```

New MWE:

````julia
using ModelingToolkit, DifferentialEquations
using TaylorDiff, ForwardDiff
using DifferentiationInterface, Enzyme, Zygote, ReverseDiff

function ADTest()
	@parameters a
	@variables t x1(t) 
	D = Differential(t)
	states = [x1]
	parameters = [a]

	@named pre_model = ODESystem([D(x1) ~ a * x1], t, states, parameters)
	model = structural_simplify(pre_model)

	ic = Dict(x1 => 1.0)
	p_true = Dict(a => 2.0)

	problem = ODEProblem{true, SciMLBase.FullSpecialize}(model, ic, [0.0, 1.0], p_true)
	soln = ModelingToolkit.solve(problem, Tsit5(), abstol = 1e-12, reltol = 1e-12)
	display(soln(0.5, idxs = [x1]))

	function different_time(new_ic, new_params, new_t)
		#newprob = ODEProblem{true, SciMLBase.FullSpecialize}(model, new_ic, [0.0, new_t*2], new_params)
		#newprob = remake(problem, u0=new_ic, tspan = [0.0, new_t], p = new_params)
		newprob = remake(problem, u0 = new_ic, tspan = [0.0, new_t], p=new_params)
        new_soln = ModelingToolkit.solve(newprob, Tsit5(), abstol = 1e-12, reltol = 1e-12)
		return (soln(new_t, idxs = [x1]))
	end

	function just_t(new_t)
		return different_time(ic, p_true, new_t)[1]
	end
	display(different_time(ic, p_true, 2e-5))
	display(just_t(0.5))

	
    g = ForwardDiff.derivative(just_t,4e-5)
	g = TaylorDiff.derivative(just_t,4e-5,1)
    value_and_gradient(just_t, AutoForwardDiff(), 1.0) 
	value_and_gradient(just_t, AutoReverseDiff(), 1.0) 	
    value_and_gradient(just_t, AutoEnzyme(Enzyme.Reverse), 1.0) 
	value_and_gradient(just_t, AutoEnzyme(Enzyme.Forward), 1.0) 
    value_and_gradient(just_t, AutoZygote(), 1.0) 
end

ADTest()
```
````

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [May 16, 2024, 1:10pm UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/4 "2024-05-16T13:10:44Z")

</div>

Remake should promote the type of of `u0` to match, i.e. equivalent to:

```julia
newprob = remake(problem, u0 = new_ic, tspan = [0.0, new_t], p=new_params)
newprob = remake(newprob, u0 = typeof(new_t).(newprob.u0))

```

I’m traveling right now, but double check that works.

Can you open an issue in ModelingToolkit.jl? @cryptic.ax take note this might be a promotion case that was missed.

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [May 16, 2024, 4:10pm UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/5 "2024-05-16T16:10:22Z")

</div>

The use of DifferentiationInterface requires creating closures which can both prevent Enzyme from differentiating code, as well as hinder performance.

What happens if you just use Enzyme.autodiff of `different_time`

---

<div class="post-metadata">

**Author:** ![orebas](https://avatars.discourse-cdn.com/v4/letter/o/bbe5ce/32.png) [@orebas](https://discourse.julialang.org/u/orebas)\
**Post date:** [May 16, 2024, 5:51pm UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/6 "2024-05-16T17:51:23Z")

</div>

> [@ChrisRackauckas](#):
>
> `newprob = remake(newprob, u0 = typeof(new_t).(newprob.u0))`

I can confirm that your fix makes ForwardDiff work (and the value is sane).

What’s your opinion on the whether I need to put FullSpecialize vs NoSpecialize vs nothing? Which is most likely to help AD packages succeed?

ReverseDiff, TaylorDiff, Enzyme, and Zygote all fail for different reasons. I’ll try and notify various places about… 5 different issues or so. (I haven’t tried Enzyme except through DifferentiationInterface.)

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [May 16, 2024, 6:01pm UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/7 "2024-05-16T18:01:57Z")

</div>

> [@orebas](#):
>
> What’s your opinion on the whether I need to put FullSpecialize vs NoSpecialize vs nothing? Which is most likely to help AD packages succeed?

You shouldn’t need to do any of that.

> [@orebas](#):
>
> ReverseDiff, TaylorDiff, Enzyme, and Zygote all fail for different reasons. I’ll try and notify various places about… 5 different issues or so. (I haven’t tried Enzyme except through DifferentiationInterface.)

ReverseDiff and TaylorDiff should be the same fix. This just needs a ModelingToolkit issue.

---

<div class="post-metadata">

**Author:** ![orebas](https://avatars.discourse-cdn.com/v4/letter/o/bbe5ce/32.png) [@orebas](https://discourse.julialang.org/u/orebas)\
**Post date:** [May 16, 2024, 9:01pm UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/8 "2024-05-16T21:01:56Z")

</div>

For bookkeeping, here’s the issue at MTK, feel free to add comments

> <https://github.com/SciML/ModelingToolkit.jl/issues/2721>
>
> \*\*Describe the bug 🐞\*\*
> 
> ForwardDiff.jl fails to differentiate a simple ODE. A… workaround was given in \[https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/6\](url). That workaround is commented out in the below MWE (search for "typeof").
> 
> 
> \*\*Minimal Reproducible Example 👇\*\*
> 
> \`\`\`
> using ModelingToolkit, DifferentialEquations
> using TaylorDiff, ForwardDiff
> using DifferentiationInterface, Enzyme, Zygote, ReverseDiff
> using SciMLSensitivity
> import Base.isnan
> function isnan(x::TaylorScalar{Float64, 2})
> return false
> end
> 
> function ADTest()
> @parameters a
> @variables t x1(t) 
> D = Differential(t)
> states = \[x1\]
> parameters = \[a\]
> 
> @named pre\_model = ODESystem(\[D(x1) ~ a \* x1\], t, states, parameters)
> model = structural\_simplify(pre\_model)
> 
> ic = Dict(x1 =\> 1.0)
> p\_true = Dict(a =\> 2.0)
> 
> problem = ODEProblem{true, SciMLBase.FullSpecialize}(model, ic, \[0.0, 1.0\], p\_true)
> soln = ModelingToolkit.solve(problem, Tsit5(), abstol = 1e-12, reltol = 1e-12)
> display(soln(0.5, idxs = \[x1\]))
> 
> function different\_time(new\_ic, new\_params, new\_t)
> #newprob = ODEProblem{true, SciMLBase.FullSpecialize}(model, new\_ic, \[0.0, new\_t\*2\], new\_params)
> 		
> newprob = remake(problem, u0=new\_ic, tspan = \[0.0, new\_t\], p = new\_params)
> 		
> #newprob = remake(problem, u0 = new\_ic, tspan = \[0.0, new\_t\], p=new\_params)
> #newprob = remake(newprob, u0 = typeof(new\_t).(newprob.u0))
>         
> new\_soln = ModelingToolkit.solve(newprob, Tsit5(), abstol = 1e-12, reltol = 1e-12)
> return (soln(new\_t, idxs = \[x1\]))
> end
> 
> function just\_t(new\_t)
> return different\_time(ic, p\_true, new\_t)\[1\]
> end
> display(different\_time(ic, p\_true, 2e-5))
> display(just\_t(0.5))
> 
> 	
> display(ForwardDiff.derivative(just\_t,1.0))
> #display(TaylorDiff.derivative(just\_t,1.0,1)) #isnan error
> #display(value\_and\_gradient(just\_t, AutoForwardDiff(), 1.0)) 
> #display(value\_and\_gradient(just\_t, AutoReverseDiff(), 1.0)) 	
> #display(value\_and\_gradient(just\_t, AutoEnzyme(Enzyme.Reverse), 1.0)) 
> #display(value\_and\_gradient(just\_t, AutoEnzyme(Enzyme.Forward), 1.0)) 
> #display(value\_and\_gradient(just\_t, AutoZygote(), 1.0)) 
> 	
> end
> 
> ADTest()
> 
> 
> \`\`\`
> 
> \*\*Error & Stacktrace ⚠️\*\*
> 
> \`\`\`ERROR: LoadError: MethodError: no method matching Float64(::ForwardDiff.Dual{ForwardDiff.Tag{var"#just\_t#6"{var"#different\_time#5"{ODESolution{…}, ODEProblem{…}, Num}, Dict{Num, Float64}, Dict{Num, Float64}}, Float64}, Float64, 1})
> 
> Closest candidates are:
> (::Type{T})(::Real, ::RoundingMode) where T\<:AbstractFloat
> @ Base rounding.jl:207
> (::Type{T})(::T) where T\<:Number
> @ Core boot.jl:792
> Float64(::IrrationalConstants.Fourinvπ)
> @ IrrationalConstants ~/.julia/packages/IrrationalConstants/vp5v4/src/macro.jl:112
> ...
> 
> Stacktrace:
> \[1\] convert(::Type{Float64}, x::ForwardDiff.Dual{ForwardDiff.Tag{var"#just\_t#6"{var"#different\_time#5"{ODESolution{…}, ODEProblem{…}, Num}, Dict{Num, Float64}, Dict{Num, Float64}}, Float64}, Float64, 1})
> @ Base ./number.jl:7
> \[2\] setindex!(A::Vector{Float64}, x::ForwardDiff.Dual{ForwardDiff.Tag{var"#just\_t#6"{var"#different\_time#5"{…}, Dict{…}, Dict{…}}, Float64}, Float64, 1}, i1::Int64)
> @ Base ./array.jl:1021
> \[3\] macro expansion
> @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/initdt.jl:119 \[inlined\]
> \[4\] macro expansion
> @ ./simdloop.jl:77 \[inlined\]
> \[5\] ode\_determine\_initdt(u0::Vector{…}, t::ForwardDiff.Dual{…}, tdir::ForwardDiff.Dual{…}, dtmax::ForwardDiff.Dual{…}, abstol::Float64, reltol::Float64, internalnorm::typeof(DiffEqBase.ODE\_DEFAULT\_NORM), prob::ODEProblem{…}, integrator::OrdinaryDiffEq.ODEIntegrator{…})
> @ OrdinaryDiffEq ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/initdt.jl:118
> \[6\] auto\_dt\_reset!
> @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/integrators/integrator\_interface.jl:453 \[inlined\]
> \[7\] handle\_dt!(integrator::OrdinaryDiffEq.ODEIntegrator{…})
> @ OrdinaryDiffEq ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:571
> \[8\] \_\_init(prob::ODEProblem{…}, alg::Tsit5{…}, timeseries\_init::Tuple{}, ts\_init::Tuple{}, ks\_init::Tuple{}, recompile::Type{…}; saveat::Tuple{}, tstops::Tuple{}, d\_discontinuities::Tuple{}, save\_idxs::Nothing, save\_everystep::Bool, save\_on::Bool, save\_start::Bool, save\_end::Nothing, callback::Nothing, dense::Bool, calck::Bool, dt::ForwardDiff.Dual{…}, dtmin::ForwardDiff.Dual{…}, dtmax::ForwardDiff.Dual{…}, force\_dtmin::Bool, adaptive::Bool, gamma::Rational{…}, abstol::Float64, reltol::Float64, qmin::Rational{…}, qmax::Int64, qsteady\_min::Int64, qsteady\_max::Int64, beta1::Nothing, beta2::Nothing, qoldinit::Rational{…}, controller::Nothing, fullnormalize::Bool, failfactor::Int64, maxiters::Int64, internalnorm::typeof(DiffEqBase.ODE\_DEFAULT\_NORM), internalopnorm::typeof(LinearAlgebra.opnorm), isoutofdomain::typeof(DiffEqBase.ODE\_DEFAULT\_ISOUTOFDOMAIN), unstable\_check::typeof(DiffEqBase.ODE\_DEFAULT\_UNSTABLE\_CHECK), verbose::Bool, timeseries\_errors::Bool, dense\_errors::Bool, advance\_to\_tstop::Bool, stop\_at\_next\_tstop::Bool, initialize\_save::Bool, progress::Bool, progress\_steps::Int64, progress\_name::String, progress\_message::typeof(DiffEqBase.ODE\_DEFAULT\_PROG\_MESSAGE), progress\_id::Symbol, userdata::Nothing, allow\_extrapolation::Bool, initialize\_integrator::Bool, alias\_u0::Bool, alias\_du0::Bool, initializealg::OrdinaryDiffEq.DefaultInit, kwargs::@Kwargs{})
> @ OrdinaryDiffEq ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:533
> \[9\] \_\_init (repeats 5 times)
> @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:11 \[inlined\]
> \[10\] #\_\_solve#787
> @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:6 \[inlined\]
> \[11\] \_\_solve
> @ ~/.julia/packages/OrdinaryDiffEq/GAgjL/src/solve.jl:1 \[inlined\]
> \[12\] solve\_call(\_prob::ODEProblem{…}, args::Tsit5{…}; merge\_callbacks::Bool, kwargshandle::Nothing, kwargs::@Kwargs{…})
> @ DiffEqBase ~/.julia/packages/DiffEqBase/X5SZr/src/solve.jl:612
> \[13\] solve\_call
> @ ~/.julia/packages/DiffEqBase/X5SZr/src/solve.jl:569 \[inlined\]
> \[14\] #solve\_up#53
> @ ~/.julia/packages/DiffEqBase/X5SZr/src/solve.jl:1080 \[inlined\]
> \[15\] solve\_up
> @ ~/.julia/packages/DiffEqBase/X5SZr/src/solve.jl:1066 \[inlined\]
> \[16\] #solve#51
> @ ~/.julia/packages/DiffEqBase/X5SZr/src/solve.jl:1003 \[inlined\]
> \[17\] (::var"#different\_time#5"{ODESolution{…}, ODEProblem{…}, Num})(new\_ic::Dict{Num, Float64}, new\_params::Dict{Num, Float64}, new\_t::ForwardDiff.Dual{ForwardDiff.Tag{…}, Float64, 1})
> @ Main ~/learning/ODETests/PLI/MWE3.jl:35
> \[18\] (::var"#just\_t#6"{var"#different\_time#5"{ODESolution{…}, ODEProblem{…}, Num}, Dict{Num, Float64}, Dict{Num, Float64}})(new\_t::ForwardDiff.Dual{ForwardDiff.Tag{var"#just\_t#6"{…}, Float64}, Float64, 1})
> @ Main ~/learning/ODETests/PLI/MWE3.jl:40
> \[19\] derivative(f::var"#just\_t#6"{var"#different\_time#5"{ODESolution{…}, ODEProblem{…}, Num}, Dict{Num, Float64}, Dict{Num, Float64}}, x::Float64)
> @ ForwardDiff ~/.julia/packages/ForwardDiff/PcZ48/src/derivative.jl:14
> \[20\] ADTest()
> @ Main ~/learning/ODETests/PLI/MWE3.jl:46
> \[21\] top-level scope
> @ ~/learning/ODETests/PLI/MWE3.jl:56
> \[22\] include(fname::String)
> @ Base.MainInclude ./client.jl:489
> \[23\] top-level scope
> @ REPL\[1\]:1
> 
> \`\`\`

I also filed an issue with TaylorDiff.jl

> <https://github.com/JuliaDiff/TaylorDiff.jl/issues/73>
>
> TaylorDiff.jl seems to throw an error when I try to differentiate a fairly simpl…e ODE solver. There is an error on the MTK side, but even after the workaround there (See https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/6) I can't get Taylor diff.jl to work.
> 
> MWE:
> \`\`\`
> using ModelingToolkit, DifferentialEquations
> using TaylorDiff, ForwardDiff
> using DifferentiationInterface, Enzyme, Zygote, ReverseDiff
> using SciMLSensitivity
> \#import Base.isnan
> \#function isnan(x::TaylorScalar{Float64, 2})
> \#	return false
> \#end
> 
> function ADTest()
> @parameters a
> @variables t x1(t) 
> D = Differential(t)
> states = \[x1\]
> parameters = \[a\]
> 
> @named pre\_model = ODESystem(\[D(x1) ~ a \* x1\], t, states, parameters)
> model = structural\_simplify(pre\_model)
> 
> ic = Dict(x1 =\> 1.0)
> p\_true = Dict(a =\> 2.0)
> 
> problem = ODEProblem{true, SciMLBase.FullSpecialize}(model, ic, \[0.0, 1.0\], p\_true)
> soln = ModelingToolkit.solve(problem, Tsit5(), abstol = 1e-12, reltol = 1e-12)
> display(soln(0.5, idxs = \[x1\]))
> 
> function different\_time(new\_ic, new\_params, new\_t)
> #newprob = ODEProblem{true, SciMLBase.FullSpecialize}(model, new\_ic, \[0.0, new\_t\*2\], new\_params)
> #newprob = remake(problem, u0=new\_ic, tspan = \[0.0, new\_t\], p = new\_params)
> newprob = remake(problem, u0 = new\_ic, tspan = \[0.0, new\_t\], p=new\_params)
> newprob = remake(newprob, u0 = typeof(new\_t).(newprob.u0))
> new\_soln = ModelingToolkit.solve(newprob, Tsit5(), abstol = 1e-12, reltol = 1e-12)
> return (soln(new\_t, idxs = \[x1\]))
> end
> 
> function just\_t(new\_t)
> return different\_time(ic, p\_true, new\_t)\[1\]
> end
> display(different\_time(ic, p\_true, 2e-5))
> display(just\_t(0.5))
> 
> 	
> #display(ForwardDiff.derivative(just\_t,1.0))
> display(TaylorDiff.derivative(just\_t,1.0,1)) #isnan error
> #display(value\_and\_gradient(just\_t, AutoForwardDiff(), 1.0)) 
> #display(value\_and\_gradient(just\_t, AutoReverseDiff(), 1.0)) 	
> #display(value\_and\_gradient(just\_t, AutoEnzyme(Enzyme.Reverse), 1.0)) 
> #display(value\_and\_gradient(just\_t, AutoEnzyme(Enzyme.Forward), 1.0)) 
> #display(value\_and\_gradient(just\_t, AutoZygote(), 1.0)) 
> 	
> end
> 
> ADTest()
> 
> \`\`\`

It is possible, as you say, that the solution for TaylorDiff lies in MTK, it’s hard for me to tell. But it looks like at a minimum it needs to handle isnan() on their type.

The error with ReverseDiff is quite complicated and I’m not 100% sure if it’s in ReverseDiff or the DifferentiationInterface layer. I’ll post it in the next comment.

---

<div class="post-metadata">

**Author:** ![orebas](https://avatars.discourse-cdn.com/v4/letter/o/bbe5ce/32.png) [@orebas](https://discourse.julialang.org/u/orebas)\
**Post date:** [May 16, 2024, 9:05pm UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/9 "2024-05-16T21:05:02Z")

</div>

Here’s the error from ReverseDiff, _after_ your workaround to cast u0.

```julia
ERROR: LoadError: MethodError: no method matching length(::ModelingToolkit.MTKParameters{Tuple{Vector{Float64}}, Tuple{}, Tuple{}, Tuple{}, Tuple{}, Nothing, Nothing})

Closest candidates are:
  length(::SymbolicUtils.Code.AtIndex)
   @ SymbolicUtils ~/.julia/packages/SymbolicUtils/JhFWV/src/utils.jl:225
  length(::RecursiveArrayTools.Chain{Tuple{}})
   @ RecursiveArrayTools ~/.julia/packages/RecursiveArrayTools/plzpk/src/utils.jl:326
  length(::MutableArithmetics.Zero)
   @ MutableArithmetics ~/.julia/packages/MutableArithmetics/iovKe/src/rewrite.jl:104
  ...

Stacktrace:
  [1] automatic_sensealg_choice(prob::ODEProblem{…}, u0::Vector{…}, p::ModelingToolkit.MTKParameters{…}, verbose::Bool)
    @ SciMLSensitivity ~/.julia/packages/SciMLSensitivity/rXkM4/src/concrete_solve.jl:84
  [2] _concrete_solve_adjoint(::ODEProblem{…}, ::Tsit5{…}, ::Nothing, ::Vector{…}, ::ModelingToolkit.MTKParameters{…}, ::SciMLBase.ReverseDiffOriginator; verbose::Bool, kwargs::@Kwargs{…})
    @ SciMLSensitivity ~/.julia/packages/SciMLSensitivity/rXkM4/src/concrete_solve.jl:218
  [3] _solve_adjoint(prob::ODEProblem{…}, sensealg::Nothing, u0::Vector{…}, p::ModelingToolkit.MTKParameters{…}, originator::SciMLBase.ReverseDiffOriginator, args::Tsit5{…}; merge_callbacks::Bool, kwargs::@Kwargs{…})
    @ DiffEqBase ~/.julia/packages/DiffEqBase/X5SZr/src/solve.jl:1537
  [4] (::DiffEqBaseReverseDiffExt.var"##solve_up#225#23"{…})(prob::ODEProblem{…}, sensealg::Nothing, u0::ReverseDiff.TrackedArray{…}, p::ModelingToolkit.MTKParameters{…}, args::Tsit5{…}; kwargs::@Kwargs{…})
    @ DiffEqBaseReverseDiffExt ~/.julia/packages/DiffEqBase/X5SZr/ext/DiffEqBaseReverseDiffExt.jl:159
  [5] track(::typeof(DiffEqBase.solve_up), prob::ODEProblem{…}, sensealg::Nothing, u0::ReverseDiff.TrackedArray{…}, p::ModelingToolkit.MTKParameters{…}, args::Tsit5{…}; kwargs::@Kwargs{…})
    @ DiffEqBaseReverseDiffExt ~/.julia/packages/ReverseDiff/p1MzG/src/macros.jl:195
  [6] solve_up(prob::ODEProblem{…}, sensealg::Nothing, u0::ReverseDiff.TrackedArray{…}, p::ModelingToolkit.MTKParameters{…}, args::Tsit5{…}; kwargs::@Kwargs{…})
    @ DiffEqBaseReverseDiffExt ~/.julia/packages/DiffEqBase/X5SZr/ext/DiffEqBaseReverseDiffExt.jl:100
  [7] solve_up(prob::ODEProblem{…}, sensealg::Nothing, u0::Vector{…}, p::ModelingToolkit.MTKParameters{…}, args::Tsit5{…}; kwargs::@Kwargs{…})
    @ DiffEqBaseReverseDiffExt ~/.julia/packages/DiffEqBase/X5SZr/ext/DiffEqBaseReverseDiffExt.jl:142
  [8] solve(prob::ODEProblem{…}, args::Tsit5{…}; sensealg::Nothing, u0::Nothing, p::Nothing, wrap::Val{…}, kwargs::@Kwargs{…})
    @ DiffEqBase ~/.julia/packages/DiffEqBase/X5SZr/src/solve.jl:1003
  [9] (::var"#different_time#9"{ODESolution{…}, ODEProblem{…}, Num})(new_ic::Dict{Num, Float64}, new_params::Dict{Num, Float64}, new_t::ReverseDiff.TrackedReal{Float64, Float64, ReverseDiff.TrackedArray{…}})
    @ Main ~/learning/ODETests/PLI/MWE3.jl:35
 [10] (::var"#just_t#10"{var"#different_time#9"{…}, Dict{…}, Dict{…}})(new_t::ReverseDiff.TrackedReal{Float64, Float64, ReverseDiff.TrackedArray{…}})
    @ Main ~/learning/ODETests/PLI/MWE3.jl:40
 [11] call_composed
    @ ./operators.jl:1044 [inlined]
 [12] (::ComposedFunction{var"#just_t#10"{var"#different_time#9"{…}, Dict{…}, Dict{…}}, typeof(only)})(x::ReverseDiff.TrackedArray{Float64, Float64, 1, Vector{Float64}, Vector{Float64}}; kw::@Kwargs{})
    @ Base ./operators.jl:1041
 [13] ComposedFunction
    @ ./operators.jl:1041 [inlined]
 [14] ReverseDiff.GradientTape(f::ComposedFunction{var"#just_t#10"{…}, typeof(only)}, input::Vector{Float64}, cfg::ReverseDiff.GradientConfig{ReverseDiff.TrackedArray{…}})
    @ ReverseDiff ~/.julia/packages/ReverseDiff/p1MzG/src/api/tape.jl:199
 [15] gradient(f::Function, input::Vector{Float64}, cfg::ReverseDiff.GradientConfig{ReverseDiff.TrackedArray{Float64, Float64, 1, Vector{Float64}, Vector{Float64}}})
    @ ReverseDiff ~/.julia/packages/ReverseDiff/p1MzG/src/api/gradients.jl:22
 [16] gradient
    @ ~/.julia/packages/ReverseDiff/p1MzG/src/api/gradients.jl:22 [inlined]
 [17] value_and_pullback(f::ComposedFunction{var"#just_t#10"{…}, typeof(only)}, ::AutoReverseDiff, x::Vector{Float64}, dy::Float64, ::DifferentiationInterface.NoPullbackExtras)
    @ DifferentiationInterfaceReverseDiffExt ~/.julia/packages/DifferentiationInterface/9POaB/ext/DifferentiationInterfaceReverseDiffExt/onearg.jl:10
 [18] value_and_pullback(f::Function, backend::AutoReverseDiff, x::Vector{Float64}, dy::Float64)
    @ DifferentiationInterface ~/.julia/packages/DifferentiationInterface/9POaB/src/pullback.jl:96
 [19] value_and_pullback
    @ ~/.julia/packages/DifferentiationInterface/9POaB/ext/DifferentiationInterfaceReverseDiffExt/onearg.jl:35 [inlined]
 [20] value_and_gradient
    @ ~/.julia/packages/DifferentiationInterface/9POaB/src/gradient.jl:57 [inlined]
 [21] value_and_gradient(f::Function, backend::AutoReverseDiff, x::Float64)
    @ DifferentiationInterface ~/.julia/packages/DifferentiationInterface/9POaB/src/gradient.jl:57
 [22] ADTest()
    @ Main ~/learning/ODETests/PLI/MWE3.jl:49
 [23] top-level scope
    @ ~/learning/ODETests/PLI/MWE3.jl:56
 [24] include(fname::String)
    @ Base.MainInclude ./client.jl:489
 [25] top-level scope
    @ REPL[1]:1
in expression starting at /home/orebas/learning/ODETests/PLI/MWE3.jl:56
Some type information was truncated. Use `show(err)` to see complete types.

```

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [May 16, 2024, 10:53pm UTC](https://discourse.julialang.org/t/error-trying-to-forwarddiff-through-an-ode-solver/114339/10 "2024-05-16T22:53:35Z")

</div>

Thanks, the two forward ones are easy.

The ReverseDiff one is just an MTK v9 thing that’s known. It’ll be solved by:

> <https://github.com/SciML/SciMLSensitivity.jl/pull/1010>
>
> This uses the SciMLStructures Tunable interface https://github.com/SciML/SciMLSt…ructures.jl in order to allow more generalized definitions of \`p\`.
> 
> \- \[\] Ensure Lux.jl is well supported (componentarrays extension in SciMLStructures
> \- \[\] Add a test for a custom SciMLStructure

Which hopefully should be in a few weeks. Zygote and Enzyme reverse will also hit similar issues to this one.
