From 57a2f3410428aff29d6e87d5363c4e39874b01fb Mon Sep 17 00:00:00 2001 From: lkdvos Date: Mon, 3 Aug 2026 16:23:33 -0400 Subject: [PATCH 1/2] Add default State and StoppingCriterionState implementations Downstream packages had to write the same three things before running a single iteration: a state struct declaring `iterate`, `iteration` and `stopping_criterion_state`, an `initialize_state` allocating it, and an `initialize_state!` resetting it. This repository showed the cost itself, with `DummyState`, `NewtonState` and `HeronState` being the same struct three times over, each with a different field order. Add `DefaultState`, holding those three expected properties next to a single generic `data` field for anything else an algorithm carries from one step to the next. `data` is opaque: nothing is forwarded to it and none of its contents are exposed as properties, so access stays type stable and the container is the algorithm's choice. Give `DefaultStoppingCriterionState` the same `data` field, defaulting to `nothing` so existing zero-argument construction is unaffected. Add defaults for `initialize_state` and `initialize_state!` at both levels, taking their arguments positionally as `(problem, algorithm, iterate, state_data, stopping_state_data)` next to the keyword form the interface documents. An algorithm that is served by these now only has to provide a `Problem`, an `Algorithm` and a `step!`; since `step!` dispatches on the `Algorithm`, sharing one state type across algorithms costs no flexibility. `StopAfterIteration` no longer carries its own initialization pair, which the criterion-level defaults now cover. Rework the documentation so all three pages lead with the default and present a state of your own as the option, keeping `HeronState` as the closing example on the interface page. Co-Authored-By: Claude Opus 5 (1M context) --- Changelog.md | 8 ++ docs/src/interface.md | 125 +++++++++++++++++------- docs/src/logging.md | 31 ++---- docs/src/stopping_criterion.md | 69 ++++++------- src/AlgorithmsInterface.jl | 2 + src/default_state.jl | 102 ++++++++++++++++++++ src/interface/interface.jl | 15 +++ src/stopping_criterion.jl | 46 ++++++--- test/state.jl | 171 +++++++++++++++++++++++++++++++++ test/stopping_criterion.jl | 64 ++++++++---- 10 files changed, 502 insertions(+), 131 deletions(-) create mode 100644 src/default_state.jl create mode 100644 test/state.jl diff --git a/Changelog.md b/Changelog.md index a185eb1..6480992 100644 --- a/Changelog.md +++ b/Changelog.md @@ -17,11 +17,19 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - `StopReasonAction`, a `LoggingAction` that reports `get_reason` at the `:Stop` context. - Defaults for `get_reason` (`nothing`) and for the type-domain `indicates_convergence` (`false`), so a criterion that implements neither no longer hits a `MethodError` from the derived convergence reporting. - Exports for `DefaultStoppingCriterionState`, `StopAfterTimePeriodState` and `GroupStoppingCriterionState`, which a downstream criterion is expected to reuse. +- `DefaultState`, a `State` storing the `iterate`, `iteration` and `stopping_criterion_state` every state is expected to provide, next to a single `data` field for anything else an algorithm has to carry from one step to the next. + The `data` field is opaque to this package and none of its contents are exposed as properties, so an algorithm reaches them through `state.data`. +- Defaults for `initialize_state` and `initialize_state!` returning and resetting a `DefaultState`, so that an algorithm whose state holds nothing beyond the expected properties only has to provide a `Problem`, an `Algorithm` and a `step!`. + Both accept their arguments positionally as `(problem, algorithm, iterate, state_data, stopping_state_data)`, next to the keyword form the interface documents. +- Defaults for `initialize_state(problem, algorithm, stopping_criterion)` and its mutating variant, returning and resetting a `DefaultStoppingCriterionState`, with a `stopping_state_data` keyword to seed its `data`. ### Changed - `indicates_convergence` without a state moved to the type domain: a new criterion implements `indicates_convergence(::Type{YourCriterion})`, and `indicates_convergence(criterion)` forwards to it. `StopWhenAll` and `StopWhenAny` combine their children in the type domain as well, so a criterion that only implements the variant taking an instance is no longer accounted for in a group. +- `DefaultStoppingCriterionState` gained a type parameter and a `data` field for whatever a criterion has to remember from one iteration to the next. + It defaults to `nothing`, so `DefaultStoppingCriterionState()` is unaffected. +- `StopAfterIteration` no longer carries its own `initialize_state` and `initialize_state!`, which the new criterion-level defaults now cover. ### Fixed diff --git a/docs/src/interface.md b/docs/src/interface.md index 461b76d..d51a180 100644 --- a/docs/src/interface.md +++ b/docs/src/interface.md @@ -48,6 +48,8 @@ They are illustrative; various performance and generality questions will be left ### Algorithm types +The [`Problem`](@ref) holds the immutable input and the [`Algorithm`](@ref) the immutable configuration: + ```@example Heron using AlgorithmsInterface @@ -58,50 +60,21 @@ end struct HeronAlgorithm <: Algorithm stopping_criterion # will be plugged in later (any StoppingCriterion) end - -mutable struct HeronState <: State - iterate::Float64 # current iterate - iteration::Int # current iteration count - stopping_criterion_state # will be plugged in later (any StoppingCriterionState) -end ``` -### Initialization +That leaves the third piece, the [`State`](@ref), which for most algorithms is entirely unremarkable: +it holds the three properties the interface asks for – the `iterate`, the `iteration` counter and the `stopping_criterion_state` – and nothing else. +This package therefore provides [`DefaultState`](@ref) for exactly that, so there is nothing to define here. +[Writing your own](@ref sec_custom_state) is the exception, covered at the end of this section. -In order to start implementing the core parts of our algorithm, we start at the very beginning. -There are two main entry points provided by the interface: - -- [`initialize_state`](@ref) constructs an entirely new state for the algorithm -- [`initialize_state!`](@ref) (in-place) reset of an existing state. - -An example implementation might look like: - -```@example Heron -function AlgorithmsInterface.initialize_state(problem::SqrtProblem, algorithm::HeronAlgorithm; kwargs...) - x0 = rand() # random initial guess - stopping_criterion_state = initialize_state(problem, algorithm, algorithm.stopping_criterion) - return HeronState(x0, 0, stopping_criterion_state) -end - -function AlgorithmsInterface.initialize_state!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState; kwargs...) - # reset the state for the algorithm - state.iterate = rand() - state.iteration = 0 - - # reset the state for the stopping criterion - state = AlgorithmsInterface.initialize_state!( - problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state - ) - return state -end -``` +The same goes for initialization: [`initialize_state`](@ref) and [`initialize_state!`](@ref), which respectively construct a fresh state and reset an existing one, both default to a [`DefaultState`](@ref). ### Iteration steps Algorithms define a mutable step via [`step!`](@ref). For Heron's method: ```@example Heron -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState) +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState) S = problem.S x = state.iterate state.iterate = 0.5 * (x + S / x) @@ -121,24 +94,102 @@ function AlgorithmsInterface.increment!(state::State) end ``` +Since [`step!`](@ref) dispatches on the [`Algorithm`](@ref), sharing one state type across algorithms costs nothing. + ### Running the algorithm +That is the whole algorithm: a [`Problem`](@ref), an [`Algorithm`](@ref) and a [`step!`](@ref). With these definitions in place you can already run (assuming you also choose a stopping criterion – added in the next section): ```@example Heron function heron_sqrt(x; maxiter = 10) prob = SqrtProblem(x) alg = HeronAlgorithm(StopAfterIteration(maxiter)) - return solve(prob, alg) # allocates & runs + return solve(prob, alg; iterate = 1.0) # allocates & runs end println("Approximate sqrt: ", heron_sqrt(16.0)) ``` +The `iterate` keyword is where the default initialization gets its starting point from, and it is required: there is no sensible guess this package could make on an algorithm's behalf. + Note that [`solve`](@ref) will default to returning `state.iterate`. If desired, this can be customized by altering [`finalize_state!`](@ref). We will refine this example with better halting logic and logging shortly. +### Carrying extra data + +Anything an algorithm has to carry from one step to the next beyond those three properties goes into the `data` field of the [`DefaultState`](@ref), which this package treats as opaque: +it is passed along and reset alongside the rest of the state, but nothing is ever read from it, so an algorithm picks whatever container suits it. +Prefer a `NamedTuple` or a small struct of your own, which keep access to the contents type stable: + +```@example Heron +struct TrackingHeronAlgorithm <: Algorithm + stopping_criterion +end + +function AlgorithmsInterface.step!(problem::SqrtProblem, ::TrackingHeronAlgorithm, state::DefaultState) + S = problem.S + x = state.iterate + state.iterate = 0.5 * (x + S / x) + push!(state.data.residuals, abs(state.iterate^2 - S)) + return state +end + +problem = SqrtProblem(16.0) +tracking = TrackingHeronAlgorithm(StopAfterIteration(5)) +state = initialize_state(problem, tracking, 1.0, (; residuals = Float64[])) +solve!(problem, tracking, state) +state.data.residuals +``` + +Note that the `NamedTuple` itself is immutable while the vector inside it is not, which is all this needs: the state keeps handing the same container back, and the algorithm mutates its contents. +Here the initial iterate and the data are passed positionally, as `initialize_state(problem, algorithm, iterate, state_data)`; the [`StoppingCriterionState`](@ref) can be given its own data as a further argument. + +### [Defining your own state](@id sec_custom_state) + +Reach for a state type of your own once `data` grows past a handful of entries, or once the algorithm wants to dispatch on the state itself. +It has to provide the three expected properties, and with it come the two initialization methods that the default was providing until now: + +```@example Heron +mutable struct HeronState <: State + iterate::Float64 # current iterate + iteration::Int # current iteration count + stopping_criterion_state # any StoppingCriterionState +end + +function AlgorithmsInterface.initialize_state(problem::SqrtProblem, algorithm::HeronAlgorithm; kwargs...) + x0 = rand() # random initial guess + stopping_criterion_state = initialize_state(problem, algorithm, algorithm.stopping_criterion) + return HeronState(x0, 0, stopping_criterion_state) +end + +function AlgorithmsInterface.initialize_state!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState; kwargs...) + # reset the state for the algorithm + state.iterate = rand() + state.iteration = 0 + + # reset the state for the stopping criterion + AlgorithmsInterface.initialize_state!( + problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state + ) + return state +end + +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState) + S = problem.S + x = state.iterate + state.iterate = 0.5 * (x + S / x) + return state +end +``` + +This `HeronAlgorithm` now determines its own initial iterate rather than being handed one, so it runs without the `iterate` keyword: + +```@example Heron +println("Approximate sqrt: ", solve(SqrtProblem(16.0), HeronAlgorithm(StopAfterIteration(10)))) +``` + ## Reference: Core interface types & functions Below are the automatic API docs for the core interface pieces. Read them after grasping the example above – the intent should now be clearer. @@ -172,7 +223,7 @@ Private = true ```@autodocs Modules = [AlgorithmsInterface] -Pages = ["interface/state.jl"] +Pages = ["interface/state.jl", "default_state.jl"] Order = [:type, :function] Private = true ``` diff --git a/docs/src/logging.md b/docs/src/logging.md index 2dcabc5..ae6830a 100644 --- a/docs/src/logging.md +++ b/docs/src/logging.md @@ -46,26 +46,7 @@ struct HeronAlgorithm <: Algorithm stopping_criterion end -mutable struct HeronState <: State - iterate::Float64 - iteration::Int - stopping_criterion_state -end - -function AlgorithmsInterface.initialize_state(problem::SqrtProblem, algorithm::HeronAlgorithm; kwargs...) - x0 = rand() - stopping_criterion_state = initialize_state(problem, algorithm, algorithm.stopping_criterion) - return HeronState(x0, 0, stopping_criterion_state) -end - -function AlgorithmsInterface.initialize_state!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState; kwargs...) - state.iterate = rand() - state.iteration = 0 - initialize_state!(problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state) - return state -end - -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState) +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState) S = problem.S x = state.iterate state.iterate = 0.5 * (x + S / x) @@ -75,11 +56,13 @@ end function heron_sqrt(x; stopping_criterion = StopAfterIteration(10)) prob = SqrtProblem(x) alg = HeronAlgorithm(stopping_criterion) - return solve(prob, alg) # allocates & runs + return solve(prob, alg; iterate = 1.0) # allocates & runs end nothing # hide ``` +Note that this leaves the state to [`DefaultState`](@ref), as the [interface section](@ref sec_interface) does, so there is neither a state type nor an `initialize_state` to be seen here. + It is already interesting to note that there are no further modifications necessary to start leveraging the logging system. ### Basic iteration printing @@ -211,7 +194,7 @@ function AlgorithmsInterface.handle_message!( action::CaptureHistory, problem::SqrtProblem, algorithm::HeronAlgorithm, - state::HeronState; + state::DefaultState; kwargs... ) push!(action.iterates, state.iterate) @@ -287,7 +270,7 @@ end StatsCollector() = StatsCollector(0, 0.0, 0.0) function AlgorithmsInterface.handle_message!( - action::StatsCollector, problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState; + action::StatsCollector, problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState; kwargs... ) action.count += 1 @@ -415,7 +398,7 @@ Here we will illustrate this by a slight adaptation of our algorithm, which coul To emit a custom logging event from within your algorithm, call [`emit_message`](@ref): ```@example Heron -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState) +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState) # Suppose we check for numerical issues if !isfinite(state.iterate) || mod(state.iteration, 10) == 0 emit_message(problem, algorithm, state, :Restart) diff --git a/docs/src/stopping_criterion.md b/docs/src/stopping_criterion.md index 944ae6d..00deb5d 100644 --- a/docs/src/stopping_criterion.md +++ b/docs/src/stopping_criterion.md @@ -27,7 +27,7 @@ The package ships several concrete [`StoppingCriterion`](@ref)s: Each criterion has an associated [`StoppingCriterionState`](@ref) storing dynamic data (iteration when met, elapsed time, etc.). -Recall our [example implementation](@ref sec_heron) for Heron's method, where we added a `stopping_criterion` to the `Algorithm`, as well as a `stopping_criterion_state` to the `State`. +Recall our [example implementation](@ref sec_heron) for Heron's method, where the `Algorithm` carries a `stopping_criterion` and the `State` a `stopping_criterion_state`. ```@example Heron using AlgorithmsInterface @@ -39,50 +39,35 @@ end struct HeronAlgorithm <: Algorithm stopping_criterion # any StoppingCriterion end - -mutable struct HeronState <: State - iterate::Float64 # current iterate - iteration::Int # current iteration count - stopping_criterion_state # any StoppingCriterionState -end ``` Here, we delve a bit deeper into the core components of what made our algorithm stop, even though we had to add very little additional functionality. ### Initialization -The first core component to enable working with stopping criteria is to extend the initialization step to include initializing a [`StoppingCriterionState`](@ref) as well. -This can conveniently be done through the same initialization functions we used for initializing the state: +The first core component to enable working with stopping criteria is that the initialization step initializes a [`StoppingCriterionState`](@ref) as well. +This happens through the same initialization functions we used for initializing the state: - [`initialize_state`](@ref) constructs an entirely new stopping state for the algorithm - [`initialize_state!`](@ref) (in-place) reset of an existing stopping state. -```@example Heron -function AlgorithmsInterface.initialize_state(problem::SqrtProblem, algorithm::HeronAlgorithm; kwargs...) - x0 = rand() # random initial guess - stopping_criterion_state = initialize_state(problem, algorithm, algorithm.stopping_criterion) - return HeronState(x0, 0, stopping_criterion_state) -end - -function AlgorithmsInterface.initialize_state!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState; kwargs...) - # reset the state for the algorithm - state.iterate = rand() - state.iteration = 0 +Since we leave the state to [`DefaultState`](@ref), this is already taken care of: the defaults pair it with the state of the algorithm's own criterion, along the lines of - # reset the state for the stopping criterion - state = AlgorithmsInterface.initialize_state!( - problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state - ) - return state +```julia +function AlgorithmsInterface.initialize_state(problem::Problem, algorithm::Algorithm; iterate, kwargs...) + stopping_criterion_state = initialize_state(problem, algorithm, algorithm.stopping_criterion; kwargs...) + return DefaultState(iterate, stopping_criterion_state) end ``` +A state of your own is where you would write that pairing out yourself. + ### Iteration During the iteration procedure, as set out by our design principles, we do not have to modify any of the code, and the stopping criteria do not show up: ```@example Heron -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState) +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState) S = problem.S x = state.iterate state.iterate = 0.5 * (x + S / x) @@ -101,8 +86,7 @@ end In other words, all of the logic is handled by the [`is_finished!`](@ref) function. The generic stopping criteria provided by this package have default implementations for this function that work out-of-the-box. -This is partially because we used conventional names for the fields in the structs. -There, `Algorithm` assumes the existence of `stopping_criterion`, while `State` assumes `iterate` and `iteration` and `stopping_criterion_state` to exist. +This is partially because everything is reached under conventional names: `Algorithm` assumes the existence of `stopping_criterion`, while `State` assumes `iterate` and `iteration` and `stopping_criterion_state` to exist — which is exactly what [`DefaultState`](@ref) provides, and what a state of your own has to provide too. ### Running the algorithm @@ -112,7 +96,7 @@ We can again combine everything into a single function, but now make the stoppin function heron_sqrt(x; stopping_criterion) prob = SqrtProblem(x) alg = HeronAlgorithm(stopping_criterion) - return solve(prob, alg) # allocates & runs + return solve(prob, alg; iterate = 1.0) # allocates & runs end heron_sqrt(2; stopping_criterion = StopAfterIteration(10)) @@ -148,6 +132,9 @@ Suppose we want to stop when successive iterates change by less than `ϵ`, we co In order to do so, we need to define our own structs and implement the required interface. Again, we split up the data into a _static_ part, the [`StoppingCriterion`](@ref), and a _dynamic_ part, the [`StoppingCriterionState`](@ref). +The dynamic part is usually not yours to write: [`DefaultStoppingCriterionState`](@ref) records the iteration at which the criterion triggered and carries a `data` field for anything else it has to remember, which covers most criteria. +We spell out a state of our own here because it shows the full picture, and because a criterion that wants its fields named and typed is exactly the case that calls for one. + ```@example Heron struct StopWhenStable <: StoppingCriterion tol::Float64 # when do we consider things converged @@ -253,8 +240,8 @@ The variant taking a criterion simply forwards to the type, and the two-argument ```@example Heron criterion = StopWhenStable(1e-8) -state = AlgorithmsInterface.initialize_state(SqrtProblem(16.0), HeronAlgorithm(criterion), criterion) -indicates_convergence(criterion), indicates_convergence(criterion, state) +criterion_state = AlgorithmsInterface.initialize_state(SqrtProblem(16.0), HeronAlgorithm(criterion), criterion) +indicates_convergence(criterion), indicates_convergence(criterion, criterion_state) ``` The criterion always *could* indicate convergence, but its fresh state has not yet seen it happen. @@ -274,7 +261,7 @@ Since [`solve`](@ref) returns only the iterate, we use [`solve!`](@ref) with a s function heron_verdict(x, criterion) problem = SqrtProblem(x) algorithm = HeronAlgorithm(criterion) - state = AlgorithmsInterface.initialize_state(problem, algorithm) + state = AlgorithmsInterface.initialize_state(problem, algorithm, 1.0) solve!(problem, algorithm, state) @@ -320,16 +307,22 @@ heron_sqrt(16.0; stopping_criterion = criterion) ### Summary -Implementing a criterion usually means defining: +Implementing a criterion means defining: 1. A subtype of [`StoppingCriterion`](@ref). -2. A state subtype of [`StoppingCriterionState`](@ref) capturing dynamic fields, including an `at_iteration` recording when the criterion triggered. -3. `initialize_state` and `initialize_state!` for setup/reset. -4. `is_finished!` (mutating) and optionally `is_finished` (non‑mutating) variants. -5. `get_reason` (return `nothing` or a string) for user feedback, gated on `is_active`. -6. `indicates_convergence(::Type{YourCriterion})` to mark if meeting it implies convergence. +2. `is_finished!` (mutating) and optionally `is_finished` (non‑mutating) variants. +3. `get_reason` (return `nothing` or a string) for user feedback, gated on `is_active`. +4. `indicates_convergence(::Type{YourCriterion})` to mark if meeting it implies convergence. The `(criterion,)` and the `(criterion, criterion_state)` variant are derived from this one and do not need to be defined. +The state is taken care of for you. +[`initialize_state`](@ref) and [`initialize_state!`](@ref) return and reset a [`DefaultStoppingCriterionState`](@ref), which records the `at_iteration` that all the reporting is built on and carries a `data` field for whatever else the criterion has to remember — the `previous_iterate` and `delta` above, for instance. + +Only a criterion that is not served by that state, as the one above wanted its fields named and typed, additionally defines: + +* A state subtype of [`StoppingCriterionState`](@ref) capturing its dynamic fields, including an `at_iteration` recording when the criterion triggered. +* `initialize_state` and `initialize_state!` for its setup and reset. + You may also implement `Base.summary(io, criterion, criterion_state)` for compact status reports, and `is_active(criterion, criterion_state)` if your state does not record its status in an `at_iteration` property. diff --git a/src/AlgorithmsInterface.jl b/src/AlgorithmsInterface.jl index 74ae2f6..3626591 100644 --- a/src/AlgorithmsInterface.jl +++ b/src/AlgorithmsInterface.jl @@ -18,12 +18,14 @@ include("interface/state.jl") include("interface/interface.jl") include("stopping_criterion.jl") +include("default_state.jl") include("logging.jl") include("test_suite.jl") # general interface export Algorithm, Problem, State +export DefaultState export initialize_state, initialize_state! export finalize_state! diff --git a/src/default_state.jl b/src/default_state.jl new file mode 100644 index 0000000..2ca3c79 --- /dev/null +++ b/src/default_state.jl @@ -0,0 +1,102 @@ +# +# +# A default state + +@doc """ + DefaultState <: State + +A [`State`](@ref) that stores the properties every [`State`](@ref) is expected to provide, +together with a single field to hold any further data an [`Algorithm`](@ref) needs. + +Together with the default [`initialize_state`](@ref) and [`initialize_state!`](@ref) methods, +this spares a downstream algorithm the definition of its own state type: +a [`Problem`](@ref), an [`Algorithm`](@ref) and a [`step!`](@ref) method suffice. +Since [`step!`](@ref) dispatches on the [`Algorithm`](@ref), reusing this state costs no +flexibility in doing so. + +# Fields + +* `iterate` stores the current iterate ``x^{(k)}``. +* `stopping_criterion_state` stores the [`StoppingCriterionState`](@ref) belonging to the + [`StoppingCriterion`](@ref) of the [`Algorithm`](@ref). +* `data` stores any further data that has to be carried from one step to the next. + It is opaque to this package, and none of its contents are exposed as properties of the + state, so an algorithm reaches them through `state.data`. + A `NamedTuple` or a struct of its own keeps access to them type stable and is the + recommended choice; `nothing` indicates that the algorithm needs no further data. +* `iteration::Int` stores the current iteration step ``k`` that is currently being performed + or was last performed. + +# Constructor + + DefaultState(iterate, stopping_criterion_state, data = nothing, iteration = 0) + +Initialize the state to start at `iterate`, carrying `data` alongside it. +""" +mutable struct DefaultState{V, S <: StoppingCriterionState, D} <: State + iterate::V + stopping_criterion_state::S + data::D + iteration::Int +end + +DefaultState(iterate, stopping_criterion_state::StoppingCriterionState, data = nothing) = + DefaultState(iterate, stopping_criterion_state, data, 0) + +# The default keeps whatever data a criterion state already carries, rather than clearing it. +# A criterion state that carries none takes `nothing`, which its `initialize_state!` ignores. +_stopping_state_data(::StoppingCriterionState) = nothing +_stopping_state_data(stopping_criterion_state::DefaultStoppingCriterionState) = + stopping_criterion_state.data + +# The signature taking an `iterate` is the most generic one there is, so this doubles as the +# fallback for any `Problem` and `Algorithm` that do not provide a state type of their own. +# Both these and their keyword counterparts below are documented by `_doc_init_state`, which +# covers exactly these signatures. +function initialize_state( + problem::Problem, algorithm::Algorithm, iterate, + state_data = nothing, stopping_state_data = nothing; + kwargs..., + ) + stopping_criterion_state = initialize_state( + problem, algorithm, algorithm.stopping_criterion; + stopping_state_data, kwargs..., + ) + return DefaultState(iterate, stopping_criterion_state, state_data) +end + +function initialize_state( + problem::Problem, algorithm::Algorithm; + iterate, state_data = nothing, stopping_state_data = nothing, kwargs..., + ) + return initialize_state( + problem, algorithm, iterate, state_data, stopping_state_data; kwargs... + ) +end + +function initialize_state!( + problem::Problem, algorithm::Algorithm, state::DefaultState, iterate, + state_data = state.data, + stopping_state_data = _stopping_state_data(state.stopping_criterion_state); + kwargs..., + ) + state.iterate = iterate + state.data = state_data + state.iteration = 0 + initialize_state!( + problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state; + stopping_state_data, kwargs..., + ) + return state +end + +function initialize_state!( + problem::Problem, algorithm::Algorithm, state::DefaultState; + iterate = state.iterate, state_data = state.data, + stopping_state_data = _stopping_state_data(state.stopping_criterion_state), + kwargs..., + ) + return initialize_state!( + problem, algorithm, state, iterate, state_data, stopping_state_data; kwargs... + ) +end diff --git a/src/interface/interface.jl b/src/interface/interface.jl index f37d513..55f4558 100644 --- a/src/interface/interface.jl +++ b/src/interface/interface.jl @@ -5,6 +5,21 @@ _doc_init_state = """ Initialize a [`State`](@ref) based on a [`Problem`](@ref) and an [`Algorithm`](@ref). The `kwargs...` should allow to initialize for example the initial point. This can be done in-place for `state`, then only values that did change have to be provided. + +Both have a default in terms of a [`DefaultState`](@ref), which stores the properties every +[`State`](@ref) is expected to provide next to a single `data` field for anything else an +algorithm needs, so an algorithm that is served by such a state does not have to provide a +state type of its own. +These defaults also take their arguments positionally: + + state = initialize_state(problem, algorithm, iterate, state_data, stopping_state_data; kwargs...) + state = initialize_state!(problem, algorithm, state, iterate, state_data, stopping_state_data; kwargs...) + +Here `state_data` is what the [`DefaultState`](@ref) carries and `stopping_state_data` what the +[`StoppingCriterionState`](@ref) does, both defaulting to `nothing` when allocating and to the +values the `state` already holds when resetting. +The remaining `kwargs...` are passed on to the corresponding function for the +[`StoppingCriterion`](@ref) of the [`Algorithm`](@ref). """ function initialize_state end diff --git a/src/stopping_criterion.jl b/src/stopping_criterion.jl index 6bed86e..eb2a3e1 100644 --- a/src/stopping_criterion.jl +++ b/src/stopping_criterion.jl @@ -3,19 +3,22 @@ An abstract type to represent a stopping criterion of an [`Algorithm`](@ref). -A concrete [`StoppingCriterion`](@ref) should also implement an -[`initialize_state(problem::Problem, algorithm::Algorithm, stopping_criterion::StoppingCriterion; kwargs...)`](@ref) function to create its accompanying -[`StoppingCriterionState`](@ref), as well as the corresponding mutating variant to reset such a [`StoppingCriterionState`](@ref). +A concrete [`StoppingCriterion`](@ref) receives its accompanying [`StoppingCriterionState`](@ref) +from a [`DefaultStoppingCriterionState`](@ref), which records the iteration at which the criterion +indicated to stop and carries a `data` field for anything else it has to remember. +A criterion is therefore free of any state bookkeeping by default. It should usually implement * [`is_finished!`](@ref)`(problem, algorithm, state, stopping_criterion, stopping_criterion_state)` * [`is_finished`](@ref)`(problem, algorithm, state, stopping_criterion, stopping_criterion_state)` -* [`initialize_state!`](@ref)`(problem, algorithm, stopping_criterion)` -* [`initialize_state`](@ref)`(problem, algorithm, stopping_criterion)` * [`get_reason`](@ref)`(stopping_criterion, stopping_criterion_state)` * [`indicates_convergence`](@ref)`(::Type{<:StoppingCriterion})` +Only a criterion that is not served by that state defines one of its own, and with it an +[`initialize_state(problem::Problem, algorithm::Algorithm, stopping_criterion::StoppingCriterion; kwargs...)`](@ref) +to create it, as well as the corresponding mutating variant to reset it. + Note that only [`indicates_convergence`](@ref) has to be implemented: it answers whether meeting this criterion *would* mean convergence, which is a static property of the criterion type alone. Both the variant taking a criterion and the one that additionally takes a [`StoppingCriterionState`](@ref), answering whether it *did* happen are derived from it. @@ -560,27 +563,46 @@ end """ DefaultStoppingCriterionState <: StoppingCriterionState -A [`StoppingCriterionState`](@ref) that does not require any information besides -storing the iteration number at which it (last) indicated to stop. +A [`StoppingCriterionState`](@ref) that stores the iteration number at which it (last) +indicated to stop, and optionally any further data its [`StoppingCriterion`](@ref) needs. # Fields * `at_iteration::Int` stores the iteration number at which this state indicated to stop. * `0` means it already indicated to stop at the start. * any negative number means that it has not yet indicated to stop. +* `data` stores any further data the criterion has to carry from one iteration to the next, + for example a value it compares against in the next one. + It is opaque to this package, and none of its contents are exposed as properties of the + state, so a criterion reaches them through `stopping_criterion_state.data`. + A mutable struct of its own is the recommended choice; `nothing`, the default, indicates + that the criterion needs no further data. + +# Constructor + + DefaultStoppingCriterionState(data = nothing) + +Initialize the state to not having indicated to stop yet, carrying `data` alongside it. """ -mutable struct DefaultStoppingCriterionState <: StoppingCriterionState +mutable struct DefaultStoppingCriterionState{D} <: StoppingCriterionState at_iteration::Int - DefaultStoppingCriterionState() = new(-1) + data::D end -initialize_state(::Problem, ::Algorithm, ::StopAfterIteration; kwargs...) = DefaultStoppingCriterionState() +DefaultStoppingCriterionState(data = nothing) = DefaultStoppingCriterionState(-1, data) + +# Fallbacks for any criterion that needs no state of its own beyond `at_iteration`, so that +# such a criterion does not have to provide these two methods at all. +initialize_state( + ::Problem, ::Algorithm, ::StoppingCriterion; stopping_state_data = nothing, kwargs... +) = DefaultStoppingCriterionState(stopping_state_data) function initialize_state!( - ::Problem, ::Algorithm, ::StopAfterIteration, + ::Problem, ::Algorithm, ::StoppingCriterion, stopping_criterion_state::DefaultStoppingCriterionState; - kwargs..., + stopping_state_data = stopping_criterion_state.data, kwargs..., ) stopping_criterion_state.at_iteration = -1 + stopping_criterion_state.data = stopping_state_data return stopping_criterion_state end diff --git a/test/state.jl b/test/state.jl new file mode 100644 index 0000000..3c04b25 --- /dev/null +++ b/test/state.jl @@ -0,0 +1,171 @@ +# Tests for the default state and the default state initialization + +using Test +using AlgorithmsInterface +using AlgorithmsInterface: Test as AIT +using AlgorithmsInterface: increment! +using Dates + +# Fixtures +# -------- + +# A problem and an algorithm that provide nothing beyond what the interface demands, so that +# every state they are solved with has to come from the defaults. Newton's method is spelled +# out with its own state in `newton.jl`; the point here is that this one has none. +struct HalvingProblem <: Problem + target::Float64 +end + +struct Halving{S <: StoppingCriterion} <: Algorithm + stopping_criterion::S +end + +function AlgorithmsInterface.step!(problem::HalvingProblem, ::Halving, state::DefaultState) + state.iterate = (state.iterate + problem.target) / 2 + return state +end + +# Carries a counter through `data` to show that an algorithm can keep data from one step to +# the next without a state type of its own. A mutable struct is what the docs recommend for +# this: the state hands the same container back every step, and only its contents change. +struct CountingHalving{S <: StoppingCriterion} <: Algorithm + stopping_criterion::S +end + +mutable struct StepCounter + steps::Int +end + +function AlgorithmsInterface.step!( + problem::HalvingProblem, ::CountingHalving, state::DefaultState, + ) + state.iterate = (state.iterate + problem.target) / 2 + state.data.steps += 1 + return state +end + +# Tests +# ----- + +@testset "DefaultState construction" begin + scs = DefaultStoppingCriterionState() + + state = DefaultState(2.0, scs) + @test state isa DefaultState{Float64, typeof(scs), Nothing} + @test state.iterate == 2.0 + @test state.stopping_criterion_state === scs + @test state.data === nothing + @test state.iteration == 0 + + # `iteration` is the fourth positional argument + @test DefaultState(2.0, scs, nothing, 3).iteration == 3 + + # the data field is opaque, so anything goes and nothing of it is exposed as a property + named_tuple_state = DefaultState(2.0, scs, (; gradient = 1.0)) + @test named_tuple_state.data.gradient == 1.0 + @test !hasproperty(named_tuple_state, :gradient) + + dict_state = DefaultState(2.0, scs, Dict{Symbol, Any}(:gradient => 1.0)) + @test dict_state.data[:gradient] == 1.0 + dict_state.data[:hessian] = 2.0 + @test dict_state.data[:hessian] == 2.0 +end + +@testset "DefaultState satisfies the State interface" begin + problem = AIT.DummyProblem() + algorithm = AIT.DummyAlgorithm(StopAfterIteration(5)) + state = DefaultState(2.0, DefaultStoppingCriterionState()) + + @test increment!(state) === state + @test state.iteration == 1 + @test finalize_state!(problem, algorithm, state) == 2.0 + @test !is_finished!(problem, algorithm, state) + @test !is_active(algorithm, state) +end + +@testset "initialize_state defaults to a DefaultState" begin + problem = AIT.DummyProblem() + algorithm = AIT.DummyAlgorithm(StopAfterIteration(5)) + + state = initialize_state(problem, algorithm, 2.0) + @test state isa DefaultState + @test state.iterate == 2.0 + @test state.iteration == 0 + @test state.data === nothing + # the accompanying criterion state comes from the algorithm's own stopping criterion + @test state.stopping_criterion_state isa DefaultStoppingCriterionState + @test state.stopping_criterion_state.at_iteration == -1 + @test state.stopping_criterion_state.data === nothing + + # `state_data` and `stopping_state_data` are positional and land on their own state + both = initialize_state(problem, algorithm, 2.0, (; a = 1), (; b = 2)) + @test both.data == (; a = 1) + @test both.stopping_criterion_state.data == (; b = 2) + + # the keyword form forwards to the same implementation + @test initialize_state( + problem, algorithm; iterate = 2.0, state_data = (; a = 1), stopping_state_data = (; b = 2), + ).data == (; a = 1) + + # and it is inferable, so reaching for the positional form is a matter of taste + @test @inferred(initialize_state(problem, algorithm, 2.0)) isa DefaultState + + # an algorithm that neither provides a state type nor an iterate to start from is a + # mistake we should not paper over + @test_throws UndefKeywordError initialize_state(problem, algorithm) + + # a criterion with a state of its own still wins over the fallback + time_algorithm = AIT.DummyAlgorithm(StopAfter(Second(1))) + @test initialize_state( + problem, time_algorithm, 2.0, + ).stopping_criterion_state isa StopAfterTimePeriodState +end + +@testset "initialize_state! resets a DefaultState" begin + problem = AIT.DummyProblem() + algorithm = AIT.DummyAlgorithm(StopAfterIteration(5)) + state = initialize_state(problem, algorithm, 2.0, (; a = 1), (; b = 2)) + + state.iteration = 4 + state.stopping_criterion_state.at_iteration = 4 + + # returns the state itself, not the stopping criterion state it also resets + @test initialize_state!(problem, algorithm, state) === state + @test state.iteration == 0 + @test state.stopping_criterion_state.at_iteration == -1 + # the iterate and both data fields are kept when they are not provided + @test state.iterate == 2.0 + @test state.data == (; a = 1) + @test state.stopping_criterion_state.data == (; b = 2) + + # positionally: iterate, then the two data fields + initialize_state!(problem, algorithm, state, 3.0, (; a = 9), (; b = 8)) + @test state.iterate == 3.0 + @test state.data == (; a = 9) + @test state.stopping_criterion_state.data == (; b = 8) + + # the keyword form forwards to the same implementation + initialize_state!(problem, algorithm, state; iterate = 4.0, state_data = (; a = 7)) + @test state.iterate == 4.0 + @test state.data == (; a = 7) + @test state.stopping_criterion_state.data == (; b = 8) +end + +@testset "an algorithm without a state type of its own can be solved" begin + problem = HalvingProblem(1.0) + algorithm = Halving(StopAfterIteration(40)) + + # no state type, no initialize_state, no initialize_state!: only `step!` above + # `solve` reaches the default through its keywords + @test solve(problem, algorithm; iterate = 100.0) ≈ 1.0 + + state = initialize_state(problem, algorithm, 100.0) + @test solve!(problem, algorithm, state) ≈ 1.0 + @test state.iteration == 40 + + # and it can carry its own data along + counting = CountingHalving(StopAfterIteration(7)) + counting_state = initialize_state(problem, counting, 100.0, StepCounter(0)) + solve!(problem, counting, counting_state) + @test counting_state.data.steps == 7 +end diff --git a/test/stopping_criterion.jl b/test/stopping_criterion.jl index 7948fe3..f40ab4f 100644 --- a/test/stopping_criterion.jl +++ b/test/stopping_criterion.jl @@ -9,18 +9,10 @@ problem = AIT.DummyProblem() # --------------- # `StopAfterIteration` and `StopAfter` both have `indicates_convergence == false`, so neither can # exercise the convergence-versus-fallback logic. This one does. +# It also leaves its state to the default, so it exercises the `initialize_state` fallbacks. struct StopWhenConverged <: StoppingCriterion at::Int end -AlgorithmsInterface.initialize_state(::Problem, ::Algorithm, ::StopWhenConverged; kwargs...) = - DefaultStoppingCriterionState() -function AlgorithmsInterface.initialize_state!( - ::Problem, ::Algorithm, ::StopWhenConverged, - stopping_criterion_state::DefaultStoppingCriterionState; kwargs..., - ) - stopping_criterion_state.at_iteration = -1 - return stopping_criterion_state -end function AlgorithmsInterface.is_finished( ::Problem, ::Algorithm, state::State, stop_when_converged::StopWhenConverged, ::DefaultStoppingCriterionState, @@ -78,18 +70,9 @@ end AlgorithmsInterface.get_reason(::CountingCriterion, ::CountingCriterionState) = nothing AlgorithmsInterface.indicates_convergence(::Type{CountingCriterion}) = false -# Indicates to stop immediately, but implements nothing beyond the bare minimum: no `get_reason` -# and no `indicates_convergence`, so it exercises the fallbacks for both. +# Indicates to stop immediately, but implements nothing beyond the bare minimum: no `get_reason`, +# no `indicates_convergence` and no state of its own, so it exercises the fallbacks for all three. struct SilentCriterion <: StoppingCriterion end -AlgorithmsInterface.initialize_state(::Problem, ::Algorithm, ::SilentCriterion; kwargs...) = - DefaultStoppingCriterionState() -function AlgorithmsInterface.initialize_state!( - ::Problem, ::Algorithm, ::SilentCriterion, - stopping_criterion_state::DefaultStoppingCriterionState; kwargs..., - ) - stopping_criterion_state.at_iteration = -1 - return stopping_criterion_state -end AlgorithmsInterface.is_finished( ::Problem, ::Algorithm, ::State, ::SilentCriterion, ::DefaultStoppingCriterionState ) = true @@ -121,6 +104,47 @@ AlgorithmsInterface.is_active( ) = stopping_criterion_state.stopped AlgorithmsInterface.indicates_convergence(::Type{UnconventionalCriterion}) = true +@testset "DefaultStoppingCriterionState" begin + scs = DefaultStoppingCriterionState() + @test scs isa StoppingCriterionState + @test scs.at_iteration == -1 + @test scs.data === nothing + + # the data field is opaque, so anything goes and nothing of it is exposed as a property + with_data = DefaultStoppingCriterionState(Dict{Symbol, Any}(:last_change => 1.0)) + @test with_data.at_iteration == -1 + @test with_data.data[:last_change] == 1.0 + @test !hasproperty(with_data, :last_change) +end + +@testset "a criterion without a state of its own falls back to the default" begin + # `SilentCriterion` provides neither `initialize_state` nor `initialize_state!` + silent = SilentCriterion() + algorithm = AIT.DummyAlgorithm(silent) + + scs = initialize_state(problem, algorithm, silent) + @test scs isa DefaultStoppingCriterionState + @test scs.at_iteration == -1 + + scs.at_iteration = 3 + @test initialize_state!(problem, algorithm, silent, scs) === scs + @test scs.at_iteration == -1 + + # `stopping_state_data` seeds the state's data, and a reset keeps it unless given a new one + seeded = initialize_state(problem, algorithm, silent; stopping_state_data = (; tol = 1.0e-8)) + @test seeded.data == (; tol = 1.0e-8) + initialize_state!(problem, algorithm, silent, seeded) + @test seeded.data == (; tol = 1.0e-8) + initialize_state!(problem, algorithm, silent, seeded; stopping_state_data = (; tol = 1.0e-4)) + @test seeded.data == (; tol = 1.0e-4) + + # a criterion that does provide them keeps its own state type + counting = CountingCriterion() + @test initialize_state( + problem, AIT.DummyAlgorithm(counting), counting, + ) isa CountingCriterionState +end + @testset "StopAfterIteration" begin s1 = StopAfterIteration(2) @test s1 isa StoppingCriterion From 0fa7cff3bb016639b9034234a6149e834248fe9d Mon Sep 17 00:00:00 2001 From: lkdvos Date: Tue, 4 Aug 2026 14:27:20 -0400 Subject: [PATCH 2/2] rework implementation --- Changelog.md | 33 +++- Project.toml | 2 +- docs/src/interface.md | 78 ++++----- docs/src/logging.md | 14 +- docs/src/stopping_criterion.md | 85 ++++++---- src/AlgorithmsInterface.jl | 5 +- src/default_state.jl | 102 ----------- src/interface/interface.jl | 62 +++---- src/interface/state.jl | 82 +++++++-- src/interface/stopping_criterion.jl | 84 +++++++++ src/stopping_criterion.jl | 255 +++++++++------------------- src/test_suite.jl | 6 - test/logging.jl | 20 +-- test/newton.jl | 29 +--- test/state.jl | 82 ++++----- test/stopping_criterion.jl | 156 ++++++++++------- 16 files changed, 525 insertions(+), 570 deletions(-) delete mode 100644 src/default_state.jl create mode 100644 src/interface/stopping_criterion.jl diff --git a/Changelog.md b/Changelog.md index 6480992..4007b62 100644 --- a/Changelog.md +++ b/Changelog.md @@ -5,7 +5,7 @@ All notable Changes to the Julia package `AlgorithmsInterface.jl` are documented The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). -## [0.1.1] unreleased +## [0.2.0] unreleased ### Added @@ -16,21 +16,36 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Convenience two-argument `get_reason(algorithm, state)`, and likewise for `is_active`, `indicates_convergence` and `get_active_stopping_criteria`, extracting the criterion and its state the way `is_finished` already did. - `StopReasonAction`, a `LoggingAction` that reports `get_reason` at the `:Stop` context. - Defaults for `get_reason` (`nothing`) and for the type-domain `indicates_convergence` (`false`), so a criterion that implements neither no longer hits a `MethodError` from the derived convergence reporting. -- Exports for `DefaultStoppingCriterionState`, `StopAfterTimePeriodState` and `GroupStoppingCriterionState`, which a downstream criterion is expected to reuse. -- `DefaultState`, a `State` storing the `iterate`, `iteration` and `stopping_criterion_state` every state is expected to provide, next to a single `data` field for anything else an algorithm has to carry from one step to the next. - The `data` field is opaque to this package and none of its contents are exposed as properties, so an algorithm reaches them through `state.data`. -- Defaults for `initialize_state` and `initialize_state!` returning and resetting a `DefaultState`, so that an algorithm whose state holds nothing beyond the expected properties only has to provide a `Problem`, an `Algorithm` and a `step!`. - Both accept their arguments positionally as `(problem, algorithm, iterate, state_data, stopping_state_data)`, next to the keyword form the interface documents. -- Defaults for `initialize_state(problem, algorithm, stopping_criterion)` and its mutating variant, returning and resetting a `DefaultStoppingCriterionState`, with a `stopping_state_data` keyword to seed its `data`. +- `StopAfterTimePeriodData`, exported, holding the clock `StopAfter` carries as the `data` of its state, which a downstream criterion working on time measurements is expected to reuse. +- Defaults for `initialize_state` and `initialize_state!` returning and resetting a `State`, so that an algorithm only has to provide a `Problem`, an `Algorithm` and a `step!`. +- Defaults for `initialize_state(problem, algorithm, stopping_criterion)` and its mutating variant, returning and resetting a `StoppingCriterionState`, with a `stopping_state_data` keyword to seed its `data`. + Every criterion honours that keyword rather than hardcoding its data: `nothing`, the default, leaves the criterion to decide, which on a reset means keeping the data the state already carries. + `StopAfter` takes the `StopAfterTimePeriodData` it is handed and otherwise makes a fresh one, and `StopWhenAll` and `StopWhenAny` read it as the states of the criteria they combine, in the order of their `criteria`. ### Changed +- `State` is no longer an abstract type but the concrete state every algorithm runs with, storing the `iterate`, the `stopping_criterion_state` and the `iteration` next to a single `data` field for anything else an algorithm has to carry from one step to the next. + The `data` field is opaque to this package and none of its contents are exposed as properties, so an algorithm reaches them through `state.data`, which is what makes one state type enough for all of them. + An algorithm that used to define a state of its own now defines nothing at all, or, if it wants to decide where to start from rather than being handed an `iterate`, only an `initialize_state` returning a `State`. +- `increment!` takes the problem and the algorithm as well, as `increment!(problem, algorithm, state)`, so that per-iteration bookkeeping beyond the iteration counter remains overloadable now that it can no longer dispatch on a state of its own. +- `initialize_state` and `initialize_state!` take their arguments positionally, as `(problem, algorithm, iterate, state_data, stopping_state_data)` and `(problem, algorithm, state, iterate, state_data, stopping_state_data)`, rather than through `iterate`, `state_data` and `stopping_state_data` keywords. + Everything beyond the `iterate` is optional, and `initialize_state!` defaults every argument to what the `state` already holds, so an in-place reset only takes what actually changes. + `solve` and `solve!` forward their trailing positional arguments there, so `solve(problem, algorithm; iterate = x)` becomes `solve(problem, algorithm, x)`. + An algorithm that decides where it starts still implements `initialize_state(problem, algorithm; kwargs...)`, which is what `solve(problem, algorithm)` reaches, and which no longer has to declare an `iterate` keyword to do so. - `indicates_convergence` without a state moved to the type domain: a new criterion implements `indicates_convergence(::Type{YourCriterion})`, and `indicates_convergence(criterion)` forwards to it. `StopWhenAll` and `StopWhenAny` combine their children in the type domain as well, so a criterion that only implements the variant taking an instance is no longer accounted for in a group. -- `DefaultStoppingCriterionState` gained a type parameter and a `data` field for whatever a criterion has to remember from one iteration to the next. - It defaults to `nothing`, so `DefaultStoppingCriterionState()` is unaffected. +- `StoppingCriterionState` is no longer an abstract type but the concrete state every criterion runs with, storing the `at_iteration` at which the criterion indicated to stop next to a single `data` field for whatever else it has to remember from one iteration to the next. + Both states put their `data` last, after the iteration they record: `StoppingCriterionState(at_iteration, data)` and `State(iterate, stopping_criterion_state, iteration, data)`, each with a shorter form omitting that iteration, as `StoppingCriterionState(data = nothing)` and `State(iterate, stopping_criterion_state, data = nothing)`. + A criterion therefore dispatches `is_finished!`, `get_reason` and the rest on itself rather than on a state of its own, and defines `initialize_state` only to fill that `data` field. + `StopAfter` keeps its clock there as a `StopAfterTimePeriodData`, and `StopWhenAll` and `StopWhenAny` the states of the criteria they combine, as a tuple in the order of their `criteria`. - `StopAfterIteration` no longer carries its own `initialize_state` and `initialize_state!`, which the new criterion-level defaults now cover. +### Removed + +- The abstract `State`: subtyping it is no longer how an algorithm provides a state, since the concrete `State` serves every algorithm. +- `AlgorithmsInterface.Test.DummyState`, which the concrete `State` replaces. +- The abstract `StoppingCriterionState`, along with `DefaultStoppingCriterionState`, `StopAfterTimePeriodState` and `GroupStoppingCriterionState`, all of which the concrete `StoppingCriterionState` replaces. + ### Fixed - `indicates_convergence(::StoppingCriterion, ::StoppingCriterionState)` had its condition inverted, claiming convergence exactly when the criterion had *not* indicated to stop. diff --git a/Project.toml b/Project.toml index 003fe00..8726ea1 100644 --- a/Project.toml +++ b/Project.toml @@ -1,7 +1,7 @@ name = "AlgorithmsInterface" uuid = "d1e3940c-cd12-4505-8585-b0a4b322527d" authors = ["Ronny Bergmann ", "Lukas Devos "] -version = "0.1.1" +version = "0.2.0" [deps] Dates = "ade2ca70-3891-5945-98fb-dc099432e06a" diff --git a/docs/src/interface.md b/docs/src/interface.md index d51a180..2255ece 100644 --- a/docs/src/interface.md +++ b/docs/src/interface.md @@ -23,13 +23,16 @@ Bundling these together loosely without forcing one giant monolithic type makes * test algorithms in isolation by constructing minimal `Problem`/`Algorithm` pairs, and * extend behavior (add new stopping criteria, new logging events) without rewriting core loops. -The interface in this package formalizes those roles with three abstract types: +The interface in this package formalizes those roles with three types: * [`Problem`](@ref): immutable, algorithm‑agnostic input data. * [`Algorithm`](@ref): immutable configuration and parameters deciding how to iterate. * [`State`](@ref): mutable data that evolves (current iterate, caches, counters, diagnostics). +The first two are abstract, and every algorithm defines its own pair. +The third is a concrete type: since the data it carries is opaque to this package, one state serves them all. + It provides a framework for decomposing iterative methods into small, composable parts: -concrete `Problem`/`Algorithm`/`State` types have to implement a minimal set of core functionality, +concrete `Problem`/`Algorithm` types have to implement a minimal set of core functionality, and this package helps to stitch everything together and provide additional helper functionality such as stopping criteria and logging functionality. ## [Concrete example: Heron's method](@id sec_heron) @@ -62,19 +65,19 @@ struct HeronAlgorithm <: Algorithm end ``` -That leaves the third piece, the [`State`](@ref), which for most algorithms is entirely unremarkable: -it holds the three properties the interface asks for – the `iterate`, the `iteration` counter and the `stopping_criterion_state` – and nothing else. -This package therefore provides [`DefaultState`](@ref) for exactly that, so there is nothing to define here. -[Writing your own](@ref sec_custom_state) is the exception, covered at the end of this section. +That leaves the third piece, the [`State`](@ref). +It holds the `iterate`, the `iteration` counter and the `stopping_criterion_state`, next to a `data` field for whatever else an algorithm might carry from one step to the next. +Since that field is opaque to this package, the same state serves every algorithm, so there is nothing to define here. -The same goes for initialization: [`initialize_state`](@ref) and [`initialize_state!`](@ref), which respectively construct a fresh state and reset an existing one, both default to a [`DefaultState`](@ref). +The same goes for initialization: [`initialize_state`](@ref) and [`initialize_state!`](@ref), which respectively construct a fresh state and reset an existing one, both have a default. +An algorithm implements them only to [decide where to start from](@ref sec_custom_init), covered at the end of this section. ### Iteration steps Algorithms define a mutable step via [`step!`](@ref). For Heron's method: ```@example Heron -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState) +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::State) S = problem.S x = state.iterate state.iterate = 0.5 * (x + S / x) @@ -85,10 +88,10 @@ end Note that we are only focussing on the actual algorithm, and *not* incrementing the iteration counter. These kinds of bookkeeping should be handled by the [`AlgorithmsInterface.increment!`](@ref) function, which will by default already increment the iteration counter. The following generic functionality is therefore enough for our purposes, and does *not* need to be defined. -Nevertheless, if additional bookkeeping would be desired, this can be achieved by overloading that function: +Nevertheless, if additional bookkeeping would be desired, this can be achieved by overloading that function for your [`Algorithm`](@ref): ```julia -function AlgorithmsInterface.increment!(state::State) +function AlgorithmsInterface.increment!(problem::Problem, algorithm::Algorithm, state::State) state.iteration += 1 return state end @@ -105,13 +108,13 @@ With these definitions in place you can already run (assuming you also choose a function heron_sqrt(x; maxiter = 10) prob = SqrtProblem(x) alg = HeronAlgorithm(StopAfterIteration(maxiter)) - return solve(prob, alg; iterate = 1.0) # allocates & runs + return solve(prob, alg, 1.0) # allocates & runs end println("Approximate sqrt: ", heron_sqrt(16.0)) ``` -The `iterate` keyword is where the default initialization gets its starting point from, and it is required: there is no sensible guess this package could make on an algorithm's behalf. +The third argument is where the default initialization gets its starting point from, and it is required: there is no sensible guess this package could make on an algorithm's behalf. Note that [`solve`](@ref) will default to returning `state.iterate`. If desired, this can be customized by altering [`finalize_state!`](@ref). @@ -119,8 +122,8 @@ We will refine this example with better halting logic and logging shortly. ### Carrying extra data -Anything an algorithm has to carry from one step to the next beyond those three properties goes into the `data` field of the [`DefaultState`](@ref), which this package treats as opaque: -it is passed along and reset alongside the rest of the state, but nothing is ever read from it, so an algorithm picks whatever container suits it. +Anything an algorithm has to carry from one step to the next beyond those three properties goes into the `data` field of the [`State`](@ref). +It is passed along and reset alongside the rest of the state, but nothing is ever read from it, so an algorithm picks whatever container suits it. Prefer a `NamedTuple` or a small struct of your own, which keep access to the contents type stable: ```@example Heron @@ -128,7 +131,7 @@ struct TrackingHeronAlgorithm <: Algorithm stopping_criterion end -function AlgorithmsInterface.step!(problem::SqrtProblem, ::TrackingHeronAlgorithm, state::DefaultState) +function AlgorithmsInterface.step!(problem::SqrtProblem, ::TrackingHeronAlgorithm, state::State) S = problem.S x = state.iterate state.iterate = 0.5 * (x + S / x) @@ -146,45 +149,28 @@ state.data.residuals Note that the `NamedTuple` itself is immutable while the vector inside it is not, which is all this needs: the state keeps handing the same container back, and the algorithm mutates its contents. Here the initial iterate and the data are passed positionally, as `initialize_state(problem, algorithm, iterate, state_data)`; the [`StoppingCriterionState`](@ref) can be given its own data as a further argument. -### [Defining your own state](@id sec_custom_state) +### [Deciding where to start](@id sec_custom_init) -Reach for a state type of your own once `data` grows past a handful of entries, or once the algorithm wants to dispatch on the state itself. -It has to provide the three expected properties, and with it come the two initialization methods that the default was providing until now: +The one thing the defaults cannot know is where an algorithm should start, which is why that argument has no default. +An algorithm that does know implements the variant without it, supplying that starting point and leaving the rest of the work to the generic methods: ```@example Heron -mutable struct HeronState <: State - iterate::Float64 # current iterate - iteration::Int # current iteration count - stopping_criterion_state # any StoppingCriterionState -end - -function AlgorithmsInterface.initialize_state(problem::SqrtProblem, algorithm::HeronAlgorithm; kwargs...) - x0 = rand() # random initial guess - stopping_criterion_state = initialize_state(problem, algorithm, algorithm.stopping_criterion) - return HeronState(x0, 0, stopping_criterion_state) -end - -function AlgorithmsInterface.initialize_state!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState; kwargs...) - # reset the state for the algorithm - state.iterate = rand() - state.iteration = 0 - - # reset the state for the stopping criterion - AlgorithmsInterface.initialize_state!( - problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state +function AlgorithmsInterface.initialize_state( + problem::SqrtProblem, algorithm::HeronAlgorithm; kwargs... ) - return state + return initialize_state(problem, algorithm, rand(); kwargs...) end -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::HeronState) - S = problem.S - x = state.iterate - state.iterate = 0.5 * (x + S / x) - return state +function AlgorithmsInterface.initialize_state!( + problem::SqrtProblem, algorithm::HeronAlgorithm, state::State; kwargs... + ) + return initialize_state!(problem, algorithm, state, rand(); kwargs...) end ``` -This `HeronAlgorithm` now determines its own initial iterate rather than being handed one, so it runs without the `iterate` keyword: +Note how both hand their guess to the generic method, which is where the [`State`](@ref) is allocated and reset. +Neither shadows that method: a caller who does pass an iterate reaches it directly and gets to say where the algorithm starts after all. +This `HeronAlgorithm` now falls back to an initial guess of its own rather than insisting on being handed one, so it runs with nothing but a problem and a criterion: ```@example Heron println("Approximate sqrt: ", solve(SqrtProblem(16.0), HeronAlgorithm(StopAfterIteration(10)))) @@ -223,7 +209,7 @@ Private = true ```@autodocs Modules = [AlgorithmsInterface] -Pages = ["interface/state.jl", "default_state.jl"] +Pages = ["interface/state.jl"] Order = [:type, :function] Private = true ``` diff --git a/docs/src/logging.md b/docs/src/logging.md index ae6830a..8b1eb2f 100644 --- a/docs/src/logging.md +++ b/docs/src/logging.md @@ -46,7 +46,7 @@ struct HeronAlgorithm <: Algorithm stopping_criterion end -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState) +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::State) S = problem.S x = state.iterate state.iterate = 0.5 * (x + S / x) @@ -56,12 +56,12 @@ end function heron_sqrt(x; stopping_criterion = StopAfterIteration(10)) prob = SqrtProblem(x) alg = HeronAlgorithm(stopping_criterion) - return solve(prob, alg; iterate = 1.0) # allocates & runs + return solve(prob, alg, 1.0) # allocates & runs end nothing # hide ``` -Note that this leaves the state to [`DefaultState`](@ref), as the [interface section](@ref sec_interface) does, so there is neither a state type nor an `initialize_state` to be seen here. +Note that this leaves the initialization to its defaults, as the [interface section](@ref sec_interface) does, so there is no `initialize_state` to be seen here and the starting iterate is handed to [`solve`](@ref). It is already interesting to note that there are no further modifications necessary to start leveraging the logging system. @@ -194,7 +194,7 @@ function AlgorithmsInterface.handle_message!( action::CaptureHistory, problem::SqrtProblem, algorithm::HeronAlgorithm, - state::DefaultState; + state::State; kwargs... ) push!(action.iterates, state.iterate) @@ -270,7 +270,7 @@ end StatsCollector() = StatsCollector(0, 0.0, 0.0) function AlgorithmsInterface.handle_message!( - action::StatsCollector, problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState; + action::StatsCollector, problem::SqrtProblem, algorithm::HeronAlgorithm, state::State; kwargs... ) action.count += 1 @@ -317,7 +317,7 @@ function solve!(problem::Problem, algorithm::Algorithm, state::State; kwargs...) while !is_finished!(problem, algorithm, state) emit_message(problem, algorithm, state, :PreStep) - increment!(state) + increment!(problem, algorithm, state) step!(problem, algorithm, state) emit_message(problem, algorithm, state, :PostStep) @@ -398,7 +398,7 @@ Here we will illustrate this by a slight adaptation of our algorithm, which coul To emit a custom logging event from within your algorithm, call [`emit_message`](@ref): ```@example Heron -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState) +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::State) # Suppose we check for numerical issues if !isfinite(state.iterate) || mod(state.iteration, 10) == 0 emit_message(problem, algorithm, state, :Restart) diff --git a/docs/src/stopping_criterion.md b/docs/src/stopping_criterion.md index 00deb5d..a6e379e 100644 --- a/docs/src/stopping_criterion.md +++ b/docs/src/stopping_criterion.md @@ -51,23 +51,23 @@ This happens through the same initialization functions we used for initializing - [`initialize_state`](@ref) constructs an entirely new stopping state for the algorithm - [`initialize_state!`](@ref) (in-place) reset of an existing stopping state. -Since we leave the state to [`DefaultState`](@ref), this is already taken care of: the defaults pair it with the state of the algorithm's own criterion, along the lines of +Since we leave the initialization to its defaults, this is already taken care of: they pair the [`State`](@ref) with the state of the algorithm's own criterion, along the lines of ```julia -function AlgorithmsInterface.initialize_state(problem::Problem, algorithm::Algorithm; iterate, kwargs...) +function AlgorithmsInterface.initialize_state(problem::Problem, algorithm::Algorithm, iterate; kwargs...) stopping_criterion_state = initialize_state(problem, algorithm, algorithm.stopping_criterion; kwargs...) - return DefaultState(iterate, stopping_criterion_state) + return State(iterate, stopping_criterion_state) end ``` -A state of your own is where you would write that pairing out yourself. +An [`initialize_state`](@ref) of your own only decides which `iterate` to pass along, and can leave that pairing to the method above. ### Iteration During the iteration procedure, as set out by our design principles, we do not have to modify any of the code, and the stopping criteria do not show up: ```@example Heron -function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::DefaultState) +function AlgorithmsInterface.step!(problem::SqrtProblem, algorithm::HeronAlgorithm, state::State) S = problem.S x = state.iterate state.iterate = 0.5 * (x + S / x) @@ -79,14 +79,14 @@ What is really going on is that behind the scenes, the loop of the iterative sol ```julia while !is_finished!(problem, algorithm, state) - increment!(state) + increment!(problem, algorithm, state) step!(problem, algorithm, state) end ``` In other words, all of the logic is handled by the [`is_finished!`](@ref) function. The generic stopping criteria provided by this package have default implementations for this function that work out-of-the-box. -This is partially because everything is reached under conventional names: `Algorithm` assumes the existence of `stopping_criterion`, while `State` assumes `iterate` and `iteration` and `stopping_criterion_state` to exist — which is exactly what [`DefaultState`](@ref) provides, and what a state of your own has to provide too. +This is partially because everything is reached under conventional names: the generic code assumes an `Algorithm` to carry a `stopping_criterion`, and reads the `iterate`, the `iteration` and the `stopping_criterion_state` of the [`State`](@ref). ### Running the algorithm @@ -96,7 +96,7 @@ We can again combine everything into a single function, but now make the stoppin function heron_sqrt(x; stopping_criterion) prob = SqrtProblem(x) alg = HeronAlgorithm(stopping_criterion) - return solve(prob, alg; iterate = 1.0) # allocates & runs + return solve(prob, alg, 1.0) # allocates & runs end heron_sqrt(2; stopping_criterion = StopAfterIteration(10)) @@ -132,24 +132,22 @@ Suppose we want to stop when successive iterates change by less than `ϵ`, we co In order to do so, we need to define our own structs and implement the required interface. Again, we split up the data into a _static_ part, the [`StoppingCriterion`](@ref), and a _dynamic_ part, the [`StoppingCriterionState`](@ref). -The dynamic part is usually not yours to write: [`DefaultStoppingCriterionState`](@ref) records the iteration at which the criterion triggered and carries a `data` field for anything else it has to remember, which covers most criteria. -We spell out a state of our own here because it shows the full picture, and because a criterion that wants its fields named and typed is exactly the case that calls for one. +The dynamic part is not yours to write: [`StoppingCriterionState`](@ref) records the iteration at which the criterion triggered, and carries a `data` field for anything else it has to remember. +What a criterion does define is that `data` — here the previous iterate to compare against, and the change we measured, which we keep around for the message: ```@example Heron struct StopWhenStable <: StoppingCriterion tol::Float64 # when do we consider things converged end -mutable struct StopWhenStableState <: StoppingCriterionState +mutable struct StableData previous_iterate::Float64 # previous value to compare to - at_iteration::Int # iteration at which stability was reached delta::Float64 # difference between the values end ``` -Note that our mutable state holds both the `previous_iterate`, which we need to compare to, -as well as the iteration at which the condition was satisfied. -This is not strictly necessary, but can be convenient to have a persistent indication that convergence was reached. +A mutable struct is the recommended container: the state keeps handing the same one back, and the criterion updates its fields in place. +Note that the iteration at which the condition was satisfied is *not* part of it — that is the `at_iteration` the state records for every criterion. ### Initialization @@ -158,20 +156,33 @@ This could be implemented as follows: ```@example Heron function AlgorithmsInterface.initialize_state(::Problem, ::Algorithm, c::StopWhenStable; kwargs...) - return StopWhenStableState(NaN, -1, NaN) + return StoppingCriterionState(StableData(NaN, NaN)) end function AlgorithmsInterface.initialize_state!( - ::Problem, ::Algorithm, stop_when::StopWhenStable, st::StopWhenStableState; + ::Problem, ::Algorithm, stop_when::StopWhenStable, st::StoppingCriterionState; kwargs... ) - st.previous_iterate = NaN + st.data.previous_iterate = NaN + st.data.delta = NaN st.at_iteration = -1 - st.delta = NaN return st end ``` +Both are handed a `stopping_state_data` keyword, which is `nothing` unless a caller wants to supply the container itself. +Ignoring it, as above, means the criterion always uses one of its own; a criterion that lets it through takes what it is given instead, the way [`StopAfter`](@ref) does with its clock: + +```julia +function AlgorithmsInterface.initialize_state( + ::Problem, ::Algorithm, c::StopWhenStable; stopping_state_data = nothing, kwargs... +) + return StoppingCriterionState( + isnothing(stopping_state_data) ? StableData(NaN, NaN) : stopping_state_data + ) +end +``` + ### Checking for convergence Then, we need to implement the logic that checks whether an algorithm has finished, which is achieved through [`is_finished`](@ref) and [`is_finished!`](@ref). @@ -179,18 +190,18 @@ Here, the mutating version alters the `stopping_criterion_state`, and should the ```@example Heron function AlgorithmsInterface.is_finished!( - ::Problem, ::Algorithm, state::State, c::StopWhenStable, st::StopWhenStableState + ::Problem, ::Algorithm, state::State, c::StopWhenStable, st::StoppingCriterionState ) k = state.iteration if k == 0 - st.previous_iterate = state.iterate + st.data.previous_iterate = state.iterate st.at_iteration = -1 return false end - st.delta = abs(state.iterate - st.previous_iterate) - st.previous_iterate = state.iterate - if st.delta < c.tol + st.data.delta = abs(state.iterate - st.data.previous_iterate) + st.data.previous_iterate = state.iterate + if st.data.delta < c.tol st.at_iteration = k return true end @@ -198,12 +209,12 @@ function AlgorithmsInterface.is_finished!( end function AlgorithmsInterface.is_finished( - ::Problem, ::Algorithm, state::State, c::StopWhenStable, st::StopWhenStableState + ::Problem, ::Algorithm, state::State, c::StopWhenStable, st::StoppingCriterionState ) k = state.iteration k == 0 && return false - Δ = abs(state.iterate - st.previous_iterate) + Δ = abs(state.iterate - st.data.previous_iterate) return Δ < c.tol end ``` @@ -217,21 +228,21 @@ There are two separate questions here, and keeping them apart is what makes the * Why, in words? This is answered by [`get_reason`](@ref), and it is for human consumption only. We get the first one for free. -The default implementation of [`is_active`](@ref) reads the `at_iteration` property of the state, and our `StopWhenStableState` has one, following the convention that a negative value means "has not (yet) indicated to stop". -Only a state that records its status some other way has to implement [`is_active`](@ref) itself. +The default implementation of [`is_active`](@ref) reads the `at_iteration` property of the state, which every criterion has, following the convention that a negative value means "has not (yet) indicated to stop". +Only a criterion that records its status in its `data` instead has to implement [`is_active`](@ref) itself. That leaves the message, plus the static statement that meeting this criterion *does* mean convergence: ```@example Heron -function AlgorithmsInterface.get_reason(c::StopWhenStable, st::StopWhenStableState) +function AlgorithmsInterface.get_reason(c::StopWhenStable, st::StoppingCriterionState) is_active(c, st) || return nothing - return "The algorithm reached an approximate stable point after $(st.at_iteration) iterations; the change $(st.delta) is less than $(c.tol).\n" + return "The algorithm reached an approximate stable point after $(st.at_iteration) iterations; the change $(st.data.delta) is less than $(c.tol).\n" end AlgorithmsInterface.indicates_convergence(::Type{StopWhenStable}) = true ``` -Note that `get_reason` gates on the *recorded* status rather than re-checking `st.delta < c.tol`. +Note that `get_reason` gates on the *recorded* status rather than re-checking `st.data.delta < c.tol`. Re-checking the predicate would make the message disappear again as soon as the state moves on, whereas `at_iteration` is a permanent record of what happened. Only the type-domain [`indicates_convergence`](@ref) needs to be defined. @@ -316,16 +327,16 @@ Implementing a criterion means defining: The `(criterion,)` and the `(criterion, criterion_state)` variant are derived from this one and do not need to be defined. The state is taken care of for you. -[`initialize_state`](@ref) and [`initialize_state!`](@ref) return and reset a [`DefaultStoppingCriterionState`](@ref), which records the `at_iteration` that all the reporting is built on and carries a `data` field for whatever else the criterion has to remember — the `previous_iterate` and `delta` above, for instance. +[`initialize_state`](@ref) and [`initialize_state!`](@ref) return and reset a [`StoppingCriterionState`](@ref), which records the `at_iteration` that all the reporting is built on and carries a `data` field for whatever else the criterion has to remember. -Only a criterion that is not served by that state, as the one above wanted its fields named and typed, additionally defines: +Only a criterion that needs such data, as the one above needed a `previous_iterate` to compare against, additionally defines: -* A state subtype of [`StoppingCriterionState`](@ref) capturing its dynamic fields, including an `at_iteration` recording when the criterion triggered. -* `initialize_state` and `initialize_state!` for its setup and reset. +* A container for it, reached as `stopping_criterion_state.data`. +* `initialize_state` and `initialize_state!` to fill and reset that container. You may also implement `Base.summary(io, criterion, criterion_state)` for compact status reports, -and `is_active(criterion, criterion_state)` if your state does not record its status in an -`at_iteration` property. +and `is_active(criterion, criterion_state)` if your criterion does not record its status in the +state's `at_iteration` property. ## Reference API diff --git a/src/AlgorithmsInterface.jl b/src/AlgorithmsInterface.jl index 3626591..fb71e1e 100644 --- a/src/AlgorithmsInterface.jl +++ b/src/AlgorithmsInterface.jl @@ -14,18 +14,17 @@ using ScopedValues include("interface/algorithm.jl") include("interface/problem.jl") +include("interface/stopping_criterion.jl") include("interface/state.jl") include("interface/interface.jl") include("stopping_criterion.jl") -include("default_state.jl") include("logging.jl") include("test_suite.jl") # general interface export Algorithm, Problem, State -export DefaultState export initialize_state, initialize_state! export finalize_state! @@ -34,7 +33,7 @@ export step!, solve, solve!, solve_loop! # stopping criteria export StoppingCriterion, StoppingCriterionState export StopAfter, StopAfterIteration, StopWhenAll, StopWhenAny -export DefaultStoppingCriterionState, StopAfterTimePeriodState, GroupStoppingCriterionState +export StopAfterTimePeriodData export is_finished, is_finished!, get_reason, is_active, indicates_convergence export get_active_stopping_criteria diff --git a/src/default_state.jl b/src/default_state.jl deleted file mode 100644 index 2ca3c79..0000000 --- a/src/default_state.jl +++ /dev/null @@ -1,102 +0,0 @@ -# -# -# A default state - -@doc """ - DefaultState <: State - -A [`State`](@ref) that stores the properties every [`State`](@ref) is expected to provide, -together with a single field to hold any further data an [`Algorithm`](@ref) needs. - -Together with the default [`initialize_state`](@ref) and [`initialize_state!`](@ref) methods, -this spares a downstream algorithm the definition of its own state type: -a [`Problem`](@ref), an [`Algorithm`](@ref) and a [`step!`](@ref) method suffice. -Since [`step!`](@ref) dispatches on the [`Algorithm`](@ref), reusing this state costs no -flexibility in doing so. - -# Fields - -* `iterate` stores the current iterate ``x^{(k)}``. -* `stopping_criterion_state` stores the [`StoppingCriterionState`](@ref) belonging to the - [`StoppingCriterion`](@ref) of the [`Algorithm`](@ref). -* `data` stores any further data that has to be carried from one step to the next. - It is opaque to this package, and none of its contents are exposed as properties of the - state, so an algorithm reaches them through `state.data`. - A `NamedTuple` or a struct of its own keeps access to them type stable and is the - recommended choice; `nothing` indicates that the algorithm needs no further data. -* `iteration::Int` stores the current iteration step ``k`` that is currently being performed - or was last performed. - -# Constructor - - DefaultState(iterate, stopping_criterion_state, data = nothing, iteration = 0) - -Initialize the state to start at `iterate`, carrying `data` alongside it. -""" -mutable struct DefaultState{V, S <: StoppingCriterionState, D} <: State - iterate::V - stopping_criterion_state::S - data::D - iteration::Int -end - -DefaultState(iterate, stopping_criterion_state::StoppingCriterionState, data = nothing) = - DefaultState(iterate, stopping_criterion_state, data, 0) - -# The default keeps whatever data a criterion state already carries, rather than clearing it. -# A criterion state that carries none takes `nothing`, which its `initialize_state!` ignores. -_stopping_state_data(::StoppingCriterionState) = nothing -_stopping_state_data(stopping_criterion_state::DefaultStoppingCriterionState) = - stopping_criterion_state.data - -# The signature taking an `iterate` is the most generic one there is, so this doubles as the -# fallback for any `Problem` and `Algorithm` that do not provide a state type of their own. -# Both these and their keyword counterparts below are documented by `_doc_init_state`, which -# covers exactly these signatures. -function initialize_state( - problem::Problem, algorithm::Algorithm, iterate, - state_data = nothing, stopping_state_data = nothing; - kwargs..., - ) - stopping_criterion_state = initialize_state( - problem, algorithm, algorithm.stopping_criterion; - stopping_state_data, kwargs..., - ) - return DefaultState(iterate, stopping_criterion_state, state_data) -end - -function initialize_state( - problem::Problem, algorithm::Algorithm; - iterate, state_data = nothing, stopping_state_data = nothing, kwargs..., - ) - return initialize_state( - problem, algorithm, iterate, state_data, stopping_state_data; kwargs... - ) -end - -function initialize_state!( - problem::Problem, algorithm::Algorithm, state::DefaultState, iterate, - state_data = state.data, - stopping_state_data = _stopping_state_data(state.stopping_criterion_state); - kwargs..., - ) - state.iterate = iterate - state.data = state_data - state.iteration = 0 - initialize_state!( - problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state; - stopping_state_data, kwargs..., - ) - return state -end - -function initialize_state!( - problem::Problem, algorithm::Algorithm, state::DefaultState; - iterate = state.iterate, state_data = state.data, - stopping_state_data = _stopping_state_data(state.stopping_criterion_state), - kwargs..., - ) - return initialize_state!( - problem, algorithm, state, iterate, state_data, stopping_state_data; kwargs... - ) -end diff --git a/src/interface/interface.jl b/src/interface/interface.jl index 55f4558..bb659e3 100644 --- a/src/interface/interface.jl +++ b/src/interface/interface.jl @@ -1,31 +1,31 @@ _doc_init_state = """ - state = initialize_state(problem::Problem, algorithm::Algorithm; kwargs...) - state = initialize_state!(problem::Problem, algorithm::Algorithm, state::State; kwargs...) - -Initialize a [`State`](@ref) based on a [`Problem`](@ref) and an [`Algorithm`](@ref). -The `kwargs...` should allow to initialize for example the initial point. -This can be done in-place for `state`, then only values that did change have to be provided. + state = initialize_state(problem::Problem, algorithm::Algorithm, iterate, state_data, stopping_state_data; kwargs...) + state = initialize_state!(problem::Problem, algorithm::Algorithm, state::State, iterate, state_data, stopping_state_data; kwargs...) -Both have a default in terms of a [`DefaultState`](@ref), which stores the properties every -[`State`](@ref) is expected to provide next to a single `data` field for anything else an -algorithm needs, so an algorithm that is served by such a state does not have to provide a -state type of its own. -These defaults also take their arguments positionally: +Allocate a [`State`](@ref) to start at `iterate`, or reset an existing one to start there again. +All arguments beyond the `iterate` are optional, and `initialize_state!` defaults every one of +them to what the `state` already holds, so resetting in place only takes what actually changes. - state = initialize_state(problem, algorithm, iterate, state_data, stopping_state_data; kwargs...) - state = initialize_state!(problem, algorithm, state, iterate, state_data, stopping_state_data; kwargs...) - -Here `state_data` is what the [`DefaultState`](@ref) carries and `stopping_state_data` what the -[`StoppingCriterionState`](@ref) does, both defaulting to `nothing` when allocating and to the -values the `state` already holds when resetting. +`state_data` is what the [`State`](@ref) carries and `stopping_state_data` what the +[`StoppingCriterionState`](@ref) does. +For the latter, `nothing` leaves the [`StoppingCriterion`](@ref) to decide which data to carry, +so it is only passed to override that decision. The remaining `kwargs...` are passed on to the corresponding function for the [`StoppingCriterion`](@ref) of the [`Algorithm`](@ref). + +An [`Algorithm`](@ref) that knows where it starts implements the variant without an `iterate`, +which is also what [`solve`](@ref) reaches when it is called without one: + + state = initialize_state(problem::Problem, algorithm::Algorithm; kwargs...) + +Since [`State`](@ref) serves every algorithm, such an implementation only picks that starting +point and hands it to the method above. """ function initialize_state end @doc "$(_doc_init_state)" -initialize_state(::Problem, ::Algorithm; kwargs...) +initialize_state(::Problem, ::Algorithm, iterate; kwargs...) function initialize_state! end @@ -44,22 +44,22 @@ finalize_state!(problem::Problem, algorithm::Algorithm, state::State) = state.it # has to be defined before used in solve but is documented alphabetically after @doc """ - solve(problem::Problem, algorithm::Algorithm; kwargs...) + solve(problem::Problem, algorithm::Algorithm, args...; kwargs...) Solve the [`Problem`](@ref) using an [`Algorithm`](@ref). -The keyword arguments `kwargs...` have to provide enough details such that -the corresponding state initialisation [`initialize_state`](@ref)`(problem, algorithm)` -returns a state. +The `args...` are passed on to [`initialize_state`](@ref)`(problem, algorithm, args...)`, and +have to provide enough detail for it to return a state: the `iterate` to start at, unless the +`algorithm` implements an [`initialize_state`](@ref) that decides that itself. By default this method continues to call [`solve!`](@ref). """ -function solve(problem::Problem, algorithm::Algorithm; kwargs...) +function solve(problem::Problem, algorithm::Algorithm, args...; kwargs...) # obtain logger once to minimize overhead from accessing ScopedValue # additionally handle logging initialization to enable stateful LoggingAction logger = algorithm_logger() # initialize the state and emit message - state = initialize_state(problem, algorithm; kwargs...) + state = initialize_state(problem, algorithm, args...; kwargs...) emit_message(logger, problem, algorithm, state, :Start) # main loop @@ -72,20 +72,22 @@ function solve(problem::Problem, algorithm::Algorithm; kwargs...) end @doc """ - solve!(problem::Problem, algorithm::Algorithm, state::State; kwargs...) + solve!(problem::Problem, algorithm::Algorithm, state::State, args...; kwargs...) Solve the [`Problem`](@ref) using an [`Algorithm`](@ref), starting from a given [`State`](@ref). The state is modified in-place. -All keyword arguments are passed to the [`initialize_state!`](@ref)`(problem, algorithm, state)` function. +The `args...` and all keyword arguments are passed to the +[`initialize_state!`](@ref)`(problem, algorithm, state, args...)` function, which resets the +`state` and defaults everything that is not given to what it already holds. """ -function solve!(problem::Problem, algorithm::Algorithm, state::State; kwargs...) +function solve!(problem::Problem, algorithm::Algorithm, state::State, args...; kwargs...) # obtain logger once to minimize overhead from accessing ScopedValue # additionally handle logging initialization to enable stateful LoggingAction logger = algorithm_logger() # initialize the state and emit message - initialize_state!(problem, algorithm, state; kwargs...) + initialize_state!(problem, algorithm, state, args...; kwargs...) emit_message(logger, problem, algorithm, state, :Start) # main loop @@ -104,7 +106,7 @@ Provide the main loop of the iterative `algorithm` for a given `problem` and sta This loop consists of: 1. Checking for convergence with [`is_finished!`](@ref) -2. Incrementing the state [`increment!`](@ref) +2. Incrementing the state [`increment!`](@ref)`(problem, algorithm, state)` 3. Performing a step [`step!`](@ref) 4. Repeat """ @@ -112,7 +114,7 @@ function solve_loop!(problem::Problem, algorithm::Algorithm, state::State) logger = algorithm_logger() while !is_finished!(problem, algorithm, state) emit_message(logger, problem, algorithm, state, :PreStep) - increment!(state) + increment!(problem, algorithm, state) step!(problem, algorithm, state) emit_message(logger, problem, algorithm, state, :PostStep) end diff --git a/src/interface/state.jl b/src/interface/state.jl index 0d61775..d0ad09d 100644 --- a/src/interface/state.jl +++ b/src/interface/state.jl @@ -1,37 +1,87 @@ @doc """ State -An abstract type to represent the state an iterative algorithm is in. +The state an iterative algorithm is in. The state consists of any information that describes the current step the algorithm is in and keeps all information needed from one step to the next. +Since the `data` field is opaque to this package, this single type serves every +[`Algorithm`](@ref): a [`Problem`](@ref), an [`Algorithm`](@ref) and a [`step!`](@ref) method +suffice to run one, and since [`step!`](@ref) dispatches on the [`Algorithm`](@ref), sharing +one state type across algorithms costs no flexibility in doing so. -## Properties +# Fields -In order to interact with the stopping criteria, the state should contain the following properties, -and provide corresponding `getproperty` and `setproperty!` methods. +* `iterate` stores the current iterate ``x^{(k)}``. +* `stopping_criterion_state` stores the [`StoppingCriterionState`](@ref) belonging to the + [`StoppingCriterion`](@ref) of the [`Algorithm`](@ref), indicating whether the + [`Algorithm`](@ref) will stop after this iteration or has stopped. +* `iteration::Int` stores the current iteration step ``k`` that is currently being performed + or was last performed. +* `data` stores any further data that has to be carried from one step to the next. + It is opaque to this package, and none of its contents are exposed as properties of the + state, so an algorithm reaches them through `state.data`. + A `NamedTuple` or a struct of its own keeps access to them type stable and is the + recommended choice; `nothing` indicates that the algorithm needs no further data. -* `iteration` – the current iteration step ``k`` that is currently being performed or was last performed. -* `stopping_criterion_state` – a [`StoppingCriterionState`](@ref) that indicates whether an [`Algorithm`](@ref) - will stop after this iteration or has stopped. -* `iterate` – the current iterate ``x^{(k)}``. +# Constructor -## Methods + State(iterate, stopping_criterion_state, data = nothing) + State(iterate, stopping_criterion_state, iteration, data) -The following methods should be implemented for a state - -* [`increment!`](@ref)(state) +Initialize the state to start at `iterate`, carrying `data` alongside it. +The second form is the one that also sets the iteration it records. """ -abstract type State end +mutable struct State{V, S <: StoppingCriterionState, D} + iterate::V + stopping_criterion_state::S + iteration::Int + data::D +end + +State(iterate, stopping_criterion_state::StoppingCriterionState, data = nothing) = + State(iterate, stopping_criterion_state, 0, data) """ - increment!(state::State) + increment!(problem::Problem, algorithm::Algorithm, state::State) Increment the current iteration that a [`State`](@ref) is currently performing or was last performing. -The default assumes that the current iteration is stored in `state.iteration`. +The default increments `state.iteration`, which is all the bookkeeping this package needs. +An [`Algorithm`](@ref) that has more of it to do overloads this function. """ -function increment!(state::State) +function increment!(::Problem, ::Algorithm, state::State) state.iteration += 1 return state end + +# These two are documented by `_doc_init_state`, which covers exactly these signatures. +# The one taking an `iterate` is the most generic there is, so it doubles as the fallback for any +# `Problem` and `Algorithm` that do not initialize a state of their own. An algorithm that does +# know where to start implements the variant without an `iterate` and leaves the rest to these. +function initialize_state( + problem::Problem, algorithm::Algorithm, iterate, + state_data = nothing, stopping_state_data = nothing; + kwargs..., + ) + stopping_criterion_state = initialize_state( + problem, algorithm, algorithm.stopping_criterion; + stopping_state_data, kwargs..., + ) + return State(iterate, stopping_criterion_state, state_data) +end + +function initialize_state!( + problem::Problem, algorithm::Algorithm, state::State, + iterate = state.iterate, state_data = state.data, stopping_state_data = nothing; + kwargs..., + ) + state.iterate = iterate + state.data = state_data + state.iteration = 0 + initialize_state!( + problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state; + stopping_state_data, kwargs..., + ) + return state +end diff --git a/src/interface/stopping_criterion.jl b/src/interface/stopping_criterion.jl new file mode 100644 index 0000000..d3514f2 --- /dev/null +++ b/src/interface/stopping_criterion.jl @@ -0,0 +1,84 @@ +@doc """ + StoppingCriterion + +An abstract type to represent a stopping criterion of an [`Algorithm`](@ref). + +A concrete [`StoppingCriterion`](@ref) is the static half of a stopping criterion, holding the +values it is configured with, while the [`StoppingCriterionState`](@ref) it is paired with +records what happens during a run. + +It should usually implement + +* [`is_finished!`](@ref)`(problem, algorithm, state, stopping_criterion, stopping_criterion_state)` +* [`is_finished`](@ref)`(problem, algorithm, state, stopping_criterion, stopping_criterion_state)` +* [`get_reason`](@ref)`(stopping_criterion, stopping_criterion_state)` +* [`indicates_convergence`](@ref)`(::Type{<:StoppingCriterion})` + +Note that only [`indicates_convergence`](@ref) has to be implemented: +it answers whether meeting this criterion *would* mean convergence, which is a static property of the criterion type alone. +Both the variant taking a criterion and the one that additionally takes a [`StoppingCriterionState`](@ref), answering whether it *did* happen are derived from it. + +A criterion that has to remember more than the iteration at which it indicated to stop +additionally implements +[`initialize_state(problem::Problem, algorithm::Algorithm, stopping_criterion::StoppingCriterion; kwargs...)`](@ref) +to fill the `data` field of its state, as well as the corresponding mutating variant to reset it. +""" +abstract type StoppingCriterion end + +@doc """ + StoppingCriterionState + +The state a [`StoppingCriterion`](@ref) is in within a [`State`](@ref). + +It records the iteration at which its [`StoppingCriterion`](@ref) indicated to stop, next to a +single field for anything else that criterion has to remember from one iteration to the next. +Since that field is opaque to this package, this single type serves every +[`StoppingCriterion`](@ref), which therefore dispatches its methods on itself rather than on a +state of its own. + +# Fields + +* `at_iteration::Int` stores the iteration number at which this state indicated to stop. + * `0` means it already indicated to stop at the start. + * any negative number means that it has not yet indicated to stop. +* `data` stores any further data the criterion has to carry from one iteration to the next, + for example a value it compares against in the next one. + It is opaque to this package, and none of its contents are exposed as properties of the + state, so a criterion reaches them through `stopping_criterion_state.data`. + A mutable struct of its own is the recommended choice; `nothing`, the default, indicates + that the criterion needs no further data. + +# Constructor + + StoppingCriterionState(data = nothing) + StoppingCriterionState(at_iteration, data) + +Initialize the state to not having indicated to stop yet, carrying `data` alongside it. +The second form is the one that also sets the iteration it records. +""" +mutable struct StoppingCriterionState{D} + at_iteration::Int + data::D +end + +StoppingCriterionState(data = nothing) = StoppingCriterionState(-1, data) + +# Throughout, `stopping_state_data = nothing` means that the criterion decides which data to +# carry, so on a reset it keeps whatever the state already holds. +_reset_data(stopping_criterion_state::StoppingCriterionState, stopping_state_data) = + isnothing(stopping_state_data) ? stopping_criterion_state.data : stopping_state_data + +# Fallbacks for any criterion that needs no data of its own, so that such a criterion does not +# have to provide these two methods at all. +initialize_state( + ::Problem, ::Algorithm, ::StoppingCriterion; stopping_state_data = nothing, kwargs... +) = StoppingCriterionState(stopping_state_data) +function initialize_state!( + ::Problem, ::Algorithm, ::StoppingCriterion, + stopping_criterion_state::StoppingCriterionState; + stopping_state_data = nothing, kwargs..., + ) + stopping_criterion_state.at_iteration = -1 + stopping_criterion_state.data = _reset_data(stopping_criterion_state, stopping_state_data) + return stopping_criterion_state +end diff --git a/src/stopping_criterion.jl b/src/stopping_criterion.jl index eb2a3e1..e602ae5 100644 --- a/src/stopping_criterion.jl +++ b/src/stopping_criterion.jl @@ -1,50 +1,3 @@ -@doc """ - StoppingCriterion - -An abstract type to represent a stopping criterion of an [`Algorithm`](@ref). - -A concrete [`StoppingCriterion`](@ref) receives its accompanying [`StoppingCriterionState`](@ref) -from a [`DefaultStoppingCriterionState`](@ref), which records the iteration at which the criterion -indicated to stop and carries a `data` field for anything else it has to remember. -A criterion is therefore free of any state bookkeeping by default. - -It should usually implement - -* [`is_finished!`](@ref)`(problem, algorithm, state, stopping_criterion, stopping_criterion_state)` -* [`is_finished`](@ref)`(problem, algorithm, state, stopping_criterion, stopping_criterion_state)` -* [`get_reason`](@ref)`(stopping_criterion, stopping_criterion_state)` -* [`indicates_convergence`](@ref)`(::Type{<:StoppingCriterion})` - -Only a criterion that is not served by that state defines one of its own, and with it an -[`initialize_state(problem::Problem, algorithm::Algorithm, stopping_criterion::StoppingCriterion; kwargs...)`](@ref) -to create it, as well as the corresponding mutating variant to reset it. - -Note that only [`indicates_convergence`](@ref) has to be implemented: -it answers whether meeting this criterion *would* mean convergence, which is a static property of the criterion type alone. -Both the variant taking a criterion and the one that additionally takes a [`StoppingCriterionState`](@ref), answering whether it *did* happen are derived from it. -""" -abstract type StoppingCriterion end - -@doc """ - StoppingCriterionState - -An abstract type to represent a stopping criterion state within a [`State`](@ref). -It represents the concrete state a [`StoppingCriterion`](@ref) is in. - -## Properties - -In order for the generic convergence reporting to work, the state should contain the following -property, and provide corresponding `getproperty` and `setproperty!` methods. - -* `at_iteration` – the iteration at which the accompanying [`StoppingCriterion`](@ref) indicated - to stop, where `0` means it already indicated to stop at the start and any negative number - means that it has not (yet) indicated to stop. - -A state that records its status differently can instead implement -[`is_active`](@ref)`(stopping_criterion, stopping_criterion_state)`. -""" -abstract type StoppingCriterionState end - @doc """ get_reason(stopping_criterion::StoppingCriterion, stopping_criterion_state::StoppingCriterionState) get_reason(algorithm::Algorithm, state::State) @@ -76,7 +29,7 @@ The second variant extracts the criterion and its state from `algorithm` and `st This is the machine-readable counterpart of [`get_reason`](@ref) and the predicate the generic convergence reporting is built on. The default implementation reads the `at_iteration` property of the state, see [`StoppingCriterionState`](@ref), -so it only has to be implemented for a state that records its status differently. +so it only has to be implemented for a criterion that records its status in the state's `data` instead. """ is_active( ::StoppingCriterion, stopping_criterion_state::StoppingCriterionState @@ -306,7 +259,7 @@ Note how this is the opposite quantifier from [`StopWhenAll`](@ref). This is deliberately pessimistic, and is why a `tolerance | budget` combination is never convergent as a criterion. To ask whether a *particular run* stopped because the convergence -criterion is what triggered, pass the accompanying [`GroupStoppingCriterionState`](@ref) as well. +criterion is what triggered, pass the accompanying [`StoppingCriterionState`](@ref) as well. """ function indicates_convergence(::Type{StopWhenAny{TCriteria}}) where {TCriteria <: Tuple} return all(indicates_convergence, fieldtypes(TCriteria)) @@ -341,27 +294,11 @@ Base.:|(s1::StoppingCriterion, s2::StopWhenAny) = StopWhenAny(s1, s2.criteria... Base.:|(s1::StopWhenAny, s2::StoppingCriterion) = StopWhenAny(s1.criteria..., s2) Base.:|(s1::StopWhenAny, s2::StopWhenAny) = StopWhenAny(s1.criteria..., s2.criteria...) -# A common state for stopping criteria working on tuples of stopping criteria -""" - GroupStoppingCriterionState <: StoppingCriterionState - -A [`StoppingCriterionState`](@ref) that groups multiple [`StoppingCriterionState`](@ref)s -internally as a tuple. -This is for example used in combination with [`StopWhenAny`](@ref) and [`StopWhenAll`](@ref). - -# Constructor - - GroupStoppingCriterionState(c::StoppingCriterionState...) -""" -mutable struct GroupStoppingCriterionState{TCriteriaStates <: Tuple} <: StoppingCriterionState - criteria_states::TCriteriaStates - at_iteration::Int - GroupStoppingCriterionState(c::StoppingCriterionState...) = new{typeof(c)}(c, -1) -end - +# A meta criterion carries the states of the criteria it combines as its `data`, as a tuple in +# the same order as its `criteria`. function get_reason( stop_when::Union{StopWhenAll, StopWhenAny}, - stopping_criterion_states::GroupStoppingCriterionState, + stopping_criterion_states::StoppingCriterionState, ) is_active(stop_when, stopping_criterion_states) || return nothing # only the children that did indicate to stop have anything to report, and of those the ones @@ -369,7 +306,7 @@ function get_reason( reasons = ( get_reason(stopping_criterion, stopping_criterion_state) for (stopping_criterion, stopping_criterion_state) in - zip(stop_when.criteria, stopping_criterion_states.criteria_states) + zip(stop_when.criteria, stopping_criterion_states.data) if is_active(stopping_criterion, stopping_criterion_state) ) reason = join(Iterators.filter(!isnothing, reasons)) @@ -379,7 +316,7 @@ function get_reason( end @doc """ - indicates_convergence(stop_when::Union{StopWhenAll, StopWhenAny}, ::GroupStoppingCriterionState) + indicates_convergence(stop_when::Union{StopWhenAll, StopWhenAny}, ::StoppingCriterionState) Return whether a group of stopping criteria stopped because of convergence. @@ -392,24 +329,24 @@ criterion is what triggered. """ function indicates_convergence( stop_when::Union{StopWhenAll, StopWhenAny}, - stopping_criterion_states::GroupStoppingCriterionState, + stopping_criterion_states::StoppingCriterionState, ) is_active(stop_when, stopping_criterion_states) || return false return any( st -> indicates_convergence(st[1], st[2]), - zip(stop_when.criteria, stopping_criterion_states.criteria_states), + zip(stop_when.criteria, stopping_criterion_states.data), ) end function get_active_stopping_criteria( stop_when::Union{StopWhenAll, StopWhenAny}, - stopping_criterion_states::GroupStoppingCriterionState, + stopping_criterion_states::StoppingCriterionState, ) pairs = Tuple{StoppingCriterion, StoppingCriterionState}[] # recurse rather than report the group itself: `&` and `|` flatten, but a mixed combination # such as `(c1 | c2) & c3` genuinely nests for (stopping_criterion, stopping_criterion_state) in - zip(stop_when.criteria, stopping_criterion_states.criteria_states) + zip(stop_when.criteria, stopping_criterion_states.data) append!( pairs, get_active_stopping_criteria(stopping_criterion, stopping_criterion_state), @@ -418,54 +355,62 @@ function get_active_stopping_criteria( return pairs end +# The `data` of a meta criterion is the states of the criteria it combines, so a +# `stopping_state_data` handed to it is read as those states, in the order of its `criteria`. +# Nothing is handed down to the children, which is what makes a nested combination work: each +# child fills and keeps its own `data` through its own initialization. function initialize_state( problem::Problem, algorithm::Algorithm, stop_when::Union{StopWhenAll, StopWhenAny}; - kwargs..., - ) - return GroupStoppingCriterionState( - ( - initialize_state(problem, algorithm, stopping_criterion; kwargs...) for - stopping_criterion in stop_when.criteria - )..., + stopping_state_data = nothing, kwargs..., ) + criteria_states = if isnothing(stopping_state_data) + map(stop_when.criteria) do stopping_criterion + initialize_state(problem, algorithm, stopping_criterion; kwargs...) + end + else + stopping_state_data + end + return StoppingCriterionState(criteria_states) end function initialize_state!( problem::Problem, algorithm::Algorithm, stop_when::Union{StopWhenAll, StopWhenAny}, - stopping_criterion_states::GroupStoppingCriterionState; - kwargs..., + stopping_criterion_states::StoppingCriterionState; + stopping_state_data = nothing, kwargs..., ) + criteria_states = _reset_data(stopping_criterion_states, stopping_state_data) for (stopping_criterion_state, stopping_criterion) in - zip(stopping_criterion_states.criteria_states, stop_when.criteria) + zip(criteria_states, stop_when.criteria) initialize_state!( problem, algorithm, stopping_criterion, stopping_criterion_state; kwargs..., ) end + stopping_criterion_states.data = criteria_states stopping_criterion_states.at_iteration = -1 return stopping_criterion_states end function is_finished( problem::Problem, algorithm::Algorithm, state::State, - stop_when_all::StopWhenAll, stopping_criterion_states::GroupStoppingCriterionState, + stop_when_all::StopWhenAll, stopping_criterion_states::StoppingCriterionState, ) # short-circuiting is fine here: unlike `is_finished!`, this may not mutate, so there is no # child left starved of an update by not being asked return all( st -> is_finished(problem, algorithm, state, st[1], st[2]), - zip(stop_when_all.criteria, stopping_criterion_states.criteria_states), + zip(stop_when_all.criteria, stopping_criterion_states.data), ) end function is_finished!( problem::Problem, algorithm::Algorithm, state::State, - stop_when_all::StopWhenAll, stopping_criterion_states::GroupStoppingCriterionState, + stop_when_all::StopWhenAll, stopping_criterion_states::StoppingCriterionState, ) k = state.iteration (k == 0) && (stopping_criterion_states.at_iteration = -1) # reset on init # `map` rather than `all`, so that every child is updated exactly once per iteration: # `all` would short-circuit and starve stateful criteria of the current iterate finished = map( - stop_when_all.criteria, stopping_criterion_states.criteria_states + stop_when_all.criteria, stopping_criterion_states.data ) do stopping_criterion, stopping_criterion_state is_finished!(problem, algorithm, state, stopping_criterion, stopping_criterion_state) end @@ -478,18 +423,18 @@ end function is_finished( problem::Problem, algorithm::Algorithm, state::State, - stop_when_any::StopWhenAny, stopping_criterion_states::GroupStoppingCriterionState, + stop_when_any::StopWhenAny, stopping_criterion_states::StoppingCriterionState, ) # short-circuiting is fine here: unlike `is_finished!`, this may not mutate, so there is no # child left starved of an update by not being asked return any( st -> is_finished(problem, algorithm, state, st[1], st[2]), - zip(stop_when_any.criteria, stopping_criterion_states.criteria_states), + zip(stop_when_any.criteria, stopping_criterion_states.data), ) end function is_finished!( problem::Problem, algorithm::Algorithm, state::State, - stop_when_any::StopWhenAny, stopping_criterion_states::GroupStoppingCriterionState, + stop_when_any::StopWhenAny, stopping_criterion_states::StoppingCriterionState, ) k = state.iteration (k == 0) && (stopping_criterion_states.at_iteration = -1) # reset on init @@ -497,7 +442,7 @@ function is_finished!( # `any` would short-circuit and starve stateful criteria of the current iterate, and # leave their `at_iteration` unset even though they did indicate to stop finished = map( - stop_when_any.criteria, stopping_criterion_states.criteria_states + stop_when_any.criteria, stopping_criterion_states.data ) do stopping_criterion, stopping_criterion_state is_finished!(problem, algorithm, state, stopping_criterion, stopping_criterion_state) end @@ -510,13 +455,13 @@ end function Base.summary( io::IO, - stop_when_any::StopWhenAny, stopping_criterion_states::GroupStoppingCriterionState, + stop_when_any::StopWhenAny, stopping_criterion_states::StoppingCriterionState, ) has_stopped = is_active(stop_when_any, stopping_criterion_states) s = has_stopped ? "reached" : "not reached" r = "Stop when _one_ of the following are fulfilled:\n" for (stopping_criterion, stopping_criterion_state) in - zip(stop_when_any.criteria, stopping_criterion_states.criteria_states) + zip(stop_when_any.criteria, stopping_criterion_states.data) t = replace(summary(stopping_criterion, stopping_criterion_state), "\n" => "\n\t") r = "$(r)\t$(t)\n" end @@ -524,13 +469,13 @@ function Base.summary( end function Base.summary( io::IO, - stop_when_all::StopWhenAll, stopping_criterion_states::GroupStoppingCriterionState, + stop_when_all::StopWhenAll, stopping_criterion_states::StoppingCriterionState, ) has_stopped = is_active(stop_when_all, stopping_criterion_states) s = has_stopped ? "reached" : "not reached" r = "Stop when _all_ of the following are fulfilled:\n" for (stopping_criterion, stopping_criterion_state) in - zip(stop_when_all.criteria, stopping_criterion_states.criteria_states) + zip(stop_when_all.criteria, stopping_criterion_states.data) t = replace(summary(stopping_criterion, stopping_criterion_state), "\n" => "\n\t") r = "$(r)\t$(t)\n" end @@ -560,64 +505,17 @@ struct StopAfterIteration <: StoppingCriterion max_iterations::Int end -""" - DefaultStoppingCriterionState <: StoppingCriterionState - -A [`StoppingCriterionState`](@ref) that stores the iteration number at which it (last) -indicated to stop, and optionally any further data its [`StoppingCriterion`](@ref) needs. - -# Fields - -* `at_iteration::Int` stores the iteration number at which this state indicated to stop. - * `0` means it already indicated to stop at the start. - * any negative number means that it has not yet indicated to stop. -* `data` stores any further data the criterion has to carry from one iteration to the next, - for example a value it compares against in the next one. - It is opaque to this package, and none of its contents are exposed as properties of the - state, so a criterion reaches them through `stopping_criterion_state.data`. - A mutable struct of its own is the recommended choice; `nothing`, the default, indicates - that the criterion needs no further data. - -# Constructor - - DefaultStoppingCriterionState(data = nothing) - -Initialize the state to not having indicated to stop yet, carrying `data` alongside it. -""" -mutable struct DefaultStoppingCriterionState{D} <: StoppingCriterionState - at_iteration::Int - data::D -end - -DefaultStoppingCriterionState(data = nothing) = DefaultStoppingCriterionState(-1, data) - -# Fallbacks for any criterion that needs no state of its own beyond `at_iteration`, so that -# such a criterion does not have to provide these two methods at all. -initialize_state( - ::Problem, ::Algorithm, ::StoppingCriterion; stopping_state_data = nothing, kwargs... -) = DefaultStoppingCriterionState(stopping_state_data) -function initialize_state!( - ::Problem, ::Algorithm, ::StoppingCriterion, - stopping_criterion_state::DefaultStoppingCriterionState; - stopping_state_data = stopping_criterion_state.data, kwargs..., - ) - stopping_criterion_state.at_iteration = -1 - stopping_criterion_state.data = stopping_state_data - return stopping_criterion_state -end - - function is_finished( ::Problem, ::Algorithm, state::State, stop_after_iteration::StopAfterIteration, - stopping_criterion_state::DefaultStoppingCriterionState, + stopping_criterion_state::StoppingCriterionState, ) return state.iteration >= stop_after_iteration.max_iterations end function is_finished!( ::Problem, ::Algorithm, state::State, stop_after_iteration::StopAfterIteration, - stopping_criterion_state::DefaultStoppingCriterionState, + stopping_criterion_state::StoppingCriterionState, ) k = state.iteration (k == 0) && (stopping_criterion_state.at_iteration = -1) @@ -629,7 +527,7 @@ function is_finished!( end function get_reason( stop_after_iteration::StopAfterIteration, - stopping_criterion_state::DefaultStoppingCriterionState, + stopping_criterion_state::StoppingCriterionState, ) if is_active(stop_after_iteration, stopping_criterion_state) return "At iteration $(stopping_criterion_state.at_iteration) the algorithm reached its maximal number of iterations ($(stop_after_iteration.max_iterations)).\n" @@ -639,7 +537,7 @@ end function Base.summary( io::IO, stop_after_iteration::StopAfterIteration, - stopping_criterion_state::DefaultStoppingCriterionState, + stopping_criterion_state::StoppingCriterionState, ) has_stopped = is_active(stop_after_iteration, stopping_criterion_state) s = has_stopped ? "reached" : "not reached" @@ -676,63 +574,70 @@ struct StopAfter <: StoppingCriterion end @doc """ - StopAfterTimePeriodState <: StoppingCriterionState + StopAfterTimePeriodData -A state for stopping criteria that are based on time measurements, -for example [`StopAfter`](@ref). +The `data` a [`StoppingCriterionState`](@ref) carries for stopping criteria that are based on +time measurements, for example [`StopAfter`](@ref). # Fields * `start` stores the starting time, recorded when the algorithm is started (the call with `k=0`). * `time` stores the elapsed time. -* `at_iteration` indicates at which iteration (including `k=0`) the stopping criterion - was fulfilled, and is `-1` while it is not fulfilled. + +# Constructor + + StopAfterTimePeriodData() + +Initialize both to zero, indicating a timer that has not been started yet. """ -mutable struct StopAfterTimePeriodState <: StoppingCriterionState +mutable struct StopAfterTimePeriodData start::Nanosecond time::Nanosecond - at_iteration::Int - function StopAfterTimePeriodState() - return new(Nanosecond(0), Nanosecond(0), -1) - end end -initialize_state(::Problem, ::Algorithm, ::StopAfter; kwargs...) = - StopAfterTimePeriodState() +StopAfterTimePeriodData() = StopAfterTimePeriodData(Nanosecond(0), Nanosecond(0)) + +initialize_state( + ::Problem, ::Algorithm, ::StopAfter; stopping_state_data = nothing, kwargs..., +) = StoppingCriterionState( + isnothing(stopping_state_data) ? StopAfterTimePeriodData() : stopping_state_data +) function initialize_state!( ::Problem, ::Algorithm, ::StopAfter, - stopping_criterion_state::StopAfterTimePeriodState; - kwargs..., + stopping_criterion_state::StoppingCriterionState; + stopping_state_data = nothing, kwargs..., ) - stopping_criterion_state.start = Nanosecond(0) - stopping_criterion_state.time = Nanosecond(0) + # a reset restarts the timer, whether it is the clock the state already held or a new one + stopping_criterion_state.data = _reset_data(stopping_criterion_state, stopping_state_data) + stopping_criterion_state.data.start = Nanosecond(0) + stopping_criterion_state.data.time = Nanosecond(0) stopping_criterion_state.at_iteration = -1 return stopping_criterion_state end function is_finished( ::Problem, ::Algorithm, state::State, - stop_after::StopAfter, stop_after_state::StopAfterTimePeriodState, + stop_after::StopAfter, stop_after_state::StoppingCriterionState, ) k = state.iteration # Read the clock rather than the `time` recorded by the last `is_finished!`, so that this # reports on the time elapsed *now*. Only the timer itself may not be (re)started here. - (k <= 0 || value(stop_after_state.start) == 0) && return false - return (Nanosecond(time_ns()) - stop_after_state.start) > Nanosecond(stop_after.threshold) + (k <= 0 || value(stop_after_state.data.start) == 0) && return false + return (Nanosecond(time_ns()) - stop_after_state.data.start) > Nanosecond(stop_after.threshold) end function is_finished!( ::Problem, ::Algorithm, state::State, - stop_after::StopAfter, stop_after_state::StopAfterTimePeriodState, + stop_after::StopAfter, stop_after_state::StoppingCriterionState, ) k = state.iteration - if value(stop_after_state.start) == 0 || k <= 0 # (re)start timer + if value(stop_after_state.data.start) == 0 || k <= 0 # (re)start timer stop_after_state.at_iteration = -1 - stop_after_state.start = Nanosecond(time_ns()) - stop_after_state.time = Nanosecond(0) + stop_after_state.data.start = Nanosecond(time_ns()) + stop_after_state.data.time = Nanosecond(0) else - stop_after_state.time = Nanosecond(time_ns()) - stop_after_state.start - if k > 0 && (stop_after_state.time > Nanosecond(stop_after.threshold)) + stop_after_state.data.time = Nanosecond(time_ns()) - stop_after_state.data.start + if k > 0 && (stop_after_state.data.time > Nanosecond(stop_after.threshold)) stop_after_state.at_iteration = k return true end @@ -741,16 +646,16 @@ function is_finished!( end function get_reason( stop_after::StopAfter, - stopping_criterion_state::StopAfterTimePeriodState, + stopping_criterion_state::StoppingCriterionState, ) if is_active(stop_after, stopping_criterion_state) - return "After iteration $(stopping_criterion_state.at_iteration) the algorithm ran for $(floor(stopping_criterion_state.time, typeof(stop_after.threshold))) (threshold: $(stop_after.threshold)).\n" + return "After iteration $(stopping_criterion_state.at_iteration) the algorithm ran for $(floor(stopping_criterion_state.data.time, typeof(stop_after.threshold))) (threshold: $(stop_after.threshold)).\n" end return nothing end function Base.summary( io::IO, - stop_after::StopAfter, stopping_criterion_state::StopAfterTimePeriodState, + stop_after::StopAfter, stopping_criterion_state::StoppingCriterionState, ) has_stopped = is_active(stop_after, stopping_criterion_state) s = has_stopped ? "reached" : "not reached" diff --git a/src/test_suite.jl b/src/test_suite.jl index e849418..2a6b7fa 100644 --- a/src/test_suite.jl +++ b/src/test_suite.jl @@ -13,10 +13,4 @@ end struct DummyProblem <: AlgorithmsInterface.Problem end -mutable struct DummyState{V, S <: AlgorithmsInterface.StoppingCriterionState} <: AlgorithmsInterface.State - iterate::V - stopping_criterion_state::S - iteration::Int -end - end diff --git a/test/logging.jl b/test/logging.jl index 09989f0..81496e4 100644 --- a/test/logging.jl +++ b/test/logging.jl @@ -6,34 +6,26 @@ struct LogDummyProblem <: Problem end struct LogDummyAlgorithm <: Algorithm stopping_criterion end -mutable struct LogDummyState{S <: StoppingCriterionState} <: State - iterate::Float64 - iteration::Int - stopping_criterion_state::S -end - -# State initialization for the dummy algorithm +# State initialization for the dummy algorithm: it starts from zero rather than from an +# iterate the caller hands it, so it implements `initialize_state` itself function AlgorithmsInterface.initialize_state(problem::LogDummyProblem, algorithm::LogDummyAlgorithm; kwargs...) sc_state = initialize_state(problem, algorithm, algorithm.stopping_criterion; kwargs...) - return LogDummyState(0.0, 0, sc_state) + return State(0.0, sc_state) end function AlgorithmsInterface.initialize_state!( problem::LogDummyProblem, algorithm::LogDummyAlgorithm, - state::LogDummyState; + state::State; kwargs... ) - initialize_state!(problem, algorithm, algorithm.stopping_criterion, state.stopping_criterion_state; kwargs...) - state.iterate = 0.0 - state.iteration = 0 - return state + return initialize_state!(problem, algorithm, state, 0.0; kwargs...) end # One trivial step per iteration (not relevant for the logging test) function AlgorithmsInterface.step!( ::LogDummyProblem, ::LogDummyAlgorithm, - state::LogDummyState, + state::State, ) state.iterate += 1.0 return state diff --git a/test/newton.jl b/test/newton.jl index f259ae0..967141f 100644 --- a/test/newton.jl +++ b/test/newton.jl @@ -16,35 +16,24 @@ struct NewtonMethod{S} <: Algorithm stopping_criterion::S end -mutable struct NewtonState{S} <: State - iteration::Int - iterate::Float64 - stopping_criterion_state::S -end - # Implementing the algorithm # -------------------------- -function initialize_state(problem::RootFindingProblem, algorithm::NewtonMethod) - scs = initialize_state(problem, algorithm, algorithm.stopping_criterion) - return NewtonState(0, 1.0, scs) # hardcode initial guess to 1.0 +# The initial guess is the algorithm's to make rather than the caller's, so both +# initializations are implemented, handing the generic `State` the iterate to start from. +function initialize_state(problem::RootFindingProblem, algorithm::NewtonMethod; kwargs...) + scs = initialize_state(problem, algorithm, algorithm.stopping_criterion; kwargs...) + return State(1.0, scs) # hardcode initial guess to 1.0 end function initialize_state!( problem::RootFindingProblem, algorithm::NewtonMethod, - state::NewtonState, + state::State; + kwargs..., ) - state.iteration = 0 - state.iterate = 1.0 - initialize_state!( - problem, - algorithm, - algorithm.stopping_criterion, - state.stopping_criterion_state, - ) - return state + return initialize_state!(problem, algorithm, state, 1.0; kwargs...) end -function step!(problem::RootFindingProblem, ::NewtonMethod, state::NewtonState) +function step!(problem::RootFindingProblem, ::NewtonMethod, state::State) state.iterate -= problem.f(state.iterate) / problem.df(state.iterate) return state end diff --git a/test/state.jl b/test/state.jl index 3c04b25..98949f6 100644 --- a/test/state.jl +++ b/test/state.jl @@ -1,4 +1,4 @@ -# Tests for the default state and the default state initialization +# Tests for the state and the default state initialization using Test using AlgorithmsInterface @@ -10,8 +10,8 @@ using Dates # -------- # A problem and an algorithm that provide nothing beyond what the interface demands, so that -# every state they are solved with has to come from the defaults. Newton's method is spelled -# out with its own state in `newton.jl`; the point here is that this one has none. +# every state they are solved with has to come from the default initialization. Newton's +# method initializes its own in `newton.jl`; the point here is that this one does not. struct HalvingProblem <: Problem target::Float64 end @@ -20,14 +20,14 @@ struct Halving{S <: StoppingCriterion} <: Algorithm stopping_criterion::S end -function AlgorithmsInterface.step!(problem::HalvingProblem, ::Halving, state::DefaultState) +function AlgorithmsInterface.step!(problem::HalvingProblem, ::Halving, state::State) state.iterate = (state.iterate + problem.target) / 2 return state end -# Carries a counter through `data` to show that an algorithm can keep data from one step to -# the next without a state type of its own. A mutable struct is what the docs recommend for -# this: the state hands the same container back every step, and only its contents change. +# Carries a counter through `data` to show how an algorithm keeps data from one step to the +# next. A mutable struct is what the docs recommend for this: the state hands the same +# container back every step, and only its contents change. struct CountingHalving{S <: StoppingCriterion} <: Algorithm stopping_criterion::S end @@ -37,7 +37,7 @@ mutable struct StepCounter end function AlgorithmsInterface.step!( - problem::HalvingProblem, ::CountingHalving, state::DefaultState, + problem::HalvingProblem, ::CountingHalving, state::State, ) state.iterate = (state.iterate + problem.target) / 2 state.data.steps += 1 @@ -47,53 +47,53 @@ end # Tests # ----- -@testset "DefaultState construction" begin - scs = DefaultStoppingCriterionState() +@testset "State construction" begin + scs = StoppingCriterionState() - state = DefaultState(2.0, scs) - @test state isa DefaultState{Float64, typeof(scs), Nothing} + state = State(2.0, scs) + @test state isa State{Float64, typeof(scs), Nothing} @test state.iterate == 2.0 @test state.stopping_criterion_state === scs @test state.data === nothing @test state.iteration == 0 - # `iteration` is the fourth positional argument - @test DefaultState(2.0, scs, nothing, 3).iteration == 3 + # `iteration` is the third positional argument, with `data` staying last + @test State(2.0, scs, 3, nothing).iteration == 3 # the data field is opaque, so anything goes and nothing of it is exposed as a property - named_tuple_state = DefaultState(2.0, scs, (; gradient = 1.0)) + named_tuple_state = State(2.0, scs, (; gradient = 1.0)) @test named_tuple_state.data.gradient == 1.0 @test !hasproperty(named_tuple_state, :gradient) - dict_state = DefaultState(2.0, scs, Dict{Symbol, Any}(:gradient => 1.0)) + dict_state = State(2.0, scs, Dict{Symbol, Any}(:gradient => 1.0)) @test dict_state.data[:gradient] == 1.0 dict_state.data[:hessian] = 2.0 @test dict_state.data[:hessian] == 2.0 end -@testset "DefaultState satisfies the State interface" begin +@testset "State interacts with the generic functionality" begin problem = AIT.DummyProblem() algorithm = AIT.DummyAlgorithm(StopAfterIteration(5)) - state = DefaultState(2.0, DefaultStoppingCriterionState()) + state = State(2.0, StoppingCriterionState()) - @test increment!(state) === state + @test increment!(problem, algorithm, state) === state @test state.iteration == 1 @test finalize_state!(problem, algorithm, state) == 2.0 @test !is_finished!(problem, algorithm, state) @test !is_active(algorithm, state) end -@testset "initialize_state defaults to a DefaultState" begin +@testset "initialize_state defaults to a State" begin problem = AIT.DummyProblem() algorithm = AIT.DummyAlgorithm(StopAfterIteration(5)) state = initialize_state(problem, algorithm, 2.0) - @test state isa DefaultState + @test state isa State @test state.iterate == 2.0 @test state.iteration == 0 @test state.data === nothing # the accompanying criterion state comes from the algorithm's own stopping criterion - @test state.stopping_criterion_state isa DefaultStoppingCriterionState + @test state.stopping_criterion_state isa StoppingCriterionState @test state.stopping_criterion_state.at_iteration == -1 @test state.stopping_criterion_state.data === nothing @@ -102,26 +102,21 @@ end @test both.data == (; a = 1) @test both.stopping_criterion_state.data == (; b = 2) - # the keyword form forwards to the same implementation - @test initialize_state( - problem, algorithm; iterate = 2.0, state_data = (; a = 1), stopping_state_data = (; b = 2), - ).data == (; a = 1) - - # and it is inferable, so reaching for the positional form is a matter of taste - @test @inferred(initialize_state(problem, algorithm, 2.0)) isa DefaultState + # it is inferable + @test @inferred(initialize_state(problem, algorithm, 2.0)) isa State - # an algorithm that neither provides a state type nor an iterate to start from is a - # mistake we should not paper over - @test_throws UndefKeywordError initialize_state(problem, algorithm) + # an algorithm that neither initializes a state itself nor is given an iterate to start + # from is a mistake we should not paper over + @test_throws MethodError initialize_state(problem, algorithm) - # a criterion with a state of its own still wins over the fallback + # a criterion that carries data of its own still wins over the fallback time_algorithm = AIT.DummyAlgorithm(StopAfter(Second(1))) @test initialize_state( problem, time_algorithm, 2.0, - ).stopping_criterion_state isa StopAfterTimePeriodState + ).stopping_criterion_state.data isa StopAfterTimePeriodData end -@testset "initialize_state! resets a DefaultState" begin +@testset "initialize_state! resets a State" begin problem = AIT.DummyProblem() algorithm = AIT.DummyAlgorithm(StopAfterIteration(5)) state = initialize_state(problem, algorithm, 2.0, (; a = 1), (; b = 2)) @@ -144,20 +139,25 @@ end @test state.data == (; a = 9) @test state.stopping_criterion_state.data == (; b = 8) - # the keyword form forwards to the same implementation - initialize_state!(problem, algorithm, state; iterate = 4.0, state_data = (; a = 7)) + # everything beyond the iterate is optional and keeps what the state already holds + initialize_state!(problem, algorithm, state, 4.0, (; a = 7)) + @test state.iterate == 4.0 + @test state.data == (; a = 7) + @test state.stopping_criterion_state.data == (; b = 8) + + initialize_state!(problem, algorithm, state) @test state.iterate == 4.0 @test state.data == (; a = 7) @test state.stopping_criterion_state.data == (; b = 8) end -@testset "an algorithm without a state type of its own can be solved" begin +@testset "an algorithm providing only a step! can be solved" begin problem = HalvingProblem(1.0) algorithm = Halving(StopAfterIteration(40)) - # no state type, no initialize_state, no initialize_state!: only `step!` above - # `solve` reaches the default through its keywords - @test solve(problem, algorithm; iterate = 100.0) ≈ 1.0 + # no initialize_state, no initialize_state!: only `step!` above + # `solve` passes the iterate it is given straight on to the default + @test solve(problem, algorithm, 100.0) ≈ 1.0 state = initialize_state(problem, algorithm, 100.0) @test solve!(problem, algorithm, state) ≈ 1.0 diff --git a/test/stopping_criterion.jl b/test/stopping_criterion.jl index f40ab4f..25abef5 100644 --- a/test/stopping_criterion.jl +++ b/test/stopping_criterion.jl @@ -15,14 +15,14 @@ struct StopWhenConverged <: StoppingCriterion end function AlgorithmsInterface.is_finished( ::Problem, ::Algorithm, state::State, - stop_when_converged::StopWhenConverged, ::DefaultStoppingCriterionState, + stop_when_converged::StopWhenConverged, ::StoppingCriterionState, ) return state.iteration >= stop_when_converged.at end function AlgorithmsInterface.is_finished!( ::Problem, ::Algorithm, state::State, stop_when_converged::StopWhenConverged, - stopping_criterion_state::DefaultStoppingCriterionState, + stopping_criterion_state::StoppingCriterionState, ) k = state.iteration (k == 0) && (stopping_criterion_state.at_iteration = -1) @@ -33,7 +33,7 @@ function AlgorithmsInterface.is_finished!( return false end function AlgorithmsInterface.get_reason( - ::StopWhenConverged, stopping_criterion_state::DefaultStoppingCriterionState, + ::StopWhenConverged, stopping_criterion_state::StoppingCriterionState, ) stopping_criterion_state.at_iteration < 0 && return nothing return "Converged at iteration $(stopping_criterion_state.at_iteration).\n" @@ -41,89 +41,89 @@ end AlgorithmsInterface.indicates_convergence(::Type{StopWhenConverged}) = true # Never indicates to stop, but records how many times it was asked, so that short-circuiting -# inside the meta criteria becomes observable. +# inside the meta criteria becomes observable. The counter is what it carries as `data`, so it +# exercises a criterion that does provide its own initialization. struct CountingCriterion <: StoppingCriterion end -mutable struct CountingCriterionState <: StoppingCriterionState - at_iteration::Int +mutable struct CallCounter calls::Int end AlgorithmsInterface.initialize_state(::Problem, ::Algorithm, ::CountingCriterion; kwargs...) = - CountingCriterionState(-1, 0) + StoppingCriterionState(CallCounter(0)) function AlgorithmsInterface.initialize_state!( ::Problem, ::Algorithm, ::CountingCriterion, - stopping_criterion_state::CountingCriterionState; kwargs..., + stopping_criterion_state::StoppingCriterionState; kwargs..., ) stopping_criterion_state.at_iteration = -1 - stopping_criterion_state.calls = 0 + stopping_criterion_state.data.calls = 0 return stopping_criterion_state end AlgorithmsInterface.is_finished( - ::Problem, ::Algorithm, ::State, ::CountingCriterion, ::CountingCriterionState + ::Problem, ::Algorithm, ::State, ::CountingCriterion, ::StoppingCriterionState ) = false function AlgorithmsInterface.is_finished!( ::Problem, ::Algorithm, ::State, ::CountingCriterion, - stopping_criterion_state::CountingCriterionState, + stopping_criterion_state::StoppingCriterionState, ) - stopping_criterion_state.calls += 1 + stopping_criterion_state.data.calls += 1 return false end -AlgorithmsInterface.get_reason(::CountingCriterion, ::CountingCriterionState) = nothing +AlgorithmsInterface.get_reason(::CountingCriterion, ::StoppingCriterionState) = nothing AlgorithmsInterface.indicates_convergence(::Type{CountingCriterion}) = false # Indicates to stop immediately, but implements nothing beyond the bare minimum: no `get_reason`, -# no `indicates_convergence` and no state of its own, so it exercises the fallbacks for all three. +# no `indicates_convergence` and no data of its own, so it exercises the fallbacks for all three. struct SilentCriterion <: StoppingCriterion end AlgorithmsInterface.is_finished( - ::Problem, ::Algorithm, ::State, ::SilentCriterion, ::DefaultStoppingCriterionState + ::Problem, ::Algorithm, ::State, ::SilentCriterion, ::StoppingCriterionState ) = true function AlgorithmsInterface.is_finished!( ::Problem, ::Algorithm, state::State, ::SilentCriterion, - stopping_criterion_state::DefaultStoppingCriterionState, + stopping_criterion_state::StoppingCriterionState, ) stopping_criterion_state.at_iteration = state.iteration return true end -# Records its status somewhere other than `at_iteration`, so it has to override +# Records its status in `data` rather than in `at_iteration`, so it has to override # `is_active` rather than rely on the default. struct UnconventionalCriterion <: StoppingCriterion end -mutable struct UnconventionalCriterionState <: StoppingCriterionState +mutable struct StoppedFlag stopped::Bool end AlgorithmsInterface.initialize_state(::Problem, ::Algorithm, ::UnconventionalCriterion; kwargs...) = - UnconventionalCriterionState(false) + StoppingCriterionState(StoppedFlag(false)) function AlgorithmsInterface.initialize_state!( ::Problem, ::Algorithm, ::UnconventionalCriterion, - stopping_criterion_state::UnconventionalCriterionState; kwargs..., + stopping_criterion_state::StoppingCriterionState; kwargs..., ) - stopping_criterion_state.stopped = false + stopping_criterion_state.data.stopped = false return stopping_criterion_state end AlgorithmsInterface.is_active( - ::UnconventionalCriterion, stopping_criterion_state::UnconventionalCriterionState -) = stopping_criterion_state.stopped + ::UnconventionalCriterion, stopping_criterion_state::StoppingCriterionState +) = stopping_criterion_state.data.stopped AlgorithmsInterface.indicates_convergence(::Type{UnconventionalCriterion}) = true -@testset "DefaultStoppingCriterionState" begin - scs = DefaultStoppingCriterionState() +@testset "StoppingCriterionState" begin + scs = StoppingCriterionState() @test scs isa StoppingCriterionState @test scs.at_iteration == -1 @test scs.data === nothing # the data field is opaque, so anything goes and nothing of it is exposed as a property - with_data = DefaultStoppingCriterionState(Dict{Symbol, Any}(:last_change => 1.0)) + with_data = StoppingCriterionState(Dict{Symbol, Any}(:last_change => 1.0)) @test with_data.at_iteration == -1 @test with_data.data[:last_change] == 1.0 @test !hasproperty(with_data, :last_change) end -@testset "a criterion without a state of its own falls back to the default" begin +@testset "a criterion without data of its own falls back to the initialization defaults" begin # `SilentCriterion` provides neither `initialize_state` nor `initialize_state!` silent = SilentCriterion() algorithm = AIT.DummyAlgorithm(silent) scs = initialize_state(problem, algorithm, silent) - @test scs isa DefaultStoppingCriterionState + @test scs isa StoppingCriterionState @test scs.at_iteration == -1 scs.at_iteration = 3 @@ -138,11 +138,41 @@ end initialize_state!(problem, algorithm, silent, seeded; stopping_state_data = (; tol = 1.0e-4)) @test seeded.data == (; tol = 1.0e-4) - # a criterion that does provide them keeps its own state type + # a criterion that does provide them gets its own data counting = CountingCriterion() @test initialize_state( problem, AIT.DummyAlgorithm(counting), counting, - ) isa CountingCriterionState + ).data isa CallCounter +end + +@testset "a criterion with data of its own takes the data it is handed" begin + # `StopAfter` defaults its data to a fresh clock rather than hardcoding one + stop_after = StopAfter(Second(1)) + algorithm = AIT.DummyAlgorithm(stop_after) + @test initialize_state(problem, algorithm, stop_after).data isa StopAfterTimePeriodData + + clock = StopAfterTimePeriodData(Nanosecond(5), Nanosecond(3)) + scs = initialize_state(problem, algorithm, stop_after; stopping_state_data = clock) + @test scs.data === clock + # a reset restarts whichever clock it ends up with + initialize_state!(problem, algorithm, stop_after, scs) + @test scs.data === clock + @test clock.start == Nanosecond(0) + replacement = StopAfterTimePeriodData(Nanosecond(7), Nanosecond(7)) + initialize_state!(problem, algorithm, stop_after, scs; stopping_state_data = replacement) + @test scs.data === replacement + @test replacement.start == Nanosecond(0) + + # a meta criterion reads the data it is handed as the states of its own criteria + group = stop_after | StopAfterIteration(2) + group_algorithm = AIT.DummyAlgorithm(group) + children = (StoppingCriterionState(StopAfterTimePeriodData()), StoppingCriterionState()) + group_state = initialize_state(problem, group_algorithm, group; stopping_state_data = children) + @test group_state.data === children + # and nothing of it reaches its children, so each of them keeps its own + @test initialize_state( + problem, group_algorithm, group, + ).data[1].data isa StopAfterTimePeriodData end @testset "StopAfterIteration" begin @@ -153,8 +183,8 @@ end algorithm = AIT.DummyAlgorithm(s1) s1_state = initialize_state(problem, algorithm, s1) @test !indicates_convergence(s1, s1_state) - state_finished = AIT.DummyState(nothing, s1_state, 2) - alg_state = AIT.DummyState(nothing, s1_state, 1) + state_finished = State(nothing, s1_state, 2, nothing) + alg_state = State(nothing, s1_state, 1, nothing) @test is_finished(problem, algorithm, state_finished) @test !is_finished(problem, algorithm, alg_state) # Fake a stop: @@ -181,7 +211,7 @@ end algorithm = AIT.DummyAlgorithm(s1) s1_state = initialize_state(problem, algorithm, s1) - alg_state = AIT.DummyState(nothing, s1_state, 0) + alg_state = State(nothing, s1_state, 0, nothing) # Iteration 0: Start timer @test !is_finished!(problem, algorithm, alg_state) @test !is_finished(problem, algorithm, alg_state) @@ -195,7 +225,7 @@ end # The non-mutating variant reads the clock rather than the `time` recorded by the last # `is_finished!`, so a stale recording does not make it report "not finished" - s1_state.time = Nanosecond(0) + s1_state.data.time = Nanosecond(0) @test is_finished(problem, algorithm, alg_state) # but it may not (re)start the timer either, so it stays quiet before the first iteration alg_state.iteration = 0 @@ -220,11 +250,11 @@ end @test contains(s1_str, "Overall: not reached") @test isnothing(AlgorithmsInterface.get_reason(s1, s1_state)) - alg_state = AIT.DummyState(nothing, s1_state, 1) + alg_state = State(nothing, s1_state, 1, nothing) @test !is_finished(problem, algorithm, alg_state) # Fake start timer - s1_state.criteria_states[2].start = Nanosecond(time_ns()) - s1_state.criteria_states[2].time = Nanosecond(7) + s1_state.data[2].data.start = Nanosecond(time_ns()) + s1_state.data[2].data.time = Nanosecond(7) # just time is not enough @test !is_finished!(problem, algorithm, alg_state) @test !is_finished(problem, algorithm, alg_state) @@ -238,7 +268,7 @@ end @test startswith(get_reason(s1, s1_state), "At iteration 2") @test alg_state.stopping_criterion_state.at_iteration > 0 AlgorithmsInterface.initialize_state!(problem, algorithm, s1, s1_state) - @test s1_state.criteria_states[1].at_iteration == -1 + @test s1_state.data[1].at_iteration == -1 # Different constructors s2 = c1 & c2 & c3 @test s1 & c3 == s2 @@ -266,12 +296,12 @@ end @test contains(s1_str, "Overall: not reached") @test isnothing(AlgorithmsInterface.get_reason(s1, s1_state)) - alg_state = AIT.DummyState(nothing, s1_state, 1) + alg_state = State(nothing, s1_state, 1, nothing) @test !is_finished!(problem, algorithm, alg_state) @test !is_finished(problem, algorithm, alg_state) # Fake two seconds of elapsed time by moving the recorded start into the past -- the # non-mutating variant derives the elapsed time from the clock, not from `time` - s1_state.criteria_states[2].start = Nanosecond(time_ns()) - Nanosecond(Second(2)) + s1_state.data[2].data.start = Nanosecond(time_ns()) - Nanosecond(Second(2)) @test is_finished(problem, algorithm, alg_state) alg_state.iteration = 2 @test is_finished(problem, algorithm, alg_state) @@ -279,7 +309,7 @@ end @test is_finished!(problem, algorithm, alg_state) @test alg_state.stopping_criterion_state.at_iteration > 0 AlgorithmsInterface.initialize_state!(problem, algorithm, s1, s1_state) - @test s1_state.criteria_states[1].at_iteration == -1 + @test s1_state.data[1].at_iteration == -1 # Different constructors s2 = c1 | c2 | c3 @test s1 | c3 == s2 @@ -297,7 +327,7 @@ end scs = initialize_state(problem, algorithm, converging) @test indicates_convergence(converging) @test !indicates_convergence(converging, scs) - state = AIT.DummyState(nothing, scs, 1) + state = State(nothing, scs, 1, nothing) @test !is_finished!(problem, algorithm, state) @test !indicates_convergence(converging, scs) state.iteration = 2 @@ -312,21 +342,21 @@ end algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 0) + state = State(nothing, scs, 0, nothing) @test !is_finished!(problem, algorithm, state) @test !indicates_convergence(stop_when, scs) state.iteration = 2 @test is_finished!(problem, algorithm, state) @test indicates_convergence(stop_when, scs) - @test indicates_convergence(converging, scs.criteria_states[1]) - @test !indicates_convergence(fallback, scs.criteria_states[2]) + @test indicates_convergence(converging, scs.data[1]) + @test !indicates_convergence(fallback, scs.data[2]) # only the fallback triggers -> stopped, but not converged stop_when = StopWhenConverged(100) | fallback algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 0) + state = State(nothing, scs, 0, nothing) @test !is_finished!(problem, algorithm, state) state.iteration = 5 @test is_finished!(problem, algorithm, state) @@ -346,30 +376,30 @@ end stop_when = StopAfterIteration(1) | counter algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 1) + state = State(nothing, scs, 1, nothing) @test is_finished!(problem, algorithm, state) - @test scs.criteria_states[2].calls == 1 + @test scs.data[2].data.calls == 1 # `StopWhenAll`: the first child already indicates *not* to stop stop_when = counter & StopAfterIteration(1) algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 1) + state = State(nothing, scs, 1, nothing) @test !is_finished!(problem, algorithm, state) - @test scs.criteria_states[1].calls == 1 + @test scs.data[1].data.calls == 1 @test !is_finished!(problem, algorithm, state) - @test scs.criteria_states[1].calls == 2 + @test scs.data[1].data.calls == 2 # a reset clears the tally again initialize_state!(problem, algorithm, stop_when, scs) - @test scs.criteria_states[1].calls == 0 + @test scs.data[1].data.calls == 0 end @testset "is_active" begin converging = StopWhenConverged(2) algorithm = AIT.DummyAlgorithm(converging) scs = initialize_state(problem, algorithm, converging) - state = AIT.DummyState(nothing, scs, 1) + state = State(nothing, scs, 1, nothing) @test !is_active(converging, scs) @test !is_finished!(problem, algorithm, state) @@ -393,7 +423,7 @@ end ucs = initialize_state(problem, algorithm, unconventional) @test !is_active(unconventional, ucs) @test !indicates_convergence(unconventional, ucs) - ucs.stopped = true + ucs.data.stopped = true @test is_active(unconventional, ucs) @test indicates_convergence(unconventional, ucs) end @@ -402,7 +432,7 @@ end silent = SilentCriterion() algorithm = AIT.DummyAlgorithm(silent) scs = initialize_state(problem, algorithm, silent) - state = AIT.DummyState(nothing, scs, 1) + state = State(nothing, scs, 1, nothing) @test is_finished!(problem, algorithm, state) @test is_active(silent, scs) @@ -419,7 +449,7 @@ end stop_when = StopWhenAny(SilentCriterion()) algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 1) + state = State(nothing, scs, 1, nothing) @test is_finished!(problem, algorithm, state) @test is_active(stop_when, scs) @@ -430,7 +460,7 @@ end stop_when = SilentCriterion() | StopWhenConverged(1) algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 1) + state = State(nothing, scs, 1, nothing) @test is_finished!(problem, algorithm, state) @test get_reason(stop_when, scs) == "Converged at iteration 1.\n" end @@ -441,17 +471,17 @@ end stop_when = StopWhenConverged(2) | StopAfterIteration(5) algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 2) + state = State(nothing, scs, 2, nothing) @test is_finished!(problem, algorithm, state) at_iteration = scs.at_iteration - child_at_iterations = map(cs -> cs.at_iteration, scs.criteria_states) + child_at_iterations = map(cs -> cs.at_iteration, scs.data) @test at_iteration == 2 state.iteration = 0 is_finished(problem, algorithm, state) @test scs.at_iteration == at_iteration - @test map(cs -> cs.at_iteration, scs.criteria_states) == child_at_iterations + @test map(cs -> cs.at_iteration, scs.data) == child_at_iterations @test indicates_convergence(stop_when, scs) end @@ -467,7 +497,7 @@ end stop_when = (converging | budget) & timer algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 0) + state = State(nothing, scs, 0, nothing) @test isempty(get_active_stopping_criteria(algorithm, state)) # iteration 0 starts the timer @@ -485,7 +515,7 @@ end stop_when = StopWhenConverged(100) | budget algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 2) + state = State(nothing, scs, 2, nothing) @test is_finished!(problem, algorithm, state) @test map(first, get_active_stopping_criteria(algorithm, state)) == [budget] end @@ -494,7 +524,7 @@ end stop_when = StopWhenConverged(2) | StopAfterIteration(5) algorithm = AIT.DummyAlgorithm(stop_when) scs = initialize_state(problem, algorithm, stop_when) - state = AIT.DummyState(nothing, scs, 2) + state = State(nothing, scs, 2, nothing) @test is_finished!(problem, algorithm, state) @test get_reason(algorithm, state) == get_reason(stop_when, scs)