diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 0bdb10497..4a80ae0fd 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -6,9 +6,9 @@ using MatrixAlgebraKit: one!, zero!, uppertriangular!, lowertriangular! using MatrixAlgebraKit: diagview, sign_safe using MatrixAlgebraKit: ROCSOLVER, LQViaTransposedQR, TruncationStrategy, NoTruncation, TruncationByValue, AbstractAlgorithm using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eigh_algorithm -import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdj! +import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdx!, gesvdj! import MatrixAlgebraKit: heevj!, heevd!, heev!, heevx! -import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback! +import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback!, complete_svd_basis! using AMDGPU using LinearAlgebra using LinearAlgebra: BlasFloat @@ -28,7 +28,7 @@ for f in (:geqrf!, :ungqr!, :unmqr!) @eval $f(::ROCSOLVER, args...) = YArocSOLVER.$f(args...) end -MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi) +MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :bisection) function gesvd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) m, n = size(A) @@ -42,6 +42,12 @@ function gesvdj!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid return MatrixAlgebraKit.svd_via_adjoint!(gesvdj!, ROCSOLVER(), A, S, U, Vᴴ; kwargs...) end +function gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) + YArocSOLVER.gesvdx!(A, S, U, Vᴴ; kwargs...) + complete_svd_basis!(U, Vᴴ, length(S)) + return S, U, Vᴴ +end + heevj!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = YArocSOLVER.heevj!(A, Dd, V; kwargs...) heevd!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index e0c5f084d..28e206660 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -175,6 +175,96 @@ for (fname, elty, relty) in end end +# `gesvdx` computes all singular values, those in the half-open interval `[vl, vu)`, or those +# with an index in `irange`. Unlike the other algorithms, it never forms the full `U` and `Vᴴ`, +# so the only job modes are `singular` and `none`. +function _gesvdx_range(::Type{T}, kwargs) where {T <: Real} + if haskey(kwargs, :irange) + irange = convert(UnitRange{Int}, kwargs[:irange]) + return rocSOLVER.rocblas_srange_index, zero(T), zero(T), first(irange), last(irange) + elseif haskey(kwargs, :vl) || haskey(kwargs, :vu) + vl = convert(T, get(kwargs, :vl, -Inf)) + vu = convert(T, get(kwargs, :vu, +Inf)) + return rocSOLVER.rocblas_srange_value, vl, vu, 0, 0 + else + return rocSOLVER.rocblas_srange_all, zero(T), zero(T), 0, 0 + end +end + +function _gesvdx_jobs(U, Vᴴ, m::Integer, n::Integer, maxnsv::Integer) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A ($m) and U ($(size(U, 1)))")) + size(U, 2) >= maxnsv || + throw(DimensionMismatch("invalid column size of U")) + jobu = rocSOLVER.rocblas_svect_singular + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A ($n) and Vᴴ ($(size(Vᴴ, 2)))")) + size(Vᴴ, 1) >= maxnsv || + throw(DimensionMismatch("invalid row size of Vᴴ")) + jobvt = rocSOLVER.rocblas_svect_singular + end + return jobu, jobvt +end + +# Wrapper for SVD via Bisection +for (fname, elty, relty) in + ( + (:rocsolver_sgesvdx, :Float32, :Float32), + (:rocsolver_dgesvdx, :Float64, :Float64), + (:rocsolver_cgesvdx, :ComplexF32, :Float32), + (:rocsolver_zgesvdx, :ComplexF64, :Float64), + ) + @eval begin + function gesvdx!( + 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)); + kwargs... + ) + chkstride1(A, U, Vᴴ, S) + m, n = size(A) + minmn = min(m, n) + srange, vl, vu, il, iu = _gesvdx_range($relty, kwargs) + maxnsv = srange == rocSOLVER.rocblas_srange_index ? iu - il + 1 : minmn + jobu, jobvt = _gesvdx_jobs(U, Vᴴ, m, n, maxnsv) + length(S) == minmn || + throw(DimensionMismatch("length mismatch between A ($minmn) and S ($(length(S)))")) + + lda = max(1, stride(A, 2)) + ldu = max(1, stride(U, 2)) + ldv = max(1, stride(Vᴴ, 2)) + ifail = ROCVector{Cint}(undef, minmn) + nsv = ROCVector{Cint}(undef, 1) + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, 1) + rocSOLVER.$fname( + dh, jobu, jobvt, srange, m, n, + A, lda, vl, vu, il, iu, nsv, + S, U, ldu, Vᴴ, ldv, ifail, + dev_info + ) + info = @allowscalar dev_info[1] + rocSOLVER.chkargsok(BlasInt(info)) + # 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))) + + AMDGPU.unsafe_free!(nsv) + AMDGPU.unsafe_free!(ifail) + AMDGPU.unsafe_free!(dev_info) + return (S, U, Vᴴ) + end + end +end + # for (jname, bname, fname, elty, relty) in # ((:sygvd!, :rocsolverDnSsygvd_bufferSize, :rocsolverDnSsygvd, :Float32, :Float32), # (:sygvd!, :rocsolverDnDsygvd_bufferSize, :rocsolverDnDsygvd, :Float64, :Float64), diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 3d20e96d4..080019e71 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -142,10 +142,16 @@ function svd_via_adjoint!(f!::F, driver::Driver, A, S, U, Vᴴ; kwargs...) where end # LAPACK -for f! in (:gesdd!, :gesvd!, :gesvdx!, :gesdvd!) +for f! in (:gesdd!, :gesvd!, :gesdvd!) @eval $f!(::LAPACK, args...; kwargs...) = YALAPACK.$f!(args...; kwargs...) end +function gesvdx!(::LAPACK, A, S, U, Vᴴ; kwargs...) + YALAPACK.gesvdx!(A, S, U, Vᴴ; kwargs...) + complete_svd_basis!(U, Vᴴ, length(S)) + return S, U, Vᴴ +end + function gesvdj!(::LAPACK, A, S, U, Vᴴ; kwargs...) m, n = size(A) m >= n && return YALAPACK.gesvdj!(A, S, U, Vᴴ) @@ -221,7 +227,27 @@ for (f, f_lapack!, Alg) in ( end supports_svd_full(::Driver, ::Symbol) = false -supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration) +supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration, :bisection) + +# Some methods (e.g. `gesvdx`) only compute the leading `min(m, n)` singular vectors. +# If `U` or `Vᴴ` is square (`svd_full!`), the remaining columns (row) of +# `U` (`Vᴴ`) need to be filled with an orthonormal basis for the complement of the +# computed singular vectors. +function complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int) + if size(U, 2) > minmn + Uc = copy_input(qr_null, view(U, :, 1:minmn)) + N = view(U, :, (minmn + 1):size(U, 2)) + N′ = qr_null!(Uc, N) + N′ === N || copyto!(N, N′) + end + if size(Vᴴ, 1) > minmn + Vc = copy_input(lq_null, view(Vᴴ, 1:minmn, :)) + Nᴴ = view(Vᴴ, (minmn + 1):size(Vᴴ, 1), :) + Nᴴ′ = lq_null!(Vc, Nᴴ) + Nᴴ′ === Nᴴ || copyto!(Nᴴ, Nᴴ′) + end + return U, Vᴴ +end function svd_trunc_no_error!(A, USVᴴ, alg::TruncatedAlgorithm) U, S, Vᴴ = svd_compact!(A, USVᴴ, alg.alg) diff --git a/src/yalapack.jl b/src/yalapack.jl index 57e974bda..f6ca7a507 100644 --- a/src/yalapack.jl +++ b/src/yalapack.jl @@ -2235,8 +2235,11 @@ for (gesvd, gesdd, gesvdx, gejsv, gesvj, elty, relty) in else size(U, 1) == m || throw(DimensionMismatch("row size mismatch between A ($m) and U ($(size(U, 1)))")) - size(U, 2) >= (range == 'I' ? iu - il + 1 : minmn) || - throw(DimensionMismatch("invalid column size of U")) + if range == 'I' + (size(U, 2) >= iu - il + 1 && size(U, 2) <= m) || throw(DimensionMismatch("invalid column size of U")) + else + (size(U, 2) == minmn || size(U, 2) == m) || throw(DimensionMismatch("invalid column size of U")) + end jobu = 'V' end if length(Vᴴ) == 0 @@ -2244,8 +2247,11 @@ for (gesvd, gesdd, gesvdx, gejsv, gesvj, elty, relty) in else size(Vᴴ, 2) == n || throw(DimensionMismatch("column size mismatch between A ($n) and Vᴴ ($(size(Vᴴ, 2)))")) - size(Vᴴ, 1) >= (range == 'I' ? iu - il + 1 : minmn) || - throw(DimensionMismatch("invalid row size of Vᴴ")) + if range == 'I' + (size(Vᴴ, 1) >= iu - il + 1 && size(Vᴴ, 1) <= n) || throw(DimensionMismatch("invalid row size of Vᴴ")) + else + (size(Vᴴ, 1) == minmn || size(Vᴴ, 1) == n) || throw(DimensionMismatch("invalid row size of Vᴴ")) + end jobvt = 'V' end length(S) == minmn || diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index c69ed3a0e..5387252a3 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -21,7 +21,7 @@ if !is_buildkite # LAPACK algorithms: for T in BLASFloats, m in (0, 54), n in (0, 37, m, 63) TestSuite.seed_rng!(123) - LAPACK_SVD_ALGS = (QRIteration(), DivideAndConquer(), SafeDivideAndConquer(; fixgauge = true)) + LAPACK_SVD_ALGS = (QRIteration(), DivideAndConquer(), SafeDivideAndConquer(; fixgauge = true), Bisection()) TestSuite.test_svd(T, (m, n)) TestSuite.test_svd_algs(T, (m, n), LAPACK_SVD_ALGS) @static if VERSION > v"1.11-" # Jacobi broken on 1.10 @@ -81,7 +81,7 @@ if AMDGPU.functional() for T in BLASFloats, m in (0, 23), n in (0, 17, m, 27) TestSuite.seed_rng!(123) TestSuite.test_svd(ROCMatrix{T}, (m, n)) - AMD_SVD_ALGS = (QRIteration(), Jacobi()) + AMD_SVD_ALGS = (QRIteration(), Jacobi(), Bisection()) TestSuite.test_svd_algs(ROCMatrix{T}, (m, n), AMD_SVD_ALGS) end