diff --git a/src/CoolPDLP.jl b/src/CoolPDLP.jl index 84da9be..3336822 100644 --- a/src/CoolPDLP.jl +++ b/src/CoolPDLP.jl @@ -65,6 +65,7 @@ export preprocess, initialize, solve, solve! export PDHG, PDLP @public Algorithm @public KKTErrors, relative +@public termination_status export is_feasible, objective_value @public Optimizer diff --git a/src/MOI_wrapper.jl b/src/MOI_wrapper.jl index 32a3dc1..2fa0b9f 100644 --- a/src/MOI_wrapper.jl +++ b/src/MOI_wrapper.jl @@ -254,7 +254,7 @@ function MOI.optimize!(dest::Optimizer{T}, fcache::MOI.Utilities.UniversalFallba dest.dual_obj_value = (max_sense ? -raw_dual_obj : raw_dual_obj) + obj_constant dest.solve_time = stats.time_elapsed - cts = stats.termination_status + cts = termination_status(stats) @assert cts != MOI.OPTIMIZE_NOT_CALLED "solve did not reach a terminal status" dest.termination_status = cts dest.primal_status, dest.dual_status = if cts == MOI.OPTIMAL diff --git a/src/algorithms/common.jl b/src/algorithms/common.jl index d749542..5647f7c 100644 --- a/src/algorithms/common.jl +++ b/src/algorithms/common.jl @@ -255,7 +255,7 @@ function try_solve_noconstraints!(state::AbstractState, milp::MILP) @. sol.x = ifelse(c > 0, lv, ifelse(c < 0, uv, clamp(zero(eltype(lv)), lv, uv))) kkt_errors!(state.stats.err, state.scratch, sol, milp) state.stats.time_elapsed = current_time() - state.stats.starting_time - state.stats.termination_status = MOI.OPTIMAL + state.stats.termination_status_code = status_code(MOI.OPTIMAL) return true else return false diff --git a/src/components/termination.jl b/src/components/termination.jl index a85b5f5..24a382c 100644 --- a/src/components/termination.jl +++ b/src/components/termination.jl @@ -40,7 +40,7 @@ end $(TYPEDFIELDS) """ -mutable struct ConvergenceStats{T <: BatchedNumber, F <: Number, I <: Number} +mutable struct ConvergenceStats{T <: BatchedNumber, F <: Number, I <: Number, S <: Number} "current KKT error" err::KKTErrors{T} "time at which the algorithm started, in seconds" @@ -49,8 +49,8 @@ mutable struct ConvergenceStats{T <: BatchedNumber, F <: Number, I <: Number} time_elapsed::F "number of multiplications by both the KKT matrix and its transpose" kkt_passes::I - "termination status (should be `MOI.OPTIMIZE_NOT_CALLED` until the algorithm actually terminates)" - termination_status::MOI.TerminationStatusCode + "integer code of the termination status (should be that of `MOI.OPTIMIZE_NOT_CALLED` until the algorithm actually terminates), read with [`termination_status`](@ref)" + termination_status_code::S "history of KKT errors, indexed by number of KKT passes" error_history::Vector{Tuple{I, KKTErrors{T}}} end @@ -60,36 +60,60 @@ function ConvergenceStats( starting_time = current_time(), time_elapsed = 0.0, kkt_passes::I = 0, - termination_status = MOI.OPTIMIZE_NOT_CALLED, + termination_status_code::S = status_code(MOI.OPTIMIZE_NOT_CALLED), error_history = [(kkt_passes, copy(err))] - ) where {T, I} + ) where {T, I, S} F = Base.promote_type(typeof(starting_time), typeof(time_elapsed)) - return ConvergenceStats{T, F, I}( + return ConvergenceStats{T, F, I, S}( err, starting_time, time_elapsed, kkt_passes, - termination_status, + termination_status_code, error_history ) end +""" + status_code(status) + +Return the integer code of a `MOI.TerminationStatusCode`, as stored in the +`termination_status_code` field of [`ConvergenceStats`](@ref). + +Unlike the enum itself, this code is an ordinary number, so Reactant can trace it: a compiled +solve can write its own status instead of freezing the trace-time one. +""" +status_code(status::MOI.TerminationStatusCode) = Int32(status) + +""" + termination_status(stats) + +Return the termination status of `stats` as a `MOI.TerminationStatusCode`. + +Decodes the `termination_status_code` field, which is stored as a plain number so that it +survives a Reactant-compiled solve. Only call this outside a compilation context: mid-trace the +code is a traced number with no value yet. +""" +function termination_status(stats::ConvergenceStats) + return MOI.TerminationStatusCode(Int32(stats.termination_status_code)) +end + function instance(stats::ConvergenceStats, i::Int) return ConvergenceStats( instance(stats.err, i); starting_time = stats.starting_time, time_elapsed = stats.time_elapsed, kkt_passes = stats.kkt_passes, - termination_status = stats.termination_status, + termination_status_code = stats.termination_status_code, error_history = [(passes, instance(err, i)) for (passes, err) in stats.error_history], ) end function Base.show(io::IO, stats::ConvergenceStats) - (; err, time_elapsed, kkt_passes, termination_status) = stats + (; err, time_elapsed, kkt_passes) = stats return print( io, - """Convergence stats with termination status $termination_status: + """Convergence stats with termination status $(termination_status(stats)): - $err - time elapsed: $time_elapsed seconds - KKT passes: $kkt_passes""", @@ -101,9 +125,13 @@ end Decide how the algorithm terminates, using `dest` as scratch space for the relative errors. -Set `stats.termination_status` and return whether the algorithm should stop. The returned -boolean is traced under Reactant, unlike the `MOI.TerminationStatusCode` enum, so it is what -the solve loops branch on. +Set `stats.termination_status_code` and return whether the algorithm should stop. + +The status is selected with nested `ifelse` calls rather than with branches, so that it survives +a Reactant-compiled solve: an assignment inside a `@trace if` only ever contributes its +trace-time value, whereas an `ifelse` over traced conditions is part of the compiled program. +The returned boolean is what the solve loops branch on, since a status code cannot drive traced +control flow. """ function set_termination_status!!( stats::ConvergenceStats, @@ -115,18 +143,20 @@ function set_termination_status!!( is_optimal = batched_all(<=(termination_reltol), relative!!(dest, err)) is_time_limit = time_elapsed >= time_limit is_iteration_limit = kkt_passes >= max_kkt_passes - # Reactant doesn't like `elseif`, see https://github.com/EnzymeAD/Reactant.jl/issues/2563#issuecomment-5584197336 - # The branches are ordered by increasing priority, so that the last write wins. - @trace if is_iteration_limit - stats.termination_status = MOI.ITERATION_LIMIT - end - @trace if is_time_limit - stats.termination_status = MOI.TIME_LIMIT - end - @trace if is_optimal - stats.termination_status = MOI.OPTIMAL - end - # `stats.termination_status` is a plain enum, so it cannot drive traced control flow. - # Return the decision as a (possibly traced) boolean instead. + # nested `ifelse` instead of `if`/`elseif`: it is a plain (traceable) computation, and it + # states the priority between simultaneous criteria in one place + stats.termination_status_code = ifelse( + is_optimal, + status_code(MOI.OPTIMAL), + ifelse( + is_time_limit, + status_code(MOI.TIME_LIMIT), + ifelse( + is_iteration_limit, + status_code(MOI.ITERATION_LIMIT), + status_code(MOI.OPTIMIZE_NOT_CALLED), + ), + ), + ) return is_optimal | is_time_limit | is_iteration_limit end diff --git a/test/algorithms/all.jl b/test/algorithms/all.jl index 9a49447..95c59bb 100644 --- a/test/algorithms/all.jl +++ b/test/algorithms/all.jl @@ -1,5 +1,6 @@ using Adapt using CoolPDLP +using CoolPDLP: termination_status using HiGHS: HiGHS using JLArrays using KernelAbstractions @@ -33,7 +34,7 @@ function test_optimizer( sol, stats = solve(milp, algo) x = sol.x - @test stats.termination_status == MOI.OPTIMAL + @test termination_status(stats) == MOI.OPTIMAL @test is_feasible(Array(x), milp; cons_tol, int_tol) @test isapprox(objective_value(jump_x, milp), objective_value(Array(x), milp); rtol = obj_rtol) return nothing diff --git a/test/algorithms/batching.jl b/test/algorithms/batching.jl index 1b05e71..5a4f185 100644 --- a/test/algorithms/batching.jl +++ b/test/algorithms/batching.jl @@ -1,6 +1,7 @@ using CoolPDLP using CoolPDLP: KKTErrors, Scratch, initialize, instance, kkt_errors!, - nbinstances, prog_showvalues, relative, restart!, restart_check!, step! + nbinstances, prog_showvalues, relative, restart!, restart_check!, step!, + termination_status using Random using Test @@ -138,7 +139,7 @@ end sol, stats = solve(milp_id, sol_id, algo) sol_single, stats_single = solve(milps[1], sols[1], algo) @test stats.kkt_passes == stats_single.kkt_passes - @test stats.termination_status == stats_single.termination_status + @test termination_status(stats) == termination_status(stats_single) for i in 1:NBATCH @test sol.x[:, i] ≈ sol_single.x @test stats.err.primal[i] ≈ stats_single.err.primal diff --git a/test/components/show.jl b/test/components/show.jl index cbe448e..f14417c 100644 --- a/test/components/show.jl +++ b/test/components/show.jl @@ -1,5 +1,5 @@ using CoolPDLP -using CoolPDLP: KKTErrors, Scratch, kkt_errors! +using CoolPDLP: KKTErrors, Scratch, kkt_errors!, termination_status using Random using SparseArrays using Test @@ -37,7 +37,7 @@ end _, stats = solve(milp, PDLP(; max_kkt_passes = 200)) str = sprint(show, stats) @test occursin("Convergence stats", str) - @test occursin(string(stats.termination_status), str) + @test occursin(string(termination_status(stats)), str) @test occursin("KKT passes: $(stats.kkt_passes)", str) @test occursin("KKT relative errors", str) end diff --git a/test/components/termination.jl b/test/components/termination.jl index 7466798..0118f0e 100644 --- a/test/components/termination.jl +++ b/test/components/termination.jl @@ -1,4 +1,5 @@ using CoolPDLP +using CoolPDLP: termination_status import MathOptInterface as MOI using Random using SparseArrays @@ -12,7 +13,7 @@ using Test milp = CoolPDLP.MILP(; c, lv, uv, A, lc, uc) algo = CoolPDLP.PDLP() sol, stats = CoolPDLP.solve(milp, algo) - @test stats.termination_status == MOI.OPTIMAL + @test termination_status(stats) == MOI.OPTIMAL end @testset "Termination statuses" begin @@ -21,11 +22,11 @@ end @testset "$alg" for alg in (PDHG, PDLP) _, stats = solve(milp, alg(; termination_reltol = 0.0, max_kkt_passes = 200)) - @test stats.termination_status == MOI.ITERATION_LIMIT + @test termination_status(stats) == MOI.ITERATION_LIMIT @test stats.kkt_passes >= 200 _, stats = solve(milp, alg(; termination_reltol = 0.0, time_limit = 0.0)) - @test stats.termination_status == MOI.TIME_LIMIT + @test termination_status(stats) == MOI.TIME_LIMIT @test stats.time_elapsed >= 0 end end @@ -41,7 +42,7 @@ end @testset "$alg" for alg in (PDHG, PDLP) sol, stats = solve(milp, alg()) - @test stats.termination_status == MOI.OPTIMAL + @test termination_status(stats) == MOI.OPTIMAL @test sol.x == clamp.(0.0, lv, uv) @test is_feasible(sol.x, milp) @test objective_value(sol.x, milp) == 0 @@ -59,7 +60,7 @@ end @testset "$alg" for alg in (PDHG, PDLP) sol, stats = solve(milp, alg()) - @test stats.termination_status == MOI.OPTIMAL + @test termination_status(stats) == MOI.OPTIMAL @test !any(isnan, sol.x) @test sol.x == [0.0, 5.0, 0.0] @test objective_value(sol.x, milp) == -5.0 @@ -82,7 +83,7 @@ end ) sol, stats = solve(milp, algo) @test !any(isnan, sol.x) && !any(isinf, sol.x) - @test stats.termination_status != MOI.OPTIMAL + @test termination_status(stats) != MOI.OPTIMAL end @testset "unbounded direction (c[1] > 0, lv[1] == -Inf)" begin @@ -91,7 +92,7 @@ end ) sol, stats = solve(milp, algo) @test !any(isnan, sol.x) && !any(isinf, sol.x) - @test stats.termination_status != MOI.OPTIMAL + @test termination_status(stats) != MOI.OPTIMAL end @testset "unbounded direction (c[1] < 0, uv[1] == Inf)" begin @@ -100,7 +101,7 @@ end ) sol, stats = solve(milp, algo) @test !any(isnan, sol.x) && !any(isinf, sol.x) - @test stats.termination_status != MOI.OPTIMAL + @test termination_status(stats) != MOI.OPTIMAL end end diff --git a/test/gpu/batching.jl b/test/gpu/batching.jl index c34ae86..4af8f2c 100644 --- a/test/gpu/batching.jl +++ b/test/gpu/batching.jl @@ -1,6 +1,6 @@ using CoolPDLP using CoolPDLP: KKTErrors, Scratch, initialize, instance, kkt_errors!, - nbinstances, preprocess, relative, step! + nbinstances, preprocess, relative, step!, termination_status using GPUArraysCore: @allowscalar using Random using Test @@ -91,7 +91,7 @@ function test_batching( algo_solve = alg(T, Int, matrix_type; backend, max_kkt_passes = 200) sol, stats = solve(milp_id, algo_solve) sol_single, stats_single = solve(milps[1], algo_solve) - @test stats.termination_status == stats_single.termination_status + @test termination_status(stats) == termination_status(stats_single) x, x_single = Array(sol.x), Array(sol_single.x) obj_single = objective_value(x_single, milps[1]) for i in 1:nbatch diff --git a/test/gpu/reactant/runtests.jl b/test/gpu/reactant/runtests.jl index 2c5940b..827fb81 100644 --- a/test/gpu/reactant/runtests.jl +++ b/test/gpu/reactant/runtests.jl @@ -1,5 +1,6 @@ using CoolPDLP -using CoolPDLP: KKTErrors +using CoolPDLP: KKTErrors, termination_status +import MathOptInterface as MOI using MathOptBenchmarkInstances using Reactant using Reactant: to_rarray @@ -113,6 +114,13 @@ configs = [ @test isapprox(err_r, err; rtol) end end + + @testset "Same termination status" begin + # the status is stored as a traced integer code, so the compiled run reports the + # criterion it actually stopped on instead of the trace-time `OPTIMIZE_NOT_CALLED` + @test termination_status(state_r.stats) != MOI.OPTIMIZE_NOT_CALLED + @test termination_status(state_r.stats) == termination_status(state.stats) + end end # Reading the host clock inside a compiled program needs a Reactant callback, which not every @@ -174,6 +182,7 @@ else @testset "The solve stops on the time limit" begin @test elapsed >= time_limit @test passes < max_kkt_passes + @test termination_status(state_r.stats) == MOI.TIME_LIMIT end end end