diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 4a80ae0fd..2d722fe77 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -6,18 +6,26 @@ 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!, gesvdx!, gesvdj! +import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesdd!, gesvdx!, gesvdj! +import MatrixAlgebraKit: gesvdj_batched!, gesdd_batched!, gesvd_batched!, gesvdx_batched! import MatrixAlgebraKit: heevj!, heevd!, heev!, heevx! -import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback!, complete_svd_basis! +import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback! using AMDGPU using LinearAlgebra using LinearAlgebra: BlasFloat include("yarocsolver.jl") -MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedROCVecOrMat{<:BlasFloat}} = ROCSOLVER() +MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedROCArray{<:BlasFloat}} = ROCSOLVER() +MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} = ROCSOLVER() -function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}} +function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCMatrix{<:BlasFloat}} + return QRIteration(; kwargs...) +end +function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCArray{<:BlasFloat, 3}} + return QRIteration(; kwargs...) +end +function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}} @@ -28,7 +36,12 @@ 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, :bisection) +MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :divide_and_conquer, :bisection) + +# rocSOLVER's `gesvd*_batched` functions take the batch as an array +# of device pointers, so a group of equally sized matrices +# doesn't need a copy into a 3D ROCArray. +MatrixAlgebraKit.supports_pointer_batch(::AbstractAlgorithm, ::ROCSOLVER, ::Type{<:StridedROCMatrix{<:BlasFloat}}) = true function gesvd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) m, n = size(A) @@ -42,11 +55,38 @@ 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...) +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ᴴ + +# rocSOLVER's batched `gesvd` requires m ≥ n, so wide matrices go through the adjoint +function gesvd_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} + m, n = size(first(As)) + m >= n && return YArocSOLVER.gesvd_batched!(As, Ss, Us, Vᴴs; kwargs...) + return MatrixAlgebraKit.batched_svd_via_adjoint!(gesvd_batched!, ROCSOLVER(), As, Ss, Us, Vᴴs; kwargs...) end +function gesvd_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} + m, n, _ = size(As) + m >= n && return YArocSOLVER.gesvd_strided_batched!(As, Ss, Us, Vᴴs; kwargs...) + return MatrixAlgebraKit.batched_svd_via_adjoint!(gesvd_batched!, ROCSOLVER(), As, Ss, Us, Vᴴs; kwargs...) +end + +gesdd_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YArocSOLVER.gesdd_batched!(As, Ss, Us, Vᴴs; kwargs...) +gesdd_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YArocSOLVER.gesdd_strided_batched!(As, Ss, Us, Vᴴs; kwargs...) + +gesvdj_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YArocSOLVER.gesvdj_batched!(As, Ss, Us, Vᴴs; kwargs...) +gesvdj_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YArocSOLVER.gesvdj_strided_batched!(As, Ss, Us, Vᴴs; kwargs...) + +gesvdx_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YArocSOLVER.gesvdx_batched!(As, Ss, Us, Vᴴs; kwargs...) +gesvdx_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YArocSOLVER.gesvdx_strided_batched!(As, Ss, Us, Vᴴs; kwargs...) + +gesdd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) = + YArocSOLVER.gesdd!(A, S, U, Vᴴ; kwargs...) heevj!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = YArocSOLVER.heevj!(A, Dd, V; kwargs...) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 6c5bd3469..cbe7730db 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -97,6 +97,425 @@ for (fname, elty, relty) in end end +# Wrappers for batched SVD via QR Iteration +for (fname, elty, relty) in + ( + (:rocsolver_sgesvd_batched, :Float32, :Float32), + (:rocsolver_dgesvd_batched, :Float64, :Float64), + (:rocsolver_cgesvd_batched, :ComplexF32, :Float32), + (:rocsolver_zgesvd_batched, :ComplexF64, :Float64), + ) + @eval begin + function gesvd_batched!( + A::AbstractVector{<:StridedROCMatrix{$elty}}, + S::StridedROCMatrix{$relty} = similar(first(A), $relty, (min(size(first(A))...), length(A))), + U::StridedROCArray{$elty, 3} = similar(first(A), $elty, size(first(A), 1), min(size(first(A))...), length(A)), + Vᴴ::StridedROCArray{$elty, 3} = similar(first(A), $elty, min(size(first(A))...), size(first(A), 2), length(A)); + check::Bool = CHECK_LIBRARY_CALLS[], + ) + for A_ in A + chkstride1(A_, U, Vᴴ, S) + end + m, n = size(first(A)) + (m < n) && throw(ArgumentError("rocSOLVER's gesvd_batched requires m ≥ n")) + minmn = min(m, n) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A and U")) + if size(U, 2) == minmn + if U === A # seems impossible? + jobu = rocSOLVER.rocblas_svect_overwrite + else + jobu = rocSOLVER.rocblas_svect_singular + end + elseif size(U, 2) == m + jobu = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid column size of U")) + end + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A and Vᴴ")) + if size(Vᴴ, 1) == minmn + if Vᴴ === A # seems impossible? + jobvt = rocSOLVER.rocblas_svect_overwrite + else + jobvt = rocSOLVER.rocblas_svect_singular + end + elseif size(Vᴴ, 1) == n + jobvt = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid row size of Vᴴ")) + end + end + size(S) == (minmn, length(A)) || + throw(DimensionMismatch("size mismatch between A and S")) + + lda = max(1, stride(first(A), 2)) + ldu = max(1, stride(U, 2)) + strideU = max(1, stride(U, 3)) + ldv = max(1, stride(Vᴴ, 2)) + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) + + strideE = minmn - 1 + E = ROCArray{$relty}(undef, length(A) * strideE) + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, length(A)) + pA = ROCVector(map(pointer, A)) + rocSOLVER.$fname( + dh, jobu, jobvt, m, n, pA, lda, + S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, + E, strideE, convert(rocSOLVER.rocblas_workmode, 'I'), + dev_info, length(A) + ) + AMDGPU.unsafe_free!(pA) + AMDGPU.unsafe_free!(E) + if check + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) + end + + return (S, U, Vᴴ) + end + end +end + +for (fname, elty, relty) in + ( + (:rocsolver_sgesvd_strided_batched, :Float32, :Float32), + (:rocsolver_dgesvd_strided_batched, :Float64, :Float64), + (:rocsolver_cgesvd_strided_batched, :ComplexF32, :Float32), + (:rocsolver_zgesvd_strided_batched, :ComplexF64, :Float64), + ) + @eval begin + function gesvd_strided_batched!( + A::StridedROCArray{$elty, 3}, + S::StridedROCMatrix{$relty} = similar(A, $relty, min(size(A, 1, size(A, 2))), size(A, 3)), + U::StridedROCArray{$elty, 3} = similar(A, $elty, size(A, 1), min(size(A, 1), size(A, 2)), size(A, 3)), + Vᴴ::StridedROCArray{$elty, 3} = similar(A, $elty, min(size(A, 1), size(A, 2)), size(A, 2), size(A, 3)); + check::Bool = CHECK_LIBRARY_CALLS[], + ) + chkstride1(A, U, Vᴴ, S) + m, n, batch_size = size(A) + (m < n) && throw(ArgumentError("rocSOLVER's gesvd_strided_batched requires m ≥ n")) + minmn = min(m, n) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + batch_size != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A and U")) + if size(U, 2) == minmn + if U === A + jobu = rocSOLVER.rocblas_svect_overwrite + else + jobu = rocSOLVER.rocblas_svect_singular + end + elseif size(U, 2) == m + jobu = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid column size of U")) + end + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + batch_size != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A and Vᴴ")) + if size(Vᴴ, 1) == minmn + if Vᴴ === A + jobvt = rocSOLVER.rocblas_svect_overwrite + else + jobvt = rocSOLVER.rocblas_svect_singular + end + elseif size(Vᴴ, 1) == n + jobvt = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid row size of Vᴴ")) + end + end + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) + + lda = max(1, stride(A, 2)) + strideA = max(1, stride(A, 3)) + ldu = max(1, stride(U, 2)) + strideU = max(1, stride(U, 3)) + ldv = max(1, stride(Vᴴ, 2)) + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) + + strideE = minmn - 1 + E = ROCArray{$relty}(undef, batch_size * strideE) + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, batch_size) + rocSOLVER.$fname( + dh, jobu, jobvt, m, n, A, lda, strideA, + S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, + E, strideE, convert(rocSOLVER.rocblas_workmode, 'I'), + dev_info, batch_size + ) + AMDGPU.unsafe_free!(E) + + if check + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) + end + + return (S, U, Vᴴ) + end + end +end + +# Wrapper for SVD via DivideAndConquer +for (fname, elty, relty) in + ( + (:rocsolver_sgesdd, :Float32, :Float32), + (:rocsolver_dgesdd, :Float64, :Float64), + (:rocsolver_cgesdd, :ComplexF32, :Float32), + (:rocsolver_zgesdd, :ComplexF64, :Float64), + ) + @eval begin + function gesdd!( + 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)); + check::Bool = CHECK_LIBRARY_CALLS[], + ) + chkstride1(A, U, Vᴴ, S) + m, n = size(A) + minmn = min(m, n) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A and U")) + if size(U, 2) == minmn + if U === A + jobu = rocSOLVER.rocblas_svect_overwrite + else + jobu = rocSOLVER.rocblas_svect_singular + end + elseif size(U, 2) == m + jobu = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid column size of U")) + end + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A and Vᴴ")) + if size(Vᴴ, 1) == minmn + if Vᴴ === A + jobvt = rocSOLVER.rocblas_svect_overwrite + else + jobvt = rocSOLVER.rocblas_svect_singular + end + elseif size(Vᴴ, 1) == n + jobvt = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid row size of Vᴴ")) + end + end + length(S) == minmn || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(A, 2)) + ldu = max(1, stride(U, 2)) + ldv = max(1, stride(Vᴴ, 2)) + + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, 1) + rocSOLVER.$fname( + dh, jobu, jobvt, m, n, + A, lda, S, U, ldu, Vᴴ, ldv, + dev_info + ) + + if check + info = @allowscalar dev_info[1] + rocSOLVER.chkargsok(BlasInt(info)) + end + return (S, U, Vᴴ) + end + end +end + +# Wrapper for batched SVD via DivideAndConquer +for (fname, elty, relty) in + ( + (:rocsolver_sgesdd_batched, :Float32, :Float32), + (:rocsolver_dgesdd_batched, :Float64, :Float64), + (:rocsolver_cgesdd_batched, :ComplexF32, :Float32), + (:rocsolver_zgesdd_batched, :ComplexF64, :Float64), + ) + @eval begin + function gesdd_batched!( + A::AbstractVector{<:StridedROCMatrix{$elty}}, + S::StridedROCMatrix{$relty} = similar(first(A), $relty, (min(size(first(A))...), length(A))), + U::StridedROCArray{$elty, 3} = similar(first(A), $elty, size(first(A), 1), min(size(first(A))...), length(A)), + Vᴴ::StridedROCArray{$elty, 3} = similar(first(A), $elty, min(size(first(A))...), size(first(A), 2), length(A)); + check::Bool = CHECK_LIBRARY_CALLS[], + ) + for A_ in A + chkstride1(A_, U, Vᴴ, S) + end + m, n = size(first(A)) + minmn = min(m, n) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A and U")) + if size(U, 2) == minmn + if U === A # seems impossible? + jobu = rocSOLVER.rocblas_svect_overwrite + else + jobu = rocSOLVER.rocblas_svect_singular + end + elseif size(U, 2) == m + jobu = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid column size of U")) + end + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A and Vᴴ")) + if size(Vᴴ, 1) == minmn + if Vᴴ === A # seems impossible? + jobvt = rocSOLVER.rocblas_svect_overwrite + else + jobvt = rocSOLVER.rocblas_svect_singular + end + elseif size(Vᴴ, 1) == n + jobvt = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid row size of Vᴴ")) + end + end + size(S) == (minmn, length(A)) || + throw(DimensionMismatch("size mismatch between A and S")) + + lda = max(1, stride(first(A), 2)) + ldu = max(1, stride(U, 2)) + strideU = max(1, stride(U, 3)) + ldv = max(1, stride(Vᴴ, 2)) + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) + + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, length(A)) + pA = ROCVector(map(pointer, A)) + rocSOLVER.$fname( + dh, jobu, jobvt, m, n, pA, lda, + S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, + dev_info, length(A) + ) + AMDGPU.unsafe_free!(pA) + if check + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) + end + + return (S, U, Vᴴ) + end + end +end + +for (fname, elty, relty) in + ( + (:rocsolver_sgesdd_strided_batched, :Float32, :Float32), + (:rocsolver_dgesdd_strided_batched, :Float64, :Float64), + (:rocsolver_cgesdd_strided_batched, :ComplexF32, :Float32), + (:rocsolver_zgesdd_strided_batched, :ComplexF64, :Float64), + ) + @eval begin + function gesdd_strided_batched!( + A::StridedROCArray{$elty, 3}, + S::StridedROCMatrix{$relty} = similar(first(A), $relty, (min(size(first(A))...), length(A))), + U::StridedROCArray{$elty, 3} = similar(first(A), $elty, size(first(A), 1), min(size(first(A))...), length(A)), + Vᴴ::StridedROCArray{$elty, 3} = similar(first(A), $elty, min(size(first(A))...), size(first(A), 2), length(A)); + check::Bool = CHECK_LIBRARY_CALLS[], + ) + chkstride1(A, U, Vᴴ, S) + m, n, batch_size = size(A) + minmn = min(m, n) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + batch_size != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A and U")) + if size(U, 2) == minmn + if U === A + jobu = rocSOLVER.rocblas_svect_overwrite + else + jobu = rocSOLVER.rocblas_svect_singular + end + elseif size(U, 2) == m + jobu = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid column size of U")) + end + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + batch_size != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A and Vᴴ")) + if size(Vᴴ, 1) == minmn + if Vᴴ === A + jobvt = rocSOLVER.rocblas_svect_overwrite + else + jobvt = rocSOLVER.rocblas_svect_singular + end + elseif size(Vᴴ, 1) == n + jobvt = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid row size of Vᴴ")) + end + end + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) + + lda = max(1, stride(A, 2)) + strideA = max(1, stride(A, 3)) + ldu = max(1, stride(U, 2)) + strideU = max(1, stride(U, 3)) + ldv = max(1, stride(Vᴴ, 2)) + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) + + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, batch_size) + rocSOLVER.$fname( + dh, jobu, jobvt, m, n, A, lda, strideA, + S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, + dev_info, batch_size + ) + if check + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) + end + return (S, U, Vᴴ) + end + end +end + # Wrapper for SVD via Jacobi for (fname, elty, relty) in ( @@ -182,6 +601,185 @@ for (fname, elty, relty) in end end +# Wrapper for batched SVD via Jacobi +for (fname, elty, relty) in + ( + (:rocsolver_sgesvdj_batched, :Float32, :Float32), + (:rocsolver_dgesvdj_batched, :Float64, :Float64), + (:rocsolver_cgesvdj_batched, :ComplexF32, :Float32), + (:rocsolver_zgesvdj_batched, :ComplexF64, :Float64), + ) + @eval begin + function gesvdj_batched!( + A::AbstractVector{<:StridedROCMatrix{$elty}}, + S::StridedROCMatrix{$relty} = similar(first(A), $relty, (min(size(first(A))...), length(A))), + U::StridedROCArray{$elty, 3} = similar(first(A), $elty, size(first(A), 1), min(size(first(A))...), length(A)), + Vᴴ::StridedROCArray{$elty, 3} = similar(first(A), $elty, min(size(first(A))...), size(first(A), 2), length(A)), + tol::$relty = eps($relty), + max_sweeps::Int = 100, + check::Bool = CHECK_LIBRARY_CALLS[], + ) + for A_ in A + chkstride1(A_, U, Vᴴ, S) + end + m, n = size(first(A)) + minmn = min(m, n) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A and U")) + if size(U, 2) == minmn + if U === A + throw(ArgumentError("overwrite mode is not supported for gesvdj")) + else + jobu = rocSOLVER.rocblas_svect_singular + end + elseif size(U, 2) == m + jobu = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid column size of U")) + end + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A and Vᴴ")) + if size(Vᴴ, 1) == minmn + if Vᴴ === A + throw(ArgumentError("overwrite mode is not supported for gesvdj")) + else + jobvt = rocSOLVER.rocblas_svect_singular + end + elseif size(Vᴴ, 1) == n + jobvt = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid row size of Vᴴ")) + end + end + size(S) == (minmn, length(A)) || + throw(DimensionMismatch("size mismatch between A and S")) + + lda = max(1, stride(first(A), 2)) + ldu = max(1, stride(U, 2)) + strideU = max(1, stride(U, 3)) + ldv = max(1, stride(Vᴴ, 2)) + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) + dev_info = ROCVector{Cint}(undef, length(A)) + dev_residual = ROCVector{$relty}(undef, length(A)) + dev_n_sweeps = ROCVector{Cint}(undef, length(A)) + + dh = rocBLAS.handle() + pA = ROCVector(map(pointer, A)) + rocSOLVER.$fname( + dh, jobu, jobvt, m, n, pA, lda, tol, + dev_residual, max_sweeps, dev_n_sweeps, + S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, + dev_info, length(A) + ) + if check + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) + end + AMDGPU.unsafe_free!(pA) + AMDGPU.unsafe_free!(dev_residual) + AMDGPU.unsafe_free!(dev_n_sweeps) + return (S, U, Vᴴ) + end + end +end + +for (fname, elty, relty) in + ( + (:rocsolver_sgesvdj_strided_batched, :Float32, :Float32), + (:rocsolver_dgesvdj_strided_batched, :Float64, :Float64), + (:rocsolver_cgesvdj_strided_batched, :ComplexF32, :Float32), + (:rocsolver_zgesvdj_strided_batched, :ComplexF64, :Float64), + ) + @eval begin + function gesvdj_strided_batched!( + A::StridedROCArray{$elty, 3}, + S::StridedROCMatrix{$relty} = similar(A, $relty, min(size(A, 1, size(A, 2))), size(A, 3)), + U::StridedROCArray{$elty, 3} = similar(A, $elty, size(A, 1), min(size(A, 1), size(A, 2)), size(A, 3)), + Vᴴ::StridedROCArray{$elty, 3} = similar(A, $elty, min(size(A, 1), size(A, 2)), size(A, 2), size(A, 3)); + tol::$relty = eps($relty), + max_sweeps::Int = 100, + check::Bool = CHECK_LIBRARY_CALLS[], + ) + chkstride1(A, U, Vᴴ, S) + m, n, batch_size = size(A) + minmn = min(m, n) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + batch_size != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A and U")) + if size(U, 2) == minmn + if U === A + throw(ArgumentError("overwrite mode is not supported for gesvdj")) + else + jobu = rocSOLVER.rocblas_svect_singular + end + elseif size(U, 2) == m + jobu = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid column size of U")) + end + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + batch_size != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A and Vᴴ")) + if size(Vᴴ, 1) == minmn + if Vᴴ === A + throw(ArgumentError("overwrite mode is not supported for gesvdj")) + else + jobvt = rocSOLVER.rocblas_svect_singular + end + elseif size(Vᴴ, 1) == n + jobvt = rocSOLVER.rocblas_svect_all + else + throw(DimensionMismatch("invalid row size of Vᴴ")) + end + end + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) + + lda = max(1, stride(A, 2)) + strideA = max(1, stride(A, 3)) + ldu = max(1, stride(U, 2)) + strideU = max(1, stride(U, 3)) + ldv = max(1, stride(Vᴴ, 2)) + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) + dev_info = ROCVector{Cint}(undef, batch_size) + dev_residual = ROCVector{$relty}(undef, batch_size) + dev_n_sweeps = ROCVector{Cint}(undef, batch_size) + + dh = rocBLAS.handle() + rocSOLVER.$fname( + dh, jobu, jobvt, m, n, A, lda, strideA, tol, + dev_residual, max_sweeps, dev_n_sweeps, + S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, + dev_info, batch_size + ) + + if check + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) + end + AMDGPU.unsafe_free!(dev_residual) + AMDGPU.unsafe_free!(dev_n_sweeps) + return (S, U, Vᴴ) + end + 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`. @@ -220,6 +818,26 @@ function _gesvdx_jobs(U, Vᴴ, m::Integer, n::Integer, maxnsv::Integer) return jobu, jobvt end +""" + _gesvdx_zero_unconverged!(S, nsv) + +Zero the entries of `S` that `gesvdx` did not write. +""" +function _gesvdx_zero_unconverged!(S::StridedROCVector, nsv::ROCVector{Cint}) + nv = @allowscalar Int(nsv[1]) + nv < length(S) && fill!(view(S, (nv + 1):length(S)), zero(eltype(S))) + return S +end +function _gesvdx_zero_unconverged!(S::StridedROCMatrix, nsv::ROCVector{Cint}) + minmn = size(S, 1) + nvs = Array(nsv) + all(==(minmn), nvs) && return S # nothing omitted, skip the per-batch fills + for (b, nv) in pairs(nvs) + nv < minmn && fill!(view(S, (nv + 1):minmn, b), zero(eltype(S))) + end + return S +end + # Wrapper for SVD via Bisection for (fname, elty, relty) in ( @@ -264,8 +882,131 @@ for (fname, elty, relty) in 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))) + _gesvdx_zero_unconverged!(S, nsv) + + AMDGPU.unsafe_free!(nsv) + AMDGPU.unsafe_free!(ifail) + AMDGPU.unsafe_free!(dev_info) + return (S, U, Vᴴ) + end + end +end + +# Wrappers for batched SVD via Bisection +for (fname, elty, relty) in + ( + (:rocsolver_sgesvdx_batched, :Float32, :Float32), + (:rocsolver_dgesvdx_batched, :Float64, :Float64), + (:rocsolver_cgesvdx_batched, :ComplexF32, :Float32), + (:rocsolver_zgesvdx_batched, :ComplexF64, :Float64), + ) + @eval begin + function gesvdx_batched!( + A::AbstractVector{<:StridedROCMatrix{$elty}}, + S::StridedROCMatrix{$relty} = similar(first(A), $relty, (min(size(first(A))...), length(A))), + U::StridedROCArray{$elty, 3} = similar(first(A), $elty, size(first(A), 1), min(size(first(A))...), length(A)), + Vᴴ::StridedROCArray{$elty, 3} = similar(first(A), $elty, min(size(first(A))...), size(first(A), 2), length(A)); + check::Bool = CHECK_LIBRARY_CALLS[], + kwargs... + ) + for A_ in A + chkstride1(A_, U, Vᴴ, S) + end + m, n = size(first(A)) + minmn = min(m, n) + batch_size = length(A) + 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) + jobu != rocSOLVER.rocblas_svect_none && batch_size != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + jobu != rocSOLVER.rocblas_svect_none && batch_size != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) + + lda = max(1, stride(first(A), 2)) + ldu = max(1, stride(U, 2)) + strideU = max(1, stride(U, 3)) + ldv = max(1, stride(Vᴴ, 2)) + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) + strideF = minmn + + dh = rocBLAS.handle() + nsv = ROCVector{Cint}(undef, batch_size) + ifail = ROCVector{Cint}(undef, strideF * batch_size) + dev_info = ROCVector{Cint}(undef, batch_size) + pA = ROCVector(map(pointer, A)) + rocSOLVER.$fname( + dh, jobu, jobvt, srange, m, n, pA, lda, + vl, vu, il, iu, nsv, + S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, + ifail, strideF, dev_info, batch_size + ) + AMDGPU.unsafe_free!(pA) + + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end + _gesvdx_zero_unconverged!(S, nsv) + + AMDGPU.unsafe_free!(nsv) + AMDGPU.unsafe_free!(ifail) + AMDGPU.unsafe_free!(dev_info) + return (S, U, Vᴴ) + end + end +end + +for (fname, elty, relty) in + ( + (:rocsolver_sgesvdx_strided_batched, :Float32, :Float32), + (:rocsolver_dgesvdx_strided_batched, :Float64, :Float64), + (:rocsolver_cgesvdx_strided_batched, :ComplexF32, :Float32), + (:rocsolver_zgesvdx_strided_batched, :ComplexF64, :Float64), + ) + @eval begin + function gesvdx_strided_batched!( + A::StridedROCArray{$elty, 3}, + S::StridedROCMatrix{$relty} = similar(A, $relty, (min(size(A, 1), size(A, 2)), size(A, 3))), + U::StridedROCArray{$elty, 3} = similar(A, $elty, size(A, 1), min(size(A, 1), size(A, 2)), size(A, 3)), + Vᴴ::StridedROCArray{$elty, 3} = similar(A, $elty, min(size(A, 1), size(A, 2)), size(A, 2), size(A, 3)); + check::Bool = CHECK_LIBRARY_CALLS[], + kwargs... + ) + chkstride1(A, U, Vᴴ, S) + m, n, batch_size = 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) + jobu != rocSOLVER.rocblas_svect_none && batch_size != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + jobu != rocSOLVER.rocblas_svect_none && batch_size != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) + + lda = max(1, stride(A, 2)) + strideA = max(1, stride(A, 3)) + ldu = max(1, stride(U, 2)) + strideU = max(1, stride(U, 3)) + ldv = max(1, stride(Vᴴ, 2)) + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) + strideF = minmn + + dh = rocBLAS.handle() + nsv = ROCVector{Cint}(undef, batch_size) + ifail = ROCVector{Cint}(undef, strideF * batch_size) + dev_info = ROCVector{Cint}(undef, batch_size) + rocSOLVER.$fname( + dh, jobu, jobvt, srange, m, n, A, lda, strideA, + vl, vu, il, iu, nsv, + S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, + ifail, strideF, dev_info, batch_size + ) + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end + _gesvdx_zero_unconverged!(S, nsv) AMDGPU.unsafe_free!(nsv) AMDGPU.unsafe_free!(ifail) diff --git a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl index 704e92e7a..8f24e3470 100644 --- a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl +++ b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl @@ -7,6 +7,7 @@ using MatrixAlgebraKit: diagview, sign_safe using MatrixAlgebraKit: CUSOLVER, LQViaTransposedQR, TruncationByValue, AbstractAlgorithm using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eig_algorithm, default_eigh_algorithm import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdp!, gesvdr!, gesvdj! +import MatrixAlgebraKit: gesvdj_batched! import MatrixAlgebraKit: heevj!, heevd!, geev! import MatrixAlgebraKit: _gpu_Xgesvdr!, _sylvester, svd_rank, svd_pullback!, eigh_pullback!, eig_pullback!, svd_pushforward! using CUDA, CUDA.cuBLAS @@ -16,11 +17,14 @@ using LinearAlgebra: BlasFloat include("yacusolver.jl") -MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedCuVecOrMat{<:BlasFloat}} = CUSOLVER() +MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedCuArray{<:BlasFloat}} = CUSOLVER() function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedCuVecOrMat{<:BlasFloat}} return QRIteration(; kwargs...) end +function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedCuArray{<:BlasFloat, 3}} + return Jacobi(; kwargs...) +end function MatrixAlgebraKit.default_eig_algorithm(::Type{T}; kwargs...) where {T <: StridedCuVecOrMat{<:BlasFloat}} return QRIteration(; kwargs...) end @@ -28,7 +32,6 @@ function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T return DivideAndConquer(; kwargs...) end - for f in (:geqrf!, :ungqr!, :unmqr!) @eval $f(::CUSOLVER, args...) = YACUSOLVER.$f(args...) end @@ -37,6 +40,9 @@ MatrixAlgebraKit.prefers_ungqr(::CUSOLVER) = true MatrixAlgebraKit.supports_svd_full(::CUSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :svd_polar) +# `cusolverDnXgesvdjBatched` only accepts matrices up to 32x32 +MatrixAlgebraKit.max_batched_blocksize(::AbstractAlgorithm, ::CUSOLVER, ::Type{<:AnyCuArray}) = 32 + function gesvd!(::CUSOLVER, A::StridedCuMatrix, S::StridedCuVector, U::StridedCuMatrix, Vᴴ::StridedCuMatrix; kwargs...) m, n = size(A) m >= n && return YACUSOLVER.gesvd!(A, S, U, Vᴴ) @@ -49,6 +55,9 @@ function gesvdj!(::CUSOLVER, A::StridedCuMatrix, S::StridedCuVector, U::StridedC return MatrixAlgebraKit.svd_via_adjoint!(gesvdj!, CUSOLVER(), A, S, U, Vᴴ; kwargs...) end +gesvdj_batched!(::CUSOLVER, As::StridedCuArray{T, 3}, Ss::StridedCuMatrix, Us::StridedCuArray{T, 3}, Vᴴs::StridedCuArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YACUSOLVER.gesvdj_batched!(As, Ss, Us, Vᴴs; kwargs...) + gesvdp!(::CUSOLVER, A::StridedCuMatrix, S::StridedCuVector, U::StridedCuMatrix, Vᴴ::StridedCuMatrix; kwargs...) = YACUSOLVER.gesvdp!(A, S, U, Vᴴ; kwargs...) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index bd19257aa..934606c49 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -276,6 +276,96 @@ for (bname, fname, elty, relty) in end end +# Wrapper for batched SVD via Jacobi +for (bname, fname, elty, relty) in + ( + (:cusolverDnSgesvdjBatched_bufferSize, :cusolverDnSgesvdjBatched, :Float32, :Float32), + (:cusolverDnDgesvdjBatched_bufferSize, :cusolverDnDgesvdjBatched, :Float64, :Float64), + (:cusolverDnCgesvdjBatched_bufferSize, :cusolverDnCgesvdjBatched, :ComplexF32, :Float32), + (:cusolverDnZgesvdjBatched_bufferSize, :cusolverDnZgesvdjBatched, :ComplexF64, :Float64), + ) + @eval begin + #! format: off + function gesvdj_batched!( + A::StridedCuArray{$elty, 3}, + S::StridedCuMatrix{$relty} = similar(A, $relty, min(size(A, 1), size(A, 2)), size(A, 3)), + U::StridedCuArray{$elty, 3} = similar(A, $elty, size(A, 1), min(size(A, 1), size(A, 2)), size(A, 3)), + Vᴴ::StridedCuArray{$elty, 3} = similar(A, $elty, min(size(A, 1), size(A, 2)), size(A, 2), size(A, 3)); + tol::$relty = eps($relty), + max_sweeps::Int = 100, + check::Bool = CHECK_LIBRARY_CALLS[], + kwargs... + ) + #! format: on + chkstride1(A, U, Vᴴ, S) + m, n, batch_size = size(A) + minmn = min(m, n) + if length(U) == 0 && length(Vᴴ) == 0 + jobz = 'N' + else + batch_size != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + batch_size != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) + jobz = 'V' + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A and U")) + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A and Vᴴ")) + if !(size(U, 2) == size(Vᴴ, 1) == minmn) && !(size(U, 2) == m && size(Vᴴ, 1) == n) + throw(DimensionMismatch("invalid column size of U or row size of Vᴴ")) + end + end + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) + + # these MUST be "full" sized + # TODO: check if U and Vᴴ already have the correct size to avoid + # some allocations + Ṽ = similar(Vᴴ, (n, n, batch_size)) + Ũ = similar(U, (m, m, batch_size)) + lda = max(1, stride(A, 2)) + ldu = max(1, stride(Ũ, 2)) + ldv = max(1, stride(Ṽ, 2)) + + params = Ref{cuSOLVER.gesvdjInfo_t}(C_NULL) + cuSOLVER.cusolverDnCreateGesvdjInfo(params) + cuSOLVER.cusolverDnXgesvdjSetTolerance(params[], tol) + cuSOLVER.cusolverDnXgesvdjSetMaxSweeps(params[], max_sweeps) + dh = cuSOLVER.dense_handle() + resize!(dh.info, batch_size) + + function bufferSize() + out = Ref{Cint}(0) + cuSOLVER.$bname( + dh, jobz, m, n, A, lda, S, Ũ, ldu, Ṽ, ldv, + out, params[], batch_size + ) + return out[] * sizeof($elty) + end + + cuSOLVER.with_workspace(dh.workspace_gpu, bufferSize) do buffer + return cuSOLVER.$fname( + dh, jobz, m, n, A, lda, S, Ũ, ldu, Ṽ, ldv, + buffer, sizeof(buffer) ÷ sizeof($elty), dh.info, + params[], batch_size + ) + end + + if check + info = collect(dh.info) + foreach(cuSOLVER.chkargsok ∘ BlasInt, info) + end + + cuSOLVER.cusolverDnDestroyGesvdjInfo(params[]) + + if jobz == 'V' + copyto!(U, view(Ũ, :, 1:size(U, 2), :)) + Vᴴ .= conj.(PermutedDimsArray(view(Ṽ, :, 1:size(Vᴴ, 1), :), (2, 1, 3))) + end + return S, U, Vᴴ + end + end +end + # Wrapper for randomized SVD function gesvdr!( A::StridedCuMatrix{T}, diff --git a/src/MatrixAlgebraKit.jl b/src/MatrixAlgebraKit.jl index 3d9662137..78d36bfe5 100644 --- a/src/MatrixAlgebraKit.jl +++ b/src/MatrixAlgebraKit.jl @@ -18,6 +18,8 @@ export qr_compact, qr_full, qr_null, lq_compact, lq_full, lq_null export qr_compact!, qr_full!, qr_null!, lq_compact!, lq_full!, lq_null! export svd_compact, svd_full, svd_vals, svd_trunc, svd_trunc_no_error export svd_compact!, svd_full!, svd_vals!, svd_trunc!, svd_trunc_no_error! +export batched_svd_compact, batched_svd_full, batched_svd_vals +export batched_svd_compact!, batched_svd_full!, batched_svd_vals! export eigh_full, eigh_vals, eigh_trunc, eigh_trunc_no_error export eigh_full!, eigh_vals!, eigh_trunc!, eigh_trunc_no_error! export eig_full, eig_vals, eig_trunc, eig_trunc_no_error @@ -99,6 +101,7 @@ include("common/utility.jl") include("yalapack.jl") include("algorithms.jl") +include("batches.jl") include("interface/projections.jl") include("interface/decompositions.jl") @@ -108,6 +111,7 @@ include("interface/matrixfunctions.jl") include("interface/qr.jl") include("interface/lq.jl") include("interface/svd.jl") +include("interface/batched_svd.jl") include("interface/eig.jl") include("interface/eigh.jl") include("interface/gen_eig.jl") @@ -121,6 +125,7 @@ include("implementations/truncation.jl") include("implementations/qr.jl") include("implementations/lq.jl") include("implementations/svd.jl") +include("implementations/batched_svd.jl") include("implementations/eig.jl") include("implementations/eigh.jl") include("implementations/gen_eig.jl") diff --git a/src/batches.jl b/src/batches.jl new file mode 100644 index 000000000..b0be7e96a --- /dev/null +++ b/src/batches.jl @@ -0,0 +1,108 @@ +_isempty_batch(A::AbstractArray{<:Any, 3}) = isempty(A) +_isempty_batch(A::AbstractVector{<:AbstractMatrix}) = all(isempty, A) + +# Adjoint of every matrix in a batch, i.e. `dst[:, :, i] = adjoint(src[:, :, i])`. +function batched_adjoint!(dst::AbstractArray{<:Any, 3}, src::AbstractArray{<:Any, 3}) + isempty(dst) && return dst + permutedims!(dst, src, (2, 1, 3)) + eltype(dst) <: Real || (dst .= conj.(dst)) + return dst +end +function batched_adjoint(A::AbstractArray{<:Any, 3}) + return batched_adjoint!(similar(A, (size(A, 2), size(A, 1), size(A, 3))), A) +end +batched_adjoint(A::AbstractVector{<:AbstractMatrix}) = map(a -> adjoint!(similar(a'), a), A) + +# Ragged batches +# -------------- + +""" + supports_ragged_batch(f!, alg::AbstractAlgorithm, driver::Driver, T::Type) -> Bool + +Whether the algorithm `alg` running on `driver` accepts a *ragged* batch of matrices +of type `T` which do not have uniform size for function `f!`. `true` by default. +""" +supports_ragged_batch(f!, alg::AbstractAlgorithm, driver::Driver, ::Type) = true +supports_ragged_batch(f!, alg::AbstractAlgorithm, ::DefaultDriver, ::Type{TA}) where {TA} = + supports_ragged_batch(f!, alg, default_driver(alg, TA), TA) + +""" + max_batched_blocksize(alg, driver::Driver, T::Type) -> Int + +Largest matrix dimension that the `driver` for the batched version of `alg` accepts for arrays +of type `T`. Larger matrices in a ragged batch are decomposed one at a time instead. +Unlimited by default. +""" +max_batched_blocksize(alg::AbstractAlgorithm, driver::Driver, ::Type) = typemax(Int) +max_batched_blocksize(alg::AbstractAlgorithm, ::DefaultDriver, ::Type{TA}) where {TA} = + max_batched_blocksize(alg, default_driver(alg, TA), TA) + +""" + supports_pointer_batch(alg, driver::Driver, T::Type) -> Bool + +Whether the low-level batched `driver` for `alg` accepts a batch of matrices of type `T` +as an `AbstractVector` of separately allocated matrices. Such a group of matrices is handed +to the driver as a vector of pointers, instead of being copied into one contiguous 3D array. +`false` by default. +""" +supports_pointer_batch(::AbstractAlgorithm, driver::Driver, ::Type) = false +supports_pointer_batch(alg::AbstractAlgorithm, ::DefaultDriver, ::Type{TA}) where {TA} = + supports_pointer_batch(alg, default_driver(alg, TA), TA) + +# Split a ragged batch into batches the driver can handle: matrices of equal size are +# batched together, and, if `pad`, whatever is left over is zero-padded into one more batch. +# Returns the batches as `(indices, (m, n))` pairs, and the indices of the matrices that have +# to be decomposed one at a time. +# TODO: should everything be padded into ONE batch? +function _ragged_batches(A::AbstractVector{<:AbstractMatrix}, alg::AbstractAlgorithm; pad::Bool = true) + batches = Tuple{Vector{Int}, Tuple{Int, Int}}[] + rest = Int[] + isempty(A) && return batches, rest + driver = get(alg.kwargs, :driver, DefaultDriver()) + batch_size_limit = max_batched_blocksize(alg, driver, typeof(first(A))) + needs_tall = requires_tall(alg) + groups = Dict{Tuple{Int, Int}, Vector{Int}}() + for i in eachindex(A) + push!(get!(Vector{Int}, groups, size(A[i])), i) + end + for ((m, n), inds) in groups + if length(inds) >= BATCHED_SVD_THRESHOLD && max(m, n) <= batch_size_limit && (!needs_tall || m >= n) + push!(batches, (inds, (m, n))) + else + append!(rest, inds) + end + end + pad || return batches, rest + # Zero padding leaves the leading `min(m, n)` singular values and vectors of every input + # untouched. Pad to a square only when the algorithm requires `m ≥ n` + # (currently only `QRIteration`). + m = maximum(i -> size(A[i], 1), rest; init = 0) + n = maximum(i -> size(A[i], 2), rest; init = 0) + padded = needs_tall ? (max(m, n), max(m, n)) : (m, n) + if length(rest) >= BATCHED_SVD_THRESHOLD && maximum(padded) <= batch_size_limit + push!(batches, (rest, padded)) + rest = Int[] + end + return batches, rest +end + +# Outputs for a batch that `_ragged_pack` produced, which is either a contiguous `(m, n, b)` +# array or, for a pointer-batch driver, a view of `b` equally sized matrices. Either way the +# outputs are packed into the 3D arrays the batched drivers write into. +_packed_output(f!, A::AbstractArray{<:Any, 3}, alg::AbstractAlgorithm) = initialize_output(f!, A, alg) + +# Gather `A[inds]` into a single `(m, n, length(inds))` batch, zero-padding where needed. +function _ragged_pack(A::AbstractVector{<:AbstractMatrix}, inds, m::Int, n::Int, alg::AbstractAlgorithm) + driver = get(alg.kwargs, :driver, DefaultDriver()) + uniform = all(i -> size(A[i]) == (m, n), inds) + uniform && supports_pointer_batch(alg, driver, typeof(A[first(inds)])) && return view(A, inds) + # `stack` can't zero-pad + # On the GPU it falls back to scalar indexing for matrices that are views + uniform && A isa AbstractVector{<:Array} && return stack(view(A, inds)) + Ab = similar(A[first(inds)], (m, n, length(inds))) + uniform || zero!(Ab) # only need to zero if matrices are ragged + for (j, i) in enumerate(inds) + copyto!(view(Ab, axes(A[i])..., j), A[i]) + end + return Ab +end diff --git a/src/common/gauge.jl b/src/common/gauge.jl index 016b60ad8..4c756b087 100644 --- a/src/common/gauge.jl +++ b/src/common/gauge.jl @@ -75,3 +75,17 @@ function gaugefix!(::Union{typeof(svd_compact!), typeof(svd_trunc!)}, U, Vᴴ) @. Vᴴ = signs_t * Vᴴ return (U, Vᴴ) end + +function gaugefix!(::typeof(batched_svd_compact!), U, Vᴴ) + for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) + gaugefix!(svd_compact!, u, vᴴ) + end + return U, Vᴴ +end + +function gaugefix!(::typeof(batched_svd_full!), U, Vᴴ) + for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) + gaugefix!(svd_full!, u, vᴴ) + end + return U, Vᴴ +end diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl new file mode 100644 index 000000000..421401556 --- /dev/null +++ b/src/implementations/batched_svd.jl @@ -0,0 +1,379 @@ +# Inputs +# ------ +copy_input(::typeof(batched_svd_full), As::AbstractVector{<:AbstractMatrix}) = map(Base.Fix1(copy_input, svd_full), As) +copy_input(::typeof(batched_svd_full), A::AbstractArray{T, 3}) where {T} = copy!(similar(A, float(T)), A) +copy_input(::typeof(batched_svd_compact), A) = copy_input(batched_svd_full, A) +copy_input(::typeof(batched_svd_vals), A) = copy_input(batched_svd_full, A) + +# ragged batches: matrices of different sizes, each with its own outputs +function check_input( + ::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, + USVᴴ::Tuple{AbstractVector{<:AbstractMatrix}, AbstractVector{<:Diagonal}, AbstractVector{<:AbstractMatrix}}, + alg::AbstractAlgorithm + ) + Us, Ss, Vᴴs = USVᴴ + length(Us) == length(Ss) == length(Vᴴs) == length(A) || + throw(DimensionMismatch("expected $(length(A)) outputs for each of U, S and Vᴴ")) + for (a, u, s, vᴴ) in zip(A, Us, Ss, Vᴴs) + check_input(svd_compact!, a, (u, s, vᴴ), alg) + end + return nothing +end +function check_input( + ::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, + USVᴴ::Tuple{AbstractVector{<:AbstractMatrix}, AbstractVector{<:AbstractMatrix}, AbstractVector{<:AbstractMatrix}}, + alg::AbstractAlgorithm + ) + Us, Ss, Vᴴs = USVᴴ + length(Us) == length(Ss) == length(Vᴴs) == length(A) || + throw(DimensionMismatch("expected $(length(A)) outputs for each of U, S and Vᴴ")) + for (a, u, s, vᴴ) in zip(A, Us, Ss, Vᴴs) + check_input(svd_full!, a, (u, s, vᴴ), alg) + end + return nothing +end +function check_input( + ::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, + S::AbstractVector{<:AbstractVector}, alg::AbstractAlgorithm + ) + length(S) == length(A) || throw(DimensionMismatch("expected $(length(A)) outputs for S")) + for (a, s) in zip(A, S) + check_input(svd_vals!, a, s, alg) + end + return nothing +end +function check_input(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) + isempty(A) && return nothing + m, n = size(first(A)) + @assert all(==((m, n)) ∘ size, A) + batch_size = length(A) + U, S, Vᴴ = USVᴴ + @assert U isa AbstractArray && S isa AbstractArray && Vᴴ isa AbstractArray + @check_size(U, (m, m, batch_size)) + @check_scalar(U, first(A)) + @check_size(S, (m, n, batch_size)) + @check_scalar(S, first(A), real) + @check_size(Vᴴ, (n, n, batch_size)) + @check_scalar(Vᴴ, first(A)) + return nothing +end +function check_input(::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) + isempty(A) && return nothing + m, n = size(first(A)) + @assert all(==((m, n)) ∘ size, A) + batch_size = length(A) + minmn = min(m, n) + U, S, Vᴴ = USVᴴ + @assert U isa AbstractArray && S isa AbstractArray && Vᴴ isa AbstractArray + @check_size(U, (m, minmn, batch_size)) + @check_scalar(U, first(A)) + @check_size(S, (minmn, batch_size)) + @check_scalar(S, first(A), real) + @check_size(Vᴴ, (minmn, n, batch_size)) + @check_scalar(Vᴴ, first(A)) + return nothing +end +function check_input(::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, S, ::AbstractAlgorithm) + isempty(A) && return nothing + m, n = size(first(A)) + @assert all(==((m, n)) ∘ size, A) + batch_size = length(A) + minmn = min(m, n) + @assert S isa AbstractMatrix + @check_size(S, (minmn, batch_size)) + @check_scalar(S, first(A), real) + return nothing +end +function check_input(::typeof(batched_svd_full!), A::AbstractArray{T, 3}, USVᴴ, ::AbstractAlgorithm) where {T} + m, n, batch_size = size(A) + U, S, Vᴴ = USVᴴ + @assert U isa AbstractArray && S isa AbstractArray && Vᴴ isa AbstractArray + @check_size(U, (m, m, batch_size)) + @check_scalar(U, A) + @check_size(S, (m, n, batch_size)) + @check_scalar(S, A, real) + @check_size(Vᴴ, (n, n, batch_size)) + @check_scalar(Vᴴ, A) + return nothing +end +function check_input(::typeof(batched_svd_compact!), A::AbstractArray{T, 3}, USVᴴ, ::AbstractAlgorithm) where {T} + m, n, batch_size = size(A) + minmn = min(m, n) + U, S, Vᴴ = USVᴴ + @assert U isa AbstractArray && S isa AbstractArray && Vᴴ isa AbstractArray + @check_size(U, (m, minmn, batch_size)) + @check_scalar(U, A) + @check_size(S, (minmn, batch_size)) + @check_scalar(S, A, real) + @check_size(Vᴴ, (minmn, n, batch_size)) + @check_scalar(Vᴴ, A) + return nothing +end +function check_input(::typeof(batched_svd_vals!), A::AbstractArray{T, 3}, S, ::AbstractAlgorithm) where {T} + m, n, batch_size = size(A) + minmn = min(m, n) + @assert S isa AbstractMatrix + @check_size(S, (minmn, batch_size)) + @check_scalar(S, A, real) + return nothing +end + +# Outputs +# ------- +# a vector of matrices, which may have different sizes, gets one output per matrix +function initialize_output(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + Us = [similar(a, (size(a, 1), size(a, 1))) for a in A] + Ss = [similar(a, real(eltype(a)), size(a)) for a in A] # TODO: Rectangular diagonal type? + Vᴴs = [similar(a, (size(a, 2), size(a, 2))) for a in A] + return (Us, Ss, Vᴴs) +end +function initialize_output(::typeof(batched_svd_full!), A::AbstractArray{T, 3}, ::AbstractAlgorithm) where {T} + m, n, batch_size = size(A) + U = similar(A, (m, m, batch_size)) + S = similar(A, real(eltype(A)), (m, n, batch_size)) + Vᴴ = similar(A, (n, n, batch_size)) + return (U, S, Vᴴ) +end +function initialize_output(::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + Us = [similar(a, (size(a, 1), minimum(size(a)))) for a in A] + Ss = [Diagonal(similar(a, real(eltype(a)), minimum(size(a)))) for a in A] + Vᴴs = [similar(a, (minimum(size(a)), size(a, 2))) for a in A] + return (Us, Ss, Vᴴs) +end +function initialize_output(::typeof(batched_svd_compact!), A::AbstractArray{T, 3}, ::AbstractAlgorithm) where {T} + m, n, batch_size = size(A) + minmn = min(m, n) + U = similar(A, (m, minmn, batch_size)) + S = similar(A, real(eltype(A)), (minmn, batch_size)) + Vᴴ = similar(A, (minmn, n, batch_size)) + return (U, S, Vᴴ) +end +function initialize_output(::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + return [similar(a, real(eltype(a)), minimum(size(a))) for a in A] +end +function initialize_output(::typeof(batched_svd_vals!), A::AbstractArray{T, 3}, ::AbstractAlgorithm) where {T} + m, n, batch_size = size(A) + return similar(A, real(eltype(A)), (min(m, n), batch_size)) +end + +for f! in (:gesdd_batched!, :gesvd_batched!, :gesvdj_batched!, :gesvdx_batched!) + @eval $f!(driver::Driver, args...) = throw(ArgumentError("$driver does not provide $($(f!))")) +end + +""" + batched_svd_via_adjoint!(f!, driver, A, S, U, Vᴴ; kwargs...) + +Compute the SVD of every matrix in the batch `A` (m × n, m < n) by computing the SVD of their +adjoints using the provided function `f!(driver, A, S, U, Vᴴ; kwargs...)`. Use this as a +building block for drivers whose batched SVD routines require m ≥ n, mirroring +[`svd_via_adjoint!`](@ref). +""" +function batched_svd_via_adjoint!(f!::F, driver::Driver, A, S, U, Vᴴ; kwargs...) where {F} + Aᴴ = batched_adjoint(A) + V = similar(Vᴴ, (size(Vᴴ, 2), size(Vᴴ, 1), size(Vᴴ, 3))) + Uᴴ = similar(U, (size(U, 2), size(U, 1), size(U, 3))) + f!(driver, Aᴴ, S, V, Uᴴ; kwargs...) + length(U) > 0 && batched_adjoint!(U, Uᴴ) + length(Vᴴ) > 0 && batched_adjoint!(Vᴴ, V) + return S, U, Vᴴ +end + +for (f, f_lapack!, Alg) in ( + (:divide_and_conquer, :gesdd_batched!, :DivideAndConquer), + (:qr_iteration, :gesvd_batched!, :QRIteration), + (:bisection, :gesvdx_batched!, :Bisection), + (:jacobi, :gesvdj_batched!, :Jacobi), + ) + svd_compact_f! = Symbol(:batched_svd_compact_, f, :!) + svd_full_f! = Symbol(:batched_svd_full_, f, :!) + svd_vals_f! = Symbol(:batched_svd_vals_, f, :!) + + # MatrixAlgebraKit wrappers + @eval begin + function batched_svd_compact!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) + check_input(batched_svd_compact!, A, USVᴴ, alg) + driver = get(alg.kwargs, :driver, default_driver(alg, eltype(A))) + supports_pointer_batch(alg, driver, eltype(A)) || + throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) + return $svd_compact_f!(A, USVᴴ...; alg.kwargs...) + end + function batched_svd_compact!(A::AbstractArray{T, 3}, USVᴴ, alg::$Alg) where {T} + check_input(batched_svd_compact!, A, USVᴴ, alg) + return $svd_compact_f!(A, USVᴴ...; alg.kwargs...) + end + function batched_svd_full!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) + check_input(batched_svd_full!, A, USVᴴ, alg) + driver = get(alg.kwargs, :driver, default_driver(alg, eltype(A))) + supports_pointer_batch(alg, driver, eltype(A)) || + throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) + return $svd_full_f!(A, USVᴴ...; alg.kwargs...) + end + function batched_svd_full!(A::AbstractArray{T, 3}, USVᴴ, alg::$Alg) where {T} + check_input(batched_svd_full!, A, USVᴴ, alg) + return $svd_full_f!(A, USVᴴ...; alg.kwargs...) + end + function batched_svd_vals!(A::AbstractVector{<:AbstractMatrix}, S, alg::$Alg) + check_input(batched_svd_vals!, A, S, alg) + driver = get(alg.kwargs, :driver, default_driver(alg, eltype(A))) + supports_pointer_batch(alg, driver, eltype(A)) || + throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) + return $svd_vals_f!(A, S; alg.kwargs...) + end + function batched_svd_vals!(A::AbstractArray{T, 3}, S, alg::$Alg) where {T} + check_input(batched_svd_vals!, A, S, alg) + return $svd_vals_f!(A, S; alg.kwargs...) + end + + # ragged batches: pack into 3D batches, see `_ragged_batches` + function batched_svd_compact!( + A::AbstractVector{<:AbstractMatrix}, + USVᴴ::Tuple{AbstractVector{<:AbstractMatrix}, AbstractVector{<:Diagonal}, AbstractVector{<:AbstractMatrix}}, + alg::$Alg + ) + check_input(batched_svd_compact!, A, USVᴴ, alg) + Us, Ss, Vᴴs = USVᴴ + batches, rest = _ragged_batches(A, alg) + for (inds, (m, n)) in batches + Ab = _ragged_pack(A, inds, m, n, alg) + Ub, Sb, Vᴴb = batched_svd_compact!(Ab, _packed_output(batched_svd_compact!, Ab, alg), alg) + for (j, i) in enumerate(inds) + copyto!(Us[i], view(Ub, axes(Us[i])..., j)) + copyto!(diagview(Ss[i]), view(Sb, axes(Ss[i], 1), j)) + copyto!(Vᴴs[i], view(Vᴴb, axes(Vᴴs[i])..., j)) + end + end + for i in rest + svd_compact!(A[i], (Us[i], Ss[i], Vᴴs[i]), alg) + end + return USVᴴ + end + function batched_svd_full!( + A::AbstractVector{<:AbstractMatrix}, + USVᴴ::Tuple{AbstractVector{<:AbstractMatrix}, AbstractVector{<:AbstractMatrix}, AbstractVector{<:AbstractMatrix}}, + alg::$Alg + ) + check_input(batched_svd_full!, A, USVᴴ, alg) + Us, Ss, Vᴴs = USVᴴ + # zero padding mixes the padded dimensions into the complements of the full + # `U` and `Vᴴ`, so only matrices of equal size are batched + batches, rest = _ragged_batches(A, alg; pad = false) + for (inds, (m, n)) in batches + Ab = _ragged_pack(A, inds, m, n, alg) + Ub, Sb, Vᴴb = batched_svd_full!(Ab, _packed_output(batched_svd_full!, Ab, alg), alg) + for (j, i) in enumerate(inds) + copyto!(Us[i], view(Ub, :, :, j)) + copyto!(Ss[i], view(Sb, :, :, j)) + copyto!(Vᴴs[i], view(Vᴴb, :, :, j)) + end + end + for i in rest + svd_full!(A[i], (Us[i], Ss[i], Vᴴs[i]), alg) + end + return USVᴴ + end + function batched_svd_vals!( + A::AbstractVector{<:AbstractMatrix}, S::AbstractVector{<:AbstractVector}, alg::$Alg + ) + check_input(batched_svd_vals!, A, S, alg) + batches, rest = _ragged_batches(A, alg) + for (inds, (m, n)) in batches + Ab = _ragged_pack(A, inds, m, n, alg) + Sb = batched_svd_vals!(Ab, _packed_output(batched_svd_vals!, Ab, alg), alg) + for (j, i) in enumerate(inds) + copyto!(S[i], view(Sb, axes(S[i], 1), j)) + end + end + for i in rest + svd_vals!(A[i], S[i], alg) + end + return S + end + end + + # driver + @eval begin + @inline $svd_compact_f!(A, U, S, Vᴴ; driver::Driver = DefaultDriver(), kwargs...) = $svd_compact_f!(driver, A, U, S, Vᴴ; kwargs...) + @inline $svd_full_f!(A, U, S, Vᴴ; driver::Driver = DefaultDriver(), kwargs...) = $svd_full_f!(driver, A, U, S, Vᴴ; kwargs...) + @inline $svd_vals_f!(A, S; driver::Driver = DefaultDriver(), kwargs...) = $svd_vals_f!(driver, A, S; kwargs...) + @inline $svd_compact_f!(::DefaultDriver, A::AbstractVector{<:AbstractMatrix}, U, S, Vᴴ; kwargs...) = $svd_compact_f!(default_driver($Alg, A), A, U, S, Vᴴ; kwargs...) + @inline $svd_compact_f!(::DefaultDriver, A::AbstractArray{<:Any, 3}, U, S, Vᴴ; kwargs...) = $svd_compact_f!(default_driver($Alg, A), A, U, S, Vᴴ; kwargs...) + @inline $svd_full_f!(::DefaultDriver, A::AbstractVector{<:AbstractMatrix}, U, S, Vᴴ; kwargs...) = $svd_full_f!(default_driver($Alg, A), A, U, S, Vᴴ; kwargs...) + @inline $svd_full_f!(::DefaultDriver, A::AbstractArray{<:Any, 3}, U, S, Vᴴ; kwargs...) = $svd_full_f!(default_driver($Alg, A), A, U, S, Vᴴ; kwargs...) + @inline $svd_vals_f!(::DefaultDriver, A::AbstractVector{<:AbstractMatrix}, S::AbstractMatrix; kwargs...) = $svd_vals_f!(default_driver($Alg, A), A, S; kwargs...) + @inline $svd_vals_f!(::DefaultDriver, A::AbstractArray{<:Any, 3}, S::AbstractMatrix; kwargs...) = $svd_vals_f!(default_driver($Alg, A), A, S; kwargs...) + end + + # Implementation + @eval begin + function $svd_compact_f!(driver::Driver, A, U, S, Vᴴ; fixgauge::Bool = true, kwargs...) + if _isempty_batch(A) + foreach(one!, eachslice(U, dims = 3)) + zero!(S) + foreach(one!, eachslice(Vᴴ, dims = 3)) + return U, S, Vᴴ + end + $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) + fixgauge && gaugefix!(batched_svd_compact!, U, Vᴴ) + return U, S, Vᴴ + end + function $svd_full_f!(driver::Driver, A, U, S, Vᴴ; fixgauge::Bool = true, kwargs...) + supports_svd_full(driver, $(QuoteNode(f))) || + throw(ArgumentError(LazyString("driver ", driver, " does not provide `$($(QuoteNode(f_lapack!)))`"))) + if _isempty_batch(A) + foreach(one!, eachslice(U, dims = 3)) + zero!(S) + foreach(one!, eachslice(Vᴴ, dims = 3)) + return U, S, Vᴴ + end + zero!(S) + m, n, batch_size = size(S) + minmn = min(m, n) + Sd = similar(S, (minmn, batch_size)) + $f_lapack!(driver, A, Sd, U, Vᴴ; kwargs...) + # `gesvdx` only computes the leading `minmn` singular vectors + if $(f === :bisection) + for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) + complete_svd_basis!(u, vᴴ, minmn) + end + end + for (s, sd) in zip(eachslice(S, dims = 3), eachslice(Sd, dims = 2)) + diagview(s) .= sd + end + fixgauge && gaugefix!(batched_svd_full!, U, Vᴴ) + return U, S, Vᴴ + end + function $svd_vals_f!(driver::Driver, A::AbstractArray{T, 3}, S::AbstractMatrix; fixgauge::Bool = true, kwargs...) where {T} + _isempty_batch(A) && return zero!(S) + U, Vᴴ = similar(A, (0, 0, size(A, 3))), similar(A, (0, 0, size(A, 3))) + $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) + return S + end + function $svd_vals_f!(driver::Driver, A::AbstractVector{<:AbstractMatrix}, S::AbstractMatrix; fixgauge::Bool = true, kwargs...) + _isempty_batch(A) && return zero!(S) + U, Vᴴ = similar(first(A), (0, 0, length(A))), similar(first(A), (0, 0, length(A))) + $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) + return S + end + end +end + +# Fewest matrices in a ragged batch that are worth a batched call +# Should this be settable by the user? +const BATCHED_SVD_THRESHOLD::Int = 4 + +function _packed_output(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + a = first(A) + m, n = size(a) + b = length(A) + return (similar(a, (m, m, b)), similar(a, real(eltype(a)), (m, n, b)), similar(a, (n, n, b))) +end +function _packed_output(::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + a = first(A) + m, n = size(a) + minmn, b = min(m, n), length(A) + return (similar(a, (m, minmn, b)), similar(a, real(eltype(a)), (minmn, b)), similar(a, (minmn, n, b))) +end +function _packed_output(::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + a = first(A) + return similar(a, real(eltype(a)), (min(size(a)...), length(A))) +end diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 080019e71..9f2b6e8ae 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -142,16 +142,10 @@ function svd_via_adjoint!(f!::F, driver::Driver, A, S, U, Vᴴ; kwargs...) where end # LAPACK -for f! in (:gesdd!, :gesvd!, :gesdvd!) +for f! in (:gesdd!, :gesvd!, :gesdvd!, :gesvdx!) @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ᴴ) @@ -212,6 +206,8 @@ for (f, f_lapack!, Alg) in ( zero!(S) minmn = min(size(A)...) $f_lapack!(driver, A, view(S, 1:minmn, 1), U, Vᴴ; kwargs...) + # `gesvdx` only computes the leading `minmn` singular vectors + $(f === :bisection) && complete_svd_basis!(U, Vᴴ, minmn) diagview(S) .= view(S, 1:minmn, 1) zero!(view(S, 2:minmn, 1)) fixgauge && gaugefix!(svd_full!, U, Vᴴ) @@ -249,6 +245,16 @@ function complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int return U, Vᴴ end +""" + requires_tall(alg) -> Bool + +Whether `alg` only accepts matrices with `m ≥ n`, as cuSOLVER's and rocSOLVER's `gesvd` do +for `QRIteration`. Single matrices work around this through the adjoint (see +`svd_via_adjoint!`), whereas ragged batches zero-pad wide matrices to a square. +""" +requires_tall(::AbstractAlgorithm) = false +requires_tall(::QRIteration) = true + function svd_trunc_no_error!(A, USVᴴ, alg::TruncatedAlgorithm) U, S, Vᴴ = svd_compact!(A, USVᴴ, alg.alg) USVᴴtrunc, ind = truncate(svd_trunc!, (U, S, Vᴴ), alg.trunc) diff --git a/src/interface/batched_svd.jl b/src/interface/batched_svd.jl new file mode 100644 index 000000000..994223fc7 --- /dev/null +++ b/src/interface/batched_svd.jl @@ -0,0 +1,66 @@ +# Batched SVD functions +# ------------- +""" + batched_svd_full(A; kwargs...) -> Us, Ss, Vᴴs + batched_svd_full(A, alg::AbstractAlgorithm) -> Us, Ss, Vᴴs + batched_svd_full!(A, [USVᴴ]; kwargs...) -> Us, Ss, Vᴴs + batched_svd_full!(A, [USVᴴ], alg::AbstractAlgorithm) -> Us, Ss, Vᴴs + +Compute the *batched* full singular value decompositions (SVD) of the rectangular +matrices `A` of size `(m, n)`, such that `A[:, :, i] = U[:, :, i] * S[:, :, i] * Vᴴ[:, :, i]`. +Here, `U[:, :, i]` and `Vᴴ[:, :, i]` are unitary matrices of size +`(m, m)` and `(n, n)` respectively, and `S[:, :, i]` is a diagonal matrix of size `(m, n)`. + +!!! note + The bang method `batched_svd_full!` optionally accepts the output structure and + possibly destroys the input matrices `A`. Always use the return value of the function + as it may not always be possible to use the provided `USVᴴ` as output. + +See also [`batched_svd_compact(!)`](@ref batched_svd_compact) and +[`batched_svd_vals(!)`](@ref batched_svd_vals). +""" +@functiondef batched_svd_full + +""" + batched_svd_compact(A; kwargs...) -> Us, Ss, Vᴴs + batched_svd_compact(A, alg::AbstractAlgorithm) -> Us, Ss, Vᴴs + batched_svd_compact!(A, [USVᴴ]; kwargs...) -> Us, Ss, Vᴴs + batched_svd_compact!(A, [USVᴴ], alg::AbstractAlgorithm) -> Us, Ss, Vᴴs + +Compute the *batched* compact singular value decomposition (SVD) of the rectangular +matrices `A` of size `(m, n)`, such that `A[:, :, i] = U[:, :, i] * S[:, :, i] * Vᴴ[:, :, i]`. +Here, `U[:, :, i]` is an isometric matrix (orthonormal columns) of size `(m, k)`, whereas +`Vᴴ[:, :, i]` is a matrix of size `(k, n)` with orthonormal rows and `S[:, :, i]` +is a square diagonal matrix of size `(k, k)`, with `k = min(m, n)`. + +!!! note + The bang method `batched_svd_compact!` optionally accepts the output structure and + possibly destroys the input matrices `A`. Always use the return value of the function + as it may not always be possible to use the provided `USVᴴ` as output. + +See also [`batched_svd_full(!)`](@ref batched_svd_full) and +[`batched_svd_vals(!)`](@ref batched_svd_vals). +""" +@functiondef batched_svd_compact + +""" + batched_svd_vals(A; kwargs...) -> Ss + batched_svd_vals(A, alg::AbstractAlgorithm) -> Ss + batched_svd_vals!(A, [S]; kwargs...) -> Ss + batched_svd_vals!(A, [S], alg::AbstractAlgorithm) -> Ss + +Compute the *batched* vector of singular values of `A`, such that for an M×N matrix `A`, +`S` is a vector of size `K = min(M, N)`, the number of kept singular values. + +See also [`batched_svd_full(!)`](@ref batched_svd_full), +[`batched_svd_compact(!)`](@ref batched_svd_compact). +""" +@functiondef batched_svd_vals + +# Algorithm selection +# ------------------- +for f in (:batched_svd_full!, :batched_svd_compact!, :batched_svd_vals!) + @eval function default_algorithm(::typeof($f), ::Type{A}; kwargs...) where {A} + return default_svd_algorithm(A; kwargs...) + end +end diff --git a/src/pullbacks/svd.jl b/src/pullbacks/svd.jl index 090012553..edacce70d 100755 --- a/src/pullbacks/svd.jl +++ b/src/pullbacks/svd.jl @@ -27,6 +27,7 @@ function check_and_prepare_svd_cotangents( ΔU₁ = zero(U₁) wtmp = similar(U₁, (r,)) utmp = similar(U₁, (m,)) + zeroj = Int[] for (j, i) in enumerate(indU) if i <= r ΔU₁[:, i] .= view(ΔU, :, j) @@ -37,9 +38,12 @@ function check_and_prepare_svd_cotangents( mul!(utmp, U₁, wtmp, -1, 1) ΔgaugeU = max(ΔgaugeU, norm(utmp)) else # remaining columns should be zero - ΔgaugeU = max(ΔgaugeU, maximum(abs, view(ΔU, :, j); init = abs(zero(eltype(ΔU))))) + push!(zeroj, j) end end + # index with a vector rather than looping over views, so wrapped GPU arrays + # (e.g. `Adjoint{<:CuArray}`) don't fall back to scalar iteration + ΔgaugeU = max(ΔgaugeU, maximum(abs, ΔU[:, zeroj]; init = abs(zero(eltype(ΔU))))) end UᴴΔU₁ = U₁' * ΔU₁ ΔU₊ = mul!(ΔU₁, U₁, UᴴΔU₁, -1, 1) @@ -59,6 +63,7 @@ function check_and_prepare_svd_cotangents( ΔV₁ᴴ = zero(V₁ᴴ) wtmp = similar(V₁ᴴ, (1, r)) vtmp = similar(V₁ᴴ, (1, n)) + zeroj = Int[] for (j, i) in enumerate(indV) if i <= r ΔV₁ᴴ[i, :] .= view(ΔVᴴ, j, :) @@ -69,9 +74,10 @@ function check_and_prepare_svd_cotangents( mul!(vtmp, wtmp, V₁ᴴ, -1, 1) ΔgaugeV = max(ΔgaugeV, norm(vtmp)) else # remaining rows should be zero - ΔgaugeV = max(ΔgaugeV, maximum(abs, view(ΔVᴴ, j, :); init = abs(zero(eltype(ΔVᴴ))))) + push!(zeroj, j) end end + ΔgaugeV = max(ΔgaugeV, maximum(abs, ΔVᴴ[zeroj, :]; init = abs(zero(eltype(ΔVᴴ))))) end VᴴΔV₁ = V₁ᴴ * ΔV₁ᴴ' ΔV₊ᴴ = mul!(ΔV₁ᴴ, VᴴΔV₁', V₁ᴴ, -1, 1) diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index 5387252a3..9d9288ddb 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -45,15 +45,25 @@ if !is_buildkite end end +batch_size = 16 + # CUDA tests # ------------ if CUDA.functional() - # LAPACK algorithms: + # CUSOLVER algorithms: for T in BLASFloats, m in (0, 23), n in (0, 17, m, 27) TestSuite.seed_rng!(123) TestSuite.test_svd(CuMatrix{T}, (m, n)) CUDA_SVD_ALGS = (QRIteration(), SVDViaPolar(), Jacobi()) TestSuite.test_svd_algs(CuMatrix{T}, (m, n), CUDA_SVD_ALGS) + + TestSuite.test_svd_batched(CuMatrix{T}, (m, n), batch_size) + CUDA_SVD_ALGS = (Jacobi(),) + TestSuite.test_svd_batched_algs(CuMatrix{T}, (m, n), batch_size, CUDA_SVD_ALGS) + end + for T in BLASFloats + TestSuite.seed_rng!(123) + TestSuite.test_svd_algs_batched_oversized(CuMatrix{T}, (Jacobi(),), batch_size) end # Randomized SVD: @@ -77,12 +87,14 @@ end # AMDGPU tests # ------------ if AMDGPU.functional() - # LAPACK algorithms: + # ROCSOLVER algorithms: 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(), Bisection()) + AMD_SVD_ALGS = (QRIteration(), Jacobi(), DivideAndConquer(), Bisection()) TestSuite.test_svd_algs(ROCMatrix{T}, (m, n), AMD_SVD_ALGS) + TestSuite.test_svd_batched(ROCMatrix{T}, (m, n), batch_size) + TestSuite.test_svd_batched_algs(ROCMatrix{T}, (m, n), batch_size, AMD_SVD_ALGS) end # Diagonal: diff --git a/test/testsuite/TestSuite.jl b/test/testsuite/TestSuite.jl index 36ae68304..653d99f80 100644 --- a/test/testsuite/TestSuite.jl +++ b/test/testsuite/TestSuite.jl @@ -42,6 +42,20 @@ instantiate_matrix(::Type{AT}, size) where {AT <: Diagonal} = Diagonal(randn(rng instantiate_matrix(::Type{AT}, size) where {T, AT <: Diagonal{T, <:CuVector}} = Diagonal(CuArray(randn(rng, eltype(AT), size))) instantiate_matrix(::Type{AT}, size) where {T, AT <: Diagonal{T, <:ROCVector}} = Diagonal(ROCArray(randn(rng, eltype(AT), size))) +# Collect a batch of separately-allocated matrices into a single contiguous batch, transferring +# the array data itself rather than a list of pointers into it. For the GPU eltypes the copies +# stay device-to-device. +function _collect_batch(As::AbstractVector{<:AbstractMatrix}) + B = similar(first(As), (size(first(As))..., length(As))) + for (i, A) in enumerate(As) + copyto!(view(B, :, :, i), A) + end + return B +end +device_batch(As::AbstractVector{<:Array}) = _collect_batch(As) +device_batch(As::AbstractVector{<:CuArray}) = _collect_batch(As) +device_batch(As::AbstractVector{<:ROCArray}) = _collect_batch(As) + precision(::Type{T}) where {T <: Number} = sqrt(eps(real(T))) precision(::Type{T}) where {T} = precision(eltype(T)) diff --git a/test/testsuite/decompositions/svd.jl b/test/testsuite/decompositions/svd.jl index 4b89d4973..d205ce63c 100644 --- a/test/testsuite/decompositions/svd.jl +++ b/test/testsuite/decompositions/svd.jl @@ -11,6 +11,16 @@ function test_svd(T::Type, sz; test_compact::Bool = true, test_full::Bool = true end end +function test_svd_batched(T::Type, sz, batch_size::Int; test_compact::Bool = true, test_full::Bool = true, test_trunc::Bool = true, kwargs...) + summary_str = testargs_summary(T, sz) + return @testset "svd batched $summary_str batch_size $batch_size" begin + test_compact && test_svd_compact_batched(T, sz, batch_size; kwargs...) + test_full && test_svd_full_batched(T, sz, batch_size; kwargs...) + # TODO + #test_trunc && test_svd_trunc(T, sz; kwargs...) + end +end + function test_svd_algs(T::Type, sz, algs; test_compact::Bool = true, test_full::Bool = true, test_trunc::Bool = true, kwargs...) summary_str = testargs_summary(T, sz) return @testset "svd algorithms $summary_str" begin @@ -20,6 +30,16 @@ function test_svd_algs(T::Type, sz, algs; test_compact::Bool = true, test_full:: end end +function test_svd_batched_algs(T::Type, sz, batch_size::Int, algs; test_compact::Bool = true, test_full::Bool = true, test_trunc::Bool = true, kwargs...) + summary_str = testargs_summary(T, sz) + return @testset "svd batched algorithms $summary_str batch_size $batch_size" begin + test_compact && test_svd_compact_algs_batched(T, sz, algs, batch_size; kwargs...) + test_full && test_svd_full_algs_batched(T, sz, algs, batch_size; kwargs...) + # TODO + #test_trunc && test_svd_trunc_algs(T, sz, algs; kwargs...) + end +end + function test_svd_compact( T::Type, sz; atol::Real = 0, rtol::Real = precision(eltype(T)), @@ -54,6 +74,47 @@ function test_svd_compact( end end +function test_svd_compact_batched( + T::Type, sz, batch_size::Int; + atol::Real = 0, rtol::Real = precision(eltype(T)), + test_vals::Bool = true, kwargs... + ) + summary_str = testargs_summary(T, sz) + return @testset "svd_compact! $summary_str batch_size $batch_size" begin + As = [instantiate_matrix(T, sz) for bi in 1:batch_size] + Ad = device_batch(As) + Ac = deepcopy(Ad) + m, n = size(first(As)) + minmn = min(m, n) + U, S, Vᴴ = @testinferred batched_svd_compact(Ad) + @test size(U) == (m, minmn, batch_size) + @test S isa AbstractMatrix{real(eltype(T))} && size(S) == (minmn, batch_size) + @test size(Vᴴ) == (minmn, n, batch_size) + for (a, u, s, vᴴ) in zip(As, eachslice(U, dims = 3), eachslice(S, dims = 2), eachslice(Vᴴ, dims = 3)) + @test u * Diagonal(s) * vᴴ ≈ a + @test isisometric(u) + @test isisometric(vᴴ; side = :right) + @test isposdef(Diagonal(s)) + end + + Sc = similar(diagview(S)) + U2, S2, V2ᴴ = @testinferred batched_svd_compact!(Ac, (U, S, Vᴴ)) + for (a, u, s, vᴴ) in zip(As, eachslice(U2, dims = 3), eachslice(S2, dims = 2), eachslice(V2ᴴ, dims = 3)) + @test u * Diagonal(s) * vᴴ ≈ a + @test isisometric(u) + @test isisometric(vᴴ; side = :right) + @test isposdef(Diagonal(s)) + end + + if test_vals + Sd = @testinferred batched_svd_vals(Ad) + for (s, sd) in zip(eachslice(S, dims = 2), eachslice(Sd, dims = 2)) + @test s ≈ sd + end + end + end +end + function test_svd_compact_algs( T::Type, sz, algs; atol::Real = 0, rtol::Real = precision(eltype(T)), @@ -87,6 +148,83 @@ function test_svd_compact_algs( end end +function test_svd_compact_algs_batched( + T::Type, sz, algs, batch_size::Int; + atol::Real = 0, rtol::Real = precision(eltype(T)), + test_vals::Bool = true, kwargs... + ) + summary_str = testargs_summary(T, sz) + return @testset "svd_compact! algorithm $alg $summary_str batch_size $batch_size" for alg in algs + As = [instantiate_matrix(T, sz) for bi in 1:batch_size] + Ad = device_batch(As) + Ac = deepcopy(Ad) + m, n = size(first(As)) + minmn = min(m, n) + U, S, Vᴴ = @testinferred batched_svd_compact(Ad; alg) + @test size(U) == (m, minmn, batch_size) + @test S isa AbstractMatrix{real(eltype(T))} && size(S) == (minmn, batch_size) + @test size(Vᴴ) == (minmn, n, batch_size) + for (a, u, s, vᴴ) in zip(As, eachslice(U, dims = 3), eachslice(S, dims = 2), eachslice(Vᴴ, dims = 3)) + @test u * Diagonal(s) * vᴴ ≈ a + @test isisometric(u) + @test isisometric(vᴴ; side = :right) + @test isposdef(Diagonal(s)) + end + + U2, S2, V2ᴴ = @testinferred batched_svd_compact!(Ac, (U, S, Vᴴ); alg) + for (a, u, s, vᴴ) in zip(As, eachslice(U2, dims = 3), eachslice(S2, dims = 2), eachslice(V2ᴴ, dims = 3)) + @test u * Diagonal(s) * vᴴ ≈ a + @test isisometric(u) + @test isisometric(vᴴ; side = :right) + @test isposdef(Diagonal(s)) + end + + if test_vals + Sd = @testinferred batched_svd_vals(Ad; alg) + for (s, sd) in zip(eachslice(S, dims = 2), eachslice(Sd, dims = 2)) + @test s ≈ sd + end + end + + # ragged batch: `As` is one group of `batch_size` equal sized matrices, + # and then `nextra` matrices of different sizes are either decomposed + # one at a time or zero-padded into one more batch + @testset "ragged with $nextra extra sizes" for nextra in (2, 5) + Ar = [As; [instantiate_matrix(T, (max(m - i % 3, 0), max(n - i % 4, 0))) for i in 1:nextra]] + Us = [similar(a, size(a, 1), minimum(size(a))) for a in Ar] + Ss = [Diagonal(similar(a, real(eltype(T)), minimum(size(a)))) for a in Ar] + Vᴴs = [similar(a, minimum(size(a)), size(a, 2)) for a in Ar] + U3, S3, V3ᴴ = @testinferred batched_svd_compact!(deepcopy(Ar), (Us, Ss, Vᴴs); alg) + @test S3 === Ss + for (a, u, s, vᴴ) in zip(Ar, U3, S3, V3ᴴ) + @test u * s * vᴴ ≈ a + @test isisometric(u) + @test isisometric(vᴴ; side = :right) + end + + U4, S4, V4ᴴ = @testinferred batched_svd_compact(Ar; alg) + @test S4 isa AbstractVector{<:Diagonal} + for (a, u, s, vᴴ) in zip(Ar, U4, S4, V4ᴴ) + @test u * s * vᴴ ≈ a + @test isisometric(u) + @test isisometric(vᴴ; side = :right) + end + + if test_vals + Sv = [similar(a, real(eltype(T)), minimum(size(a))) for a in Ar] + Sv2 = @testinferred batched_svd_vals!(deepcopy(Ar), Sv; alg) + for (s, s2) in zip(S3, Sv2) + @test collect(diagview(s)) ≈ collect(s2) + end + Sv3 = @testinferred batched_svd_vals(Ar; alg) + for (s, s3) in zip(S3, Sv3) + @test collect(diagview(s)) ≈ collect(s3) + end + end + end + end +end + function test_svd_full( T::Type, sz; atol::Real = 0, rtol::Real = precision(eltype(T)), @@ -120,6 +258,45 @@ function test_svd_full( end end +function test_svd_full_batched( + T::Type, sz, batch_size::Int; + atol::Real = 0, rtol::Real = precision(eltype(T)), + kwargs... + ) + summary_str = testargs_summary(T, sz) + return @testset "svd_full! $summary_str batch_size $batch_size" begin + As = [instantiate_matrix(T, sz) for bi in 1:batch_size] + Ad = device_batch(As) + Ac = deepcopy(Ad) + m, n = size(first(As)) + minmn = min(m, n) + U, S, Vᴴ = @testinferred batched_svd_full(Ad) + @test size(U) == (m, m, batch_size) + @test S isa AbstractArray{real(eltype(T)), 3} && size(S) == (m, n, batch_size) + @test size(Vᴴ) == (n, n, batch_size) + for (a, u, s, vᴴ) in zip(As, eachslice(U, dims = 3), eachslice(S, dims = 3), eachslice(Vᴴ, dims = 3)) + @test u * s * vᴴ ≈ a + @test isunitary(u) + @test isunitary(vᴴ) + @test all(isposdef, diagview(s)) + end + + U2, S2, V2ᴴ = @testinferred batched_svd_full!(Ac, (U, S, Vᴴ)) + for (a, u, s, vᴴ) in zip(As, eachslice(U2, dims = 3), eachslice(S2, dims = 3), eachslice(V2ᴴ, dims = 3)) + @test u * s * vᴴ ≈ a + @test isunitary(u) + @test isunitary(vᴴ) + @test all(isposdef, diagview(s)) + end + + Sc = similar(first(As), real(eltype(T)), min(m, n), batch_size) + Sc2 = @testinferred batched_svd_vals!(copy!(Ac, Ad), Sc) + for (s, s2) in zip(eachslice(S, dims = 3), eachslice(Sc, dims = 2)) + @test collect(diagview(s)) ≈ collect(s2) + end + end +end + function test_svd_full_algs( T::Type, sz, algs; atol::Real = 0, rtol::Real = precision(eltype(T)), @@ -153,6 +330,116 @@ function test_svd_full_algs( end end +function test_svd_full_algs_batched( + T::Type, sz, algs, batch_size::Int; + atol::Real = 0, rtol::Real = precision(eltype(T)), + kwargs... + ) + summary_str = testargs_summary(T, sz) + return @testset "svd_full! algorithm $alg $summary_str batch_size $batch_size" for alg in algs + As = [instantiate_matrix(T, sz) for bi in 1:batch_size] + Ad = device_batch(As) + Ac = deepcopy(Ad) + m, n = size(first(As)) + minmn = min(m, n) + U, S, Vᴴ = @testinferred batched_svd_full(Ad; alg) + @test size(U) == (m, m, batch_size) + @test S isa AbstractArray{real(eltype(T)), 3} && size(S) == (m, n, batch_size) + @test size(Vᴴ) == (n, n, batch_size) + for (a, u, s, vᴴ) in zip(As, eachslice(U, dims = 3), eachslice(S, dims = 3), eachslice(Vᴴ, dims = 3)) + @test u * s * vᴴ ≈ a + @test isunitary(u) + @test isunitary(vᴴ) + @test all(isposdef, diagview(s)) + end + + U2, S2, V2ᴴ = @testinferred batched_svd_full!(Ac, (U, S, Vᴴ); alg) + for (a, u, s, vᴴ) in zip(As, eachslice(U2, dims = 3), eachslice(S2, dims = 3), eachslice(V2ᴴ, dims = 3)) + @test u * s * vᴴ ≈ a + @test isunitary(u) + @test isunitary(vᴴ) + @test all(isposdef, diagview(s)) + end + + Sc = similar(first(As), real(eltype(T)), min(m, n), batch_size) + Sc2 = @testinferred batched_svd_vals!(copy!(Ac, Ad), Sc; alg) + for (s, s2) in zip(eachslice(S, dims = 3), eachslice(Sc, dims = 2)) + @test collect(diagview(s)) ≈ collect(s2) + end + + # ragged batch: `As` is one group of `batch_size` equal sized matrices, + # and then `nextra` matrices of different sizes are either decomposed + # one at a time or zero-padded into one more batch + @testset "ragged with $nextra extra sizes" for nextra in (2, 5) + Ar = [As; [instantiate_matrix(T, (max(m - i % 3, 0), max(n - i % 4, 0))) for i in 1:nextra]] + Us = [similar(a, size(a, 1), size(a, 1)) for a in Ar] + Ss = [similar(a, real(eltype(T)), size(a)) for a in Ar] + Vᴴs = [similar(a, size(a, 2), size(a, 2)) for a in Ar] + U3, S3, V3ᴴ = @testinferred batched_svd_full!(deepcopy(Ar), (Us, Ss, Vᴴs); alg) + for (a, u, s, vᴴ) in zip(Ar, U3, S3, V3ᴴ) + @test size(u) == (size(a, 1), size(a, 1)) + @test size(s) == size(a) + @test size(vᴴ) == (size(a, 2), size(a, 2)) + @test u * s * vᴴ ≈ a + @test isunitary(u) + @test isunitary(vᴴ) + end + + U4, S4, V4ᴴ = @testinferred batched_svd_full(Ar; alg) + for (a, u, s, vᴴ) in zip(Ar, U4, S4, V4ᴴ) + @test u * s * vᴴ ≈ a + @test isunitary(u) + @test isunitary(vᴴ) + end + + Sv = [similar(a, real(eltype(T)), minimum(size(a))) for a in Ar] + Sv2 = @testinferred batched_svd_vals!(deepcopy(Ar), Sv; alg) + for (s, s2) in zip(S3, Sv2) + @test collect(diagview(s)) ≈ collect(s2) + end + end + end +end + +# Ragged batches containing matrices larger than the driver's batched size limit must +# be split off and decomposed one at a time rather than handed to the driver. +function test_svd_algs_batched_oversized( + T::Type, algs, batch_size::Int; + atol::Real = 0, rtol::Real = precision(eltype(T)), + kwargs... + ) + summary_str = testargs_summary(T) + return @testset "batched svd over the size limit, algorithm $alg $summary_str" for alg in algs + limit = MatrixAlgebraKit.max_batched_blocksize(alg, MatrixAlgebraKit.default_driver(alg, T), T) + if limit < typemax(Int) # nothing to do if the driver + algo combo has no limit + sizes = ((limit + 5, limit + 3), (limit + 1, 5), (limit - 2, limit - 4)) + Ar = [instantiate_matrix(T, sz) for sz in sizes for _ in 1:batch_size] + batches, _ = MatrixAlgebraKit._ragged_batches(Ar, alg) + @test all(((inds, mn),) -> maximum(mn) <= limit, batches) + @test !isempty(batches) + + U, S, Vᴴ = @testinferred batched_svd_compact(Ar; alg) + for (a, u, s, vᴴ) in zip(Ar, U, S, Vᴴ) + @test u * s * vᴴ ≈ a + @test isisometric(u) + @test isisometric(vᴴ; side = :right) + end + + Uf, Sf, Vfᴴ = @testinferred batched_svd_full(Ar; alg) + for (a, u, s, vᴴ) in zip(Ar, Uf, Sf, Vfᴴ) + @test u * s * vᴴ ≈ a + @test isunitary(u) + @test isunitary(vᴴ) + end + + Sv = @testinferred batched_svd_vals(Ar; alg) + for (s, sv) in zip(S, Sv) + @test collect(diagview(s)) ≈ collect(sv) + end + end + end +end + function test_svd_trunc( T::Type, sz; atol::Real = 0, rtol::Real = precision(eltype(T)),