diff --git a/Project.toml b/Project.toml index b95d8154383..3dededbfb98 100644 --- a/Project.toml +++ b/Project.toml @@ -62,14 +62,14 @@ ForwardDiff = "1.3.3" FunctionWrappersWrappers = "1" NonlinearSolve = "4.20.3" OrdinaryDiffEqBDF = "2.4.5" -OrdinaryDiffEqCore = "4.15.1" +OrdinaryDiffEqCore = "4.15.3, 4.16" OrdinaryDiffEqDefault = "2" OrdinaryDiffEqRosenbrock = "2.6.3" OrdinaryDiffEqTsit5 = "2.1.4" OrdinaryDiffEqVerner = "2" PreallocationTools = "1.1.2" RecursiveArrayTools = "4.2.0" -SciMLBase = "3.46" +SciMLBase = "3.56" SciMLLogging = "2.0.0" SciMLOperators = "1.24.4" SciMLTesting = "2.8" @@ -140,4 +140,4 @@ Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" Unitful = "1986cc42-f94f-5a68-af5c-568840ba703d" [targets] -test = ["ADTypes", "ArrayInterface", "ComponentArrays", "AlgebraicMultigrid", "DiffEqCallbacks", "DifferentiationInterface", "DiffEqDevTools", "ExplicitImports", "ForwardDiff", "FunctionWrappersWrappers", "IncompleteLU", "InteractiveUtils", "LinearAlgebra", "LinearSolve", "ODEProblemLibrary", "OrdinaryDiffEqAdamsBashforthMoulton", "OrdinaryDiffEqDifferentiation", "OrdinaryDiffEqExplicitRK", "OrdinaryDiffEqExplicitTableaus", "OrdinaryDiffEqExponentialRK", "OrdinaryDiffEqExtrapolation", "OrdinaryDiffEqFIRK", "OrdinaryDiffEqFeagin", "OrdinaryDiffEqFunctionMap", "OrdinaryDiffEqHighOrderRK", "OrdinaryDiffEqIMEXMultistep", "OrdinaryDiffEqLinear", "OrdinaryDiffEqLowOrderRK", "OrdinaryDiffEqLowStorageRK", "NonlinearSolve", "OrdinaryDiffEqNonlinearSolve", "OrdinaryDiffEqNordsieck", "OrdinaryDiffEqPDIRK", "OrdinaryDiffEqPRK", "OrdinaryDiffEqQPRK", "OrdinaryDiffEqRKN", "OrdinaryDiffEqSDIRK", "OrdinaryDiffEqSSPRK", "OrdinaryDiffEqStabilizedIRK", "OrdinaryDiffEqStabilizedRK", "OrdinaryDiffEqSymplecticRK", "ElasticArrays", "JLArrays", "Random", "SafeTestsets", "SciMLOperators", "SciMLTesting", "StableRNGs", "StructArrays", "Test", "Unitful", "Pkg", "PreallocationTools", "RecursiveArrayTools", "RecursiveFactorization", "SparseArrays", "SparseConnectivityTracer", "SparseMatrixColorings", "StaticArrays", "Statistics"] +test = ["ADTypes", "ArrayInterface", "ComponentArrays", "AlgebraicMultigrid", "DiffEqCallbacks", "DifferentiationInterface", "DiffEqDevTools", "ExplicitImports", "ForwardDiff", "FunctionWrappersWrappers", "IncompleteLU", "InteractiveUtils", "LinearAlgebra", "LinearSolve", "ODEProblemLibrary", "OrdinaryDiffEqAdamsBashforthMoulton", "OrdinaryDiffEqDifferentiation", "OrdinaryDiffEqExplicitRK", "OrdinaryDiffEqExplicitTableaus", "OrdinaryDiffEqExponentialRK", "OrdinaryDiffEqExtrapolation", "OrdinaryDiffEqFIRK", "OrdinaryDiffEqFeagin", "OrdinaryDiffEqFunctionMap", "OrdinaryDiffEqHighOrderRK", "OrdinaryDiffEqIMEXMultistep", "OrdinaryDiffEqLinear", "OrdinaryDiffEqLowOrderRK", "OrdinaryDiffEqLowStorageRK", "NonlinearSolve", "OrdinaryDiffEqNonlinearSolve", "OrdinaryDiffEqNordsieck", "OrdinaryDiffEqPDIRK", "OrdinaryDiffEqPRK", "OrdinaryDiffEqQPRK", "OrdinaryDiffEqRKN", "OrdinaryDiffEqSDIRK", "OrdinaryDiffEqSSPRK", "OrdinaryDiffEqStabilizedIRK", "OrdinaryDiffEqStabilizedRK", "OrdinaryDiffEqSymplecticRK", "ElasticArrays", "JLArrays", "Random", "SafeTestsets", "SciMLOperators", "SciMLTesting", "StableRNGs", "StructArrays", "Test", "Unitful", "Pkg", "PreallocationTools", "RecursiveArrayTools", "RecursiveFactorization", "SparseArrays", "SparseConnectivityTracer", "SparseMatrixColorings", "StaticArrays", "Statistics"] diff --git a/docs/Project.toml b/docs/Project.toml index ac12ce6c39e..e5d808042a7 100644 --- a/docs/Project.toml +++ b/docs/Project.toml @@ -164,6 +164,6 @@ StochasticDiffEqMilstein = "2" StochasticDiffEqROCK = "2" StochasticDiffEqRODE = "2" StochasticDiffEqWeak = "2" -SciMLBase = "3.39" +SciMLBase = "3.56" SciMLLogging = "2.0.0" SciMLOperators = "1.24" diff --git a/docs/src/assets/Project.toml b/docs/src/assets/Project.toml index ac12ce6c39e..e5d808042a7 100644 --- a/docs/src/assets/Project.toml +++ b/docs/src/assets/Project.toml @@ -164,6 +164,6 @@ StochasticDiffEqMilstein = "2" StochasticDiffEqROCK = "2" StochasticDiffEqRODE = "2" StochasticDiffEqWeak = "2" -SciMLBase = "3.39" +SciMLBase = "3.56" SciMLLogging = "2.0.0" SciMLOperators = "1.24" diff --git a/docs/src/devtools/internals/public_api.md b/docs/src/devtools/internals/public_api.md index 12d7a690e31..c2e21cbddba 100644 --- a/docs/src/devtools/internals/public_api.md +++ b/docs/src/devtools/internals/public_api.md @@ -34,6 +34,7 @@ DiffEqBase.get_condition DiffEqBase.get_tstops DiffEqBase.get_tstops_array DiffEqBase.get_tstops_max +DiffEqBase.has_callbacks DiffEqBase.initialize! DiffEqBase.max_vector_callback_length DiffEqBase.max_vector_callback_length_int diff --git a/docs/src/usage.md b/docs/src/usage.md index 029c74138ba..07f5b1da1f4 100644 --- a/docs/src/usage.md +++ b/docs/src/usage.md @@ -67,6 +67,28 @@ sol2 = solve(prob, KahanLi8(), dt = 1 / 10); Other refined forms are IMEX and semi-linear ODEs (for exponential integrators). +## Reactant compilation + +An explicit ODE solve can be part of a [`Reactant.@jit`](https://enzymead.github.io/Reactant.jl/stable/api/#Reactant.@jit) compiled function. Both adaptive and fixed-step solver loops are staged as device-side loops, so the compiled executable can be reused with new state and parameter values: + +```julia +using OrdinaryDiffEq, Reactant + +f(u, p, t) = p .* u + +function compiled_solve(u, p) + prob = ODEProblem(f, u, (0.0f0, 1.0f0), p) + return solve(prob, Tsit5()) +end + +u0 = Reactant.to_rarray(Float32[1, 2]) +p = Reactant.to_rarray(Float32[-1]) +sol = Reactant.@jit compiled_solve(u0, p) +Array(sol.u[end]) +``` + +Reactant requires statically shaped outputs. A compiled solve therefore returns an endpoint-only `ODESolution`: `sol.u` and `sol.t` contain the final state and time, while `sol.prob`, `sol.stats`, and `sol.interp` are `nothing`. Adaptive solves are tested with `IController`, `PIController`, and `PIDController`; the tested algorithms are the `Tsit5` and Verner explicit Runge–Kutta families. Implicit algorithms, saving intermediate or partial states (`saveat` or `save_idxs`), callbacks, user `tstops`, discontinuity handling, `force_dtmin`, progress reporting, custom domain or instability checks, and step limiters are not currently supported inside Reactant compilation and produce an `ArgumentError` instead of silently changing the solve. + ## Available Solvers For the list of available solvers, please refer to the [DifferentialEquations.jl ODE Solvers](https://docs.sciml.ai/DiffEqDocs/stable/solvers/ode_solve/), [Dynamical ODE Solvers](https://docs.sciml.ai/DiffEqDocs/stable/solvers/dynamical_solve/), and the [Split ODE Solvers](https://docs.sciml.ai/DiffEqDocs/stable/solvers/split_ode_solve/) pages. diff --git a/lib/DiffEqBase/src/DiffEqBase.jl b/lib/DiffEqBase/src/DiffEqBase.jl index fd24e06c64c..9bac5979c86 100644 --- a/lib/DiffEqBase/src/DiffEqBase.jl +++ b/lib/DiffEqBase/src/DiffEqBase.jl @@ -206,6 +206,7 @@ export AutoDespecialize, AutoRespecialize, AutoDePSpecialize :public, :get_tstops, :get_tstops_array, :get_tstops_max, :ExplicitRKTableau, :ImplicitRKTableau, :DECostFunction, :merge_problem_kwargs, + :has_callbacks, # Callback API (DiffEqBase-owned shared functionality used by downstream solvers) :apply_callback!, :apply_discrete_callback!, :CallbackCache, :find_first_continuous_callback, :find_callback_time, diff --git a/lib/DiffEqBase/src/solve.jl b/lib/DiffEqBase/src/solve.jl index 0baa2915f88..44cceea3886 100644 --- a/lib/DiffEqBase/src/solve.jl +++ b/lib/DiffEqBase/src/solve.jl @@ -19,9 +19,13 @@ NO_TSPAN_PROBS = Union{ } """ - has_callbacks(kwargs) + has_callbacks(kwargs) -> Bool -Check if there are any callbacks in the kwargs. Returns `true` if callbacks are present. +Return `true` if `kwargs` contains a nonzero callback. + +An absent or `nothing` `:callback` is treated as no callback. An empty +`CallbackSet` (including the empty erasure set injected on Julia ≥ 1.12) is also +treated as no callback; any other callback value is treated as present. """ function has_callbacks(kwargs) cb = get(kwargs, :callback, nothing) diff --git a/lib/DiffEqBase/test/problem_kwargs_merging.jl b/lib/DiffEqBase/test/problem_kwargs_merging.jl index 1842944f761..9e5db83c963 100644 --- a/lib/DiffEqBase/test/problem_kwargs_merging.jl +++ b/lib/DiffEqBase/test/problem_kwargs_merging.jl @@ -100,3 +100,16 @@ import SciMLBase kwargs_out = DiffEqBase.merge_problem_kwargs(prob_full; kwargs_in...) @test kwargs_out.callback === cb1 end + +@testset "has_callbacks API" begin + @static if VERSION >= v"1.11" + @test Base.ispublic(DiffEqBase, :has_callbacks) + end + @test !DiffEqBase.has_callbacks((;)) + @test !DiffEqBase.has_callbacks((; abstol = 1.0e-6)) + @test !DiffEqBase.has_callbacks((; callback = nothing)) + @test !DiffEqBase.has_callbacks((; callback = CallbackSet())) + cb = DiscreteCallback((u, t, integrator) -> false, integrator -> nothing) + @test DiffEqBase.has_callbacks((; callback = cb)) + @test DiffEqBase.has_callbacks((; callback = CallbackSet(cb))) +end diff --git a/lib/GlobalDiffEq/src/companion.jl b/lib/GlobalDiffEq/src/companion.jl index df6198bdc0a..4a74a636a94 100644 --- a/lib/GlobalDiffEq/src/companion.jl +++ b/lib/GlobalDiffEq/src/companion.jl @@ -26,7 +26,8 @@ function _validate_estimation_problem(prob, name) prob.f.mass_matrix == LinearAlgebra.I || throw(ArgumentError("$name currently requires the standard mass matrix")) problem_kwargs = values(prob.kwargs) - if haskey(problem_kwargs, :callback) && problem_kwargs.callback !== nothing + if haskey(problem_kwargs, :callback) && + DiffEqBase.has_callbacks((; callback = problem_kwargs.callback)) throw(ArgumentError("$name does not currently support callbacks")) end return nothing @@ -264,7 +265,7 @@ function _companion_error_estimate( make_rhs, name, prob, inner_alg, companion_alg, args...; abstol, reltol, companion_abstol, companion_reltol, kwargs... ) - haskey(kwargs, :callback) && + DiffEqBase.has_callbacks(kwargs) && throw(ArgumentError("$name does not currently support callbacks")) _validate_estimation_problem(prob, name) solve_kwargs = merge((; kwargs...), _DENSE_SOLVE_KWARGS) @@ -285,7 +286,7 @@ function _companion_error_estimate_streaming( make_rhs, name, prob, inner_alg, companion_alg, args...; abstol, reltol, companion_abstol, companion_reltol, kwargs... ) - haskey(kwargs, :callback) && + DiffEqBase.has_callbacks(kwargs) && throw(ArgumentError("$name does not currently support callbacks")) _validate_estimation_problem(prob, name) integrator = SciMLBase.init( diff --git a/lib/GlobalDiffEq/src/estimation.jl b/lib/GlobalDiffEq/src/estimation.jl index 36599c720de..9390efcbc33 100644 --- a/lib/GlobalDiffEq/src/estimation.jl +++ b/lib/GlobalDiffEq/src/estimation.jl @@ -134,7 +134,7 @@ function SciMLBase.__solve( ArgumentError("GlobalErrorEstimation requires a positive `gtol` constructor keyword") ) _validate_tolerances(abstol, reltol, "local") - haskey(kwargs, :callback) && + DiffEqBase.has_callbacks(kwargs) && throw(ArgumentError("GlobalErrorEstimation does not currently support callbacks")) estimator = (local_abstol, local_reltol) -> global_error_estimate( prob, alg, args...; diff --git a/lib/OrdinaryDiffEqCore/AGENTS.md b/lib/OrdinaryDiffEqCore/AGENTS.md new file mode 100644 index 00000000000..b7a19443cd5 --- /dev/null +++ b/lib/OrdinaryDiffEqCore/AGENTS.md @@ -0,0 +1,6 @@ +# Solver compatibility + +- Keep Reactant support in the ordinary solver, controller, and initial-step paths. Use shared traceable control flow; reserve backend checks for tracing boundaries and host-only diagnostics. +- Keep temporary estimates local to traced branches with `let` scopes; avoid exposing logging-macro temporaries as branch outputs. +- Implement dependency-owned methods in the owning package. In particular, specialization policy for SciMLBase function types belongs in SciMLBase. +- `get_fsalfirstlast` initializes storage for composite/default caches; their active FSAL buffers are held by the integrator after algorithm selection. Preserve this distinction when updating derivatives. diff --git a/lib/OrdinaryDiffEqCore/Project.toml b/lib/OrdinaryDiffEqCore/Project.toml index d0e0a69f7e4..0df9897d02f 100644 --- a/lib/OrdinaryDiffEqCore/Project.toml +++ b/lib/OrdinaryDiffEqCore/Project.toml @@ -30,6 +30,7 @@ PrecompileTools = "aea7be01-6a6a-4083-8856-8a6e6704d82a" Preferences = "21216c6a-2e73-6563-6e65-726566657250" Printf = "de0858da-6303-5e67-8744-51eddeeeb8d7" Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c" +ReactantCore = "a3311ec8-5e00-46d5-b541-4f83e724a433" RecursiveArrayTools = "731186ca-8d62-57ce-b412-fbd966d074cd" Reexport = "189a3867-3050-52da-a836-e630ba90ab69" SciMLBase = "0bca4576-84f4-4d90-8ffe-ffa030f20462" @@ -87,10 +88,11 @@ PrecompileTools = "1.2.1, 1.3" Preferences = "1.5.0" Printf = "1.9" Random = "<0.0.1, 1" +ReactantCore = "0.1.21" RecursiveArrayTools = "4.2.0" Reexport = "1.2.2" SafeTestsets = "0.1.0" -SciMLBase = "3.53" +SciMLBase = "3.56" SciMLLogging = "2" SciMLOperators = "1.24.3" SciMLStructures = "1.7" diff --git a/lib/OrdinaryDiffEqCore/src/OrdinaryDiffEqCore.jl b/lib/OrdinaryDiffEqCore/src/OrdinaryDiffEqCore.jl index 837325b377a..66c9f74fe91 100644 --- a/lib/OrdinaryDiffEqCore/src/OrdinaryDiffEqCore.jl +++ b/lib/OrdinaryDiffEqCore/src/OrdinaryDiffEqCore.jl @@ -115,6 +115,7 @@ import SymbolicIndexingInterface: parameter_values using EnumX: @enumx import EnzymeCore +using ReactantCore: ReactantCore """ Predictor @@ -408,6 +409,7 @@ include("disco.jl") include("dense/generic_dense.jl") include("iterator_interface.jl") +include("reactant.jl") include("solve.jl") include("initdt.jl") include("interp_func.jl") diff --git a/lib/OrdinaryDiffEqCore/src/initdt.jl b/lib/OrdinaryDiffEqCore/src/initdt.jl index f8c68d0d0bb..cd5783c071a 100644 --- a/lib/OrdinaryDiffEqCore/src/initdt.jl +++ b/lib/OrdinaryDiffEqCore/src/initdt.jl @@ -6,12 +6,41 @@ # d₂ uses max(|Δf±ΔgMax|)/sk instead of Δf/sk # ============================================================================= -# Coerce `a == b` to a scalar `Bool`. Some array wrappers (notably PyCall -# `PyObject` of arrays — JuliaPy/PyCall.jl#900) return `Vector{Bool}` from `==`, -# which is not valid in a boolean context. See OrdinaryDiffEq.jl#1402. +# PyCall wrappers may return arrays from `==` (https://github.com/JuliaPy/PyCall.jl/issues/900). +# Reduce these arrays while preserving scalar Boolean types. @inline function _bool_equal(a, b) r = a == b - return r isa Bool ? r : all(r) + return r isa Number ? r : all(r) +end + +# Prefer DiffEqBase.NAN_CHECK when a method exists (Number / AbstractArray / +# ArrayPartition / …). Custom non-array states without a method keep master's +# reduction-only path. (NAN_CHECK on Vector{<:Dual} only sees NaN values, not +# NaN Dual partials — that is DiffEqBase's behavior.) +@inline function _ode_nan_check(x) + return applicable(DiffEqBase.NAN_CHECK, x) ? DiffEqBase.NAN_CHECK(x) : false +end + +@muladd function _initdt_euler_step!(u₁, u0, dt, f₀) + if u0 isa Array + @inbounds @simd ivdep for i in eachindex(u0) + u₁[i] = u0[i] + dt * f₀[i] + end + else + @.. broadcast = false u₁ = u0 + dt * f₀ + end + return u₁ +end + +@muladd function _initdt_scaled_diff!(tmp, u0, f₁, f₀, sk, oneunit_tType) + if u0 isa Array + @inbounds @simd ivdep for i in eachindex(u0) + tmp[i] = (f₁[i] - f₀[i]) / sk[i] * oneunit_tType + end + else + @.. broadcast = false tmp = (f₁ - f₀) / sk * oneunit_tType + end + return tmp end @muladd function _ode_initdt_iip( @@ -24,7 +53,7 @@ end oneunit_tType = oneunit(t) dtmax_tdir = tdir * dtmax - dtmin = nextfloat(max(integrator.opts.dtmin, convert(_tType, oneunit_tType * eps(SciMLBase.value(t))))) + dtmin = _initial_dtmin(t, integrator.opts.dtmin) smalldt = max(dtmin, convert(_tType, oneunit_tType * 1 // 10^(6))) if integrator.isdae @@ -143,10 +172,11 @@ end =# ftmp = nothing - if !_is_identity_massmatrix(prob.f.mass_matrix) && ( - !(prob.f isa DynamicalODEFunction) || - any(!_is_identity_massmatrix, prob.f.mass_matrix) - ) + has_mass_matrix = !_is_identity_massmatrix(prob.f.mass_matrix) && ( + !(prob.f isa DynamicalODEFunction) || + any(!_is_identity_massmatrix, prob.f.mass_matrix) + ) + if has_mass_matrix ftmp = zero(f₀) try integrator.alg.linsolve(ftmp, copy(prob.f.mass_matrix), f₀, true) @@ -186,115 +216,118 @@ end d₁ = internalnorm(tmp, t) end - # Better than checking any(x->any(isnan, x), f₀) - # because it also checks if partials are NaN + # Prefer DiffEqBase.NAN_CHECK when a method exists; otherwise keep master's + # reduction-only path via isnan(d₁) so custom non-array states are not forced + # to add a new method. Complements the fast-math norm which can hide NaNs + # from a subsequent scalar isnan check: # https://discourse.julialang.org/t/incorporating-forcing-functions-in-the-ode-model/70133/26 - if isnan(d₁) + has_nan = _ode_nan_check(f₀) | isnan(d₁) + warn_initial_dt = !ReactantCore.within_compile() + if warn_initial_dt && has_nan @SciMLMessage( "First function call produced NaNs. Exiting. Double check that none of the initial conditions, parameters, or timespan values are NaN.", integrator.opts.verbose, :init_NaN ) - return tdir * dtmin end - - dt₀ = ifelse( - (d₀ < 1 // 10^(5)) | - (d₁ < 1 // 10^(5)), smalldt, - convert( - _tType, - oneunit_tType * SciMLBase.value( - (d₀ / d₁) / - 100 - ) - ) - ) - # if d₀ < 1//10^(5) || d₁ < 1//10^(5) - # dt₀ = smalldt - # else - # dt₀ = convert(_tType,oneunit_tType*(d₀/d₁)/100) - # end - dt₀ = min(dt₀, dtmax_tdir) - - if typeof(one(_tType)) <: AbstractFloat && dt₀ < 10eps(_tType) * oneunit(_tType) - # This catches Andreas' non-singular example - # should act like it's singular - result_dt = tdir * max(smalldt, dtmin) - @SciMLMessage( - lazy"Initial timestep too small (near machine epsilon), using default: dt = $(result_dt)", - integrator.opts.verbose, :dt_epsilon - ) - return result_dt - end - - dt₀_tdir = tdir * dt₀ - - u₁ = zero(u0) # required by DEDataArray - - if u0 isa Array - @inbounds @simd ivdep for i in eachindex(u0) - u₁[i] = u0[i] + dt₀_tdir * f₀[i] - end + ReactantCore.@trace track_numbers = false if has_nan + result_dt = tdir * dtmin else - @.. broadcast = false u₁ = u0 + dt₀_tdir * f₀ - end - f₁ = zero(f₀) - f(f₁, u₁, p, t + dt₀_tdir) - - if !_is_identity_massmatrix(prob.f.mass_matrix) && ( - !(prob.f isa DynamicalODEFunction) || - any(!_is_identity_massmatrix, prob.f.mass_matrix) + dt₀ = ifelse( + (d₀ < 1 // 10^(5)) | + (d₁ < 1 // 10^(5)), smalldt, + convert( + _tType, + oneunit_tType * SciMLBase.value( + (d₀ / d₁) / + 100 + ) + ) ) - integrator.alg.linsolve(ftmp, prob.f.mass_matrix, f₁, false) - copyto!(f₁, ftmp) - end - - # Constant zone before callback - # Just return first guess - # Avoids AD issues. - # `==` is not guaranteed to return `Bool` (e.g. PyCall `PyObject` arrays - # return `Vector{Bool}` — JuliaPy/PyCall.jl#900 / OrdinaryDiffEq.jl#1402). - # Coerce array-valued equality so the boolean context always receives a Bool. - length(u0) > 0 && _bool_equal(f₀, f₁) && return tdir * max(dtmin, 100dt₀) + # if d₀ < 1//10^(5) || d₁ < 1//10^(5) + # dt₀ = smalldt + # else + # dt₀ = convert(_tType,oneunit_tType*(d₀/d₁)/100) + # end + dt₀ = min(dt₀, dtmax_tdir) + + if (eltype(prob.tspan) <: AbstractFloat) && dt₀ < 10eps(eltype(prob.tspan)) * oneunit_tType + result_dt = tdir * max(smalldt, dtmin) - # d₂: fold in diffusion terms when g !== nothing - if g !== nothing - if noise_prototype !== nothing - g₁ = zero(noise_prototype) else - g₁ = zero(u0) - end - g(g₁, u₁, p, t + dt₀_tdir) - g₁ .*= 3 - ΔgMax = max.(internalnorm.(g₀ .- g₁, t), internalnorm.(g₀ .+ g₁, t)) - d₂ = internalnorm( - max.(internalnorm.(f₁ .- f₀ .+ ΔgMax, t), internalnorm.(f₁ .- f₀ .- ΔgMax, t)) ./ sk, - t - ) / dt₀ - else - if u0 isa Array - @inbounds @simd ivdep for i in eachindex(u0) - tmp[i] = (f₁[i] - f₀[i]) / sk[i] * oneunit_tType + result_dt = let result_dt, tmp = tmp + dt₀_tdir = tdir * dt₀ + + u₁ = zero(u0) # required by DEDataArray + + u₁ = _initdt_euler_step!(u₁, u0, dt₀_tdir, f₀) + f₁ = zero(f₀) + f(f₁, u₁, p, t + dt₀_tdir) + + if has_mass_matrix + integrator.alg.linsolve(ftmp, prob.f.mass_matrix, f₁, false) + copyto!(f₁, ftmp) + end + + # Constant zone before callback + # Just return first guess + # Avoids AD issues. + # `==` is not guaranteed to return `Bool` (e.g. PyCall `PyObject` arrays + # return `Vector{Bool}` — JuliaPy/PyCall.jl#900 / OrdinaryDiffEq.jl#1402). + # Coerce array-valued equality so the boolean context always receives a Bool. + if _bool_equal(f₀, f₁) + result_dt = tdir * max(dtmin, 100dt₀) + else + result_dt = let + # d₂: fold in diffusion terms when g !== nothing + if g !== nothing + if noise_prototype !== nothing + g₁ = zero(noise_prototype) + else + g₁ = zero(u0) + end + g(g₁, u₁, p, t + dt₀_tdir) + g₁ .*= 3 + ΔgMax = max.(internalnorm.(g₀ .- g₁, t), internalnorm.(g₀ .+ g₁, t)) + d₂ = internalnorm( + max.(internalnorm.(f₁ .- f₀ .+ ΔgMax, t), internalnorm.(f₁ .- f₀ .- ΔgMax, t)) ./ sk, + t + ) / dt₀ + else + tmp = _initdt_scaled_diff!(tmp, u0, f₁, f₀, sk, oneunit_tType) + d₂ = internalnorm(tmp, t) / dt₀ * oneunit_tType + end + # Hairer has d₂ = sqrt(sum(abs2,tmp))/dt₀, note the lack of norm correction + + max_d₁d₂ = max(d₁, d₂) + if max_d₁d₂ <= 1 // Int64(10)^(15) + dt₁ = max(convert(_tType, oneunit_tType * 1 // 10^(6)), dt₀ * 1 // 10^(3)) + else + dt₁ = convert( + _tType, + oneunit_tType * + SciMLBase.value( + 10.0^(-(2 + log10(max_d₁d₂)) / order) + ) + ) + end + tdir * max(dtmin, min(100dt₀, dt₁, dtmax_tdir)) + end + end + result_dt end - else - @.. broadcast = false tmp = (f₁ - f₀) / sk * oneunit_tType end - d₂ = internalnorm(tmp, t) / dt₀ * oneunit_tType - end - # Hairer has d₂ = sqrt(sum(abs2,tmp))/dt₀, note the lack of norm correction - - max_d₁d₂ = max(d₁, d₂) - if max_d₁d₂ <= 1 // Int64(10)^(15) - dt₁ = max(convert(_tType, oneunit_tType * 1 // 10^(6)), dt₀ * 1 // 10^(3)) - else - dt₁ = convert( - _tType, - oneunit_tType * - SciMLBase.value( - 10.0^(-(2 + log10(max_d₁d₂)) / order) + # Keep this warning inside the non-NaN branch: `dt₀` is only assigned there, + # and `ReactantCore.@trace if` confuses JET's definite-assignment analysis. + if warn_initial_dt && + (eltype(prob.tspan) <: AbstractFloat) && + dt₀ < 10eps(eltype(prob.tspan)) * oneunit_tType + @SciMLMessage( + lazy"Initial timestep too small (near machine epsilon), using default: dt = $(result_dt)", + integrator.opts.verbose, :dt_epsilon ) - ) + end end - return tdir * max(dtmin, min(100dt₀, dt₁, dtmax_tdir)) + return result_dt end # ODE iip entry point @@ -353,7 +386,7 @@ end oneunit_tType = oneunit(t) dtmax_tdir = tdir * dtmax - dtmin = nextfloat(max(integrator.opts.dtmin, convert(_tType, oneunit_tType * eps(SciMLBase.value(t))))) + dtmin = _initial_dtmin(t, integrator.opts.dtmin) smalldt = max(dtmin, convert(_tType, oneunit_tType * 1 // 10^(6))) if integrator.isdae @@ -371,17 +404,28 @@ end f₀ = f(u0, p, t) - # Use the overloadable DiffEqBase.NAN_CHECK hook (same intent as the IIP - # isnan(d₁) path) rather than nested any(isnan, ·), which custom array / - # field types cannot sensibly overload (OrdinaryDiffEq #1404). - if DiffEqBase.NAN_CHECK(f₀) + # Same NAN_CHECK-when-applicable policy as the IIP path (OrdinaryDiffEq #1404). + f0_nan = _ode_nan_check(f₀) + warn_initial_dt = !ReactantCore.within_compile() + if warn_initial_dt && f0_nan @SciMLMessage( "First function call produced NaNs. Exiting. Double check that none of the initial conditions, parameters, or timespan values are NaN.", integrator.opts.verbose, :init_NaN ) - return tdir * dtmin end + ReactantCore.@trace track_numbers = false if f0_nan + result_dt = tdir * dtmin + else + result_dt = _ode_initdt_oop_after_f0(prob, u0, t, tdir, sk, f₀, g, order, integrator, dtmin, smalldt, dtmax_tdir, d₀, internalnorm) + end + return result_dt +end +@muladd function _ode_initdt_oop_after_f0(prob, u0, t, tdir, sk, f₀, g, order, integrator, dtmin, smalldt, dtmax_tdir, d₀, internalnorm) + f = prob.f + p = prob.p + _tType = eltype(t) + oneunit_tType = oneunit(t) inferredtype = Base.promote_op(/, typeof(u0), typeof(oneunit(t))) if !(f₀ isa inferredtype) throw(TypeNotConstantError(inferredtype, typeof(f₀))) @@ -392,7 +436,7 @@ end g₀ = nothing if g !== nothing g₀ = 3g(u0, p, t) - if DiffEqBase.NAN_CHECK(g₀) + if _ode_nan_check(g₀) @SciMLMessage( "First function call for g produced NaNs. Exiting.", integrator.opts.verbose, :init_NaN @@ -406,56 +450,65 @@ end d₁ = internalnorm(f₀ ./ sk .* oneunit_tType, t) end - # Also catch NaN AD partials that NAN_CHECK on values may miss (matches IIP). - if isnan(d₁) + # Match IIP: also reject when the norm itself is NaN (fast-math can hide + # elementwise NaNs from NAN_CHECK on some array types). + warn_initial_dt = !ReactantCore.within_compile() + if warn_initial_dt && isnan(d₁) @SciMLMessage( "First function call produced NaNs. Exiting. Double check that none of the initial conditions, parameters, or timespan values are NaN.", integrator.opts.verbose, :init_NaN ) - return tdir * dtmin end - - if d₀ < 1 // 10^(5) || d₁ < 1 // 10^(5) - dt₀ = smalldt - else - dt₀ = convert(_tType, oneunit_tType * SciMLBase.value((d₀ / d₁) / 100)) - end - dt₀ = min(dt₀, dtmax_tdir) - dt₀_tdir = tdir * dt₀ - - u₁ = @.. broadcast = false u0 + dt₀_tdir * f₀ - f₁ = f(u₁, p, t + dt₀_tdir) - - # Constant zone before callback - # Just return first guess - # Avoids AD issues. - # See iip path: coerce array-valued `==` (OrdinaryDiffEq.jl#1402). - _bool_equal(f₀, f₁) && return tdir * max(dtmin, 100dt₀) - - # d₂: fold in diffusion terms when g !== nothing - if g !== nothing - g₁ = 3g(u₁, p, t + dt₀_tdir) - ΔgMax = max.(internalnorm.(g₀ .- g₁, t), internalnorm.(g₀ .+ g₁, t)) - d₂ = internalnorm( - max.(internalnorm.(f₁ .- f₀ .+ ΔgMax, t), internalnorm.(f₁ .- f₀ .- ΔgMax, t)) ./ sk, - t - ) / dt₀ - else - d₂ = internalnorm((f₁ .- f₀) ./ sk .* oneunit_tType, t) / dt₀ * oneunit_tType - end - - max_d₁d₂ = max(d₁, d₂) - if max_d₁d₂ <= 1 // Int64(10)^(15) - dt₁ = max(smalldt, dt₀ * 1 // 10^(3)) + ReactantCore.@trace track_numbers = false if isnan(d₁) + result_dt = tdir * dtmin else - dt₁ = _tType( - oneunit_tType * - SciMLBase.value( - 10^(-(2 + log10(max_d₁d₂)) / order) - ) - ) + if d₀ < 1 // 10^(5) || d₁ < 1 // 10^(5) + dt₀ = smalldt + else + dt₀ = convert(_tType, oneunit_tType * SciMLBase.value((d₀ / d₁) / 100)) + end + dt₀ = min(dt₀, dtmax_tdir) + dt₀_tdir = tdir * dt₀ + + u₁ = @.. broadcast = false u0 + dt₀_tdir * f₀ + f₁ = f(u₁, p, t + dt₀_tdir) + + # Constant zone before callback + # Just return first guess + # Avoids AD issues. + # See iip path: coerce array-valued `==` (OrdinaryDiffEq.jl#1402). + if _bool_equal(f₀, f₁) + result_dt = tdir * max(dtmin, 100dt₀) + else + result_dt = let + # d₂: fold in diffusion terms when g !== nothing + if g !== nothing + g₁ = 3g(u₁, p, t + dt₀_tdir) + ΔgMax = max.(internalnorm.(g₀ .- g₁, t), internalnorm.(g₀ .+ g₁, t)) + d₂ = internalnorm( + max.(internalnorm.(f₁ .- f₀ .+ ΔgMax, t), internalnorm.(f₁ .- f₀ .- ΔgMax, t)) ./ sk, + t + ) / dt₀ + else + d₂ = internalnorm((f₁ .- f₀) ./ sk .* oneunit_tType, t) / dt₀ * oneunit_tType + end + + max_d₁d₂ = max(d₁, d₂) + if max_d₁d₂ <= 1 // Int64(10)^(15) + dt₁ = max(smalldt, dt₀ * 1 // 10^(3)) + else + dt₁ = _tType( + oneunit_tType * + SciMLBase.value( + 10^(-(2 + log10(max_d₁d₂)) / order) + ) + ) + end + tdir * max(dtmin, min(100dt₀, dt₁, dtmax_tdir)) + end + end end - return tdir * max(dtmin, min(100dt₀, dt₁, dtmax_tdir)) + return result_dt end # ODE oop entry point @@ -523,3 +576,8 @@ function ode_determine_initdt( prob, g, effective_order, integrator ) end + +function _initial_dtmin(t, dtmin) + T = eltype(t) + return nextfloat(max(dtmin, convert(T, oneunit(t) * eps(SciMLBase.value(t))))) +end diff --git a/lib/OrdinaryDiffEqCore/src/integrators/controllers.jl b/lib/OrdinaryDiffEqCore/src/integrators/controllers.jl index 68436f37b72..c4572fb2074 100644 --- a/lib/OrdinaryDiffEqCore/src/integrators/controllers.jl +++ b/lib/OrdinaryDiffEqCore/src/integrators/controllers.jl @@ -286,10 +286,7 @@ on normal steps but 10^4 on the first step. See also: https://github.com/SciML/DifferentialEquations.jl/issues/299 """ @inline function get_current_qmax(integrator, qmax) - if integrator.success_iter == 0 - return get_qmax_first_step(integrator) - end - return qmax + return ifelse(iszero(integrator.success_iter), get_qmax_first_step(integrator), qmax) end """ @@ -657,7 +654,7 @@ mutable struct IControllerCache{T, E, NLPType} <: AbstractControllerCache end function setup_controller_cache(alg, cache, controller::IController, ::Type{E}, disco_probs) where {E} - QT = _resolved_QT(controller.basic) + QT = typeof(_maybe_traced(zero(_resolved_QT(controller.basic)))) resolved = IController(resolve_basic(controller.basic, alg, QT; disco_probs)) T = QT return IControllerCache{T, E, eltype(disco_probs)}(resolved, T(1 // 10^4), oneunit(E)) @@ -668,10 +665,10 @@ end qmax = get_current_qmax(integrator, qmax) EEst = SciMLBase.value(get_EEst(integrator)) - if iszero(EEst) + expo = 1 / (get_current_adaptive_order(alg, integrator.cache) + 1) + ReactantCore.@trace track_numbers = false if iszero(EEst) q = inv(qmax) else - expo = 1 / (get_current_adaptive_order(alg, integrator.cache) + 1) qtmp = fastpower(EEst, expo) / gamma @fastmath q = SciMLBase.value(max(inv(qmax), min(inv(qmin), qtmp))) # TODO: Shouldn't this be in `step_accept_controller!` as for the PI controller? @@ -686,9 +683,7 @@ function step_accept_controller!(integrator, cache::IControllerCache, alg, q) t = integrator.t dt = integrator.dt - if qsteady_min <= q <= qsteady_max - q = one(q) - end + q = ifelse((qsteady_min <= q) & (q <= qsteady_max), one(q), q) return handle_disco_accept!(integrator, cache.controller.basic, t, dt / q) end @@ -791,7 +786,7 @@ mutable struct PIControllerCache{T, E, NLPType} <: AbstractControllerCache end function setup_controller_cache(alg, cache, controller::PIController, ::Type{E}, disco_probs) where {E} - QT = _resolved_QT(controller.basic) + QT = typeof(_maybe_traced(zero(_resolved_QT(controller.basic)))) basic = resolve_basic(controller.basic, alg, QT; disco_probs) resolved = PIController{typeof(basic), QT}( basic, QT(controller.beta1), QT(controller.beta2), QT(controller.qoldinit) @@ -809,7 +804,7 @@ end (; beta1, beta2) = controller EEst = SciMLBase.value(get_EEst(integrator)) - if iszero(EEst) + ReactantCore.@trace track_numbers = false if iszero(EEst) q = inv(qmax) else q11 = fastpower(EEst, beta1) @@ -828,9 +823,7 @@ function step_accept_controller!(integrator, cache::PIControllerCache, alg, q) t = integrator.t dt = integrator.dt - if qsteady_min <= q <= qsteady_max - q = one(q) - end + q = ifelse((qsteady_min <= q) & (q <= qsteady_max), one(q), q) cache.errold = max(EEst, qoldinit) return handle_disco_accept!(integrator, controller.basic, t, dt / q) end @@ -986,24 +979,24 @@ the resolved `controller`, its limiter, the error history, and the scalar `EEst` """ mutable struct PIDControllerCache{T, Limiter, E, NLPType} <: AbstractControllerCache controller::PIDController{CommonControllerOptions{T, NLPType}, T, Limiter} - err::Vector{T} # history of the error estimates + err::NTuple{3, T} # history of the error estimates dt_factor::T EEst::E end function reinit_controller!(integrator::SciMLBase.DEIntegrator, cache::PIDControllerCache{T}) where {T} - cache.err = ones(T, 3) + cache.err = (one(T), one(T), one(T)) cache.dt_factor = one(T) return nothing end function setup_controller_cache(alg, cache, controller::PIDController, ::Type{E}, disco_probs) where {E} - QT = _resolved_QT(controller.basic) + QT = typeof(_maybe_traced(zero(_resolved_QT(controller.basic)))) basic = resolve_basic(controller.basic, alg, QT; disco_probs) resolved = PIDController{typeof(basic), QT, typeof(controller.limiter)}( basic, map(QT, controller.beta), QT(controller.accept_safety), controller.limiter, ) - err = ones(QT, 3) + err = (one(QT), one(QT), one(QT)) return PIDControllerCache{QT, typeof(controller.limiter), E, eltype(disco_probs)}( resolved, err, one(QT), oneunit(E), ) @@ -1033,12 +1026,12 @@ end # ``` EEst = max(EEst, EEst_min) - cache.err[1] = inv(EEst) + cache.err = (inv(EEst), cache.err[2], cache.err[3]) err1, err2, err3 = cache.err k = min(alg_order(alg), alg_adaptive_order(alg)) + 1 dt_factor = err1^(beta1 / k) * err2^(beta2 / k) * err3^(beta3 / k) - if isnan(dt_factor) + if !ReactantCore.within_compile() && isnan(dt_factor) @warn "unlimited dt_factor" dt_factor err1 err2 err3 beta1 beta2 beta3 k end cache.dt_factor = controller.limiter(dt_factor) @@ -1062,13 +1055,11 @@ function step_accept_controller!(integrator, cache::PIDControllerCache, alg, dt_ t = integrator.t dt = integrator.dt - if qsteady_min <= inv(dt_factor) <= qsteady_max - dt_factor = one(dt_factor) - end - @inbounds begin - cache.err[3] = cache.err[2] - cache.err[2] = cache.err[1] - end + dt_factor = ifelse( + (qsteady_min <= inv(dt_factor)) & (inv(dt_factor) <= qsteady_max), + one(dt_factor), dt_factor + ) + cache.err = (cache.err[1], cache.err[1], cache.err[2]) return handle_disco_accept!(integrator, controller.basic, t, dt * dt_factor) end diff --git a/lib/OrdinaryDiffEqCore/src/integrators/integrator_utils.jl b/lib/OrdinaryDiffEqCore/src/integrators/integrator_utils.jl index 23adbb3cbfa..e404352ef6c 100644 --- a/lib/OrdinaryDiffEqCore/src/integrators/integrator_utils.jl +++ b/lib/OrdinaryDiffEqCore/src/integrators/integrator_utils.jl @@ -94,27 +94,8 @@ function loopheader!(integrator) end # Accept or reject the step - if integrator.iter > 0 - if (!integrator.force_stepfail) && - ( - !integrator.opts.adaptive || integrator.accept_step || - isaposteriori(integrator.alg) - ) - # ACCEPT - @SciMLMessage( - lazy"Step accepted: t = $(integrator.t), dt = $(integrator.dt), EEst = $(get_EEst(integrator))", - integrator.opts.verbose, :step_accepted - ) - integrator.success_iter += 1 - apply_step!(integrator) - elseif ( - integrator.opts.adaptive && !integrator.accept_step && - !isaposteriori(integrator.alg) - ) || - integrator.force_stepfail - # REJECT - handle_step_rejection!(integrator) - end + ReactantCore.@trace track_numbers = false if integrator.iter > 0 + _accept_or_reject_step!(integrator) end integrator.iter += 1 @@ -126,6 +107,30 @@ function loopheader!(integrator) return nothing end +function _accept_or_reject_step!(integrator) + ReactantCore.@trace track_numbers = false if (!integrator.force_stepfail) && + ( + !integrator.opts.adaptive || integrator.accept_step || + isaposteriori(integrator.alg) + ) + # ACCEPT + @SciMLMessage( + lazy"Step accepted: t = $(integrator.t), dt = $(integrator.dt), EEst = $(get_EEst(integrator))", + integrator.opts.verbose, :step_accepted + ) + integrator.success_iter += 1 + apply_step!(integrator) + elseif ( + integrator.opts.adaptive && !integrator.accept_step && + !isaposteriori(integrator.alg) + ) || + integrator.force_stepfail + # REJECT + handle_step_rejection!(integrator) + end + return nothing +end + # Handles step rejection in loopheader: adjust dt, reject noise, and call post_step_reject!. function handle_step_rejection!(integrator) @SciMLMessage( @@ -215,6 +220,41 @@ end return nothing end +# Host path: always use the integrator's FSAL buffers (what `alg_cache` / +# init wired). Calling `get_fsalfirstlast` here is wrong for caches whose +# accessor allocates throwaways each call (ExpRK: `zero(cache.rtmp)`). +# +# Under Reactant compile, `_dealias_traced!` may break integrator↔cache aliasing +# so `perform_step!` mutates cache fields while `integrator.fsalfirst` is a +# detached copy. Resolve via `get_fsalfirstlast` once and accept those buffers +# only when *both* are already owned by the cache (some caches return an owned +# first buffer and a fresh last one); otherwise keep the integrator buffers. +@inline function _cache_owns_buffer(cache, buf) + @inbounds for i in 1:nfields(cache) + getfield(cache, i) === buf && return true + end + return false +end + +function _fsal_copy_buffers(integrator) + fsalfirst = integrator.fsalfirst + fsallast = integrator.fsallast + if !ReactantCore.within_compile() + return fsalfirst, fsallast + end + cache = integrator.cache + cache isa Union{CompositeCache, DefaultCache, OrdinaryDiffEqConstantCache} && + return fsalfirst, fsallast + c_first, c_last = get_fsalfirstlast(cache, integrator.u) + if c_first === fsalfirst || c_last === fsallast + return fsalfirst, fsallast + end + if _cache_owns_buffer(cache, c_first) && _cache_owns_buffer(cache, c_last) + return c_first, c_last + end + return fsalfirst, fsallast +end + function update_fsal!(integrator) if has_discontinuity(integrator) && first_discontinuity(integrator) == integrator.tdir * integrator.t @@ -232,7 +272,8 @@ function update_fsal!(integrator) reset_fsal!(integrator) else # Do not reeval_fsal, instead copyto! over if isinplace(integrator.sol.prob) - recursivecopy!(integrator.fsalfirst, integrator.fsallast) + fsalfirst, fsallast = _fsal_copy_buffers(integrator) + recursivecopy!(fsalfirst, fsallast) else integrator.fsalfirst = integrator.fsallast end @@ -257,14 +298,14 @@ end _get_next_step_tstop(integrator::ODEIntegrator) = integrator.next_step_tstop _get_next_step_tstop(integrator) = false -function _set_tstop_flag!(integrator::ODEIntegrator, is_tstop::Bool, target = nothing) +function _set_tstop_flag!(integrator::ODEIntegrator, is_tstop, target = nothing) integrator.next_step_tstop = is_tstop - if is_tstop && target !== nothing - integrator.tstop_target = target + if target !== nothing + integrator.tstop_target = ifelse(is_tstop, target, integrator.tstop_target) end return nothing end -_set_tstop_flag!(integrator, is_tstop::Bool, target = nothing) = nothing +_set_tstop_flag!(integrator, is_tstop, target = nothing) = nothing _get_tstop_target(integrator::ODEIntegrator) = integrator.tstop_target @@ -277,46 +318,42 @@ function modify_dt_for_tstops!(integrator) # distance_to_tstop to within rounding still triggers the tstop # branch. Without this, accumulated `t + dt + dt + …` can drift # just past the last tstop and produce a spurious micro-step. - tstop_tol = if integrator.t isa AbstractFloat && isfinite(tdir_tstop) && - isfinite(integrator.t) - 100 * eps( - float( - max(abs(integrator.t), abs(tdir_tstop)) / - oneunit(integrator.t) - ) - ) * oneunit(integrator.t) - else - zero(distance_to_tstop) + # Match DelayDiffEq's sync check and pre-Reactant Core: `100 * eps(mag)`, + # not `100 * eps(T) * mag`. The latter is larger near typical times and + # lets the tstop snap exceed DelayDiffEq's tolerance ("unexpected time + # discrepancy"). Use `eps` of the magnitude so Unitful times keep the + # oneunit scaling of the previous host-only formula. + tstop_tol = zero(distance_to_tstop) + if eltype(integrator.sol.prob.tspan) <: AbstractFloat + ReactantCore.@trace track_numbers = false if isfinite(tdir_tstop) & isfinite(integrator.t) + t_mag = max(abs(integrator.t), abs(tdir_tstop)) + tstop_tol = 100 * eps(float(t_mag / oneunit(integrator.t))) * + oneunit(integrator.t) + end end if integrator.opts.adaptive original_dt = abs(integrator.dt) integrator.dtpropose = integrator.tdir * original_dt - if original_dt + tstop_tol < distance_to_tstop - _set_tstop_flag!(integrator, false) - else - _set_tstop_flag!( - integrator, true, integrator.tdir * tdir_tstop - ) - end - integrator.dt = integrator.tdir * min(original_dt, distance_to_tstop) - elseif iszero(integrator.dtcache) && integrator.dtchangeable - integrator.dt = integrator.tdir * distance_to_tstop _set_tstop_flag!( - integrator, true, integrator.tdir * tdir_tstop + integrator, !(original_dt + tstop_tol < distance_to_tstop), + integrator.tdir * tdir_tstop ) - elseif integrator.dtchangeable && !integrator.force_stepfail - # always try to step! with dtcache, but lower if a tstop - # however, if force_stepfail then don't set to dtcache, and no tstop worry - if abs(integrator.dtcache) + tstop_tol < distance_to_tstop - _set_tstop_flag!(integrator, false) - else - _set_tstop_flag!( - integrator, true, integrator.tdir * tdir_tstop + integrator.dt = integrator.tdir * min(original_dt, distance_to_tstop) + elseif integrator.dtchangeable + zero_dt = iszero(integrator.dtcache) + if integrator.force_stepfail + integrator.dt = ifelse( + zero_dt, integrator.tdir * distance_to_tstop, integrator.dt ) + is_tstop = zero_dt + else + original_dt = ifelse(zero_dt, distance_to_tstop, abs(integrator.dtcache)) + integrator.dt = integrator.tdir * min(original_dt, distance_to_tstop) + is_tstop = !(original_dt + tstop_tol < distance_to_tstop) end - integrator.dt = integrator.tdir * - min(abs(integrator.dtcache), distance_to_tstop) + integrator.dtpropose = ifelse(zero_dt, integrator.dt, integrator.dtpropose) + _set_tstop_flag!(integrator, is_tstop, integrator.tdir * tdir_tstop) else _set_tstop_flag!(integrator, false) end @@ -612,40 +649,16 @@ function _loopfooter!(integrator) elseif integrator.opts.adaptive q = stepsize_controller!(integrator, integrator.alg) integrator.isout = integrator.opts.isoutofdomain(integrator.u, integrator.p, ttmp) - integrator.accept_step = ( - !integrator.isout && - accept_step_controller( - integrator, - integrator.alg - ) - ) || - ( - integrator.opts.force_dtmin && - abs(integrator.dt) <= timedepentdtmin(integrator) - ) - if integrator.accept_step # Accept - increment_accept!(integrator.stats) - apply_solve_step_limiter!(integrator, ttmp) - integrator.last_stepfail = false - integrator.tprev = integrator.t - - if _get_next_step_tstop(integrator) - # Step controller dt is overly pessimistic, since dt = time to tstop. - # Restore the original dt so the controller proposes a reasonable next step. - integrator.dt = integrator.dtpropose - end - integrator.t = fixed_t_for_tstop_error!(integrator, ttmp) - - dtnew = SciMLBase.value( - step_accept_controller!( - integrator, - integrator.alg, - q - ) - ) * - oneunit(integrator.dt) - calc_dt_propose!(integrator, dtnew) - handle_callbacks!(integrator) + ReactantCore.@trace track_numbers = false if !integrator.isout + integrator.accept_step = accept_step_controller(integrator, integrator.alg) + else + integrator.accept_step = false + end + if integrator.opts.force_dtmin + integrator.accept_step = integrator.accept_step | (abs(integrator.dt) <= timedepentdtmin(integrator)) + end + ReactantCore.@trace track_numbers = false if integrator.accept_step # Accept + _accept_step!(integrator, q, ttmp) else # Reject increment_reject!(integrator.stats) end @@ -679,6 +692,32 @@ function _loopfooter!(integrator) return nothing end +function _accept_step!(integrator, q, ttmp) + increment_accept!(integrator.stats) + apply_solve_step_limiter!(integrator, ttmp) + integrator.last_stepfail = false + integrator.tprev = integrator.t + + ReactantCore.@trace track_numbers = false if _get_next_step_tstop(integrator) + # Step controller dt is overly pessimistic, since dt = time to tstop. + # Restore the original dt so the controller proposes a reasonable next step. + integrator.dt = integrator.dtpropose + end + integrator.t = fixed_t_for_tstop_error!(integrator, ttmp) + + dtnew = SciMLBase.value( + step_accept_controller!( + integrator, + integrator.alg, + q + ) + ) * + oneunit(integrator.dt) + calc_dt_propose!(integrator, dtnew) + handle_callbacks!(integrator) + return nothing +end + # Trait: is this a composite algorithm cache? Override to include SDE composite caches. """ is_composite_cache(cache) -> Bool @@ -738,6 +777,9 @@ handle_force_stepfail!(integrator) = post_newton_controller!(integrator, integra Increment the accepted-step counter `stats.naccept` by one. """ function increment_accept!(stats) + # DEStats counters are plain `Int`; Reactant forbids mutating untraced fields + # inside `@trace if` (EnzymeAD/Reactant.jl#3271). Compiled solves strip stats. + ReactantCore.within_compile() && return nothing return stats.naccept += 1 end @@ -747,6 +789,7 @@ end Increment the rejected-step counter `stats.nreject` by one. """ function increment_reject!(stats) + ReactantCore.within_compile() && return nothing return stats.nreject += 1 end @@ -1028,12 +1071,13 @@ function SciMLBase.log_numerical_instability(integrator::ODEIntegrator; jacobian end function fixed_t_for_tstop_error!(integrator, ttmp) - if _get_next_step_tstop(integrator) + ReactantCore.@trace track_numbers = false if _get_next_step_tstop(integrator) _set_tstop_flag!(integrator, false) - return _get_tstop_target(integrator) + target = _get_tstop_target(integrator) else - return ttmp + target = ttmp end + return target end # Type-stable check: did the callback that fired have maybe_discontinuity = true? diff --git a/lib/OrdinaryDiffEqCore/src/integrators/type.jl b/lib/OrdinaryDiffEqCore/src/integrators/type.jl index a3c63ceb918..6dfdb6a072d 100644 --- a/lib/OrdinaryDiffEqCore/src/integrators/type.jl +++ b/lib/OrdinaryDiffEqCore/src/integrators/type.jl @@ -157,7 +157,7 @@ mutable struct ODEIntegrator{ uType, duType, tType, pType, eigenType, tdirType, ksEltype, SolType, F, CacheType, O, FSALType, EventErrorType, CallbackCacheType, IA, DV, CC, RNGType, WType, PType, SqdtType, - NoiseType, CType, RCType, + NoiseType, CType, RCType, CountType, FlagType, } <: SciMLBase.AbstractODEIntegrator{algType, IIP, uType, tType} sol::SolType @@ -179,23 +179,23 @@ mutable struct ODEIntegrator{ tdir::tdirType eigen_est::eigenType controller_cache::CC - success_iter::Int - iter::Int + success_iter::CountType + iter::CountType saveiter::Int saveiter_dense::Int cache::CacheType callback_cache::CallbackCacheType kshortsize::Int force_stepfail::Bool - last_stepfail::Bool + last_stepfail::FlagType just_hit_tstop::Bool - next_step_tstop::Bool + next_step_tstop::FlagType tstop_target::tType do_error_check::Bool event_last_time::Int vector_event_last_time::Int last_event_error::EventErrorType - accept_step::Bool + accept_step::FlagType isout::Bool reeval_fsal::Bool derivative_discontinuity::Bool diff --git a/lib/OrdinaryDiffEqCore/src/reactant.jl b/lib/OrdinaryDiffEqCore/src/reactant.jl new file mode 100644 index 00000000000..f93d97d6e5c --- /dev/null +++ b/lib/OrdinaryDiffEqCore/src/reactant.jl @@ -0,0 +1,44 @@ +_maybe_traced(x) = ReactantCore.within_compile() ? ReactantCore.promote_to_traced(x) : x + +# Reactant requires every loop-carried path to own its traced value at the loop boundary. +function _dealias_traced!(x) + ReactantCore.within_compile() || return x + x isa Union{Type, Function, Module, AbstractString, Symbol, SciMLBase.AbstractSciMLFunction} && return x + if x isa DenseArray + if eltype(x) <: Number + return x .* one(eltype(x)) + end + y = copy(x) + for i in eachindex(x) + y[i] = _dealias_traced!(x[i]) + end + return y + end + if x isa Number + return x * one(x) + end + x isa Union{Tuple, NamedTuple} && return map(_dealias_traced!, x) + T = typeof(x) + isbitstype(T) && return x + if ismutable(x) + for name in fieldnames(T) + (isdefined(x, name) && !isconst(T, name)) || continue + setfield!(x, name, _dealias_traced!(getfield(x, name))) + end + return x + end + names = fieldnames(T) + isempty(names) && return x + values = map(name -> _dealias_traced!(getfield(x, name)), names) + return ConstructionBase.setproperties(x, NamedTuple{names}(values)) +end + +function _traced_finalize_solution(integrator, retcode) + return ConstructionBase.setproperties( + integrator.sol, + (; + u = [integrator.u], t = [integrator.t], k = nothing, prob = nothing, + interp = nothing, dense = false, stats = nothing, retcode, + ) + ) +end diff --git a/lib/OrdinaryDiffEqCore/src/solve.jl b/lib/OrdinaryDiffEqCore/src/solve.jl index 854030bf5d1..13c0a7c368c 100644 --- a/lib/OrdinaryDiffEqCore/src/solve.jl +++ b/lib/OrdinaryDiffEqCore/src/solve.jl @@ -7,8 +7,7 @@ Base.@constprop :aggressive function SciMLBase.__solve( kwargs... ) integrator = SciMLBase.__init(prob, alg, args...; kwargs...) - solve!(integrator) - return integrator.sol + return ReactantCore.within_compile() ? solve!(integrator) : (solve!(integrator); integrator.sol) end determine_controller_datatype(u::AbstractVector{<:Number}, internalnorm, ts::Tuple{<:Number, <:Number}) = promote_type(typeof(SciMLBase.value(internalnorm(u, ts[1]))), typeof(SciMLBase.value(internalnorm(u, ts[2]))), eltype(SciMLBase.value.(ts))) @@ -347,6 +346,31 @@ Base.@constprop :aggressive function _ode_init( seed = UInt64(0), kwargs... ) + if ReactantCore.within_compile() + prob isa SciMLBase.AbstractODEProblem || + throw(ArgumentError("only ODEProblem is supported inside Reactant compilation")) + isimplicit(alg) && + throw(ArgumentError("implicit algorithms are not supported inside Reactant compilation")) + !adaptive && isnothing(dt) && + throw(ArgumentError("dt is required for fixed-step solves inside Reactant compilation")) + isempty(saveat) || throw(ArgumentError("saveat is not supported inside Reactant compilation")) + isempty(tstops) || throw(ArgumentError("tstops are not supported inside Reactant compilation")) + isempty(d_discontinuities) || throw(ArgumentError("d_discontinuities are not supported inside Reactant compilation")) + isnothing(callback) || throw(ArgumentError("callbacks are not supported inside Reactant compilation")) + isnothing(save_idxs) || throw(ArgumentError("save_idxs is not supported inside Reactant compilation")) + isoutofdomain === ODE_DEFAULT_ISOUTOFDOMAIN || + throw(ArgumentError("isoutofdomain is not supported inside Reactant compilation")) + unstable_check === ODE_DEFAULT_UNSTABLE_CHECK || + throw(ArgumentError("unstable_check is not supported inside Reactant compilation")) + force_dtmin && throw(ArgumentError("force_dtmin is not supported inside Reactant compilation")) + progress && throw(ArgumentError("progress is not supported inside Reactant compilation")) + save_on = false + save_everystep = false + save_start = false + save_end = true + dense = false + calck = false + end # ODE/DAE-specific validation (skip for RODE/SDE problems) if !(prob isa SciMLBase.AbstractRODEProblem) if prob isa SciMLBase.AbstractDAEProblem && alg isa OrdinaryDiffEqAlgorithm @@ -385,6 +409,9 @@ Base.@constprop :aggressive function _ode_init( stage_limiter!, step_limiter! = resolve_stage_step_limiters( alg, stage_limiter, step_limiter, verbose_spec ) + if ReactantCore.within_compile() && step_limiter! !== trivial_limiter! + throw(ArgumentError("step_limiter is not supported inside Reactant compilation")) + end if alg isa OrdinaryDiffEqRosenbrockAdaptiveAlgorithm && # https://github.com/SciML/OrdinaryDiffEq.jl/pull/2079 fixes this for Rosenbrock23 and 32 @@ -841,21 +868,22 @@ Base.@constprop :aggressive function _ode_init( # rate/state = (state/time)/state = 1/t units, internalnorm drops units # we don't want to differentiate through eigenvalue estimation eigen_est = inv(one(tType)) + t = _maybe_traced(t) tprev = t - dtcache = tType(_dt) - dtpropose = tType(_dt) - iter = 0 + dtcache = _maybe_traced(tType(_dt)) + dtpropose = _maybe_traced(tType(_dt)) + iter = _maybe_traced(0) kshortsize = 0 reeval_fsal = false derivative_discontinuity = false EEst = oneunit(EEstT) # https://github.com/JuliaPhysics/Measurements.jl/pull/135 just_hit_tstop = false - next_step_tstop = false + next_step_tstop = _maybe_traced(false) tstop_target = zero(t) isout = false - accept_step = false + accept_step = _maybe_traced(false) force_stepfail = false - last_stepfail = false + last_stepfail = _maybe_traced(false) do_error_check = true event_last_time = 0 vector_event_last_time = 1 @@ -865,7 +893,7 @@ Base.@constprop :aggressive function _ode_init( 0.0 ) dtchangeable = isdtchangeable(_alg) - success_iter = 0 + success_iter = _maybe_traced(0) reinitialize = true saveiter = 0 # Starts at 0 so first save is at 1 saveiter_dense = 0 @@ -923,7 +951,7 @@ Base.@constprop :aggressive function _ode_init( integrator = ODEIntegrator{ typeof(_alg), isinplace(prob), uType, typeof(du), - tType, typeof(p), typeof(eigen_est), + typeof(t), typeof(p), typeof(eigen_est), typeof(tdir), typeof(k), SolType, FType, cacheType, typeof(opts), typeof(fsalfirst), @@ -931,9 +959,9 @@ Base.@constprop :aggressive function _ode_init( typeof(initializealg), typeof(differential_vars), typeof(controller_cache), typeof(_rng), typeof(W), typeof(P), typeof(sqdt), - typeof(noise), typeof(c), typeof(rate_constants), + typeof(noise), typeof(c), typeof(rate_constants), typeof(iter), typeof(accept_step), }( - sol, u, du, k, t, tType(_dt), f, p, + sol, u, du, k, t, dtcache, f, p, uprev, uprev2, duprev, tprev, _alg, dtcache, dtchangeable, dtpropose, tdir, eigen_est, @@ -1062,26 +1090,26 @@ end function SciMLBase.solve!(integrator::ODEIntegrator) @inbounds while !isempty(integrator.opts.tstops) first_tstop = first(integrator.opts.tstops) - while integrator.tdir * integrator.t < first_tstop - loopheader!(integrator) - if integrator.do_error_check && check_error!(integrator) != ReturnCode.Success - return integrator.sol - end - - # Use special tstop handling if flag is set, otherwise normal stepping - if integrator.next_step_tstop - handle_tstop_step!(integrator) - else - perform_step!(integrator, integrator.cache) - end - - should_exit = integrator.next_step_tstop - - loopfooter!(integrator) - if isempty(integrator.opts.tstops) || should_exit - break + stop = _maybe_traced(false) + errored = _maybe_traced(false) + maxiters = ReactantCore.within_compile() ? integrator.opts.maxiters : typemax(Int) + _dealias_traced!(integrator) + ReactantCore.@trace track_numbers = false while (integrator.tdir * integrator.t < first_tstop) & !stop & + (integrator.iter < maxiters) + stop, errored = _solve_step!(integrator) + _dealias_traced!(integrator) + end + if ReactantCore.within_compile() + ReactantCore.@trace track_numbers = false if !integrator.accept_step + integrator.u = integrator.uprev end + retcode = ifelse( + integrator.tdir * integrator.t >= first_tstop, + ReturnCode.Success, ReturnCode.MaxIters + ) + return _traced_finalize_solution(integrator, retcode) end + errored && return integrator.sol handle_tstop!(integrator) end postamble!(integrator) @@ -1103,10 +1131,35 @@ function SciMLBase.solve!(integrator::ODEIntegrator) return integrator.sol = SciMLBase.solution_new_retcode(integrator.sol, ReturnCode.Success) end +function _solve_step!(integrator) + loopheader!(integrator) + if !ReactantCore.within_compile() && integrator.do_error_check && + check_error!(integrator) != ReturnCode.Success + return true, true + end + ReactantCore.@trace track_numbers = false if integrator.next_step_tstop + handle_tstop_step!(integrator) + else + perform_step!(integrator, integrator.cache) + end + should_exit = integrator.next_step_tstop + loopfooter!(integrator) + return isempty(integrator.opts.tstops) | (should_exit & integrator.accept_step), false +end + # Helpers function handle_dt!(integrator) - return if iszero(integrator.dt) && integrator.opts.adaptive + # 1-arg form is used by DelayDiffEq; ODE init uses the 2-arg form below. + # During Reactant compilation, skip the traced `iszero(dt)` gate and only apply + # the host-side tdir sign fix (auto-dt reset is handled by the 2-arg ODE path). + if ReactantCore.within_compile() + if integrator.opts.adaptive && integrator.tdir < 0 + integrator.dt = abs(integrator.dt) * integrator.tdir + end + return nothing + end + if iszero(integrator.dt) && integrator.opts.adaptive auto_dt_reset!(integrator) if sign(integrator.dt) != integrator.tdir && !iszero(integrator.dt) && !isnan(integrator.dt) @@ -1118,28 +1171,33 @@ function handle_dt!(integrator) integrator.opts.verbose, :dt_NaN ) end - elseif integrator.opts.adaptive && integrator.dt > zero(integrator.dt) && - integrator.tdir < 0 - integrator.dt *= integrator.tdir # Allow positive dt, but auto-convert + elseif integrator.opts.adaptive && integrator.tdir < 0 + integrator.dt = abs(integrator.dt) * integrator.tdir end + return nothing end function handle_dt!(integrator, dt) - return if isnothing(dt) && iszero(integrator.dt) && integrator.opts.adaptive + # Always run auto_dt_reset! and the tdir sign fix; only skip host diagnostics during + # Reactant compilation (traced comparisons cannot drive ordinary `if`). + if isnothing(dt) && integrator.opts.adaptive && + (ReactantCore.within_compile() || iszero(integrator.dt)) auto_dt_reset!(integrator) - if sign(integrator.dt) != integrator.tdir && !iszero(integrator.dt) && - !isnan(integrator.dt) - error("Automatic dt setting has the wrong sign. Exiting. Please report this error.") - end - if isnan(integrator.dt) - @SciMLMessage( - "Automatic dt set the starting dt as NaN, causing instability. Exiting.", - integrator.opts.verbose, :dt_NaN - ) + if !ReactantCore.within_compile() + if sign(integrator.dt) != integrator.tdir && !iszero(integrator.dt) && + !isnan(integrator.dt) + error("Automatic dt setting has the wrong sign. Exiting. Please report this error.") + end + if isnan(integrator.dt) + @SciMLMessage( + "Automatic dt set the starting dt as NaN, causing instability. Exiting.", + integrator.opts.verbose, :dt_NaN + ) + end end - elseif integrator.opts.adaptive && integrator.dt > zero(integrator.dt) && - integrator.tdir < 0 - integrator.dt *= integrator.tdir # Allow positive dt, but auto-convert + elseif integrator.opts.adaptive && integrator.tdir < 0 + integrator.dt = abs(integrator.dt) * integrator.tdir end + return nothing end """ diff --git a/lib/OrdinaryDiffEqExponentialRK/test/qa/allocation_tests.jl b/lib/OrdinaryDiffEqExponentialRK/test/qa/allocation_tests.jl index b011bea060c..1a27539c117 100644 --- a/lib/OrdinaryDiffEqExponentialRK/test/qa/allocation_tests.jl +++ b/lib/OrdinaryDiffEqExponentialRK/test/qa/allocation_tests.jl @@ -181,3 +181,51 @@ end # essentially flat, so a generous 3x bound cleanly separates the two. @test large < 3 * small end + +# FSAL refresh must not call allocating `get_fsalfirstlast` factories (ExpRK +# returns fresh `zero(cache.rtmp)` each call). Assert `update_fsal!` stays at 0 +# bytes. Under `Pkg.test`'s `--check-bounds=yes`, ETDRK2 `step!` allocates a +# small constant amount, so guard size-independence instead of absolute zero. +@testset "ExpRK FSAL refresh does not allocate on warmed in-place steps" begin + f!(du, u, p, t) = (du .= .-u; nothing) + function warmed_etdrk2(n) + prob = SplitODEProblem( + MatrixOperator(Diagonal(fill(-1.0, n))), + (du, u, p, t) -> fill!(du, 0), + ones(n), + (0.0, 100.0) + ) + integrator = init( + prob, ETDRK2(krylov = true); + dt = 0.01, adaptive = false, save_everystep = false, dense = false + ) + for _ in 1:5 + step!(integrator) + end + return integrator + end + for n in (10, 10000), alg in (LawsonEuler(krylov = true), ETDRK2(krylov = true)) + prob = SplitODEProblem( + MatrixOperator(Diagonal(fill(-1.0, n))), + (du, u, p, t) -> fill!(du, 0), + ones(n), + (0.0, 100.0) + ) + integrator = init( + prob, alg; dt = 0.01, adaptive = false, save_everystep = false, dense = false + ) + for _ in 1:5 + step!(integrator) # warm up + end + @test (@allocated OrdinaryDiffEqCore.update_fsal!(integrator)) == 0 + end + small = begin + i = warmed_etdrk2(10) + @allocated step!(i) + end + large = begin + i = warmed_etdrk2(10000) + @allocated step!(i) + end + @test small == large +end diff --git a/lib/OrdinaryDiffEqLowOrderRK/src/low_order_rk_perform_step.jl b/lib/OrdinaryDiffEqLowOrderRK/src/low_order_rk_perform_step.jl index fe3d3b1face..9881defd7d9 100644 --- a/lib/OrdinaryDiffEqLowOrderRK/src/low_order_rk_perform_step.jl +++ b/lib/OrdinaryDiffEqLowOrderRK/src/low_order_rk_perform_step.jl @@ -1109,7 +1109,7 @@ end integrator.u = u end -get_fsalfirstlast(cache::RKMCache, u) = (cache.k1, zero(cache.k1)) +get_fsalfirstlast(cache::RKMCache, u) = (cache.k1, cache.fsalfirst) function initialize!(integrator, cache::RKMCache) (; k, fsalfirst) = cache integrator.kshortsize = 6 @@ -1151,7 +1151,9 @@ end dt * (β1 * k1 + β2 * k2 + β3 * k3 + β4 * k4 + β6 * k6) stage_limiter!(u, integrator, p, t + dt) - f(integrator.fsallast, u, p, t + dt) + # Write the FSAL into the cache-owned last buffer (`cache.fsalfirst` via + # `get_fsalfirstlast`), not a possibly-dealiased `integrator.fsallast`. + f(fsalfirst, u, p, t + dt) OrdinaryDiffEqCore.increment_nf!(integrator.stats, 6) return nothing end diff --git a/lib/OrdinaryDiffEqSDIRK/src/OrdinaryDiffEqSDIRK.jl b/lib/OrdinaryDiffEqSDIRK/src/OrdinaryDiffEqSDIRK.jl index 5df40e74a22..e32321e8a00 100644 --- a/lib/OrdinaryDiffEqSDIRK/src/OrdinaryDiffEqSDIRK.jl +++ b/lib/OrdinaryDiffEqSDIRK/src/OrdinaryDiffEqSDIRK.jl @@ -33,7 +33,7 @@ using SciMLBase: SciMLBase, SplitFunction, ODEProblem, _vec, _reshape, _unwrap_v # `calculate_residuals`/`calculate_residuals!` are only called. import DiffEqBase: initialize! using DiffEqBase: calculate_residuals, calculate_residuals! -using LinearAlgebra: mul!, diag, I +using LinearAlgebra: mul!, diag, I, UniformScaling import OrdinaryDiffEqCore using OrdinaryDiffEqDifferentiation: dolinsolve diff --git a/lib/OrdinaryDiffEqSDIRK/src/generic_imex_perform_step.jl b/lib/OrdinaryDiffEqSDIRK/src/generic_imex_perform_step.jl index f4c5bb14523..814e107dcbc 100644 --- a/lib/OrdinaryDiffEqSDIRK/src/generic_imex_perform_step.jl +++ b/lib/OrdinaryDiffEqSDIRK/src/generic_imex_perform_step.jl @@ -56,7 +56,15 @@ end @inline _mmmul(z, d) = d * z function _mmdiag(tab, mass_matrix) - return (mass_matrix === I || !tab.explicit_first_stage) ? nothing : diag(mass_matrix) + return (mass_matrix === I || !tab.explicit_first_stage) ? nothing : + _mmdiag_values(mass_matrix) +end +# λ·I without a dense layout: UniformScaling (λ ≠ 1) and ScalarOperator +# (`axes == ()`). `_mmdiv`/`_mmmul` accept a scalar the same way. +_mmdiag_values(mass_matrix::UniformScaling) = mass_matrix.λ +function _mmdiag_values(mass_matrix) + isempty(axes(mass_matrix)) && return convert(Number, mass_matrix) + return diag(mass_matrix) end # =========================================================================== diff --git a/test/InterfaceI/controllers.jl b/test/InterfaceI/controllers.jl index ab6437cbaab..1580601f12d 100644 --- a/test/InterfaceI/controllers.jl +++ b/test/InterfaceI/controllers.jl @@ -106,3 +106,17 @@ end @test integ.controller_cache.errold == errold_before @test integ.controller_cache.q11 == q11_before end + +@testset "PID history on reinit!" begin + integ = init( + ODEProblem((u, p, t) -> -20u, 1.0, (0.0, 1.0)), Tsit5(); + controller = PIDController(0.7, -0.4) + ) + step!(integ) + history = Tuple(integ.controller_cache.err) + @test any(!isone, history) + reinit!(integ; reinit_controller = false) + @test Tuple(integ.controller_cache.err) == history + reinit!(integ) + @test all(isone, integ.controller_cache.err) +end diff --git a/test/InterfaceI/ode_initdt_tests.jl b/test/InterfaceI/ode_initdt_tests.jl index e191af13b19..1eb60c2a442 100644 --- a/test/InterfaceI/ode_initdt_tests.jl +++ b/test/InterfaceI/ode_initdt_tests.jl @@ -151,3 +151,11 @@ end end end end + +@testset "IIP initdt NaN fallback" for T in (Float32, Float64) + f_nan!(du, u, p, t) = (du .= p .* u; nothing) + prob = ODEProblem(f_nan!, ones(T, 2), (zero(T), one(T)), T[NaN]) + dtmin = T(1.0e-5) + integrator = init(prob, Tsit5(); dtmin) + @test integrator.dt == nextfloat(dtmin) +end diff --git a/test/Reactant/reactant_tests.jl b/test/Reactant/reactant_tests.jl new file mode 100644 index 00000000000..9214ae2ac9a --- /dev/null +++ b/test/Reactant/reactant_tests.jl @@ -0,0 +1,90 @@ +using OrdinaryDiffEqCore: IController, PIController, PIDController +using OrdinaryDiffEq +using Reactant +using SciMLBase +using Test + +f(u, p, t) = p .* u + +function f!(du, u, p, t) + du .= p .* u + return nothing +end + +struct CompiledODESolve{F, A, K} + f::F + alg::A + kwargs::K +end + +function (s::CompiledODESolve)(u, p) + prob = ODEProblem(s.f, u, (0.0f0, 1.0f0), p) + return solve(prob, s.alg; s.kwargs...) +end + +u0 = Reactant.to_rarray(Float32[1, 2]) +p0 = Reactant.to_rarray(Float32[-1]) + +@testset "Rejected first step" for rhs in (f, f!) + maxiters_solver = CompiledODESolve( + rhs, Tsit5(), (; maxiters = 1, dt = 1.0f0, abstol = eps(Float32), reltol = eps(Float32)) + ) + maxiters_sol = Reactant.@jit maxiters_solver(u0, p0) + @test maxiters_sol.retcode == ReturnCode.MaxIters + @test !SciMLBase.successful_retcode(maxiters_sol) + @test maxiters_sol.t == Float32[0] + @test Array(maxiters_sol.u[end]) == Float32[1, 2] +end + +implicit_solver = CompiledODESolve(f, Rosenbrock23(), (;)) +@test_throws ArgumentError Reactant.@jit implicit_solver(u0, p0) +fixed_solver_without_dt = CompiledODESolve(f, Tsit5(), (; adaptive = false)) +@test_throws ArgumentError Reactant.@jit fixed_solver_without_dt(u0, p0) + +solver_cases = ( + ("PIDController Tsit5", CompiledODESolve(f, Tsit5(), (; controller = PIDController(0.7, -0.4), abstol = 1.0f-7, reltol = 1.0f-5))), + ("adaptive Tsit5", CompiledODESolve(f, Tsit5(), (;))), + ("fixed Tsit5", CompiledODESolve(f, Tsit5(), (; adaptive = false, dt = 0.1f0))), + ("in-place adaptive Tsit5", CompiledODESolve(f!, Tsit5(), (;))), + ("custom PIController Tsit5", CompiledODESolve(f, Tsit5(), (; controller = PIController(0.14, 0.08)))), + ("IController Tsit5", CompiledODESolve(f, Tsit5(), (; controller = IController(), abstol = 1.0f-7, reltol = 1.0f-5))), + ("in-place fixed Tsit5", CompiledODESolve(f!, Tsit5(), (; adaptive = false, dt = 0.1f0))), + ("adaptive Vern7", CompiledODESolve(f, Vern7(), (;))), +) + +@testset "$name" for (name, solver) in solver_cases + compiled = Reactant.compile(solver, (u0, p0)) + for rate in (0.0f0, -1.0f0, -3.0f0) + sol = compiled( + Reactant.to_rarray(Float32[1, 2]), Reactant.to_rarray(Float32[rate]) + ) + @test sol.u[end] isa Reactant.ConcreteRArray + @test Array(sol.u[end]) ≈ Float32[exp(rate), 2exp(rate)] rtol = 5.0f-4 + @test sol.t == Float32[1] + @test sol.retcode == ReturnCode.Success + @test SciMLBase.successful_retcode(sol) + @test sol.prob === nothing + @test sol.stats === nothing + @test sol.interp === nothing + end +end + +@testset "Fixed-step endpoint clipping" for rhs in (f, f!), direction in (1.0f0, -1.0f0) + function step_count(u, p) + integrator = init( + ODEProblem(rhs, u, (0.0f0, direction), p), Tsit5(); + adaptive = false, dt = direction * 0.01f0, save_everystep = false + ) + solve!(integrator) + return integrator.iter + end + @test Reactant.@jit(step_count(u0, p0)) == step_count(Float32[1, 2], Float32[-1]) +end + +@testset "Initial-step NaN fallback" for rhs in (f, f!) + function first_dt(u, p) + return init(ODEProblem(rhs, u, (0.0f0, 1.0f0), p), Tsit5(); dtmin = 1.0f-5).dt + end + nan_p = Reactant.to_rarray(Float32[NaN]) + @test Reactant.@jit(first_dt(u0, nan_p)) == first_dt(Float32[1, 2], Float32[NaN]) +end diff --git a/test/runtests.jl b/test/runtests.jl index 50f64fe0bfc..5e3e352238e 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -1,4 +1,5 @@ using Pkg + using SafeTestsets, Test using SciMLTesting @@ -203,6 +204,22 @@ function qa_group() return @time @safetestset "Quality Assurance Tests" include("qa/qa_tests.jl") end +function reactant_group() + is_APPVEYOR && return + # Reactant v0.2.289 includes traced enums (EnzymeAD/Reactant.jl#3232). Keep the + # Reactant suite self-contained here rather than listing Reactant in the root + # [extras]/[targets] test environment. + withenv("JULIA_PKG_PRECOMPILE_AUTO" => "0") do + Pkg.add( + [ + PackageSpec(name = "Reactant", version = v"0.2.289"), + PackageSpec(name = "ReactantCore", version = v"0.1.23"), + ] + ) + end + return @time @safetestset "Reactant Tests" include("Reactant/reactant_tests.jl") +end + function activate_gpu_env() Pkg.activate(joinpath(@__DIR__, "gpu")) Pkg.develop(PackageSpec(path = dirname(@__DIR__))) @@ -314,6 +331,7 @@ end "AD" => ad_group, "ODEInterfaceRegression" => odeinterface_group, "GPU" => gpu_group, + "Reactant" => reactant_group, ), # QA runs in the root test environment (no per-group Project.toml); # its body is the ExplicitImports testset, not the standard diff --git a/test/test_groups.toml b/test/test_groups.toml index 10c7f76f34b..099be6492af 100644 --- a/test/test_groups.toml +++ b/test/test_groups.toml @@ -41,3 +41,5 @@ versions = ["lts"] versions = ["1"] runner = ["self-hosted", "Linux", "X64", "gpu"] timeout = 120 +[Reactant] +versions = ["1"]