Skip to content
Open
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
156 changes: 89 additions & 67 deletions ext/HSL/ma57_workspace.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,14 +2,17 @@ mutable struct PenaltyMA57Workspace{
WP<:Ma57,
K2<:AbstractMatrix,
V<:AbstractVector,
VI<:Union{Nothing,AbstractVector},
T<:Real,
} <: AbstractHSLWorkspace
M::WP
H::K2
x::V
work::V
_qn_work::V
_info::Base.RefValue{BlasInt} # For CompactBFGS LU factorization
dx::V
_ipiv::VI # For CompactBFGS LU factorization
σ::T
n::Int
m::Int
Expand Down Expand Up @@ -42,6 +45,8 @@ function construct_ma57_workspace(
similar(u1, 4*length(u1)),
similar(u1, 0),
similar(u1),
nothing,
Ref{BlasInt}(0),
zero(T),
n,
m,
Expand Down Expand Up @@ -72,6 +77,8 @@ function construct_ma57_workspace(
similar(u1, 4*length(u1)),
similar(u1, 2*length(u1)*H.B._mem),
similar(u1),
Vector{LinearAlgebra.BlasInt}(undef, 2 * H.B._mem),
Ref{BlasInt}(0),
zero(T),
n,
m,
Expand Down Expand Up @@ -214,12 +221,12 @@ function solve_system!(
B = workspace.H.B
n, m = workspace.n, workspace.m
p = min(B._insert - 1, B._mem)
x1, x2, x3, y1, y2 = H.x1, H.x2, H.x3, H.y1, H.y2
x1, x2, x3, y1 = H.x1, H.x2, H.x3, H.y1
Z1, Z2 = H.Z1, H.Z2
Uk = @view B.Uk[:, 1:p]
Vk = @view B.Vk[:, 1:p]

# Step 0: Write (#TODO: we can use easily use QRMumps instead of LDLFactorization here...)
# Step 0: Write
# [B Aᵀ] = [σI+ξI Aᵀ] + [-U V]([U V])ᵀ
# [A -αI] = [A -αI] + [ 0 0]([0 0])
# Hence,
Expand Down Expand Up @@ -273,79 +280,94 @@ function solve_system!(
return
end

# Step 3: Compute
# y₁ = Fᵀx₁ = [Uᵀx₁(1:n)]
# y₁ = Fᵀx₁ = [Vᵀx₁(1:n)]
@views mul!(y1[1:p], Uk', x1[1:n])
@views mul!(y1[(p+1):(2*p)], Vk', x1[1:n])
if p > 0

# Step 3: Compute
# y₁ = Fᵀx₁ = [Uᵀx₁(1:n)]
# y₁ = Fᵀx₁ = [Vᵀx₁(1:n)]
@views mul!(y1[1:p], Uk', x1[1:n])
@views mul!(y1[(p+1):(2*p)], Vk', x1[1:n])

# Step 4: Assemble Schur complement (I + Fᵀ [σI+ξI Aᵀ]⁻¹ E )
# ( [A -αI] )
# Step 4.1: Compute
# Z₁ = [σI+ξI Aᵀ]⁻¹ E = [σI+ξI Aᵀ]⁻¹[-U V]
# Z₁ = [A -αI] E = [A -αI] [ 0 0]
Z1 .= 0

@views Z1[1:n, 1:p] .= Uk .* (-1)
@views Z1[1:n, (p+1):(2*p)] .= Vk
try
ma57_solve!(workspace.M, Z1, workspace._qn_work)
catch e
!(e isa HSL.Ma57Exception) && rethrow(e)
end
if any(isnan, Z1) || workspace.M.info.info[1] < 0
workspace.status = :failed
return
end
# Step 4: Assemble Schur complement (I + Fᵀ [σI+ξI Aᵀ]⁻¹ E )
# ( [A -αI] )
# Step 4.1: Compute
# Z₁ = [σI+ξI Aᵀ]⁻¹ E = [σI+ξI Aᵀ]⁻¹[-U V]
# Z₁ = [A -αI] E = [A -αI] [ 0 0]
Z1 .= 0

# Step 4.2: Compute
# Z₂ = FᵀZ₁ = UᵀZ₁[1:n]
# Z₂ = FᵀZ₁ = VᵀZ₁[1:n]
Z2 .= 0
@views mul!(Z2[1:p, 1:(2*p)], Uk', Z1[1:n, (1:(2*p))])
@views mul!(Z2[(p+1):(2*p), 1:(2*p)], Vk', Z1[1:n, (1:(2*p))])

# Step 4.3: Compute
# Z₂ = I + Z₂
for i = 1:(2*p)
Z2[i, i] += 1
end
@views Z1[1:n, 1:p] .= Uk .* (-1)
@views Z1[1:n, (p+1):(2*p)] .= Vk
try
ma57_solve!(workspace.M, Z1, workspace._qn_work)
catch e
!(e isa HSL.Ma57Exception) && rethrow(e)
end
if any(isnan, Z1) || workspace.M.info.info[1] < 0
workspace.status = :failed
return
end

# Step 5: Solve
# (I + Fᵀ [σI+ξI Aᵀ]⁻¹ E )⁻¹[y₁]
# ( [A -αI] ) [y₁]
# using Julia LinearALgebra's lu!
F = lu!(Z2[1:(2*p), 1:(2*p)], check = false) # FIXME ?
@views ldiv!(y2[1:(2*p)], F, y1[1:(2*p)])
if any(isnan, y2)
workspace.status = :failed
return
end
# Step 4.2: Compute
# Z₂ = FᵀZ₁ = UᵀZ₁[1:n]
# Z₂ = FᵀZ₁ = VᵀZ₁[1:n]
Z2 .= 0
@views mul!(Z2[1:p, 1:(2*p)], Uk', Z1[1:n, (1:(2*p))])
@views mul!(Z2[(p+1):(2*p), 1:(2*p)], Vk', Z1[1:n, (1:(2*p))])

# Step 4.3: Compute
# Z₂ = I + Z₂
for i = 1:(2*p)
Z2[i, i] += 1
end

# Step 5: Solve
# (I + Fᵀ [σI+ξI Aᵀ]⁻¹ E )⁻¹[y₁]
# ( [A -αI] ) [y₁]
# using LAPACK
@views info_f = getrf!(BlasInt(2p), BlasInt(2p), Z2, stride(Z2, 2), workspace._ipiv, workspace._info)
if info_f != 0
workspace.status = :failed
return
end

# Step 6: Compute
# x₂ = E[y₂] = [-U V][y₂] = [-Uy₂ + Vy₂]
# x₂ = E[y₂] = [ 0 0][y₂] = [0]
@views mul!(x2[1:n], Vk, y2[(p+1):(2*p)])
@views mul!(x2[1:n], Uk, y2[1:p], -one(eltype(y2)), one(eltype(y2)))
@views info_s = getrs!('N', BlasInt(2p), BlasInt(1), Z2[1:(2p), 1:(2p)], stride(Z2, 2),
workspace._ipiv, y1[1:(2p)], BlasInt(2p), workspace._info)
if info_s != 0
workspace.status = :failed
return
end
if any(isnan, @view y1[1:(2*p)])
workspace.status = :failed
return
end

# Step 7: Solve
# [x₃] = [σI+ξI Aᵀ]⁻¹[x₂]
# [x₃] = [A -αI] [x₂]
try
ma57_solve!(workspace.M, x2, x3, workspace.dx, workspace.work, 10)
catch e
!(e isa HSL.Ma57Exception) && rethrow(e)
end
if any(isnan, x3) || workspace.M.info.info[1] < 0
workspace.status = :failed
return
end
# Step 6: Compute
# x₂ = E[y₂] = [-U V][y₂] = [-Uy₂ + Vy₂]
# x₂ = E[y₂] = [ 0 0][y₂] = [0]
@views mul!(x2[1:n], Vk, y1[(p+1):(2*p)])
@views mul!(x2[1:n], Uk, y1[1:p], -one(eltype(y1)), one(eltype(y1)))

# Step 8:
# [B Aᵀ]⁻¹[u] = x₁ - x₃
# [A -αI] [u] = x₁ - x₃
workspace.x .= x1 .- x3
# Step 7: Solve
# [x₃] = [σI+ξI Aᵀ]⁻¹[x₂]
# [x₃] = [A -αI] [x₂]
try
ma57_solve!(workspace.M, x2, x3, workspace.dx, workspace.work, 10)
catch e
!(e isa HSL.Ma57Exception) && rethrow(e)
end
if any(isnan, x3) || workspace.M.info.info[1] < 0
workspace.status = :failed
return
end

# Step 8:
# [B Aᵀ]⁻¹[u] = x₁ - x₃
# [A -αI] [u] = x₁ - x₃
workspace.x .= x1 .- x3
else
workspace.x .= x1
end
end

function get_solution!(
Expand Down
Loading
Loading