diff --git a/docs/src/index.md b/docs/src/index.md index bae6ad1..b54932c 100644 --- a/docs/src/index.md +++ b/docs/src/index.md @@ -48,6 +48,36 @@ prob = SteadyStateProblem((u, p, t) -> 1 .- u, [0.0]) sol = solve(prob, SICNM(Rodas3d())) ``` +## Initialization and stepping + +`DynamicSS` and `SICNM` also support the `init`/`solve!` interface. `init(prob, alg; +kwargs...)` accepts the same keyword arguments as `solve` and returns a cache, and +`solve!(cache)` returns the same solution as `solve(prob, alg; kwargs...)`: its return +code comes from the termination condition (for example `ReturnCode.Unstable` from a +safe termination mode, or a failure when the time span ends before steady state is +reached), and `save_idxs` selects components of the final state. + +For a problem that is integrated as a whole, `cache.integrator` is the underlying ODE +integrator, and `step!(cache)` advances it. For `SICNM` that integrator solves the +extended continuous-Newton system, so its state is `[y; z]`: the first +`length(prob.u0)` components are the steady-state variables `y` and the rest are the +Newton direction `z`. Its intermediate states are therefore not states of `prob`; +`solve!` returns only `y`. + +When `prob` carries an `SCCNonlinearProblem` lowering, the blocks are solved one after +another and there is no single integrator to step. `solve!` runs that sequential solve. + +```julia +using SciMLBase: SteadyStateProblem, init, solve!, step! +using SteadyStateDiffEq +using OrdinaryDiffEqRosenbrock: Rodas5P + +prob = SteadyStateProblem((u, p, t) -> 1 .- u, [0.0]) +cache = init(prob, SICNM(Rodas5P())) +step!(cache) +sol = solve!(cache) +``` + ## API ```@docs diff --git a/src/SteadyStateDiffEq.jl b/src/SteadyStateDiffEq.jl index e508efc..2d13eb0 100644 --- a/src/SteadyStateDiffEq.jl +++ b/src/SteadyStateDiffEq.jl @@ -10,7 +10,7 @@ using LinearSolve: LinearSolve using SciMLPublic: @public using SciMLBase: SciMLBase, CallbackSet, LinearProblem, NonlinearProblem, ODEProblem, NonlinearSolution, ReturnCode, SteadyStateProblem, SteadyStateSolution, get_du, init, - isinplace, remake, solve, successful_retcode + isinplace, remake, solve, solve!, step!, successful_retcode using SymbolicIndexingInterface: parameter_values const infnorm = Base.Fix2(norm, Inf) diff --git a/src/solve.jl b/src/solve.jl index acacbee..461be71 100644 --- a/src/solve.jl +++ b/src/solve.jl @@ -98,16 +98,21 @@ function SciMLBase.solve( return SciMLBase.build_solution(prob, alg, u, resid; retcode, original = sols) end -# A `SteadyStateProblem` that records an `SCCNonlinearProblem` lowering is -# solved block-sequentially in the lowering's ordering instead of one -# monolithic integration. Returns `nothing` when there is no SCC lowering. -function __solve_scc_lowering( - prob, alg, args...; save_idxs = nothing, kwargs... - ) +# The `SCCNonlinearProblem` lowering recorded on a `SteadyStateProblem`, or +# `nothing` when there is none. +function __scc_lowering(prob) lp = prob isa SteadyStateProblem ? prob.lowered_problem : nothing lp === nothing && return nothing lp isa SciMLBase.AbstractSciMLProblem || (lp = lp(prob)) - lp isa SciMLBase.SCCNonlinearProblem || return nothing + return lp isa SciMLBase.SCCNonlinearProblem ? lp : nothing +end + +# A `SteadyStateProblem` that records an `SCCNonlinearProblem` lowering `lp` is +# solved block-sequentially in the lowering's ordering instead of one +# monolithic integration. +function __solve_scc_lowering( + prob, lp, alg, args...; save_idxs = nothing, kwargs... + ) sccsol = solve(lp, alg, args...; kwargs...) save_idxs === nothing && return sccsol return SciMLBase.build_solution( @@ -149,6 +154,39 @@ function SciMLBase.solve(prob::SteadyStateProblem, args...; kwargs...) ) end +# `init(prob, ::DynamicSS/SICNM)` returns one of these caches; `solve!` finishes +# it with the same steady-state finalization as `solve`. +@concrete struct SteadyStateSCCCache + prob + lowered_problem + alg + args + kwargs +end + +function SciMLBase.solve!(cache::SteadyStateSCCCache) + return __solve_scc_lowering( + cache.prob, cache.lowered_problem, cache.alg, cache.args...; cache.kwargs... + ) +end + +@concrete struct SteadyStateODECache + prob + alg + integrator + setup + save_idxs +end + +function SciMLBase.solve!(cache::SteadyStateODECache) + odesol = solve!(cache.integrator) + return __steady_state_solution( + cache.prob, cache.alg, cache.setup, odesol, cache.save_idxs + ) +end + +SciMLBase.step!(cache::SteadyStateODECache, args...) = step!(cache.integrator, args...) + __get_tspan(u0, alg::Union{DynamicSS, SICNM}) = __get_tspan(u0, alg.tspan) __get_tspan(u0, tspan::Tuple) = tspan function __get_tspan(u0, tspan::Number) @@ -161,18 +199,13 @@ function __without_verbose(kwargs) return (; (name => value for (name, value) in pairs(kwargs) if name !== :verbose)...) end -function SciMLBase.__solve( - prob::SciMLBase.AbstractSteadyStateProblem, alg::DynamicSS, - args...; abstol = 1.0e-8, reltol = 1.0e-6, odesolve_kwargs = (;), - save_idxs = nothing, termination_condition = NonlinearSolveBase.NormTerminationMode(infnorm), - alias = SciMLBase.NonlinearAliasSpecifier(), kwargs... +# Shared `DynamicSS` setup for `__solve` and `__init`: the ODE problem integrating +# `prob`'s residual to steady state and its termination callback. SCC lowerings are +# handled by `__solve_scc_lowering` before this is reached. +function __dynamicss_ode_setup( + prob::SciMLBase.AbstractSteadyStateProblem, alg::DynamicSS; + abstol, reltol, odesolve_kwargs, termination_condition, alias, kwargs... ) - sccsol = __solve_scc_lowering( - prob, alg, args...; abstol, reltol, odesolve_kwargs, - termination_condition, alias, save_idxs, kwargs... - ) - sccsol !== nothing && return sccsol - tspan = __get_tspan(prob.u0, alg) f = if prob isa SteadyStateProblem @@ -216,19 +249,39 @@ function SciMLBase.__solve( haskey(kwargs, :callback) && (callback = CallbackSet(callback, kwargs[:callback])) haskey(odesolve_kwargs, :callback) && (callback = CallbackSet(callback, odesolve_kwargs[:callback])) - kwargs = pairs(__without_verbose(kwargs)) - # Construct and solve the ODEProblem + run_kwargs = pairs(__without_verbose(kwargs)) odeprob = ODEProblem{isinplace(prob), true}(f, prob.u0, tspan, prob.p) + odealias = SciMLBase.ODEAliasSpecifier(; + alias_p = alias.alias_p, alias_f = alias.alias_f, alias_u0 = alias.alias_u0 + ) + return (; odeprob, tc_cache, abstol, reltol, callback, run_kwargs, odealias) +end + +function SciMLBase.__solve( + prob::SciMLBase.AbstractSteadyStateProblem, alg::DynamicSS, + args...; abstol = 1.0e-8, reltol = 1.0e-6, odesolve_kwargs = (;), + save_idxs = nothing, termination_condition = NonlinearSolveBase.NormTerminationMode(infnorm), + alias = SciMLBase.NonlinearAliasSpecifier(), kwargs... + ) + lp = __scc_lowering(prob) + lp !== nothing && return __solve_scc_lowering( + prob, lp, alg, args...; abstol, reltol, odesolve_kwargs, + termination_condition, alias, save_idxs, kwargs... + ) + + setup = __dynamicss_ode_setup( + prob, alg; abstol, reltol, odesolve_kwargs, termination_condition, alias, kwargs... + ) odesol = solve( - odeprob, alg.alg, args...; abstol, reltol, kwargs..., - odesolve_kwargs..., callback, save_end = true, - alias = SciMLBase.ODEAliasSpecifier(; - alias_p = alias.alias_p, - alias_f = alias.alias_f, alias_u0 = alias.alias_u0 - ) + setup.odeprob, alg.alg, args...; setup.abstol, setup.reltol, + setup.run_kwargs..., odesolve_kwargs..., setup.callback, save_end = true, + alias = setup.odealias ) + return __steady_state_solution(prob, alg, setup, odesol, save_idxs) +end - resid, u, retcode = __get_result_from_sol(tc_cache, odesol) +function __steady_state_solution(prob, alg::DynamicSS, setup, odesol, save_idxs) + resid, u, retcode = __get_result_from_sol(setup.tc_cache, odesol) if save_idxs !== nothing u = u[save_idxs] @@ -241,6 +294,31 @@ function SciMLBase.__solve( ) end +function SciMLBase.__init( + prob::SciMLBase.AbstractSteadyStateProblem, alg::DynamicSS, + args...; abstol = 1.0e-8, reltol = 1.0e-6, odesolve_kwargs = (;), + save_idxs = nothing, termination_condition = NonlinearSolveBase.NormTerminationMode(infnorm), + alias = SciMLBase.NonlinearAliasSpecifier(), kwargs... + ) + lp = __scc_lowering(prob) + lp !== nothing && return SteadyStateSCCCache( + prob, lp, alg, args, (; + abstol, reltol, odesolve_kwargs, termination_condition, alias, + save_idxs, kwargs..., + ) + ) + + setup = __dynamicss_ode_setup( + prob, alg; abstol, reltol, odesolve_kwargs, termination_condition, alias, kwargs... + ) + integrator = init( + setup.odeprob, alg.alg, args...; setup.abstol, setup.reltol, + setup.run_kwargs..., odesolve_kwargs..., setup.callback, save_end = true, + alias = setup.odealias + ) + return SteadyStateODECache(prob, alg, integrator, setup, save_idxs) +end + # SICNM: Semi-Implicit Continuous Newton Method # Solves 0 = g(y) by integrating the DAE ẏ = z, 0 = J(y)z + g(y) to steady state, # where J is the Jacobian of g. See the SICNM docstring for details and references. @@ -270,19 +348,13 @@ function __sicnm_g_and_jvp!(gval, jvp, g!::G, y, z) where {G} return nothing end -function SciMLBase.__solve( - prob::SciMLBase.AbstractSteadyStateProblem, alg::SICNM, - args...; abstol = 1.0e-8, reltol = 1.0e-6, odesolve_kwargs = (;), - save_idxs = nothing, - termination_condition = NonlinearSolveBase.AbsNormTerminationMode(infnorm), - alias = SciMLBase.NonlinearAliasSpecifier(), kwargs... - ) - sccsol = __solve_scc_lowering( - prob, alg, args...; abstol, reltol, odesolve_kwargs, - termination_condition, alias, save_idxs, kwargs... +# Shared `SICNM` setup for `__solve` and `__init`: builds the extended DAE ODE +# problem whose continuous-Newton flow drives `g(y) = 0`, along with the +# termination callback based on the residual `g`. Mirrors `__dynamicss_ode_setup`. +function __sicnm_ode_setup( + prob::SciMLBase.AbstractSteadyStateProblem, alg::SICNM; + abstol, reltol, odesolve_kwargs, termination_condition, kwargs... ) - sccsol !== nothing && return sccsol - prob.u0 isa AbstractVector || throw(ArgumentError("SICNM currently only supports `AbstractVector` initial conditions")) tspan = __get_tspan(prob.u0, alg) @@ -384,17 +456,45 @@ function SciMLBase.__solve( odefun = SciMLBase.ODEFunction{iip, SciMLBase.FullSpecialize}(fext; mass_matrix) odeprob = ODEProblem{iip}(odefun, u0, tspan, p) + run_kwargs = pairs(__without_verbose(kwargs)) + + return (; + odeprob, tc_cache, n, g, gbuf, iip, + ode_abstol, ode_reltol, callback, run_kwargs, + ) +end + +function SciMLBase.__solve( + prob::SciMLBase.AbstractSteadyStateProblem, alg::SICNM, + args...; abstol = 1.0e-8, reltol = 1.0e-6, odesolve_kwargs = (;), + save_idxs = nothing, + termination_condition = NonlinearSolveBase.AbsNormTerminationMode(infnorm), + alias = SciMLBase.NonlinearAliasSpecifier(), kwargs... + ) + lp = __scc_lowering(prob) + lp !== nothing && return __solve_scc_lowering( + prob, lp, alg, args...; abstol, reltol, odesolve_kwargs, + termination_condition, alias, save_idxs, kwargs... + ) + + setup = __sicnm_ode_setup( + prob, alg; abstol, reltol, odesolve_kwargs, termination_condition, kwargs... + ) odesol = solve( - odeprob, alg.alg, args...; abstol = ode_abstol, reltol = ode_reltol, - kwargs..., odesolve_kwargs..., callback, save_end = true + setup.odeprob, alg.alg, args...; abstol = setup.ode_abstol, + reltol = setup.ode_reltol, setup.run_kwargs..., odesolve_kwargs..., + setup.callback, save_end = true ) + return __steady_state_solution(prob, alg, setup, odesol, save_idxs) +end - u, retcode = __sicnm_result(tc_cache, odesol, n) - resid = if iip - g(gbuf, u) - gbuf +function __steady_state_solution(prob, alg::SICNM, setup, odesol, save_idxs) + u, retcode = __sicnm_result(setup.tc_cache, odesol, setup.n) + resid = if setup.iip + setup.g(setup.gbuf, u) + setup.gbuf else - g(u) + setup.g(u) end if save_idxs !== nothing @@ -408,6 +508,32 @@ function SciMLBase.__solve( ) end +function SciMLBase.__init( + prob::SciMLBase.AbstractSteadyStateProblem, alg::SICNM, + args...; abstol = 1.0e-8, reltol = 1.0e-6, odesolve_kwargs = (;), + save_idxs = nothing, + termination_condition = NonlinearSolveBase.AbsNormTerminationMode(infnorm), + alias = SciMLBase.NonlinearAliasSpecifier(), kwargs... + ) + lp = __scc_lowering(prob) + lp !== nothing && return SteadyStateSCCCache( + prob, lp, alg, args, (; + abstol, reltol, odesolve_kwargs, termination_condition, alias, + save_idxs, kwargs..., + ) + ) + + setup = __sicnm_ode_setup( + prob, alg; abstol, reltol, odesolve_kwargs, termination_condition, kwargs... + ) + integrator = init( + setup.odeprob, alg.alg, args...; abstol = setup.ode_abstol, + reltol = setup.ode_reltol, setup.run_kwargs..., odesolve_kwargs..., + setup.callback, save_end = true + ) + return SteadyStateODECache(prob, alg, integrator, setup, save_idxs) +end + function __sicnm_result(tc_cache, odesol, n) u, _, retcode = termination_condition_result( tc_cache, last(odesol.u)[1:n], last(odesol.t), odesol.retcode diff --git a/test/scc.jl b/test/scc.jl index 4780d4f..31ec16b 100644 --- a/test/scc.jl +++ b/test/scc.jl @@ -2,8 +2,9 @@ using SteadyStateDiffEq, NonlinearSolve, OrdinaryDiffEq, Test using ModelingToolkit using ModelingToolkit: t_nounits as t, D_nounits as D using SCCNonlinearSolve: SCCAlg +using NonlinearSolve.NonlinearSolveBase: AbsNormSafeTerminationMode using SciMLBase: HomotopyProblem, LinearProblem, NonlinearProblem, SCCNonlinearProblem, - SteadyStateSolution + SteadyStateSolution, step! function coupled_scc_problem(iip, use_vector) f = if iip @@ -435,3 +436,103 @@ end @test sol.prob isa SCCNonlinearProblem @test sol.original isa Tuple{SciMLBase.LinearSolution, NonlinearSolution} end + +# `solve!(init(prob, alg))` goes through the same steady-state finalization as +# `solve(prob, alg)`: termination-condition retcodes, best-state selection, +# `save_idxs`, and the sequential SCC solve of a stored lowering. +@testset "init/solve! on DynamicSS/SICNM matches solve" begin + function test_matches_solve(prob, alg; kwargs...) + sol = solve!(init(prob, alg; kwargs...)) + ref = solve(prob, alg; kwargs...) + @test sol isa NonlinearSolution + @test sol.retcode == ref.retcode + @test sol.u ≈ ref.u + @test sol.resid ≈ ref.resid + return sol + end + + @testset "plain SteadyStateProblem, alg=$(nameof(typeof(alg)))" for alg in ( + DynamicSS(Tsit5()), SICNM(Rodas5P()), + ) + prob = SteadyStateProblem((u, p, t) -> [1, 2] .- u, [0.0, 0.0]) + sol = test_matches_solve(prob, alg; abstol = 1.0e-10, reltol = 1.0e-10) + @test successful_retcode(sol) + @test sol.u ≈ [1.0, 2.0] atol = 1.0e-6 + + sol = test_matches_solve( + prob, alg; save_idxs = [2], abstol = 1.0e-10, reltol = 1.0e-10 + ) + @test sol.u ≈ [2.0] atol = 1.0e-6 + + cache = init(prob, alg; abstol = 1.0e-10, reltol = 1.0e-10) + step!(cache) + @test cache.integrator.t > 0 + @test successful_retcode(solve!(cache)) + end + + @testset "finite-time nonconvergence, alg=$(nameof(typeof(alg)))" for alg in ( + DynamicSS(Tsit5(); tspan = 1.0e-3), SICNM(Rodas5P(); tspan = 1.0e-3), + ) + prob = SteadyStateProblem((u, p, t) -> [1, 2] .- u, [0.0, 0.0]) + sol = test_matches_solve(prob, alg; abstol = 1.0e-10, reltol = 1.0e-10) + @test !successful_retcode(sol) + end + + @testset "protective termination" begin + # `u' = u` grows past the protective threshold. + prob = SteadyStateProblem((u, p, t) -> u, [1.0]) + tc = AbsNormSafeTerminationMode(u -> maximum(abs, u); protective_threshold = 1.01) + sol = test_matches_solve( + prob, DynamicSS(Tsit5()); abstol = 1.0e-10, reltol = 1.0e-10, + termination_condition = tc + ) + @test sol.retcode == ReturnCode.Unstable + end + + @testset "manually-built SCC lowering, alg=$(nameof(typeof(alg)))" for (alg, shortalg) in ( + (DynamicSS(Tsit5()), DynamicSS(Tsit5(); tspan = 1.0e-3)), + (SICNM(Rodas5P()), SICNM(Rodas5P(); tspan = 1.0e-3)), + ) + sccprob = dynamicss_scc_problem(false, false) + prob = SteadyStateProblem( + (u, p, t) -> 1 .- u, [0.0, 0.0]; lowered_problem = sccprob + ) + sol = test_matches_solve(prob, alg; abstol = 1.0e-10, reltol = 1.0e-10) + @test successful_retcode(sol) + @test sol.u ≈ [1, 2, 1, 2] atol = 1.0e-8 + + sol = test_matches_solve( + prob, alg; save_idxs = [2, 4], abstol = 1.0e-10, reltol = 1.0e-10 + ) + @test sol.u ≈ [2, 2] atol = 1.0e-8 + + # The block cannot reach steady state within the short time span. + failing = SCCNonlinearProblem( + (NonlinearProblem((u, p) -> 1 .- u, [0.0]),), (Returns(nothing),) + ) + prob = SteadyStateProblem( + (u, p, t) -> 1 .- u, [0.0]; lowered_problem = failing + ) + sol = test_matches_solve(prob, shortalg; abstol = 1.0e-10, reltol = 1.0e-10) + @test !successful_retcode(sol) + end + + @testset "ModelingToolkit SCC decomposition, alg=$(nameof(typeof(alg)))" for alg in ( + DynamicSS(Tsit5()), SICNM(Rodas5P()), + ) + @variables a(t) b(t) x(t) [irreducible = true] + @named model = System( + [D(a) ~ 5 - 3a - b, D(b) ~ 5 - a - 2b, D(x) ~ a + b - x^3], t + ) + sys = mtkcompile(model) + prob = SteadyStateProblem(sys, [a => 0.8, b => 1.8, x => 0.8]) + + sol = test_matches_solve(prob, alg; abstol = 1.0e-10, reltol = 1.0e-10) + @test successful_retcode(sol) + @test sol[[a, b, x]] ≈ [1, 2, cbrt(3)] atol = 1.0e-8 + + test_matches_solve( + prob, alg; save_idxs = [1], abstol = 1.0e-10, reltol = 1.0e-10 + ) + end +end