diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index a8ca3333e..7fc474e04 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -108,7 +108,7 @@ jobs: fail-fast: false matrix: julia-version: ['1'] - gpu: [cuda, amdgpu] + gpu: [cuda, amdgpu, oneapi] runs-on: - self-hosted - ${{ matrix.gpu }} @@ -118,13 +118,30 @@ jobs: with: channel: ${{ matrix.julia-version }} - uses: julia-actions/cache@v1 - - name: test MadNLPGPU - shell: julia --project=./lib/MadNLPGPU --color=yes {0} + - name: Setup test environment + shell: julia --project=./lib/MadNLPGPU/test --color=yes {0} run: | using Pkg Pkg.Registry.update() Pkg.develop(path=".") - Pkg.test("MadNLPGPU", coverage=true) + Pkg.develop(path="./lib/MadNLPGPU") + Pkg.develop(path="./lib/MadNLPTests") + - name: Add AMDGPU + if: matrix.gpu == 'amdgpu' + shell: julia --project=./lib/MadNLPGPU/test --color=yes {0} + run: | + using Pkg + Pkg.add("AMDGPU") + - name: Add oneAPI + if: matrix.gpu == 'oneapi' + shell: julia --project=./lib/MadNLPGPU/test --color=yes {0} + run: | + using Pkg + Pkg.add("oneAPI") + - name: test MadNLPGPU + env: + MADNLP_GPU_BACKEND: ${{ matrix.gpu }} + run: julia --project=./lib/MadNLPGPU/test --color=yes -e 'include("lib/MadNLPGPU/test/runtests.jl")' - uses: julia-actions/julia-processcoverage@v1 with: directories: lib/MadNLPGPU/src diff --git a/lib/MadNLPGPU/Project.toml b/lib/MadNLPGPU/Project.toml index c5c1ddabb..10a52c7c8 100644 --- a/lib/MadNLPGPU/Project.toml +++ b/lib/MadNLPGPU/Project.toml @@ -6,6 +6,7 @@ version = "0.7.18" AMD = "14f7f29c-3bd6-536c-9a0b-7339e30b5a3e" CUDA = "052768ef-5323-5732-b1bb-66c8b64840ba" CUDSS = "45b445bb-4962-46a0-9369-b4df9d0f772e" +GPUArraysCore = "46192b85-c4d5-4398-a991-12ede77f4527" KernelAbstractions = "63c18a36-062a-441e-b654-da1e3ab1ce7c" LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" MadNLP = "2621e9c9-9eb4-46b1-8089-e8c72242dfb6" @@ -14,24 +15,20 @@ SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf" [weakdeps] AMDGPU = "21141c5a-9bdb-4563-92ae-f87d6854732e" +oneAPI = "8f75cd03-7ff8-4ecb-9b8f-daf728133b1b" [extensions] MadNLPGPUAMDGPUExt = "AMDGPU" +MadNLPGPUOneAPIExt = "oneAPI" [compat] AMD = "0.5" AMDGPU = "2" CUDA = "5.4.0" CUDSS = "0.6.4" +GPUArraysCore = "0.2" KernelAbstractions = "0.9" MadNLP = "0.8.12" -MadNLPTests = "0.5.3" Metis = "1" julia = "1.10" - -[extras] -MadNLPTests = "b52a2a03-04ab-4a5f-9698-6a2deff93217" -Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" - -[targets] -test = ["Test", "MadNLPTests", "AMDGPU"] +oneAPI = "2.6.0" diff --git a/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/MadNLPGPUOneAPIExt.jl b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/MadNLPGPUOneAPIExt.jl new file mode 100644 index 000000000..e40cb9b60 --- /dev/null +++ b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/MadNLPGPUOneAPIExt.jl @@ -0,0 +1,26 @@ +module MadNLPGPUOneAPIExt + +import LinearAlgebra +import SparseArrays: SparseMatrixCSC, nonzeros, nnz +import LinearAlgebra: Symmetric + +import MadNLP +import MadNLPGPU + +import KernelAbstractions: synchronize +import GPUArraysCore: @allowscalar + +using oneAPI +using oneAPI.oneMKL, oneAPI.Support + +function __init__() + setglobal!(MadNLPGPU, :LapackOneMKLSolver, LapackOneMKLSolver) + return +end + +include("oneapi_dense.jl") +include("oneapi_sparse.jl") +include("onemkl.jl") +include("oneapi.jl") + +end diff --git a/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi.jl b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi.jl new file mode 100644 index 000000000..5b4fc7e28 --- /dev/null +++ b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi.jl @@ -0,0 +1,106 @@ +#= + MadNLP.MadNLPOptions +=# + +function MadNLP.MadNLPOptions{T}( + nlp::MadNLP.AbstractNLPModel{T, VT}; + dense_callback = MadNLP.is_dense_callback(nlp), + callback = dense_callback ? MadNLP.DenseCallback : MadNLP.SparseCallback, + kkt_system = dense_callback ? MadNLP.DenseCondensedKKTSystem : MadNLP.SparseCondensedKKTSystem, + linear_solver = MadNLPGPU.LapackOneMKLSolver, + tol = MadNLP.get_tolerance(T, kkt_system), + bound_relax_factor = tol, + ) where {T, VT <: oneVector{T}} + return MadNLP.MadNLPOptions{T}( + tol = tol, + callback = callback, + kkt_system = kkt_system, + linear_solver = linear_solver, + bound_relax_factor = bound_relax_factor, + ) +end + +#= + SparseMatrixCSC to oneSparseMatrixCSC +=# + +function oneMKL.oneSparseMatrixCSC{Tv, Ti}(A::SparseMatrixCSC{Tv, Ti}) where {Tv, Ti} + return oneMKL.oneSparseMatrixCSC{Tv, Ti}( + oneVector(A.colptr), + oneVector(A.rowval), + oneVector(A.nzval), + size(A), + ) +end + +#= + oneSparseMatrixCSC to oneMatrix +=# + +function MadNLPGPU.gpu_transfer!(y::oneMatrix{T}, x::oneMKL.oneSparseMatrixCSC{T}) where {T} + n = size(y, 2) + fill!(y, zero(T)) + backend = oneAPIBackend() + MadNLPGPU._csc_to_dense_kernel!(backend)(y, x.colPtr, x.rowVal, x.nzVal, ndrange = n) + synchronize(backend) + return +end + +#= + MadNLP._syr! +=# + +MadNLP._syr!(uplo::Char, alpha::T, x::oneVector{T}, A::oneMatrix{T}) where {T} = oneMKL.syr!(uplo, alpha, x, A) + +#= + MadNLP._symv! +=# + +MadNLP._symv!(uplo::Char, alpha::T, A::oneMatrix{T}, x::oneVector{T}, beta::T, y::oneVector{T}) where {T} = oneMKL.symv!(uplo, alpha, A, x, beta, y) + +#= + MadNLP._syrk! +=# + +MadNLP._syrk!(uplo::Char, trans::Char, alpha::T, A::oneMatrix{T}, beta::T, C::oneMatrix{T}) where {T} = oneMKL.syrk!(uplo, trans, alpha, A, beta, C) + +#= + MadNLP._trsm! +=# + +MadNLP._trsm!(side::Char, uplo::Char, transa::Char, diag::Char, alpha::T, A::oneMatrix{T}, B::oneMatrix{T}) where {T} = oneMKL.trsm!(side, uplo, transa, diag, alpha, A, B) + +#= + MadNLP._dgmm! +=# + +MadNLP._dgmm!(side::Char, A::oneMatrix{T}, x::oneVector{T}, B::oneMatrix{T}) where {T} = oneMKL.dgmm!(side, A, x, B) + +#= + LinearAlgebra.norm for oneAPI arrays. + GPUArrays' norm uses LinearAlgebra.norm as a map function in mapreduce, + which fails on oneAPI because norm on scalars generates jl_f_throw_methoderror + calls that can't be compiled to SPIRV. Replace with abs-based implementation. + The SubArray method also avoids scalar indexing on views. +=# + +function _onenorm(v, p::Real) + isempty(v) && return float(zero(eltype(v))) + if p == Inf + return maximum(abs, v) + elseif p == -Inf + return minimum(abs, v) + elseif p == 1 + return mapreduce(abs, +, v) + elseif p == 2 + return sqrt(mapreduce(abs2, +, v)) + elseif p == 0 + return float(count(!iszero, v)) + else + spp = float(p) + return mapreduce(x -> abs(x)^spp, +, v)^inv(spp) + end +end + +LinearAlgebra.norm(v::oneArray{<:Number}, p::Real = 2) = _onenorm(v, p) +LinearAlgebra.norm(v::SubArray{<:Number, <:Any, <:oneArray}, p::Real = 2) = _onenorm(v, p) diff --git a/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi_dense.jl b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi_dense.jl new file mode 100644 index 000000000..555e730ce --- /dev/null +++ b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi_dense.jl @@ -0,0 +1,149 @@ +######################################################################## +##### oneAPI wrappers for DenseKKTSystem / DenseCondensedKKTSystem ##### +######################################################################## + +#= + MadNLP._ger! +=# + +MadNLP._ger!(alpha::T, x::oneVector{T}, y::oneVector{T}, A::oneMatrix{T}) where {T} = oneMKL.ger!(alpha, x, y, A) + +#= + MadNLP._madnlp_unsafe_wrap +=# + +function MadNLP._madnlp_unsafe_wrap(vec::VT, n, shift = 1) where {T, VT <: oneVector{T}} + return view(vec, shift:(shift + n - 1)) +end + +#= + MadNLP.diag! +=# + +function MadNLP.diag!(dest::oneVector{T}, src::oneMatrix{T}) where {T} + @assert length(dest) == size(src, 1) + backend = oneAPIBackend() + MadNLPGPU._copy_diag_kernel!(backend)(dest, src, ndrange = length(dest)) + synchronize(backend) + return +end + +#= + MadNLP.diag_add! +=# + +function MadNLP.diag_add!(dest::oneMatrix, src1::oneVector, src2::oneVector) + backend = oneAPIBackend() + MadNLPGPU._add_diagonal_kernel!(backend)(dest, src1, src2, ndrange = size(dest, 1)) + synchronize(backend) + return +end + +#= + MadNLP._set_diag! +=# + +function MadNLP._set_diag!(A::oneMatrix, inds, a) + if !isempty(inds) + backend = oneAPIBackend() + MadNLPGPU._set_diag_kernel!(backend)(A, inds, a; ndrange = length(inds)) + synchronize(backend) + end + return +end + +#= + MadNLP._build_dense_kkt_system! +=# + +function MadNLP._build_dense_kkt_system!( + dest::oneMatrix, + hess::oneMatrix, + jac::oneMatrix, + pr_diag::oneVector, + du_diag::oneVector, + diag_hess::oneVector, + ind_ineq::AbstractVector, + n, + m, + ns, + ) + ind_ineq_gpu = oneVector(ind_ineq) + ndrange = (n + m + ns, n) + backend = oneAPIBackend() + MadNLPGPU._build_dense_kkt_system_kernel!(backend)( + dest, + hess, + jac, + pr_diag, + du_diag, + diag_hess, + ind_ineq_gpu, + n, + m, + ns, + ndrange = ndrange, + ) + synchronize(backend) + return +end + +#= + MadNLP._build_ineq_jac! +=# + +function MadNLP._build_ineq_jac!( + dest::oneMatrix, + jac::oneMatrix, + diag_buffer::oneVector, + ind_ineq::AbstractVector, + n, + m_ineq, + ) + (m_ineq == 0) && return # nothing to do if no ineq. constraints + ind_ineq_gpu = oneVector(ind_ineq) + ndrange = (m_ineq, n) + backend = oneAPIBackend() + MadNLPGPU._build_jacobian_condensed_kernel!(backend)( + dest, + jac, + diag_buffer, + ind_ineq_gpu, + m_ineq, + ndrange = ndrange, + ) + synchronize(backend) + return +end + +#= + MadNLP._build_condensed_kkt_system! +=# + +function MadNLP._build_condensed_kkt_system!( + dest::oneMatrix, + hess::oneMatrix, + jac::oneMatrix, + pr_diag::oneVector, + du_diag::oneVector, + ind_eq::AbstractVector, + n, + m_eq, + ) + ind_eq_gpu = oneVector(ind_eq) + ndrange = (n + m_eq, n) + backend = oneAPIBackend() + MadNLPGPU._build_condensed_kkt_system_kernel!(backend)( + dest, + hess, + jac, + pr_diag, + du_diag, + ind_eq_gpu, + n, + m_eq, + ndrange = ndrange, + ) + synchronize(backend) + return +end diff --git a/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi_sparse.jl b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi_sparse.jl new file mode 100644 index 000000000..aaa1dffea --- /dev/null +++ b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/oneapi_sparse.jl @@ -0,0 +1,494 @@ +######################################################## +##### oneAPI wrappers for SparseCondensedKKTSystem ##### +######################################################## + +#= + SparseMatrixCOO to oneSparseMatrixCSC +=# + +function MadNLP.transfer!( + dest::oneMKL.oneSparseMatrixCSC, + src::MadNLP.SparseMatrixCOO, + map, + ) + return copyto!(view(dest.nzVal, map), src.V) +end + +#= + MadNLP.mul! with SparseCondensedKKTSystem + oneAPI +=# + +function MadNLP.mul!( + w::MadNLP.AbstractKKTVector{T, VT}, + kkt::MadNLP.SparseCondensedKKTSystem, + x::MadNLP.AbstractKKTVector, + alpha = one(T), + beta = zero(T), + ) where {T, VT <: oneVector{T}} + n = size(kkt.hess_com, 1) + m = size(kkt.jt_csc, 2) + + # Decompose results + xx = view(MadNLP.full(x), 1:n) + xs = view(MadNLP.full(x), (n + 1):(n + m)) + xz = view(MadNLP.full(x), (n + m + 1):(n + 2 * m)) + + # Decompose buffers + wx = view(MadNLP.full(w), 1:n) + ws = view(MadNLP.full(w), (n + 1):(n + m)) + wz = view(MadNLP.full(w), (n + m + 1):(n + 2 * m)) + + MadNLP.mul!(wx, kkt.hess_com, xx, alpha, beta) + MadNLP.mul!(wx, kkt.hess_com', xx, alpha, one(T)) + MadNLP.mul!(wx, kkt.jt_csc, xz, alpha, beta) + if !isempty(kkt.ext.diag_map_to) + backend = oneAPIBackend() + MadNLPGPU._diag_operation_kernel!(backend)( + wx, + kkt.hess_com.nzVal, + xx, + alpha, + kkt.ext.diag_map_to, + kkt.ext.diag_map_fr; + ndrange = length(kkt.ext.diag_map_to), + ) + synchronize(backend) + end + + MadNLP.mul!(wz, kkt.jt_csc', xx, alpha, one(T)) + MadNLP.axpy!(-alpha, xz, ws) + MadNLP.axpy!(-alpha, xs, wz) + return MadNLP._kktmul!( + w, + x, + kkt.reg, + kkt.du_diag, + kkt.l_lower, + kkt.u_lower, + kkt.l_diag, + kkt.u_diag, + alpha, + beta, + ) +end + +function MadNLP.mul_hess_blk!( + wx::VT, + kkt::Union{MadNLP.SparseKKTSystem, MadNLP.SparseCondensedKKTSystem}, + t, + ) where {T, VT <: oneVector{T}} + n = size(kkt.hess_com, 1) + wxx = @view(wx[1:n]) + tx = @view(t[1:n]) + + MadNLP.mul!(wxx, kkt.hess_com, tx, one(T), zero(T)) + MadNLP.mul!(wxx, kkt.hess_com', tx, one(T), one(T)) + if !isempty(kkt.ext.diag_map_to) + backend = oneAPIBackend() + MadNLPGPU._diag_operation_kernel!(backend)( + wxx, + kkt.hess_com.nzVal, + tx, + one(T), + kkt.ext.diag_map_to, + kkt.ext.diag_map_fr; + ndrange = length(kkt.ext.diag_map_to), + ) + synchronize(backend) + end + + fill!(@view(wx[(n + 1):end]), 0) + wx .+= t .* kkt.pr_diag + return +end + +function MadNLP.get_tril_to_full(csc::oneMKL.oneSparseMatrixCSC{Tv, Ti}) where {Tv, Ti} + cscind = MadNLP.SparseMatrixCSC{Int, Ti}( + Symmetric( + MadNLP.SparseMatrixCSC{Int, Ti}( + size(csc)..., + Array(csc.colPtr), + Array(csc.rowVal), + collect(1:MadNLP.nnz(csc)), + ), + :L, + ), + ) + return oneMKL.oneSparseMatrixCSC{Tv, Ti}( + oneArray(cscind.colptr), + oneArray(cscind.rowval), + oneVector{Tv}(undef, MadNLP.nnz(cscind)), + size(csc), + ), + view(csc.nzVal, oneArray(cscind.nzval)) +end + +function MadNLP.get_sparse_condensed_ext( + ::Type{VT}, + hess_com, + jptr, + jt_map, + hess_map, + ) where {T, VT <: oneVector{T}} + zvals = oneVector{Int}(1:length(hess_map)) + hess_com_ptr = map((i, j) -> (i, j), hess_map, zvals) + if length(hess_com_ptr) > 0 # otherwise error is thrown + sort!(hess_com_ptr) + end + + jvals = oneVector{Int}(1:length(jt_map)) + jt_csc_ptr = map((i, j) -> (i, j), jt_map, jvals) + if length(jt_csc_ptr) > 0 # otherwise error is thrown + sort!(jt_csc_ptr) + end + + by = (i, j) -> i[1] != j[1] + jptrptr = MadNLP.getptr(jptr, by = by) + hess_com_ptrptr = MadNLP.getptr(hess_com_ptr, by = by) + jt_csc_ptrptr = MadNLP.getptr(jt_csc_ptr, by = by) + + diag_map_to, diag_map_fr = get_diagonal_mapping(hess_com.colPtr, hess_com.rowVal) + + return ( + jptrptr = jptrptr, + hess_com_ptr = hess_com_ptr, + hess_com_ptrptr = hess_com_ptrptr, + jt_csc_ptr = jt_csc_ptr, + jt_csc_ptrptr = jt_csc_ptrptr, + diag_map_to = diag_map_to, + diag_map_fr = diag_map_fr, + ) +end + +function get_diagonal_mapping(colptr, rowval) + nnz = length(rowval) + if nnz == 0 + return similar(colptr, 0), similar(colptr, 0) + end + inds1 = findall( + map( + (x, y) -> ((x <= nnz) && (x != y)), + @view(colptr[1:(end - 1)]), + @view(colptr[2:end]) + ), + ) + if length(inds1) == 0 + return similar(rows, 0), similar(ptrs, 0) + end + ptrs = colptr[inds1] + rows = rowval[ptrs] + inds2 = findall(inds1 .== rows) + if length(inds2) == 0 + return similar(rows, 0), similar(ptrs, 0) + end + + return rows[inds2], ptrs[inds2] +end + +function MadNLP._sym_length(Jt::oneMKL.oneSparseMatrixCSC) + return mapreduce( + (x, y) -> begin + z = x - y + div(z^2 + z, 2) + end, + +, + @view(Jt.colPtr[2:end]), + @view(Jt.colPtr[1:(end - 1)]) + ) +end + +function MadNLP._first_and_last_col(sym2::oneVector, ptr2) + @allowscalar begin + first = sym2[1][2] + last = sym2[ptr2[end]][2] + end + return (first, last) +end + +MadNLP.nzval(H::oneMKL.oneSparseMatrixCSC) = H.nzVal + +function MadNLP._get_sparse_csc(dims, colptr::oneVector, rowval, nzval) + return oneMKL.oneSparseMatrixCSC(colptr, rowval, nzval, dims) +end + +function getij(idx, n) + j = ceil(Int, ((2n + 1) - sqrt((2n + 1)^2 - 8 * idx)) / 2) + i = idx - div((j - 1) * (2n - j), 2) + return (i, j) +end + +#= + MadNLP._set_colptr! +=# + +function MadNLP._set_colptr!(colptr::oneVector, ptr2, sym2, guide) + if length(ptr2) > 1 # otherwise error is thrown + backend = oneAPIBackend() + MadNLPGPU._set_colptr_kernel!(backend)( + colptr, + sym2, + ptr2, + guide; + ndrange = length(ptr2) - 1, + ) + synchronize(backend) + end + return +end + + +#= + MadNLP.tril_to_full! +=# + +function MadNLP.tril_to_full!(dense::oneMatrix{T}) where {T} + n = size(dense, 1) + backend = oneAPIBackend() + MadNLPGPU._tril_to_full_kernel!(backend)(dense; ndrange = div(n^2 + n, 2)) + synchronize(backend) + return +end + +#= + MadNLP.force_lower_triangular! +=# + +function MadNLP.force_lower_triangular!(I::oneVector{T}, J) where {T} + if !isempty(I) + backend = oneAPIBackend() + MadNLPGPU._force_lower_triangular_kernel!(backend)(I, J; ndrange = length(I)) + synchronize(backend) + end + return +end + +#= + MadNLP.coo_to_csc +=# + +function MadNLP.coo_to_csc( + coo::MadNLP.SparseMatrixCOO{T, I, VT, VI}, + ) where {T, I, VT <: oneArray, VI <: oneArray} + zvals = oneVector{Int}(1:length(coo.I)) + coord = map((i, j, k) -> ((i, j), k), coo.I, coo.J, zvals) + if length(coord) > 0 + sort!(coord, lt = (((i, j), k), ((n, m), l)) -> (j, i) < (m, n)) + end + + mapptr = MadNLP.getptr(coord; by = ((x1, x2), (y1, y2)) -> x1 != y1) + + colptr = similar(coo.I, size(coo, 2) + 1) + + coord_csc = coord[@view(mapptr[1:(end - 1)])] + + backend = oneAPIBackend() + if length(coord_csc) > 0 + MadNLPGPU._set_coo_to_colptr_kernel!(backend)( + colptr, + coord_csc, + ndrange = length(coord_csc), + ) + synchronize(backend) + else + fill!(colptr, one(Int)) + end + + rowval = map(x -> x[1][1], coord_csc) + nzval = similar(rowval, T) + csc = oneMKL.oneSparseMatrixCSC(colptr, rowval, nzval, size(coo)) + + cscmap = similar(coo.I, Int) + if length(mapptr) > 1 + MadNLPGPU._set_coo_to_csc_map_kernel!(backend)( + cscmap, + mapptr, + coord, + ndrange = length(mapptr) - 1, + ) + synchronize(backend) + end + + return csc, cscmap +end + +#= + MadNLP.build_condensed_aug_coord! +=# + +function MadNLP.build_condensed_aug_coord!( + kkt::MadNLP.AbstractCondensedKKTSystem{T, VT, MT}, + ) where {T, VT, MT <: oneMKL.oneSparseMatrixCSC{T}} + fill!(kkt.aug_com.nzVal, zero(T)) + backend = oneAPIBackend() + if length(kkt.hptr) > 0 + MadNLPGPU._transfer_hessian_kernel!(backend)( + kkt.aug_com.nzVal, + kkt.hptr, + kkt.hess_com.nzVal; + ndrange = length(kkt.hptr), + ) + synchronize(backend) + end + if length(kkt.dptr) > 0 + MadNLPGPU._transfer_hessian_kernel!(backend)( + kkt.aug_com.nzVal, + kkt.dptr, + kkt.pr_diag; + ndrange = length(kkt.dptr), + ) + synchronize(backend) + end + if length(kkt.ext.jptrptr) > 1 # otherwise error is thrown + MadNLPGPU._transfer_jtsj_kernel!(backend)( + kkt.aug_com.nzVal, + kkt.jptr, + kkt.ext.jptrptr, + kkt.jt_csc.nzVal, + kkt.diag_buffer; + ndrange = length(kkt.ext.jptrptr) - 1, + ) + synchronize(backend) + end + return +end + +#= + MadNLP.compress_hessian! / MadNLP.compress_jacobian! +=# + +function MadNLP.compress_hessian!( + kkt::MadNLP.AbstractSparseKKTSystem{T, VT, MT}, + ) where {T, VT, MT <: oneMKL.oneSparseMatrixCSC{T, Int32}} + fill!(kkt.hess_com.nzVal, zero(T)) + backend = oneAPIBackend() + if length(kkt.ext.hess_com_ptrptr) > 1 + MadNLPGPU._transfer_to_csc_kernel!(backend)( + kkt.hess_com.nzVal, + kkt.ext.hess_com_ptr, + kkt.ext.hess_com_ptrptr, + kkt.hess_raw.V; + ndrange = length(kkt.ext.hess_com_ptrptr) - 1, + ) + synchronize(backend) + end + return +end + +function MadNLP.compress_jacobian!( + kkt::MadNLP.SparseCondensedKKTSystem{T, VT, MT}, + ) where {T, VT, MT <: oneMKL.oneSparseMatrixCSC{T, Int32}} + fill!(kkt.jt_csc.nzVal, zero(T)) + backend = oneAPIBackend() + if length(kkt.ext.jt_csc_ptrptr) > 1 # otherwise error is thrown + MadNLPGPU._transfer_to_csc_kernel!(backend)( + kkt.jt_csc.nzVal, + kkt.ext.jt_csc_ptr, + kkt.ext.jt_csc_ptrptr, + kkt.jt_coo.V; + ndrange = length(kkt.ext.jt_csc_ptrptr) - 1, + ) + synchronize(backend) + end + return +end + +#= + MadNLP._set_con_scale_sparse! +=# + +function MadNLP._set_con_scale_sparse!( + con_scale::VT, + jac_I, + jac_buffer, + ) where {T, VT <: oneVector{T}} + ind_jac = oneVector{Int}(1:length(jac_I)) + inds = map((i, j) -> (i, j), jac_I, ind_jac) + !isempty(inds) && sort!(inds) + ptr = MadNLP.getptr(inds; by = ((x1, x2), (y1, y2)) -> x1 != y1) + if length(ptr) > 1 + backend = oneAPIBackend() + MadNLPGPU._set_con_scale_sparse_kernel!(backend)( + con_scale, + ptr, + inds, + jac_I, + jac_buffer; + ndrange = length(ptr) - 1, + ) + synchronize(backend) + end + return +end + +#= + MadNLP._build_condensed_aug_symbolic_hess +=# + +function MadNLP._build_condensed_aug_symbolic_hess( + H::oneMKL.oneSparseMatrixCSC{Tv, Ti}, + sym, + sym2, + ) where {Tv, Ti} + if size(H, 2) > 0 + backend = oneAPIBackend() + MadNLPGPU._build_condensed_aug_symbolic_hess_kernel!(backend)( + sym, + sym2, + H.colPtr, + H.rowVal; + ndrange = size(H, 2), + ) + synchronize(backend) + end + return +end + +#= + MadNLP._build_condensed_aug_symbolic_jt +=# + +function MadNLP._build_condensed_aug_symbolic_jt( + Jt::oneMKL.oneSparseMatrixCSC{Tv, Ti}, + sym, + sym2, + ) where {Tv, Ti} + if size(Jt, 2) > 0 + _offsets = map( + (i, j) -> div((j - i)^2 + (j - i), 2), + @view(Jt.colPtr[1:(end - 1)]), + @view(Jt.colPtr[2:end]) + ) + offsets = cumsum(_offsets) + backend = oneAPIBackend() + MadNLPGPU._build_condensed_aug_symbolic_jt_kernel!(backend)( + sym, + sym2, + Jt.colPtr, + Jt.rowVal, + offsets; + ndrange = size(Jt, 2), + ) + synchronize(backend) + end + return +end + +#= + MadNLP._build_scale_augmented_system_coo! +=# + +function MadNLP._build_scale_augmented_system_coo!(dest, src, scaling::oneArray, n, m) + backend = oneAPIBackend() + MadNLPGPU._scale_augmented_system_coo_kernel!(backend)( + dest.V, + src.I, + src.J, + src.V, + scaling, + n, + m; + ndrange = nnz(src), + ) + synchronize(backend) + return +end diff --git a/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/onemkl.jl b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/onemkl.jl new file mode 100644 index 000000000..fd76e665f --- /dev/null +++ b/lib/MadNLPGPU/ext/MadNLPGPUOneAPIExt/onemkl.jl @@ -0,0 +1,345 @@ +mutable struct LapackOneMKLSolver{T, MT} <: MadNLP.AbstractLinearSolver{T} + A::MT + fact::oneMatrix{T} + n::Int64 + sol::oneVector{T} + tau::oneVector{T} + Λ::oneVector{T} + info::oneVector{Cint} + ipiv::oneVector{Int64} + scratchpad::oneVector{T} + scratchpad_size::Int64 + device_queue::SYCL.syclQueue_t + alpha::Base.RefValue{T} + beta::Base.RefValue{T} + opt::MadNLP.LapackOptions + logger::MadNLP.MadNLPLogger + + function LapackOneMKLSolver( + A::MT; + option_dict::Dict{Symbol, Any} = Dict{Symbol, Any}(), + opt = MadNLP.LapackOptions(), + logger = MadNLP.MadNLPLogger(), + kwargs..., + ) where {MT <: AbstractMatrix} + # Reclaim GPU resources from previous solvers. Julia's GC doesn't see GPU + # memory pressure, so old oneArray/MKL handles accumulate. + # Synchronize FIRST to ensure all pending GPU operations complete, making + # it safe for GC finalizers to free GPU buffers. Then GC to collect stale + # objects, flush deferred MKL sparse handles, and GC again. + oneAPI.synchronize() + GC.gc(true) + oneAPI.oneL0._run_reclaim_callbacks() + GC.gc(true) + oneAPI.synchronize() + MadNLP.set_options!(opt, option_dict, kwargs...) + T = eltype(A) + m, n = size(A) + @assert m == n + fact = oneMatrix{T}(undef, m, n) + sol = oneVector{T}(undef, 0) + tau = oneVector{T}(undef, 0) + Λ = oneVector{T}(undef, 0) + info = oneVector{Cint}(undef, 1) + ipiv = oneVector{Int64}(undef, 0) + scratchpad = oneVector{T}(undef, 0) + scratchpad_size = 0 + # Get the device queue from the oneAPI context + queue = oneAPI.global_queue(oneAPI.context(fact), oneAPI.device(fact)) + device_queue = oneAPI.sycl_queue(queue) + alpha = Ref{T}(1) + beta = Ref{T}(0) + solver = new{T, MT}(A, fact, n, sol, tau, Λ, info, ipiv, scratchpad, scratchpad_size, device_queue, alpha, beta, opt, logger) + setup!(solver) + return solver + end +end + +MadNLP.improve!(M::LapackOneMKLSolver) = false +MadNLP.is_inertia(M::LapackOneMKLSolver) = (M.opt.lapack_algorithm == MadNLP.CHOLESKY) || (M.opt.lapack_algorithm == MadNLP.EVD) +function MadNLP.inertia(M::LapackOneMKLSolver) + return if M.opt.lapack_algorithm == MadNLP.CHOLESKY + sum(M.info) == 0 ? (M.n, 0, 0) : (0, M.n, 0) + elseif M.opt.lapack_algorithm == MadNLP.EVD + numpos = count(λ -> λ > 0, M.Λ) + numneg = count(λ -> λ < 0, M.Λ) + numzero = M.n - numpos - numneg + (numpos, numzero, numneg) + else + error(M.logger, "Invalid lapack_algorithm") + end +end + +MadNLP.input_type(::Type{LapackOneMKLSolver}) = :dense +MadNLP.default_options(::Type{LapackOneMKLSolver}) = MadNLP.LapackOptions(MadNLP.EVD) +MadNLP.introduce(M::LapackOneMKLSolver) = "OneAPI -- ($(M.opt.lapack_algorithm))" +# MadNLP.introduce(M::LapackOneMKLSolver) = "OneAPI v$(oneAPI.version()) -- ($(M.opt.lapack_algorithm))" + +function setup!(M::LapackOneMKLSolver) + return if M.opt.lapack_algorithm == MadNLP.LU + setup_lu!(M) + elseif M.opt.lapack_algorithm == MadNLP.QR + setup_qr!(M) + elseif M.opt.lapack_algorithm == MadNLP.CHOLESKY + setup_cholesky!(M) + elseif M.opt.lapack_algorithm == MadNLP.EVD + setup_evd!(M) + else + error(M.logger, "Invalid lapack_algorithm") + end +end + +function MadNLP.factorize!(M::LapackOneMKLSolver) + MadNLPGPU.gpu_transfer!(M.fact, M.A) + return if M.opt.lapack_algorithm == MadNLP.LU + MadNLP.tril_to_full!(M.fact) + factorize_lu!(M) + elseif M.opt.lapack_algorithm == MadNLP.QR + MadNLP.tril_to_full!(M.fact) + factorize_qr!(M) + elseif M.opt.lapack_algorithm == MadNLP.CHOLESKY + factorize_cholesky!(M) + elseif M.opt.lapack_algorithm == MadNLP.EVD + factorize_evd!(M) + else + error(M.logger, "Invalid lapack_algorithm") + end +end + +for T in (:Float32, :Float64) + @eval begin + function MadNLP.solve!(M::LapackOneMKLSolver{$T}, x::oneVector{$T}) + return if M.opt.lapack_algorithm == MadNLP.LU + solve_lu!(M, x) + elseif M.opt.lapack_algorithm == MadNLP.QR + solve_qr!(M, x) + elseif M.opt.lapack_algorithm == MadNLP.CHOLESKY + solve_cholesky!(M, x) + elseif M.opt.lapack_algorithm == MadNLP.EVD + solve_evd!(M, x) + else + error(M.logger, "Invalid lapack_algorithm") + end + end + + MadNLP.is_supported(::Type{LapackOneMKLSolver}, ::Type{$T}) = true + end +end + +function MadNLP.solve!(M::LapackOneMKLSolver, x::AbstractVector) + isempty(M.sol) && resize!(M.sol, M.n) + copyto!(M.sol, x) + MadNLP.solve!(M, M.sol) + copyto!(x, M.sol) + return x +end + +for (potrf, potrf_buffer, potrs, potrs_buffer, T) in + ( + (:onemklDpotrf, :onemklDpotrf_scratchpad_size, :onemklDpotrs, :onemklDpotrs_scratchpad_size, :Float64), + (:onemklSpotrf, :onemklSpotrf_scratchpad_size, :onemklSpotrs, :onemklSpotrs_scratchpad_size, :Float32), + ) + @eval begin + function setup_cholesky!(M::LapackOneMKLSolver{$T}) + potrf_scratchpad_size = Support.$potrf_buffer(M.device_queue, 'L', M.n, M.n) + potrs_scratchpad_size = Support.$potrs_buffer(M.device_queue, 'L', M.n, one(Int64), M.n, M.n) + M.scratchpad_size = max(potrf_scratchpad_size, potrs_scratchpad_size) + resize!(M.scratchpad, M.scratchpad_size) + return M + end + + function factorize_cholesky!(M::LapackOneMKLSolver{$T}) + Support.$potrf( + M.device_queue, + 'L', + M.n, + M.fact, + M.n, + M.scratchpad, + M.scratchpad_size, + ) + return M + end + + function solve_cholesky!(M::LapackOneMKLSolver{$T}, x::oneVector{$T}) + Support.$potrs( + M.device_queue, + 'L', + M.n, + one(Int64), + M.fact, + M.n, + x, + M.n, + M.scratchpad, + M.scratchpad_size, + ) + return x + end + end +end + +for (getrf, getrf_buffer, getrs, getrs_buffer, T) in + ( + (:onemklDgetrf, :onemklDgetrf_scratchpad_size, :onemklDgetrs, :onemklDgetrs_scratchpad_size, :Float64), + (:onemklSgetrf, :onemklSgetrf_scratchpad_size, :onemklSgetrs, :onemklSgetrs_scratchpad_size, :Float32), + ) + @eval begin + function setup_lu!(M::LapackOneMKLSolver{$T}) + resize!(M.ipiv, M.n) + getrf_scratchpad_size = Support.$getrf_buffer(M.device_queue, M.n, M.n, M.n) + getrs_scratchpad_size = Support.$getrs_buffer(M.device_queue, 'N', M.n, one(Int64), M.n, M.n) + M.scratchpad_size = max(getrf_scratchpad_size, getrs_scratchpad_size) + resize!(M.scratchpad, M.scratchpad_size) + return M + end + + function factorize_lu!(M::LapackOneMKLSolver{$T}) + Support.$getrf( + M.device_queue, + M.n, + M.n, + M.fact, + M.n, + M.ipiv, + M.scratchpad, + M.scratchpad_size, + ) + return M + end + + function solve_lu!(M::LapackOneMKLSolver{$T}, x::oneVector{$T}) + Support.$getrs( + M.device_queue, + 'N', + M.n, + one(Int64), + M.fact, + M.n, + M.ipiv, + x, + M.n, + M.scratchpad, + M.scratchpad_size, + ) + return x + end + end +end + +for (geqrf, geqrf_buffer, ormqr, ormqr_buffer, trsv, T) in + ( + (:onemklDgeqrf, :onemklDgeqrf_scratchpad_size, :onemklDormqr, :onemklDormqr_scratchpad_size, :onemklDtrsv, :Float64), + (:onemklSgeqrf, :onemklSgeqrf_scratchpad_size, :onemklSormqr, :onemklSormqr_scratchpad_size, :onemklStrsv, :Float32), + ) + @eval begin + function setup_qr!(M::LapackOneMKLSolver{$T}) + resize!(M.tau, M.n) + geqrf_scratchpad_size = Support.$geqrf_buffer(M.device_queue, M.n, M.n, M.n) + ormqr_scratchpad_size = Support.$ormqr_buffer(M.device_queue, 'L', 'T', M.n, one(Int64), M.n, M.n, M.n) + M.scratchpad_size = max(geqrf_scratchpad_size, ormqr_scratchpad_size) + resize!(M.scratchpad, M.scratchpad_size) + return M + end + + function factorize_qr!(M::LapackOneMKLSolver{$T}) + Support.$geqrf( + M.device_queue, + M.n, + M.n, + M.fact, + M.n, + M.tau, + M.scratchpad, + M.scratchpad_size, + ) + return M + end + + function solve_qr!(M::LapackOneMKLSolver{$T}, x::oneVector{$T}) + # Apply Q^T to x: x = Q^T * x + Support.$ormqr( + M.device_queue, + 'L', # side (left multiplication) + 'T', # trans (transpose) + M.n, # m + one(Int64), # n (single RHS) + M.n, # k + M.fact, # A + M.n, # lda + M.tau, # tau + x, # c (the RHS vector) + M.n, # ldc + M.scratchpad, + M.scratchpad_size, + ) + # Solve R*x = Q^T*b using triangular solve + oneMKL.trsv!('U', 'N', 'N', M.fact, x) # upper, no-trans, non-unit diagonal + return x + end + end +end + +for (syevd, syevd_buffer, gemv, T) in + ( + (:onemklDsyevd, :onemklDsyevd_scratchpad_size, :onemklDgemv, :Float64), + (:onemklSsyevd, :onemklSsyevd_scratchpad_size, :onemklSgemv, :Float32), + ) + @eval begin + function setup_evd!(M::LapackOneMKLSolver{$T}) + resize!(M.tau, M.n) + resize!(M.Λ, M.n) + M.scratchpad_size = Support.$syevd_buffer(M.device_queue, 'V', 'L', M.n, M.n) + resize!(M.scratchpad, M.scratchpad_size) + return M + end + + function factorize_evd!(M::LapackOneMKLSolver{$T}) + Support.$syevd( + M.device_queue, + 'V', + 'L', + M.n, + M.fact, + M.n, + M.Λ, + M.scratchpad, + M.scratchpad_size, + ) + return M + end + + function solve_evd!(M::LapackOneMKLSolver{$T}, x::oneVector{$T}) + Support.$gemv( + M.device_queue, + 'T', + M.n, + M.n, + M.alpha, + M.fact, + M.n, + x, + one(Int64), + M.beta, + M.tau, + one(Int64), + ) + M.tau ./= M.Λ + Support.$gemv( + M.device_queue, + 'N', + M.n, + M.n, + M.alpha, + M.fact, + M.n, + M.tau, + one(Int64), + M.beta, + x, + one(Int64), + ) + return x + end + end +end diff --git a/lib/MadNLPGPU/src/MadNLPGPU.jl b/lib/MadNLPGPU/src/MadNLPGPU.jl index 60667f0ba..6fbcf2540 100644 --- a/lib/MadNLPGPU/src/MadNLPGPU.jl +++ b/lib/MadNLPGPU/src/MadNLPGPU.jl @@ -39,7 +39,8 @@ include("LinearSolvers/cudss.jl") include("cuda.jl") global LapackROCmSolver -export LapackCUDASolver, CUDSSSolver, LapackROCmSolver +global LapackOneMKLSolver +export LapackCUDASolver, CUDSSSolver, LapackROCmSolver, LapackOneMKLSolver # re-export MadNLP, including deprecated names for name in names(MadNLP, all=true) diff --git a/lib/MadNLPGPU/test/Project.toml b/lib/MadNLPGPU/test/Project.toml new file mode 100644 index 000000000..c1c2a45b0 --- /dev/null +++ b/lib/MadNLPGPU/test/Project.toml @@ -0,0 +1,9 @@ +[deps] +CUDA = "052768ef-5323-5732-b1bb-66c8b64840ba" +MadNLP = "2621e9c9-9eb4-46b1-8089-e8c72242dfb6" +MadNLPGPU = "d72a61cc-809d-412f-99be-fd81f4b8a598" +MadNLPTests = "b52a2a03-04ab-4a5f-9698-6a2deff93217" +Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" + +[compat] +MadNLPTests = "0.5.3" diff --git a/lib/MadNLPGPU/test/densekkt_oneapi.jl b/lib/MadNLPGPU/test/densekkt_oneapi.jl new file mode 100644 index 000000000..1c137dc5c --- /dev/null +++ b/lib/MadNLPGPU/test/densekkt_oneapi.jl @@ -0,0 +1,89 @@ +using oneAPI +using MadNLPTests + +function _compare_oneapi_with_cpu(KKTSystem, n, m, ind_fixed) + for (T, tol, atol) in [ + (Float32, 1.0e-4, 1.0e1), + (Float64, 1.0e-8, 1.0e-6), + ] + madnlp_options = Dict{Symbol, Any}( + :callback => MadNLP.DenseCallback, + :kkt_system => KKTSystem, + :linear_solver => LapackOneMKLSolver, + :lapack_algorithm => MadNLP.QR, + :print_level => MadNLP.ERROR, + :tol => tol + ) + + # Host evaluator + nlph = MadNLPTests.DenseDummyQP(zeros(T, n); m = m, fixed_variables = ind_fixed) + # Device evaluator + nlpd = MadNLPTests.DenseDummyQP(oneAPI.zeros(T, n); m = m, fixed_variables = oneArray(ind_fixed)) + + # Solve on CPU + h_solver = MadNLPSolver(nlph; madnlp_options...) + results_cpu = MadNLP.solve!(h_solver) + + # Solve on GPU + d_solver = MadNLPSolver(nlpd; madnlp_options...) + results_gpu = MadNLP.solve!(d_solver) + + @test isa(d_solver.kkt, KKTSystem{T}) + # Check that both results match exactly + if T == Float64 + @test h_solver.cnt.k == d_solver.cnt.k + @test results_cpu.objective ≈ results_gpu.objective + @test results_cpu.solution ≈ Array(results_gpu.solution) atol = atol + @test results_cpu.multipliers ≈ Array(results_gpu.multipliers) atol = atol + end + end + return +end + +@testset "MadNLPGPU -- LapackOneMKLSolver -- ($(kkt_system))" for kkt_system in [ + MadNLP.DenseKKTSystem, + MadNLP.DenseCondensedKKTSystem, + ] + @testset "Size: ($n, $m)" for (n, m) in [(10, 0), (10, 5), (50, 10)] + _compare_oneapi_with_cpu(kkt_system, n, m, Int[]) + end + @testset "Fixed variables" for (n, m) in [(10, 0), (10, 5), (50, 10)] + _compare_oneapi_with_cpu(kkt_system, n, m, Int[1, 2]) + end +end + +@testset "MadNLP -- LapackOneMKLSolver: $QN + $KKT" for QN in [ + MadNLP.BFGS, + MadNLP.DampedBFGS, + ], KKT in [ + MadNLP.DenseKKTSystem, + MadNLP.DenseCondensedKKTSystem, + ] + @testset "Size: ($n, $m)" for (n, m) in [(10, 0), (10, 5), (50, 10)] + nlp = MadNLPTests.DenseDummyQP(zeros(Float64, n); m = m) + solver_exact = MadNLPSolver( + nlp; + callback = MadNLP.DenseCallback, + print_level = MadNLP.ERROR, + kkt_system = KKT, + linear_solver = LapackOneMKLSolver, + ) + results_ref = MadNLP.solve!(solver_exact) + + nlp = MadNLPTests.DenseDummyQP(oneAPI.zeros(Float64, n); m = m) + solver_qn = MadNLPSolver( + nlp; + callback = MadNLP.DenseCallback, + print_level = MadNLP.ERROR, + kkt_system = KKT, + hessian_approximation = QN, + linear_solver = LapackOneMKLSolver, + ) + results_qn = MadNLP.solve!(solver_qn) + + @test results_qn.status == MadNLP.SOLVE_SUCCEEDED + @test results_qn.objective ≈ results_ref.objective atol = 1.0e-6 + @test Array(results_qn.solution) ≈ Array(results_ref.solution) atol = 1.0e-6 + @test solver_qn.cnt.lag_hess_cnt == 0 + end +end diff --git a/lib/MadNLPGPU/test/madnlpgpu_test.jl b/lib/MadNLPGPU/test/madnlpgpu_test.jl index b9573a83d..21e701507 100644 --- a/lib/MadNLPGPU/test/madnlpgpu_test.jl +++ b/lib/MadNLPGPU/test/madnlpgpu_test.jl @@ -175,8 +175,37 @@ rocm_testset = [ ], ] +oneapi_testset = [ + [ + "LapackOneMKLSolver-LU", + () -> MadNLP.Optimizer( + linear_solver = LapackOneMKLSolver, + lapack_algorithm = MadNLP.LU, + print_level = MadNLP.ERROR, + ), + [], + ], + # [ + # "LapackOneMKLSolver-QR", + # ()->MadNLP.Optimizer( + # linear_solver=LapackOneMKLSolver, + # lapack_algorithm=MadNLP.QR, + # print_level=MadNLP.ERROR, + # ), + # [], + # ], + # [ + # "LapackOneMKLSolver-CHOLESKY", + # ()->MadNLP.Optimizer( + # linear_solver=LapackOneMKLSolver, + # lapack_algorithm=MadNLP.CHOLESKY, + # print_level=MadNLP.ERROR, + # ), + # ["infeasible", "lootsma", "eigmina", "lp_examodels_issue75"], # KKT system not PD + # ], +] @testset "MadNLPGPU test" begin - if CUDA.functional() + if CUDA.functional() && GPU_BACKEND == "cuda" MadNLPTests.test_linear_solver(LapackCUDASolver,Float32) MadNLPTests.test_linear_solver(LapackCUDASolver,Float64) # Test LapackGPU wrapper @@ -184,11 +213,18 @@ rocm_testset = [ test_madnlp(name,optimizer_constructor,exclude; Arr=CuArray) end end - if AMDGPU.functional() + if GPU_BACKEND == "amdgpu" && AMDGPU.functional() MadNLPTests.test_linear_solver(LapackROCmSolver,Float32) MadNLPTests.test_linear_solver(LapackROCmSolver,Float64) for (name,optimizer_constructor,exclude) in rocm_testset test_madnlp(name,optimizer_constructor,exclude; Arr=ROCArray) end end + if GPU_BACKEND == "oneapi" && oneAPI.functional() + MadNLPTests.test_linear_solver(LapackOneMKLSolver, Float32) + MadNLPTests.test_linear_solver(LapackOneMKLSolver, Float64) + for (name, optimizer_constructor, exclude) in oneapi_testset + test_madnlp(name, optimizer_constructor, exclude; Arr = oneArray) + end + end end diff --git a/lib/MadNLPGPU/test/runtests.jl b/lib/MadNLPGPU/test/runtests.jl index 8f754f280..683321385 100644 --- a/lib/MadNLPGPU/test/runtests.jl +++ b/lib/MadNLPGPU/test/runtests.jl @@ -1,15 +1,28 @@ -using Test, CUDA, AMDGPU, MadNLP, MadNLPGPU, MadNLPTests +using Test, CUDA, MadNLP, MadNLPGPU, MadNLPTests + +# Get backend from environment +const GPU_BACKEND = get(ENV, "MADNLP_GPU_BACKEND", "cuda") + +# Conditionally load GPU backends +if GPU_BACKEND == "amdgpu" + using AMDGPU +elseif GPU_BACKEND == "oneapi" + using oneAPI +end @testset "MadNLPGPU test" begin include("madnlpgpu_test.jl") - if CUDA.functional() + if GPU_BACKEND == "cuda" && CUDA.functional() include("densekkt_cuda.jl") # Need to add support for CompactLBFGS in SparseCondensedKKTSystem (Issue #563) # include("sparsekkt_cuda.jl") end - if AMDGPU.functional() + if GPU_BACKEND == "amdgpu" && AMDGPU.functional() include("densekkt_rocm.jl") # Need to add support for CompactLBFGS in SparseCondensedKKTSystem (Issue #563) # include("sparsekkt_rocm.jl") end + if GPU_BACKEND == "oneapi" && oneAPI.functional() + include("densekkt_oneapi.jl") + end end