From 28fdf080501695930d0a8f1395a921b4f57f3c8e Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 17 Aug 2026 04:21:11 -0400 Subject: [PATCH 1/6] Allow turning off allowscalar copies for yacusolver --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 127 ++++++++++++++-------- 1 file changed, 79 insertions(+), 48 deletions(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 0cfa64c3c..b07ae6f32 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -26,7 +26,8 @@ for (bname, fname, elty, relty) in A::StridedCuMatrix{$elty}, S::StridedCuVector{$relty} = similar(A, $relty, min(size(A)...)), U::StridedCuMatrix{$elty} = similar(A, $elty, size(A, 1), min(size(A)...)), - Vᴴ::StridedCuMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)) + Vᴴ::StridedCuMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); + check::Bool = false, ) chkstride1(A, U, Vᴴ, S) m, n = size(A) @@ -90,8 +91,10 @@ for (bname, fname, elty, relty) in end CUDA.unsafe_free!(rwork) - info = @allowscalar dh.info[1] - cuSOLVER.chkargsok(BlasInt(info)) + if check + info = @allowscalar dh.info[1] + cuSOLVER.chkargsok(BlasInt(info)) + end return (S, U, Vᴴ) end @@ -103,7 +106,8 @@ function gesvdp!( S::StridedCuVector = similar(A, real(T), min(size(A)...)), U::StridedCuMatrix{T} = similar(A, T, size(A, 1), min(size(A)...)), Vᴴ::StridedCuMatrix{T} = similar(A, T, min(size(A)...), size(A, 2)); - tol = norm(A) * eps(real(T)) + tol = norm(A) * eps(real(T)), + check::Bool = false, ) where {T <: BlasFloat} chkstride1(A, U, S, Vᴴ) m, n = size(A) @@ -166,8 +170,10 @@ function gesvdp!( err = h_err_sigma[] err > tol && @warn "gesvdp! did not attain the requested tolerance: error = $err > tolerance = $tol" - flag = @allowscalar dh.info[1] - cuSOLVER.chklapackerror(BlasInt(flag)) + if check + flag = @allowscalar dh.info[1] + cuSOLVER.chklapackerror(BlasInt(flag)) + end if Ũ !== U && length(U) > 0 U .= view(Ũ, 1:m, 1:size(U, 2)) end @@ -196,6 +202,7 @@ for (bname, fname, elty, relty) in U::StridedCuMatrix{$elty} = similar(A, $elty, size(A, 1), min(size(A)...)), Vᴴ::StridedCuMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); tol::$relty = eps($relty), + check::Bool = false, max_sweeps::Int = 100, kwargs... ) @@ -253,8 +260,10 @@ for (bname, fname, elty, relty) in ) end - info = @allowscalar dh.info[1] - cuSOLVER.chkargsok(BlasInt(info)) + if check + info = @allowscalar dh.info[1] + cuSOLVER.chkargsok(BlasInt(info)) + end cuSOLVER.cusolverDnDestroyGesvdjInfo(params[]) @@ -272,6 +281,7 @@ function gesvdr!( S::StridedCuVector = similar(A, real(T), min(size(A)...)), U::StridedCuMatrix{T} = similar(A, T, size(A, 1), min(size(A)...)), Vᴴ::StridedCuMatrix{T} = similar(A, T, min(size(A)...), size(A, 2)); + check::Bool = false, k::Int = length(S), p::Int = min(size(A)...) - k - 1, niters::Int = 1 @@ -319,8 +329,10 @@ function gesvdr!( ) end - flag = @allowscalar dh.info[1] - cuSOLVER.chklapackerror(BlasInt(flag)) + if check + flag = @allowscalar dh.info[1] + cuSOLVER.chklapackerror(BlasInt(flag)) + end if Ũ !== U && length(U) > 0 U .= view(Ũ, 1:m, 1:size(U, 2)) end @@ -336,7 +348,7 @@ end # Wrapper for general eigensolver for (celty, elty) in ((:ComplexF32, :Float32), (:ComplexF64, :Float64), (:ComplexF32, :ComplexF32), (:ComplexF64, :ComplexF64)) @eval begin - function Xgeev!(A::StridedCuMatrix{$elty}, D::StridedCuVector{$celty}, V::StridedCuMatrix{$celty}) + function Xgeev!(A::StridedCuMatrix{$elty}, D::StridedCuVector{$celty}, V::StridedCuMatrix{$celty}; check::Bool = false) require_one_based_indexing(A, V, D) chkstride1(A, V, D) n = checksquare(A) @@ -391,8 +403,10 @@ for (celty, elty) in ((:ComplexF32, :Float32), (:ComplexF64, :Float64), (:Comple sizeof(buffer_gpu), buffer_cpu, sizeof(buffer_cpu), dh.info ) end - flag = @allowscalar dh.info[1] - cuSOLVER.chkargsok(BlasInt(flag)) + if check + flag = @allowscalar dh.info[1] + cuSOLVER.chkargsok(BlasInt(flag)) + end if eltype(A) <: Real work = CuVector{$elty}(undef, n) DR = view(D2, 1:n) @@ -414,7 +428,8 @@ end # jobz::Char, # uplo::Char, # A::StridedCuMatrix{$elty}, -# B::StridedCuMatrix{$elty}) +# B::StridedCuMatrix{$elty}; +# check::Bool = false) # chkuplo(uplo) # nA, nB = checksquare(A, B) # if nB != nA @@ -436,9 +451,11 @@ end # return $fname(dh, itype, jobz, uplo, n, A, lda, B, ldb, W, # buffer, sizeof(buffer) ÷ sizeof($elty), dh.info) # end - -# info = @allowscalar dh.info[1] -# chkargsok(BlasInt(info)) +# +# if check +# info = @allowscalar dh.info[1] +# chkargsok(BlasInt(info)) +# end # if jobz == 'N' # return W @@ -460,6 +477,7 @@ end # uplo::Char, # A::StridedCuMatrix{$elty}, # B::StridedCuMatrix{$elty}; +# check::Bool = false, # tol::$relty=eps($relty), # max_sweeps::Int=100) # chkuplo(uplo) @@ -488,10 +506,10 @@ end # return $fname(dh, itype, jobz, uplo, n, A, lda, B, ldb, W, # buffer, sizeof(buffer) ÷ sizeof($elty), dh.info, params[]) # end - -# info = @allowscalar dh.info[1] -# chkargsok(BlasInt(info)) - +# if check +# info = @allowscalar dh.info[1] +# chkargsok(BlasInt(info)) +# end # cusolverDnDestroySyevjInfo(params[]) # if jobz == 'N' @@ -516,6 +534,7 @@ end # function $jname(jobz::Char, # uplo::Char, # A::StridedCuArray{$elty}; +# check::Bool = false, # tol::$relty=eps($relty), # max_sweeps::Int=100) @@ -548,14 +567,15 @@ end # sizeof(buffer) ÷ sizeof($elty), dh.info, params[], batchSize) # end -# # Copy the solver info and delete the device memory -# info = @allowscalar collect(dh.info) +# if check +# # Copy the solver info and delete the device memory +# info = collect(dh.info) -# # Double check the solver's exit status -# for i in 1:batchSize -# chkargsok(BlasInt(info[i])) +# # Double check the solver's exit status +# for i in 1:batchSize +# chkargsok(BlasInt(info[i])) +# end # end - # cusolverDnDestroySyevjInfo(params[]) # # Return eigenvalues (in W) and possibly eigenvectors (in A) @@ -575,7 +595,8 @@ end # @eval begin # function potrsBatched!(uplo::Char, # A::Vector{<:StridedCuMatrix{$elty}}, -# B::Vector{<:StridedCuVecOrMat{$elty}}) +# B::Vector{<:StridedCuVecOrMat{$elty}}; +# check::Bool = false,) # if length(A) != length(B) # throw(DimensionMismatch("")) # end @@ -602,10 +623,11 @@ end # # Run the solver # $fname(dh, uplo, n, nrhs, Aptrs, lda, Bptrs, ldb, dh.info, batchSize) -# # Copy the solver info and delete the device memory -# info = @allowscalar dh.info[1] -# chklapackerror(BlasInt(info)) - +# if check +# # Copy the solver info and delete the device memory +# info = @allowscalar dh.info[1] +# chklapackerror(BlasInt(info)) +# end # return B # end # end @@ -616,7 +638,7 @@ end # (:cusolverDnCpotrfBatched, :ComplexF32), # (:cusolverDnZpotrfBatched, :ComplexF64)) # @eval begin -# function potrfBatched!(uplo::Char, A::Vector{<:StridedCuMatrix{$elty}}) +# function potrfBatched!(uplo::Char, A::Vector{<:StridedCuMatrix{$elty}}; check::Bool = false) # # Set up information for the solver arguments # chkuplo(uplo) @@ -632,14 +654,15 @@ end # # Run the solver # $fname(dh, uplo, n, Aptrs, lda, dh.info, batchSize) -# # Copy the solver info and delete the device memory -# info = @allowscalar collect(dh.info) +# if check +# # Copy the solver info and delete the device memory +# info = collect(dh.info) -# # Double check the solver's exit status -# for i in 1:batchSize -# chkargsok(BlasInt(info[i])) +# # Double check the solver's exit status +# for i in 1:batchSize +# chkargsok(BlasInt(info[i])) +# end # end - # # info[i] > 0 means the leading minor of order info[i] is not positive definite # # LinearAlgebra.LAPACK does not throw Exception here # # to simplify calls to isposdef! and factorize @@ -649,7 +672,8 @@ end # end # # gesv -# function gesv!(X::CuVecOrMat{T}, A::CuMatrix{T}, B::CuVecOrMat{T}; fallback::Bool=true, +# function gesv!(X::CuVecOrMat{T}, A::CuMatrix{T}, B::CuVecOrMat{T}; +# fallback::Bool=true, check::Bool = false, # residual_history::Bool=false, irs_precision::String="AUTO", # refinement_solver::String="CLASSICAL", # maxiters::Int=0, maxiters_inner::Int=0, tol::Float64=0.0, @@ -703,10 +727,11 @@ end # X, ldx, buffer, sizeof(buffer), niters, dh.info) # end -# # Copy the solver flag and delete the device memory -# flag = @allowscalar dh.info[1] -# chklapackerror(BlasInt(flag)) - +# if check +# # Copy the solver flag and delete the device memory +# flag = @allowscalar dh.info[1] +# chklapackerror(BlasInt(flag)) +# end # return X, info # end @@ -721,6 +746,7 @@ for (bname, fname, elty, relty) in ( A::StridedCuMatrix{$elty}, W::StridedCuVector{$relty}, V::StridedCuMatrix{$elty}; + check::Bool = false, uplo::Char = 'U', tol::$relty = eps($relty), max_sweeps::Int = 100 @@ -752,8 +778,10 @@ for (bname, fname, elty, relty) in ( ) end - info = @allowscalar dh.info[1] - chkargsok(BlasInt(info)) + if check + info = @allowscalar dh.info[1] + chkargsok(BlasInt(info)) + end if jobz == 'V' && V !== A copy!(V, A) @@ -767,6 +795,7 @@ function heevd!( A::StridedCuMatrix{T}, W::StridedCuVector{Tr}, V::StridedCuMatrix{T}; + check::Bool = false, uplo::Char = 'U' ) where {T <: BlasFloat, Tr <: BlasReal} chkuplo(uplo) @@ -800,8 +829,10 @@ function heevd!( ) end - info = @allowscalar dh.info[1] - chkargsok(BlasInt(info)) + if check + info = @allowscalar dh.info[1] + chkargsok(BlasInt(info)) + end if jobz == 'V' && V !== A copy!(V, A) From 5ac35c3f8bc740fc90e0db30cfba1aff842e78a0 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 17 Sep 2026 11:50:30 +0200 Subject: [PATCH 2/6] Make the checks on by default and add a docstring --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 26 +++++++++++------------ src/MatrixAlgebraKit.jl | 15 +++++++++++++ 2 files changed, 28 insertions(+), 13 deletions(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index b07ae6f32..19fa3a6ef 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -27,7 +27,7 @@ for (bname, fname, elty, relty) in S::StridedCuVector{$relty} = similar(A, $relty, min(size(A)...)), U::StridedCuMatrix{$elty} = similar(A, $elty, size(A, 1), min(size(A)...)), Vᴴ::StridedCuMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); - check::Bool = false, + check::Bool = CHECK_LIBRARY_CALLS[], ) chkstride1(A, U, Vᴴ, S) m, n = size(A) @@ -107,7 +107,7 @@ function gesvdp!( U::StridedCuMatrix{T} = similar(A, T, size(A, 1), min(size(A)...)), Vᴴ::StridedCuMatrix{T} = similar(A, T, min(size(A)...), size(A, 2)); tol = norm(A) * eps(real(T)), - check::Bool = false, + check::Bool = CHECK_LIBRARY_CALLS[], ) where {T <: BlasFloat} chkstride1(A, U, S, Vᴴ) m, n = size(A) @@ -202,7 +202,7 @@ for (bname, fname, elty, relty) in U::StridedCuMatrix{$elty} = similar(A, $elty, size(A, 1), min(size(A)...)), Vᴴ::StridedCuMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); tol::$relty = eps($relty), - check::Bool = false, + check::Bool = CHECK_LIBRARY_CALLS[], max_sweeps::Int = 100, kwargs... ) @@ -281,7 +281,7 @@ function gesvdr!( S::StridedCuVector = similar(A, real(T), min(size(A)...)), U::StridedCuMatrix{T} = similar(A, T, size(A, 1), min(size(A)...)), Vᴴ::StridedCuMatrix{T} = similar(A, T, min(size(A)...), size(A, 2)); - check::Bool = false, + check::Bool = CHECK_LIBRARY_CALLS[], k::Int = length(S), p::Int = min(size(A)...) - k - 1, niters::Int = 1 @@ -348,7 +348,7 @@ end # Wrapper for general eigensolver for (celty, elty) in ((:ComplexF32, :Float32), (:ComplexF64, :Float64), (:ComplexF32, :ComplexF32), (:ComplexF64, :ComplexF64)) @eval begin - function Xgeev!(A::StridedCuMatrix{$elty}, D::StridedCuVector{$celty}, V::StridedCuMatrix{$celty}; check::Bool = false) + function Xgeev!(A::StridedCuMatrix{$elty}, D::StridedCuVector{$celty}, V::StridedCuMatrix{$celty}; check::Bool = CHECK_LIBRARY_CALLS[]) require_one_based_indexing(A, V, D) chkstride1(A, V, D) n = checksquare(A) @@ -429,7 +429,7 @@ end # uplo::Char, # A::StridedCuMatrix{$elty}, # B::StridedCuMatrix{$elty}; -# check::Bool = false) +# check::Bool = CHECK_LIBRARY_CALLS[]) # chkuplo(uplo) # nA, nB = checksquare(A, B) # if nB != nA @@ -477,7 +477,7 @@ end # uplo::Char, # A::StridedCuMatrix{$elty}, # B::StridedCuMatrix{$elty}; -# check::Bool = false, +# check::Bool = CHECK_LIBRARY_CALLS[], # tol::$relty=eps($relty), # max_sweeps::Int=100) # chkuplo(uplo) @@ -534,7 +534,7 @@ end # function $jname(jobz::Char, # uplo::Char, # A::StridedCuArray{$elty}; -# check::Bool = false, +# check::Bool = CHECK_LIBRARY_CALLS[], # tol::$relty=eps($relty), # max_sweeps::Int=100) @@ -596,7 +596,7 @@ end # function potrsBatched!(uplo::Char, # A::Vector{<:StridedCuMatrix{$elty}}, # B::Vector{<:StridedCuVecOrMat{$elty}}; -# check::Bool = false,) +# check::Bool = CHECK_LIBRARY_CALLS[],) # if length(A) != length(B) # throw(DimensionMismatch("")) # end @@ -638,7 +638,7 @@ end # (:cusolverDnCpotrfBatched, :ComplexF32), # (:cusolverDnZpotrfBatched, :ComplexF64)) # @eval begin -# function potrfBatched!(uplo::Char, A::Vector{<:StridedCuMatrix{$elty}}; check::Bool = false) +# function potrfBatched!(uplo::Char, A::Vector{<:StridedCuMatrix{$elty}}; check::Bool = CHECK_LIBRARY_CALLS[]) # # Set up information for the solver arguments # chkuplo(uplo) @@ -673,7 +673,7 @@ end # # gesv # function gesv!(X::CuVecOrMat{T}, A::CuMatrix{T}, B::CuVecOrMat{T}; -# fallback::Bool=true, check::Bool = false, +# fallback::Bool=true, check::Bool = CHECK_LIBRARY_CALLS[], # residual_history::Bool=false, irs_precision::String="AUTO", # refinement_solver::String="CLASSICAL", # maxiters::Int=0, maxiters_inner::Int=0, tol::Float64=0.0, @@ -746,7 +746,7 @@ for (bname, fname, elty, relty) in ( A::StridedCuMatrix{$elty}, W::StridedCuVector{$relty}, V::StridedCuMatrix{$elty}; - check::Bool = false, + check::Bool = CHECK_LIBRARY_CALLS[], uplo::Char = 'U', tol::$relty = eps($relty), max_sweeps::Int = 100 @@ -795,7 +795,7 @@ function heevd!( A::StridedCuMatrix{T}, W::StridedCuVector{Tr}, V::StridedCuMatrix{T}; - check::Bool = false, + check::Bool = CHECK_LIBRARY_CALLS[], uplo::Char = 'U' ) where {T <: BlasFloat, Tr <: BlasReal} chkuplo(uplo) diff --git a/src/MatrixAlgebraKit.jl b/src/MatrixAlgebraKit.jl index 4d4e0084e..a3e007df2 100644 --- a/src/MatrixAlgebraKit.jl +++ b/src/MatrixAlgebraKit.jl @@ -145,4 +145,19 @@ include("pushforwards/svd.jl") include("precompile.jl") +""" + CHECK_LIBRARY_CALLS::Ref{Bool} + +Library setting which controls whether GPU solver libraries (CUSOLVER, rocSOLVER) +perform expensive checking operations after library calls. These libraries implement +the check by doing an `@allowscalar` read of a GPU array, which forces device +synchronization. Default is `Ref(true)` (checks are performed), but this can be +changed by setting `MatrixAlgebraKit.CHECK_LIBRARY_CALLS[] = false`. + +!!! warning + Disabling the checks means you may encounter errors later in program execution, + far from their original source, which will make debugging substantially more difficult. +""" +const CHECK_LIBRARY_CALLS = Ref(true) + end From 4f992008aaeda92046fd2a59fe8afb9b19cdf441 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 17 Sep 2026 12:04:06 +0200 Subject: [PATCH 3/6] Missing import --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 1 + 1 file changed, 1 insertion(+) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 19fa3a6ef..245076a81 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -7,6 +7,7 @@ using LinearAlgebra.LAPACK: chkargsok, chklapackerror, chktrans, chkside, chkdia using CUDA using CUDA: @allowscalar, i32 using CUDA.cuSOLVER +using ..MatrixAlgebraKit: CHECK_LIBRARY_CALLS # QR methods are implemented with full access to allocated arrays, so we do not need to redo this: using CUDA.cuSOLVER: geqrf!, ormqr!, orgqr! From dbf81ba57eecf6a32d8348956585394d2b463e6a Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 17 Sep 2026 12:51:06 +0200 Subject: [PATCH 4/6] Formatter --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 245076a81..bd19257aa 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -7,7 +7,7 @@ using LinearAlgebra.LAPACK: chkargsok, chklapackerror, chktrans, chkside, chkdia using CUDA using CUDA: @allowscalar, i32 using CUDA.cuSOLVER -using ..MatrixAlgebraKit: CHECK_LIBRARY_CALLS +using ..MatrixAlgebraKit: CHECK_LIBRARY_CALLS # QR methods are implemented with full access to allocated arrays, so we do not need to redo this: using CUDA.cuSOLVER: geqrf!, ormqr!, orgqr! From 68132deebedf5b34b8e8a9696bb1fc56b01cc23c Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 23 Sep 2026 10:44:42 +0200 Subject: [PATCH 5/6] Add the flags for YAROCSOLVER too --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 129 ++++++++++++------- 1 file changed, 81 insertions(+), 48 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 28e206660..8cbfeed16 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -8,6 +8,7 @@ using AMDGPU using AMDGPU: @allowscalar using AMDGPU.rocSOLVER using AMDGPU.rocBLAS +using ..MatrixAlgebraKit: CHECK_LIBRARY_CALLS # QR methods are implemented with full access to allocated arrays, so we do not need to redo this: using AMDGPU.rocSOLVER: geqrf!, ormqr!, orgqr! @@ -27,7 +28,8 @@ for (fname, elty, relty) in A::StridedROCMatrix{$elty}, S::StridedROCVector{$relty} = similar(A, $relty, min(size(A)...)), U::StridedROCMatrix{$elty} = similar(A, $elty, size(A, 1), min(size(A)...)), - Vᴴ::StridedROCMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)) + Vᴴ::StridedROCMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); + check::Bool = CHECK_LIBRARY_CALLS[], ) chkstride1(A, U, Vᴴ, S) m, n = size(A) @@ -85,8 +87,10 @@ for (fname, elty, relty) in ) AMDGPU.unsafe_free!(rwork) - info = @allowscalar dev_info[1] - rocSOLVER.chkargsok(BlasInt(info)) + if check + info = @allowscalar dev_info[1] + rocSOLVER.chkargsok(BlasInt(info)) + end return (S, U, Vᴴ) end @@ -109,6 +113,7 @@ for (fname, elty, relty) in Vᴴ::StridedROCMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); tol::$relty = eps($relty), max_sweeps::Int = 100, + check::Bool = CHECK_LIBRARY_CALLS[], ) chkstride1(A, U, Vᴴ, S) m, n = size(A) @@ -165,8 +170,10 @@ for (fname, elty, relty) in S, U, ldu, Vᴴ, ldv, dev_info, ) - info = @allowscalar dev_info[1] - rocSOLVER.chkargsok(BlasInt(info)) + if check + info = @allowscalar dev_info[1] + rocSOLVER.chkargsok(BlasInt(info)) + end AMDGPU.unsafe_free!(dev_residual) AMDGPU.unsafe_free!(dev_n_sweeps) @@ -275,7 +282,9 @@ end # jobz::Char, # uplo::Char, # A::StridedROCMatrix{$elty}, -# B::StridedROCMatrix{$elty}) +# B::StridedROCMatrix{$elty}; +# check::Bool = CHECK_LIBRARY_CALLS[], +# ) # chkuplo(uplo) # nA, nB = checksquare(A, B) # if nB != nA @@ -298,9 +307,10 @@ end # buffer, sizeof(buffer) ÷ sizeof($elty), dh.info) # end -# info = @allowscalar dh.info[1] -# chkargsok(BlasInt(info)) - +# if check +# info = @allowscalar dh.info[1] +# chkargsok(BlasInt(info)) +# end # if jobz == 'N' # return W # elseif jobz == 'V' @@ -322,7 +332,9 @@ end # A::StridedROCMatrix{$elty}, # B::StridedROCMatrix{$elty}; # tol::$relty=eps($relty), -# max_sweeps::Int=100) +# max_sweeps::Int=100; +# check::Bool = CHECK_LIBRARY_CALLS[], +# ) # chkuplo(uplo) # nA, nB = checksquare(A, B) # if nB != nA @@ -349,9 +361,10 @@ end # return $fname(dh, itype, jobz, uplo, n, A, lda, B, ldb, W, # buffer, sizeof(buffer) ÷ sizeof($elty), dh.info, params[]) # end - -# info = @allowscalar dh.info[1] -# chkargsok(BlasInt(info)) +# if check +# info = @allowscalar dh.info[1] +# chkargsok(BlasInt(info)) +# end # rocsolverDnDestroySyevjInfo(params[]) @@ -378,7 +391,9 @@ end # uplo::Char, # A::StridedROCArray{$elty}; # tol::$relty=eps($relty), -# max_sweeps::Int=100) +# max_sweeps::Int=100; +# check::Bool = CHECK_LIBRARY_CALLS[], +# ) # # Set up information for the solver arguments # chkuplo(uplo) @@ -409,14 +424,15 @@ end # sizeof(buffer) ÷ sizeof($elty), dh.info, params[], batchSize) # end -# # Copy the solver info and delete the device memory -# info = @allowscalar collect(dh.info) - -# # Double check the solver's exit status -# for i in 1:batchSize -# chkargsok(BlasInt(info[i])) -# end - +# if check +# # Copy the solver info and delete the device memory +# info = collect(dh.info) +# +# # Double check the solver's exit status +# for i in 1:batchSize +# chkargsok(BlasInt(info[i])) +# end +# end # rocsolverDnDestroySyevjInfo(params[]) # # Return eigenvalues (in W) and possibly eigenvectors (in A) @@ -436,7 +452,9 @@ end # @eval begin # function potrsBatched!(uplo::Char, # A::Vector{<:StridedROCMatrix{$elty}}, -# B::Vector{<:StridedROCVecOrMat{$elty}}) +# B::Vector{<:StridedROCVecOrMat{$elty}}; +# check::Bool = CHECK_LIBRARY_CALLS[], +# ) # if length(A) != length(B) # throw(DimensionMismatch("")) # end @@ -462,11 +480,11 @@ end # # Run the solver # $fname(dh, uplo, n, nrhs, Aptrs, lda, Bptrs, ldb, dh.info, batchSize) - -# # Copy the solver info and delete the device memory -# info = @allowscalar dh.info[1] -# chklapackerror(BlasInt(info)) - +# if check +# # Copy the solver info and delete the device memory +# info = @allowscalar dh.info[1] +# chklapackerror(BlasInt(info)) +# end # return B # end # end @@ -477,7 +495,7 @@ end # (:rocsolverDnCpotrfBatched, :ComplexF32), # (:rocsolverDnZpotrfBatched, :ComplexF64)) # @eval begin -# function potrfBatched!(uplo::Char, A::Vector{<:StridedROCMatrix{$elty}}) +# function potrfBatched!(uplo::Char, A::Vector{<:StridedROCMatrix{$elty}}; check::Bool = CHECK_LIBRARY_CALLS[],) # # Set up information for the solver arguments # chkuplo(uplo) @@ -493,14 +511,15 @@ end # # Run the solver # $fname(dh, uplo, n, Aptrs, lda, dh.info, batchSize) -# # Copy the solver info and delete the device memory -# info = @allowscalar collect(dh.info) +# if check +# # Copy the solver info and delete the device memory +# info = collect(dh.info) -# # Double check the solver's exit status -# for i in 1:batchSize -# chkargsok(BlasInt(info[i])) +# # Double check the solver's exit status +# for i in 1:batchSize +# chkargsok(BlasInt(info[i])) +# end # end - # # info[i] > 0 means the leading minor of order info[i] is not positive definite # # LinearAlgebra.LAPACK does not throw Exception here # # to simplify calls to isposdef! and factorize @@ -510,11 +529,13 @@ end # end # # gesv -# function gesv!(X::CuVecOrMat{T}, A::CuMatrix{T}, B::CuVecOrMat{T}; fallback::Bool=true, +# function gesv!(X::ROCVecOrMat{T}, A::ROCMatrix{T}, B::ROCVecOrMat{T}; fallback::Bool=true, # residual_history::Bool=false, irs_precision::String="AUTO", # refinement_solver::String="CLASSICAL", # maxiters::Int=0, maxiters_inner::Int=0, tol::Float64=0.0, -# tol_inner=Float64 = 0.0) where {T<:BlasFloat} +# tol_inner=Float64 = 0.0, +# check::Bool = CHECK_LIBRARY_CALLS[], +# ) where {T<:BlasFloat} # params = CuSolverIRSParameters() # info = CuSolverIRSInformation() # n = checksquare(A) @@ -583,7 +604,8 @@ for (heevd, heev, heevx, heevj, elty, relty) in A::StridedROCMatrix{$elty}, W::StridedROCVector{$relty}, V::StridedROCMatrix{$elty}; - uplo::Char = 'U' + uplo::Char = 'U', + check::Bool = CHECK_LIBRARY_CALLS[], ) chkuplo(uplo) n = checksquare(A) @@ -601,8 +623,10 @@ for (heevd, heev, heevx, heevj, elty, relty) in roc_uplo = convert(rocSOLVER.rocblas_fill, uplo) $heevd(dh, jobz, roc_uplo, n, A, lda, W, work, dev_info) - info = @allowscalar dev_info[1] - chkargsok(BlasInt(info)) + if check + info = @allowscalar dev_info[1] + chkargsok(BlasInt(info)) + end if jobz == rocSOLVER.rocblas_evect_original && V !== A copy!(V, A) @@ -613,7 +637,8 @@ for (heevd, heev, heevx, heevj, elty, relty) in A::StridedROCMatrix{$elty}, W::StridedROCVector{$relty}, V::StridedROCMatrix{$elty}; - uplo::Char = 'U' + uplo::Char = 'U', + check::Bool = CHECK_LIBRARY_CALLS[], ) chkuplo(uplo) n = checksquare(A) @@ -631,8 +656,10 @@ for (heevd, heev, heevx, heevj, elty, relty) in roc_uplo = convert(rocSOLVER.rocblas_fill, uplo) $heev(dh, jobz, roc_uplo, n, A, lda, W, work, dev_info) - info = @allowscalar dev_info[1] - chkargsok(BlasInt(info)) + if check + info = @allowscalar dev_info[1] + chkargsok(BlasInt(info)) + end if jobz == rocSOLVER.rocblas_evect_original && V !== A copy!(V, A) @@ -644,6 +671,7 @@ for (heevd, heev, heevx, heevj, elty, relty) in W::StridedROCVector{$relty}, V::StridedROCMatrix{$elty}; uplo::Char = 'U', + check::Bool = CHECK_LIBRARY_CALLS[], kwargs... ) chkuplo(uplo) @@ -680,8 +708,10 @@ for (heevd, heev, heevx, heevj, elty, relty) in roc_uplo = convert(rocSOLVER.rocblas_fill, uplo) $heevx(dh, jobz, range, roc_uplo, n, A, lda, vl, vu, il, iu, abstol, nev, W, V, ldv, ifail, dev_info) - info = @allowscalar dev_info[1] - chkargsok(BlasInt(info)) + if check + info = @allowscalar dev_info[1] + chkargsok(BlasInt(info)) + end m = @allowscalar nev[1] return W, V, m end @@ -692,7 +722,8 @@ for (heevd, heev, heevx, heevj, elty, relty) in uplo::Char = 'U', tol::$relty = eps($relty), max_sweeps::Int = 100, - sort::Char = 'N' + sort::Char = 'N', + check::Bool = CHECK_LIBRARY_CALLS[], ) chkuplo(uplo) n = checksquare(A) @@ -712,8 +743,10 @@ for (heevd, heev, heevx, heevj, elty, relty) in roc_sort = sort == 'N' ? rocSOLVER.rocblas_esort_none : rocSOLVER.rocblas_esort_ascending $heevj(dh, roc_sort, jobz, roc_uplo, n, A, lda, tol, residual, max_sweeps, n_sweeps, W, dev_info) - info = @allowscalar dev_info[1] - chkargsok(BlasInt(info)) + if check + info = @allowscalar dev_info[1] + chkargsok(BlasInt(info)) + end if jobz == rocSOLVER.rocblas_evect_original && V !== A copy!(V, A) From d8cf28cdba64226776257928b4e13c4d1639e8e2 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 23 Sep 2026 11:15:29 +0200 Subject: [PATCH 6/6] Add the check for gesvdx too --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 8cbfeed16..6c5bd3469 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -234,6 +234,7 @@ for (fname, elty, relty) in S::StridedROCVector{$relty} = similar(A, $relty, min(size(A)...)), U::StridedROCMatrix{$elty} = similar(A, $elty, size(A, 1), min(size(A)...)), Vᴴ::StridedROCMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); + check::Bool = CHECK_LIBRARY_CALLS[], kwargs... ) chkstride1(A, U, Vᴴ, S) @@ -258,8 +259,10 @@ for (fname, elty, relty) in S, U, ldu, Vᴴ, ldv, ifail, dev_info ) - info = @allowscalar dev_info[1] - rocSOLVER.chkargsok(BlasInt(info)) + if check + info = @allowscalar dev_info[1] + rocSOLVER.chkargsok(BlasInt(info)) + end # Zero the entries of `S` that `gesvdx` did not write. nv = @allowscalar Int(nsv[1]) nv < length(S) && fill!(view(S, (nv + 1):length(S)), zero(eltype(S)))