Skip to content

Commit 578a1b7

Browse files
committed
more carefully designated supported drivers
1 parent c31d563 commit 578a1b7

3 files changed

Lines changed: 8 additions & 4 deletions

File tree

ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,9 @@ for f in (:geqrf!, :ungqr!, :unmqr!)
2727
@eval $f(::ROCSOLVER, args...) = YArocSOLVER.$f(args...)
2828
end
2929

30+
MatrixAlgebraKit.supports_svd(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi)
31+
MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi)
32+
3033
function gesvd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...)
3134
m, n = size(A)
3235
m >= n && return YArocSOLVER.gesvd!(A, S, U, Vᴴ)
@@ -38,6 +41,7 @@ function gesvdj!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid
3841
m >= n && return YArocSOLVER.gesvdj!(A, S, U, Vᴴ; kwargs...)
3942
return MatrixAlgebraKit.svd_via_adjoint!(gesvdj!, ROCSOLVER(), A, S, U, Vᴴ; kwargs...)
4043
end
44+
4145
_gpu_heevj!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
4246
YArocSOLVER.heevj!(A, Dd, V; kwargs...)
4347
_gpu_heevd!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =

ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,10 +27,14 @@ function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T
2727
return CUSOLVER_DivideAndConquer(; kwargs...)
2828
end
2929

30+
3031
for f in (:geqrf!, :ungqr!, :unmqr!)
3132
@eval $f(::CUSOLVER, args...) = YACUSOLVER.$f(args...)
3233
end
3334

35+
MatrixAlgebraKit.supports_svd(::CUSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :svd_polar)
36+
MatrixAlgebraKit.supports_svd_full(::CUSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :svd_polar)
37+
3438
function gesvd!(::CUSOLVER, A::StridedCuMatrix, S::StridedCuVector, U::StridedCuMatrix, Vᴴ::StridedCuMatrix; kwargs...)
3539
m, n = size(A)
3640
m >= n && return YACUSOLVER.gesvd!(A, S, U, Vᴴ)

src/implementations/svd.jl

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -216,13 +216,9 @@ end
216216
supports_svd(::Driver, ::Symbol) = false
217217
supports_svd(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration, :bisection, :jacobi)
218218
supports_svd(::GLA, f::Symbol) = f === :qr_iteration
219-
supports_svd(::CUSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :svd_polar)
220-
supports_svd(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi)
221219
supports_svd_full(::Driver, ::Symbol) = false
222220
supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration)
223221
supports_svd_full(::GLA, f::Symbol) = f === :qr_iteration
224-
supports_svd_full(::CUSOLVER, f::Symbol) = f === :qr_iteration
225-
supports_svd_full(::ROCSOLVER, f::Symbol) = f === :qr_iteration
226222

227223
function svd_trunc_no_error!(A, USVᴴ, alg::TruncatedAlgorithm)
228224
U, S, Vᴴ = svd_compact!(A, USVᴴ, alg.alg)

0 commit comments

Comments
 (0)