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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/CoolPDLP.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/MOI_wrapper.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/algorithms/common.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
82 changes: 56 additions & 26 deletions src/components/termination.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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
Expand All @@ -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""",
Expand All @@ -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,
Expand All @@ -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
3 changes: 2 additions & 1 deletion test/algorithms/all.jl
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
using Adapt
using CoolPDLP
using CoolPDLP: termination_status
using HiGHS: HiGHS
using JLArrays
using KernelAbstractions
Expand Down Expand Up @@ -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
Expand Down
5 changes: 3 additions & 2 deletions test/algorithms/batching.jl
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions test/components/show.jl
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
17 changes: 9 additions & 8 deletions test/components/termination.jl
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
using CoolPDLP
using CoolPDLP: termination_status
import MathOptInterface as MOI
using Random
using SparseArrays
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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

Expand Down
4 changes: 2 additions & 2 deletions test/gpu/batching.jl
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down
11 changes: 10 additions & 1 deletion test/gpu/reactant/runtests.jl
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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