Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 30 additions & 0 deletions docs/src/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/SteadyStateDiffEq.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
218 changes: 172 additions & 46 deletions src/solve.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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]
Expand All @@ -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.
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
Loading
Loading