Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ KernelAbstractions = "0.9"
LDLFactorizations = "0.10.1"
LinearAlgebra = "1.10"
MadNLP = "0.8.12"
MadNLPGPU = "0.7.15"
MadNLPGPU = "0.7"
MadNLPTests = "0.5.3"
MathOptInterface = "1.42"
NLPModels = "0.21.5"
Expand Down
3 changes: 2 additions & 1 deletion ext/MadIPMCUDAExt/MadIPMCUDAExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -136,5 +136,6 @@ function Base.convert(::Type{QuadraticModel{T, S}}, qp::QuadraticModel{T}) where
)
end

end
include("UniformBatch/UniformBatch.jl")

end
31 changes: 31 additions & 0 deletions ext/MadIPMCUDAExt/UniformBatch/UniformBatch.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
## for each KKT system:
get_rhs_size(kkt::MadNLP.AbstractReducedKKTSystem, vec::MadNLP.UnreducedKKTVector) = length(get_rhs(typeof(kkt), vec))
get_rhs(::Type{<:MadNLP.AbstractReducedKKTSystem}, vec::MadNLP.UnreducedKKTVector) = MadNLP.primal_dual(vec)

# TODO: better modularize; below combine MadIPM.solve_system! and MadNLP.solve!(::AbstractReducedKKTSystem)
pre_solve!(solver::MadIPM.MPCSolver{T,VT,VI,KKTSystem}) where {
T,VT,VI,KKTSystem<:MadNLP.AbstractReducedKKTSystem,
} = begin
copyto!(MadNLP.full(solver.d), MadNLP.full(solver.p))
MadNLP.reduce_rhs!(solver.kkt, solver.d)
return
end
post_solve!(solver::MadIPM.MPCSolver{T,VT,VI,KKTSystem}) where {
T,VT,VI,KKTSystem<:MadNLP.AbstractReducedKKTSystem,
} = begin
MadNLP.finish_aug_solve!(solver.kkt, solver.d)
MadIPM.post_solve!(solver.d, solver, solver.p)
end

## dummy solver to make sure batch solve uses batch solver only
Comment thread
klamike marked this conversation as resolved.
struct NoLinearSolver{T} <: MadNLP.AbstractLinearSolver{T} end
NoLinearSolver(A; kwargs...) = NoLinearSolver{Float64}()
MadNLP.default_options(::Type{NoLinearSolver}) = nothing
MadNLP.set_options!(::Nothing, x) = x
MadNLP.is_supported(::Type{NoLinearSolver}, ::Type{T}) where {T<:AbstractFloat} = true


include("kkt.jl")
include("structure.jl")
include("broadcast.jl")
include("solver.jl")
24 changes: 24 additions & 0 deletions ext/MadIPMCUDAExt/UniformBatch/broadcast.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
struct ActiveSolversStyle <: Base.Broadcast.BroadcastStyle end
struct ActiveSolversIterator{S}
batch_solver::UniformBatchSolver{S}
end

Base.Broadcast.BroadcastStyle(::Type{<:UniformBatchSolver}) = ActiveSolversStyle()
Base.Broadcast.BroadcastStyle(::ActiveSolversStyle, ::Base.Broadcast.DefaultArrayStyle{0}) = ActiveSolversStyle()
Base.Broadcast.instantiate(bc::Base.Broadcast.Broadcasted{ActiveSolversStyle}) = bc
Base.broadcastable(batch_solver::UniformBatchSolver) = ActiveSolversIterator(batch_solver)
Base.axes(iter::ActiveSolversIterator) = (Base.OneTo(iter.batch_solver.bkkt.active_batch_size[]),)
Base.ndims(::Type{<:ActiveSolversIterator}) = 1
Base.getindex(iter::ActiveSolversIterator, i::Int) = iter.batch_solver.solvers[iter.batch_solver.bkkt.batch_map_rev[i]]

# FIXME: this works but does not seem to be the correct way to do it
@inline function Base.Broadcast.materialize(bc::Base.Broadcast.Broadcasted{ActiveSolversStyle})
iter = bc.args[1]
bkkt = iter.batch_solver.bkkt
active_i = 1
while (solver_idx = bkkt.batch_map_rev[active_i]) != 0
bc.f(iter.batch_solver.solvers[solver_idx])
active_i += 1
Comment thread
klamike marked this conversation as resolved.
Outdated
end
return nothing
end
144 changes: 144 additions & 0 deletions ext/MadIPMCUDAExt/UniformBatch/kkt.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,144 @@
abstract type AbstractBatchKKTSystem{KKTSystem,LS} end
Comment thread
klamike marked this conversation as resolved.

Comment thread
klamike marked this conversation as resolved.
struct UniformBatchKKTSystem{ # NOTE: move to MadIPM/MadNLP
Comment thread
klamike marked this conversation as resolved.
KKTSystem<:MadNLP.AbstractSparseKKTSystem, LS, T, VT<:AbstractVector{T}, VI
} <: AbstractBatchKKTSystem{KKTSystem,LS}
kkts::Vector{KKTSystem}
vecs::Vector{MadNLP.UnreducedKKTVector{T,VT,VI}}
linear_solver::LS
batch_rhs::VT
batch_nzVal::VT
rhs_slices::Vector{VT}
nzVal_slices::Vector{VT}
rhs_size::Int
nzVal_size::Int
batch_map::Vector{Int}
batch_map_rev::Vector{Int}
batch_size::Int
active_batch_size::Base.RefValue{Int}
is_active::BitVector
active_rhs::Base.RefValue{VT}
end

all_done(bkkt::UniformBatchKKTSystem) = !any(bkkt.is_active)
is_active(bkkt::UniformBatchKKTSystem, i) = bkkt.is_active[i]

function UniformBatchKKTSystem( # NOTE: move to MadNLPGPU
kkts::Vector{KKTSystem},
vecs::Vector{KKTVector},
linear_solver::Type{MadNLPGPU.CUDSSSolver};
opt_linear_solver=MadNLP.default_options(linear_solver),
) where {T,VT,KKTSystem<:MadNLP.AbstractSparseKKTSystem,KKTVector<:MadNLP.UnreducedKKTVector{T,VT}}
kkt1 = first(kkts)

for kkt in kkts
@assert kkt.aug_com.colPtr == kkt1.aug_com.colPtr "Cannot use UniformBatchKKTSystem when KKTSystems do not share sparsity structure (colPtr)."
@assert kkt.aug_com.rowVal == kkt1.aug_com.rowVal "Cannot use UniformBatchKKTSystem when KKTSystems do not share sparsity structure (rowVal)."
end
Comment thread
klamike marked this conversation as resolved.

vec1 = first(vecs)
rhs_size = get_rhs_size(kkt1, vec1)
nzVal_size = length(nonzeros(kkt1.aug_com))

batch_size = length(kkts)
batch_rhs = fill!(VT(undef, rhs_size * batch_size), zero(T))
batch_nzVal = fill!(VT(undef, nzVal_size * batch_size), zero(T))

rhs_slices = [MadNLP._madnlp_unsafe_wrap(batch_rhs, rhs_size, (j - 1) * rhs_size + 1) for j in 1:batch_size]
nzVal_slices = [MadNLP._madnlp_unsafe_wrap(batch_nzVal, nzVal_size, (j - 1) * nzVal_size + 1) for j in 1:batch_size]

batch_aug_com = similar(kkt1.aug_com)
batch_aug_com.nzVal = batch_nzVal
_linear_solver = linear_solver(
batch_aug_com; opt = opt_linear_solver
)

bkkt = UniformBatchKKTSystem(
kkts, vecs, _linear_solver,
batch_rhs, batch_nzVal,
rhs_slices, nzVal_slices,
rhs_size, nzVal_size,
collect(1:batch_size), collect(1:batch_size),
batch_size, Ref(batch_size),
trues(batch_size),
Ref(MadNLP._madnlp_unsafe_wrap(batch_rhs, rhs_size * batch_size, 1)),
)

update_pointers!(bkkt)

return bkkt
end

function update_batch!(bkkt::UniformBatchKKTSystem{KKTSystem,LS}) where {KKTSystem,LS<:MadNLPGPU.CUDSSSolver}
# NOTE: only called if an update is needed
active_pos = 0
for i in 1:bkkt.batch_size
if bkkt.is_active[i]
active_pos += 1
bkkt.batch_map[i] = active_pos
bkkt.batch_map_rev[active_pos] = i
else
bkkt.batch_map[i] = 0
end
end

for j in (active_pos + 1):bkkt.batch_size
bkkt.batch_map_rev[j] = 0
end

bkkt.active_batch_size[] = active_pos
bkkt.active_rhs[] = MadNLP._madnlp_unsafe_wrap(bkkt.batch_rhs, active_pos * bkkt.rhs_size, 1)
bkkt.linear_solver.tril.nzVal = MadNLP._madnlp_unsafe_wrap(bkkt.batch_nzVal, active_pos * bkkt.nzVal_size, 1)

update_pointers!(bkkt)

MadNLPGPU.CUDSS.cudss_set(bkkt.linear_solver.inner, "ubatch_size", active_pos)
return
end

function update_pointers!(bkkt::UniformBatchKKTSystem{KKTSystem,LS}) where {KKTSystem,LS<:MadNLPGPU.CUDSSSolver}
for (i, kkt_i) in enumerate(bkkt.kkts)
batch_pos = bkkt.batch_map[i]
if batch_pos > 0
kkt_i.aug_com.nzVal = bkkt.nzVal_slices[batch_pos]
end
end
return
end

function batch_factorize!(bkkt::UniformBatchKKTSystem{KKTSystem,LS}) where {KKTSystem,LS<:MadNLPGPU.CUDSSSolver}
MadNLP.factorize!(bkkt.linear_solver)
return
end

function batch_solve!(bkkt::UniformBatchKKTSystem{KKTSystem,LS}) where {KKTSystem,LS<:MadNLPGPU.CUDSSSolver}
copy_batch_rhs!(bkkt)
MadNLP.solve!(bkkt.linear_solver, bkkt.active_rhs[])
copy_batch_solution!(bkkt)
return
end

# TODO: can be better (each copyto syncs)
Comment thread
klamike marked this conversation as resolved.
function copy_batch_rhs!(bkkt::UniformBatchKKTSystem{KKTSystem,LS}) where {KKTSystem,LS<:MadNLPGPU.CUDSSSolver}
for active_i in 1:bkkt.active_batch_size[]
dest = bkkt.rhs_slices[active_i]
vec = get_active_vec(bkkt, active_i)
src = get_rhs(KKTSystem, vec)
copyto!(dest, src)
end
return
end

function copy_batch_solution!(bkkt::UniformBatchKKTSystem{KKTSystem,LS}) where {KKTSystem,LS<:MadNLPGPU.CUDSSSolver}
Comment thread
klamike marked this conversation as resolved.
for active_i in 1:bkkt.active_batch_size[]
vec = get_active_vec(bkkt, active_i)
dest = get_rhs(KKTSystem, vec)
src = bkkt.rhs_slices[active_i]
copyto!(dest, src)
end
return
end

function get_active_vec(bkkt::UniformBatchKKTSystem{KKTSystem,LS}, active_i) where {KKTSystem,LS<:MadNLPGPU.CUDSSSolver}
return bkkt.vecs[bkkt.batch_map_rev[active_i]]
end
94 changes: 94 additions & 0 deletions ext/MadIPMCUDAExt/UniformBatch/solver.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
batch_init_starting_point_solve!(batch_solver::AbstractBatchSolver) = begin
MadIPM.set_initial_regularization!.(batch_solver)
MadNLP.build_kkt!.(batch_solver)
batch_factorize!(batch_solver.bkkt)

MadIPM.set_initial_primal_rhs!.(batch_solver)
batch_solve_system!(batch_solver)
MadIPM.update_primal_start!.(batch_solver)

MadIPM.set_initial_dual_rhs!.(batch_solver)
batch_solve_system!(batch_solver)
return
end
batch_initialize!(batch_solver::AbstractBatchSolver) = begin
MadIPM.pre_initialize!.(batch_solver)
batch_init_starting_point_solve!(batch_solver)
MadIPM.post_initialize!.(batch_solver)
return
end
batch_factorize_regularized_system!(batch_solver::AbstractBatchSolver) = begin
MadIPM.set_aug_diagonal_reg!.(batch_solver)
MadNLP.build_kkt!.(batch_solver)
batch_factorize!(batch_solver.bkkt)
return
end
batch_solve_system!(batch_solver::AbstractBatchSolver) = begin
pre_solve!.(batch_solver)
batch_solve!(batch_solver.bkkt)
post_solve!.(batch_solver)
return
end

function batch_mpc!(batch_solver::AbstractBatchSolver)
while true
# Check termination criteria
MadNLP.print_iter.(batch_solver)
MadIPM.update_termination_criteria!.(batch_solver)
update_batch!(batch_solver)
all_done(batch_solver) && return

# Factorize KKT system
MadIPM.update_regularization!.(batch_solver)
batch_factorize_regularized_system!(batch_solver)

# Affine direction
MadIPM.set_predictive_rhs!.(batch_solver)
batch_solve_system!(batch_solver)

# Prediction step size
MadIPM.prediction_step_size!.(batch_solver)

# Mehrotra's Correction direction
MadIPM.set_correction_rhs!.(batch_solver)
batch_solve_system!(batch_solver)

# Gondzio's additional correction direction FIXME
Comment thread
klamike marked this conversation as resolved.
Outdated
# batch_gondzio_correction_direction!(batch_solver)

# Update step size
MadIPM.update_step_size!.(batch_solver)

# Apply step
MadIPM.apply_step!.(batch_solver)

# Evaluate model at new iterate
MadIPM.evaluate_model!.(batch_solver)
end
end


function MadIPM.solve!(batch_solver::AbstractBatchSolver)
batch_stats = [MadNLP.MadNLPExecutionStats(solver) for solver in batch_solver] # TODO: BatchExecutionStats?

try
MadNLP.@notice(first(batch_solver).logger,"This is MadIPM, running with $(MadNLP.introduce(batch_solver.bkkt.linear_solver)), batch size $(length(batch_solver))\n")
batch_initialize!(batch_solver)
batch_mpc!(batch_solver)
catch e
rethrow(e) # FIXME
finally
for (stats, solver) in zip(batch_stats, batch_solver.solvers)
MadIPM.finalize!(stats, solver)
end
end

return batch_stats
end


function MadIPM.madipm(ms::AbstractVector{NLPModel}; kwargs...) where {NLPModel <: NLPModels.AbstractNLPModel}
solvers = MadIPM.MPCSolver.(ms; linear_solver = NoLinearSolver, kwargs...) # TODO: special constructor to share kkt/cb memory/set NoLinearSolver
Comment thread
klamike marked this conversation as resolved.
batch_solver = UniformBatchSolver(solvers, linear_solver = MadNLPGPU.CUDSSSolver) # TODO: add some detection for the best BatchSolver to use (for now we only have UniformBatchSolver anyway)
return MadIPM.solve!(batch_solver)
end
41 changes: 41 additions & 0 deletions ext/MadIPMCUDAExt/UniformBatch/structure.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,41 @@
abstract type AbstractBatchSolver end
Comment thread
klamike marked this conversation as resolved.

struct UniformBatchSolver{VS} <: AbstractBatchSolver
solvers::VS
bkkt::UniformBatchKKTSystem

Comment thread
klamike marked this conversation as resolved.
Outdated
function UniformBatchSolver(solvers::Vector{Solver}; linear_solver::Type) where {Solver<:MadIPM.MPCSolver}
batch_size = length(solvers)
solver1 = first(solvers)
kkt1 = solver1.kkt
vec1 = solver1.d

Comment thread
klamike marked this conversation as resolved.
Outdated
kkts = Vector{typeof(kkt1)}(undef, batch_size)
vecs = Vector{typeof(vec1)}(undef, batch_size)
for i in 1:batch_size
solver_i = solvers[i]
kkts[i] = solver_i.kkt
vecs[i] = solver_i.d
end

Comment thread
klamike marked this conversation as resolved.
Outdated
return new{Vector{Solver}}(solvers, UniformBatchKKTSystem(kkts, vecs, linear_solver))
end
end
Comment thread
klamike marked this conversation as resolved.

all_done(batch_solver::UniformBatchSolver) = all_done(batch_solver.bkkt)
is_active(batch_solver::UniformBatchSolver, i) = is_active(batch_solver.bkkt, i)
Base.length(batch_solver::UniformBatchSolver) = length(batch_solver.solvers)
Base.iterate(batch_solver::UniformBatchSolver, i=1) = iterate(batch_solver.solvers, i)
Base.getindex(batch_solver::UniformBatchSolver, i) = batch_solver.solvers[i]

update_batch!(batch_solver::UniformBatchSolver) = begin
needs_update = false
for (i, solver) in enumerate(batch_solver)
if is_active(batch_solver, i) && MadIPM.is_done(solver)
needs_update = true
batch_solver.bkkt.is_active[i] = false
end
end
needs_update && update_batch!(batch_solver.bkkt)
return
end
Comment thread
klamike marked this conversation as resolved.
Outdated
11 changes: 9 additions & 2 deletions src/linear_solver.jl
Original file line number Diff line number Diff line change
Expand Up @@ -3,10 +3,12 @@
Interface to direct solver for solving KKT system
=#

MadNLP.build_kkt!(solver::MadIPM.MPCSolver) = MadNLP.build_kkt!(solver.kkt)
set_aug_diagonal_reg!(solver) = set_aug_diagonal_reg!(solver.kkt, solver)
function factorize_regularized_system!(solver)
max_trials = 3
for ntrial in 1:max_trials
set_aug_diagonal_reg!(solver.kkt, solver)
set_aug_diagonal_reg!(solver)
MadNLP.factorize_wrapper!(solver)
if is_factorized(solver.kkt.linear_solver)
break
Expand All @@ -21,9 +23,14 @@ function solve_system!(
solver::MadNLP.AbstractMadNLPSolver{T},
p::MadNLP.UnreducedKKTVector{T},
) where T
opt = solver.opt
copyto!(MadNLP.full(d), MadNLP.full(p))
MadNLP.solve!(solver.kkt, d)
post_solve!(d, solver, p)
return d
end

function post_solve!(d::MadNLP.UnreducedKKTVector{T}, solver, p) where T
opt = solver.opt

# Check residual
w = solver._w1
Expand Down
Loading
Loading