From 7475023fed0af46ba6d42b1703ef9b75ce41fc47 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 21 Aug 2026 16:41:22 +0200 Subject: [PATCH 01/46] Batched SVD support for ROCSOLVER and CUSOLVER --- .../MatrixAlgebraKitAMDGPUExt.jl | 43 +- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 573 ++++++++++++++++++ .../MatrixAlgebraKitCUDAExt.jl | 9 +- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 86 +++ src/MatrixAlgebraKit.jl | 7 +- src/implementations/svd.jl | 224 ++++++- src/interface/decompositions.jl | 104 +++- test/decompositions/svd.jl | 17 +- test/testsuite/TestSuite.jl | 14 + test/testsuite/decompositions/svd.jl | 179 ++++++ 10 files changed, 1240 insertions(+), 16 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 4a80ae0fd..a47a5ea70 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -6,7 +6,8 @@ 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! using AMDGPU @@ -15,10 +16,17 @@ 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}} - return QRIteration(; kwargs...) +function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCMatrix{<:BlasFloat}} + return DivideAndConquer(; kwargs...) +end +function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCArray{<:BlasFloat, 3}} + return DivideAndConquerBatched(; kwargs...) +end +function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} + return DivideAndConquerBatched(; kwargs...) end function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}} return DivideAndConquer(; kwargs...) @@ -28,7 +36,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, :bisection) +MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :divide_and_conquer, :bisection, :qr_iteration_batched, :jacobi_batched, :divide_and_conquer_batched, :bisection_batched) function gesvd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) m, n = size(A) @@ -48,6 +56,31 @@ function gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid return S, U, Vᴴ end +gesvd_batched!(::ROCSOLVER, As::Vector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YArocSOLVER.gesvd_batched!(As, Ss, Us, Vᴴs; kwargs...) +gesvd_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = + YArocSOLVER.gesvd_strided_batched!(As, Ss, Us, Vᴴs; kwargs...) + +gesdd_batched!(::ROCSOLVER, As::Vector{<: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::Vector{<: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::Vector{<: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...) + +gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) = + YArocSOLVER.gesvdx!(A, S, U, Vᴴ; 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...) heevd!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 6c5bd3469..5d43b061a 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -97,6 +97,405 @@ 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::StridedROCVector{<: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)), + ) + 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 + 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 + 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 + length(S) == minmn * length(A) || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(first(A), 2)) + ldu = max(1, stride(U, 2)) + strideU = ldu * size(U, 2) + ldv = max(1, stride(Vᴴ, 2)) + strideV = ldv * n + strideS = minmn + + strideE = minmn - 1 + E = ROCArray{$relty}(undef, length(A) * strideE) + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, length(A)) + pA = 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) + + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + + 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)), + ) + 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 + 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 * batch_size || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(A, 2)) + strideA = lda * n + ldu = max(1, stride(U, 2)) + strideU = ldu * size(U, 2) + ldv = max(1, stride(Vᴴ, 2)) + strideV = ldv * n + strideS = minmn + + 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) + + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + + 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)) + ) + 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 + ) + + info = @allowscalar dev_info[1] + rocSOLVER.chkargsok(BlasInt(info)) + + 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::StridedROCVector{<: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)), + ) + 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 + 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 + 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 + length(S) == minmn * length(A) || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(first(A), 2)) + ldu = max(1, stride(U, 2)) + strideU = ldu * size(U, 2) + ldv = max(1, stride(Vᴴ, 2)) + strideV = ldv * n + strideS = minmn + + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, length(A)) + pA = 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) + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + + 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)), + ) + 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 + 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 * batch_size || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(A, 2)) + strideA = lda * n + ldu = max(1, stride(U, 2)) + strideU = ldu * size(U, 2) + ldv = max(1, stride(Vᴴ, 2)) + strideV = ldv * n + strideS = minmn + + 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 + ) + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + + return (S, U, Vᴴ) + end + end +end + # Wrapper for SVD via Jacobi for (fname, elty, relty) in ( @@ -182,6 +581,180 @@ 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::StridedROCVector{<: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, + ) + 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 + 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 + 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 + length(S) == minmn * length(A) || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(first(A), 2)) + ldu = max(1, stride(U, 2)) + strideU = ldu * size(U, 2) + ldv = max(1, stride(Vᴴ, 2)) + strideV = ldv * n + strideS = minmn + 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 = 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) + ) + + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + + 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, + ) + 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 + 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 + 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 + length(S) == minmn * batch_size || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(A, 2)) + strideA = lda * n + ldu = max(1, stride(U, 2)) + strideU = ldu * size(U, 2) + ldv = max(1, stride(Vᴴ, 2)) + strideV = ldv * n + strideS = minmn + 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 + ) + + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + + 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`. diff --git a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl index 704e92e7a..a5df7079e 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 @@ -21,6 +22,9 @@ MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedCuVecOrMat{<:Bla 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 JacobiBatched(; kwargs...) +end function MatrixAlgebraKit.default_eig_algorithm(::Type{T}; kwargs...) where {T <: StridedCuVecOrMat{<:BlasFloat}} return QRIteration(; kwargs...) end @@ -35,7 +39,7 @@ end MatrixAlgebraKit.prefers_ungqr(::CUSOLVER) = true -MatrixAlgebraKit.supports_svd_full(::CUSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :svd_polar) +MatrixAlgebraKit.supports_svd_full(::CUSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :svd_polar, :jacobi_batched) function gesvd!(::CUSOLVER, A::StridedCuMatrix, S::StridedCuVector, U::StridedCuMatrix, Vᴴ::StridedCuMatrix; kwargs...) m, n = size(A) @@ -49,6 +53,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_strided_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..8f991c4cd 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -276,6 +276,92 @@ 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, + 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' + econ = 0 + else + 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 + econ = 1 + elseif size(U, 2) == m && size(Vᴴ, 1) == n + econ = 0 + else + throw(DimensionMismatch("invalid column size of U or row size of Vᴴ")) + end + end + length(S) == minmn || + throw(DimensionMismatch("length mismatch between A and S")) + + Ṽ = (jobz == 'V') ? similar(Vᴴ') : similar(Vᴴ, (n, minmn)) + Ũ = (jobz == 'V') ? U : similar(U, (m, minmn)) + 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() + + function bufferSize() + out = Ref{Cint}(0) + cuSOLVER.$bname( + dh, jobz, econ, m, n, A, lda, S, Ũ, ldu, Ṽ, ldv, + out, params[] + ) + return out[] * sizeof($elty) + end + + cuSOLVER.with_workspace(dh.workspace_gpu, bufferSize) do buffer + return cuSOLVER.$fname( + dh, jobz, econ, m, n, A, lda, S, Ũ, ldu, Ṽ, ldv, + buffer, sizeof(buffer) ÷ sizeof($elty), dh.info, + params[], batch_size + ) + end + + info = collect(dh.info) + cuSOLVER.chkargsok.(BlasInt.(info)) + + cuSOLVER.cusolverDnDestroyGesvdjInfo(params[]) + + if jobz == 'V' + adjoint!(Vᴴ, Ṽ) + 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..59f9cf42e 100644 --- a/src/MatrixAlgebraKit.jl +++ b/src/MatrixAlgebraKit.jl @@ -34,7 +34,7 @@ export exponential, exponential! export Householder, Native_HouseholderQR, Native_HouseholderLQ export DivideAndConquer, SafeDivideAndConquer, QRIteration, Bisection, Jacobi, SVDViaPolar -export RobustRepresentations +export RobustRepresentations, DivideAndConquerBatched, QRIterationBatched, BisectionBatched, JacobiBatched export LAPACK_HouseholderQR, LAPACK_HouseholderLQ, LAPACK_Simple, LAPACK_Expert, LAPACK_QRIteration, LAPACK_Bisection, LAPACK_MultipleRelativelyRobustRepresentations, LAPACK_DivideAndConquer, LAPACK_Jacobi, LAPACK_SafeDivideAndConquer @@ -46,9 +46,10 @@ export DefaultAlgorithm export DiagonalAlgorithm export NativeBlocked export CUSOLVER_Simple, CUSOLVER_HouseholderQR, CUSOLVER_QRIteration, CUSOLVER_SVDPolar, - CUSOLVER_Jacobi, CUSOLVER_Randomized, CUSOLVER_DivideAndConquer + CUSOLVER_Jacobi, CUSOLVER_Randomized, CUSOLVER_DivideAndConquer, CUSOLVER_JacobiBatched export ROCSOLVER_HouseholderQR, ROCSOLVER_QRIteration, ROCSOLVER_Jacobi, - ROCSOLVER_DivideAndConquer, ROCSOLVER_Bisection + ROCSOLVER_DivideAndConquer, ROCSOLVER_Bisection, ROCSOLVER_QRIterationBatched, ROCSOLVER_JacobiBatched, + ROCSOLVER_DivideAndConquerBatched, ROCSOLVER_BisectionBatched export notrunc, truncrank, trunctol, truncerror, truncfilter diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 080019e71..18a54a7ba 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -1,5 +1,7 @@ # Input # ------ +copy_input(::typeof(svd_full), As::AbstractVector{<:AbstractMatrix}) = map(A -> copy!(similar(A, float(eltype(A))), A), As) +copy_input(::typeof(svd_full), A::AbstractArray{T, 3}) where {T} = copy!(similar(A, float(T)), A) copy_input(::typeof(svd_full), A::AbstractMatrix) = copy!(similar(A, float(eltype(A))), A) copy_input(::typeof(svd_compact), A) = copy_input(svd_full, A) copy_input(::typeof(svd_vals), A) = copy_input(svd_full, A) @@ -42,6 +44,80 @@ function check_input(::typeof(svd_vals!), A::AbstractMatrix, S, ::AbstractAlgori return nothing end +# batched varieties +function check_input(::typeof(svd_full!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(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(svd_compact!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(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(svd_vals!), A::AbstractVector{<:AbstractMatrix}, S, ::AbstractAlgorithm) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(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(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(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(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 + function check_input(::typeof(svd_full!), A::AbstractMatrix, USVᴴ, ::DiagonalAlgorithm) m, n = size(A) @assert m == n && isdiag(A) @@ -92,6 +168,47 @@ end function initialize_output(::Union{typeof(svd_trunc!), typeof(svd_trunc_no_error!)}, A, alg::TruncatedAlgorithm) return initialize_output(svd_compact!, A, alg.alg) end +# batched versions +function initialize_output(::typeof(svd_full!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + m, n = size(first(A)) + U = similar(first(A), (m, m, length(A))) + S = similar(first(A), real(eltype(first(A))), (m, n, length(A))) # TODO: Rectangular diagonal type? + Vᴴ = similar(first(A), (n, n, length(A))) + return (U, S, Vᴴ) +end +function initialize_output(::typeof(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(svd_compact!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(A)) + minmn = min(m, n) + U = similar(first(A), (m, minmn, length(A))) + S = similar(first(A), real(eltype(first(A))), minmn, length(A)) + Vᴴ = similar(first(A), (minmn, n, length(A))) + return (U, S, Vᴴ) +end +function initialize_output(::typeof(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(svd_vals!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(A)) + return similar(first(A), real(eltype(first(A))), (min(m, n), length(A))) +end +function initialize_output(::typeof(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 function initialize_output(::typeof(svd_full!), A::Diagonal, ::DiagonalAlgorithm) TA = eltype(A) @@ -120,10 +237,17 @@ end # IMPLEMENTATIONS # ========================== -for f! in (:gesdd!, :gesvd!, :gesvdj!, :gesvdp!, :gesvdx!, :gesvdr!, :gesdvd!) +for f! in (:gesdd!, :gesvd!, :gesvdj!, :gesvdp!, :gesvdx!, :gesvdr!, :gesdvd!, :gesdd_batched!, :gesvd_batched!, :gesvdj_batched!, :gesvdx_batched!) @eval $f!(driver::Driver, args...) = throw(ArgumentError("$driver does not provide $($(f!))")) end +# declare these as dummies so the GPU extensions can import them safely +function gesvd_batched! end +function gesdd_batched! end +function gesvdj_batched! end +function gesvdx_batched! end + + """ svd_via_adjoint!(f!, driver, A, S, U, Vᴴ; kwargs...) @@ -226,6 +350,104 @@ for (f, f_lapack!, Alg) in ( end end +# batched varieties +for (f, f_lapack!, Alg) in ( + (:divide_and_conquer_batched, :gesdd_batched!, :DivideAndConquerBatched), + (:qr_iteration_batched, :gesvd_batched!, :QRIterationBatched), + (:bisection_batched, :gesvdx_batched!, :BisectionBatched), + (:jacobi_batched, :gesvdj_batched!, :JacobiBatched), + ) + svd_compact_f! = Symbol(:svd_compact_, f, :!) + svd_full_f! = Symbol(:svd_full_, f, :!) + svd_vals_f! = Symbol(:svd_vals_, f, :!) + + # MatrixAlgebraKit wrappers + @eval begin + function svd_compact!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) + check_input(svd_compact!, A, USVᴴ, alg) + return $svd_compact_f!(A, USVᴴ...; alg.kwargs...) + end + function svd_compact!(A::AbstractArray{T, 3}, USVᴴ, alg::$Alg) where {T} + check_input(svd_compact!, A, USVᴴ, alg) + return $svd_compact_f!(A, USVᴴ...; alg.kwargs...) + end + function svd_full!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) + check_input(svd_full!, A, USVᴴ, alg) + return $svd_full_f!(A, USVᴴ...; alg.kwargs...) + end + function svd_full!(A::AbstractArray{T, 3}, USVᴴ, alg::$Alg) where {T} + check_input(svd_full!, A, USVᴴ, alg) + return $svd_full_f!(A, USVᴴ...; alg.kwargs...) + end + function svd_vals!(A::AbstractVector{<:AbstractMatrix}, S, alg::$Alg) + check_input(svd_vals!, A, S, alg) + return $svd_vals_f!(A, S; alg.kwargs...) + end + function svd_vals!(A::AbstractArray{T, 3}, S, alg::$Alg) where {T} + check_input(svd_vals!, A, S, alg) + return $svd_vals_f!(A, S; alg.kwargs...) + 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...) + isempty(A) && return one!(U), zero!(S), one!(Vᴴ) + $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) + if fixgauge + for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) + gaugefix!(svd_compact!, u, vᴴ) + end + end + 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!)))`"))) + isempty(A) && return one!(U), zero!(S), one!(Vᴴ) + 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...) + for (s, sd) in zip(eachslice(S, dims = 3), eachslice(Sd, dims = 2)) + diagview(s) .= sd + end + if fixgauge + for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) + gaugefix!(svd_full!, u, vᴴ) + end + end + return U, S, Vᴴ + end + function $svd_vals_f!(driver::Driver, A::AbstractArray{T, 3}, S::AbstractMatrix; fixgauge::Bool = true, kwargs...) where {T} + isempty(A) && return zero!(S) + U, Vᴴ = similar(A, (0, 0, 0)), similar(A, (0, 0, 0)) + $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(A) && return zero!(S) + U, Vᴴ = similar(first(A), (0, 0, 0)), similar(first(A), (0, 0, 0)) + $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) + return S + end + end +end + supports_svd_full(::Driver, ::Symbol) = false supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration, :bisection) diff --git a/src/interface/decompositions.jl b/src/interface/decompositions.jl index 1d423564a..90999fb3c 100644 --- a/src/interface/decompositions.jl +++ b/src/interface/decompositions.jl @@ -100,6 +100,17 @@ The optional `driver` keyword can be used to choose between different implementa """ @algdef DivideAndConquer +""" + DivideAndConquerBatched(; [driver], fixgauge = default_fixgauge()) + +Algorithm type for computing the *batched* eigenvalue decompositions of a set of Hermitian matrices, +or the singular value decompositions of a set of general matrices using the divide-and-conquer algorithm. + +$_fixgauge_docs +The optional `driver` keyword can be used to choose between different implementations of this algorithm. +""" +@algdef DivideAndConquerBatched + """ SafeDivideAndConquer(; [driver], fixgauge = default_fixgauge()) @@ -162,6 +173,49 @@ The optional `driver` keyword can be used to choose between different implementa """ @algdef Jacobi +""" + QRIterationBatched(; [driver], fixgauge = default_fixgauge(), kwargs...) + +Algorithm type for computing the *batched* eigenvalue, Schur or singular value decompositions of a set of matrices via QR iteration. + +## Keyword arguments + +Various customizations are available, depending on the type of decomposition this algorithm is used for. + +Schur decompositions are not yet supported. + +For non-Hermitian eigenvalue decompositions there is `permute = true` and `scale = true` to control whether +or not to balance the input matrix before starting the QR iterations. + +For the singular value and eigenvalue decompositions, there is residual freedom in the outputs that can be resolved. +$_fixgauge_docs + +In all cases, the optional `driver` keyword can be used to choose between different implementations of this algorithm. +""" +@algdef QRIterationBatched + +""" + BisectionBatched(; [driver], fixgauge = default_fixgauge()) + +Algorithm type for computing the *batched* eigenvalue decompositions of a set of Hermitian matrices +via the bisection algorithm, or the singular value decompositions of a set of general matrices. + +$_fixgauge_docs +The optional `driver` keyword can be used to choose between different implementations of this algorithm. +""" +@algdef BisectionBatched + +""" + JacobiBatched(; [driver], fixgauge = default_fixgauge()) + +Algorithm type for computing the *batched* eigenvalue decompositions of a set of Hermitian matrices, +or the singular value decompositions of a set of general matrices using the Jacobi algorithm. + +$_fixgauge_docs +The optional `driver` keyword can be used to choose between different implementations of this algorithm. +""" +@algdef JacobiBatched + """ RobustRepresentations(; [driver], fixgauge = default_fixgauge()) @@ -397,6 +451,15 @@ $_fixgauge_docs """ @algdef CUSOLVER_Jacobi +""" + CUSOLVER_JacobiBatched(; fixgauge = default_fixgauge()) + +Algorithm type to denote the CUSOLVER driver for computing the *batched singular value decompositions +of a set of general matrices using the Jacobi algorithm. +$_fixgauge_docs +""" +@algdef CUSOLVER_JacobiBatched + """ CUSOLVER_Randomized(; k, p, niters) @@ -485,6 +548,45 @@ $_fixgauge_docs """ @algdef ROCSOLVER_DivideAndConquer +""" + ROCSOLVER_QRIterationBatched(; fixgauge = default_fixgauge()) + +Algorithm type to denote the ROCSOLVER driver for computing the *batched* eigenvalue decompositions of a +set of Hermitian matrices, or the singular value decompositions of a set of general matrices using the +QR Iteration algorithm. +$_fixgauge_docs +""" +@algdef ROCSOLVER_QRIterationBatched + +""" + ROCSOLVER_JacobiBatched(; fixgauge = default_fixgauge()) + +Algorithm type to denote the ROCSOLVER driver for computing the *batched* singular value decompositions of +a set of general matrices using the Jacobi algorithm. +$_fixgauge_docs +""" +@algdef ROCSOLVER_JacobiBatched + +""" + ROCSOLVER_BisectionBatched(; fixgauge = default_fixgauge()) + +Algorithm type to denote the ROCSOLVER driver for computing the *batched* eigenvalue decomposition of a +set of Hermitian matrices, or the singular value decompositions of a set of general matrices using the +Bisection algorithm. +$_fixgauge_docs +""" +@algdef ROCSOLVER_BisectionBatched + +""" + ROCSOLVER_DivideAndConquerBatched(; fixgauge = default_fixgauge()) + +Algorithm type to denote the ROCSOLVER driver for computing the *batched* eigenvalue decomposition of a +set of Hermitian matrices, or the singular value decompositions of a set of general matrices using the +Divide and Conquer algorithm. +$_fixgauge_docs +""" +@algdef ROCSOLVER_DivideAndConquerBatched + # Various consts and unions # ------------------------- @@ -509,7 +611,6 @@ const CUSOLVER_SVDAlgorithm = Union{ CUSOLVER_QRIteration, CUSOLVER_SVDPolar, CUSOLVER_Jacobi, CUSOLVER_Randomized, } const GPU_SVDAlgorithm = Union{CUSOLVER_SVDAlgorithm, ROCSOLVER_SVDAlgorithm} - const LAPACK_EighAlgorithm = Union{ LAPACK_QRIteration, LAPACK_Bisection, @@ -524,7 +625,6 @@ const LAPACK_EigAlgorithm = Union{LAPACK_Simple, LAPACK_Expert} const CUSOLVER_EigAlgorithm = Union{CUSOLVER_Simple} const GPU_EigAlgorithm = Union{GPU_Simple} - # List of available algorithms - for docs and convenience purposes const SVDAlgorithms = Union{ SafeDivideAndConquer, diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index 5387252a3..8a30c03d8 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -45,15 +45,21 @@ 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 = (JacobiBatched(),) + TestSuite.test_svd_batched_algs(CuMatrix{T}, (m, n), batch_size, CUDA_SVD_ALGS) end # Randomized SVD: @@ -77,12 +83,15 @@ end # AMDGPU tests # ------------ if AMDGPU.functional() - # LAPACK algorithms: - for T in BLASFloats, m in (0, 23), n in (0, 17, m, 27) + # ROCSOLVER algorithms: + for T in BLASFloats, m in (23,), n in (17, m) 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) + AMD_SVD_ALGS = (QRIterationBatched(), JacobiBatched(), DivideAndConquerBatched(), BisectionBatched()) + 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..1b020de86 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 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 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 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,46 @@ 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 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 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 svd_vals(Ad; alg) + for (s, sd) in zip(eachslice(S, dims = 2), eachslice(Sd, dims = 2)) + @test s ≈ sd + end + end + end +end + function test_svd_full( T::Type, sz; atol::Real = 0, rtol::Real = precision(eltype(T)), @@ -120,6 +221,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 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 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 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 +293,45 @@ 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 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 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 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 + end +end + function test_svd_trunc( T::Type, sz; atol::Real = 0, rtol::Real = precision(eltype(T)), From 606c42bf5da090feddb2d93c95ba300a19412381 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sat, 22 Aug 2026 20:43:16 +0200 Subject: [PATCH 02/46] Go back to QRIteration default --- ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl | 6 +++--- test/decompositions/svd.jl | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index a47a5ea70..b00034d48 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -20,13 +20,13 @@ MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedROCArray{<:BlasF MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} = ROCSOLVER() function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCMatrix{<:BlasFloat}} - return DivideAndConquer(; kwargs...) + return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCArray{<:BlasFloat, 3}} - return DivideAndConquerBatched(; kwargs...) + return QRIterationBatched(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} - return DivideAndConquerBatched(; kwargs...) + return QRIterationBatched(; kwargs...) end function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}} return DivideAndConquer(; kwargs...) diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index 8a30c03d8..3e7c1d05e 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -84,7 +84,7 @@ end # ------------ if AMDGPU.functional() # ROCSOLVER algorithms: - for T in BLASFloats, m in (23,), n in (17, m) + 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(), DivideAndConquer(), Bisection()) From 57b5bc91da0b2a6c8b2bff9e052621cf2fa48b8d Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 23 Aug 2026 00:45:08 -0400 Subject: [PATCH 03/46] Fix default driver for CUDA --- ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl index a5df7079e..178022ff1 100644 --- a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl +++ b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl @@ -17,7 +17,7 @@ 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...) From 7ff48b39c378e0d03dde4843f5e7710e2ef672ca Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 23 Aug 2026 01:16:43 -0400 Subject: [PATCH 04/46] Fix typo --- ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl index 178022ff1..8e2bc3fc3 100644 --- a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl +++ b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl @@ -54,7 +54,7 @@ function gesvdj!(::CUSOLVER, A::StridedCuMatrix, S::StridedCuVector, U::StridedC 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_strided_batched!(As, Ss, Us, Vᴴs; kwargs...) + 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...) From e626f2ba55701e2c62e9e543a68d3b28790cd174 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 23 Aug 2026 01:25:37 -0400 Subject: [PATCH 05/46] And fix bad check in yacusolver --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 8f991c4cd..52b0086f2 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -317,7 +317,7 @@ for (bname, fname, elty, relty) in throw(DimensionMismatch("invalid column size of U or row size of Vᴴ")) end end - length(S) == minmn || + length(S) == minmn * batch_size || throw(DimensionMismatch("length mismatch between A and S")) Ṽ = (jobz == 'V') ? similar(Vᴴ') : similar(Vᴴ, (n, minmn)) From b68177983926c62aeb322c294c18616a89b1229e Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 23 Aug 2026 01:33:56 -0400 Subject: [PATCH 06/46] One more API fix --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 13 ++++--------- 1 file changed, 4 insertions(+), 9 deletions(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 52b0086f2..038d56fb1 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -302,18 +302,13 @@ for (bname, fname, elty, relty) in if length(U) == 0 && length(Vᴴ) == 0 jobz = 'N' - econ = 0 else 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 - econ = 1 - elseif size(U, 2) == m && size(Vᴴ, 1) == n - econ = 0 - else + 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 @@ -335,15 +330,15 @@ for (bname, fname, elty, relty) in function bufferSize() out = Ref{Cint}(0) cuSOLVER.$bname( - dh, jobz, econ, m, n, A, lda, S, Ũ, ldu, Ṽ, ldv, - out, params[] + 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, econ, m, n, A, lda, S, Ũ, ldu, Ṽ, ldv, + dh, jobz, m, n, A, lda, S, Ũ, ldu, Ṽ, ldv, buffer, sizeof(buffer) ÷ sizeof($elty), dh.info, params[], batch_size ) From 4e59aaff9431cb030c5afc4f09fc2d59d5ddfd0c Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 23 Aug 2026 02:26:51 -0400 Subject: [PATCH 07/46] Fix size of U and V bar --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 038d56fb1..91460d610 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -315,8 +315,8 @@ for (bname, fname, elty, relty) in length(S) == minmn * batch_size || throw(DimensionMismatch("length mismatch between A and S")) - Ṽ = (jobz == 'V') ? similar(Vᴴ') : similar(Vᴴ, (n, minmn)) - Ũ = (jobz == 'V') ? U : similar(U, (m, minmn)) + Ṽ = (jobz == 'V') ? similar(Vᴴ') : similar(Vᴴ, (n, minmn, batch_size)) + Ũ = (jobz == 'V') ? U : similar(U, (m, minmn, batch_size)) lda = max(1, stride(A, 2)) ldu = max(1, stride(Ũ, 2)) ldv = max(1, stride(Ṽ, 2)) From f8f46f2bb3e4fd7d6531ba892f1470c961e9c7bc Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 23 Aug 2026 09:44:02 -0400 Subject: [PATCH 08/46] More API fixes --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 91460d610..7c0472618 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -315,8 +315,9 @@ for (bname, fname, elty, relty) in length(S) == minmn * batch_size || throw(DimensionMismatch("length mismatch between A and S")) - Ṽ = (jobz == 'V') ? similar(Vᴴ') : similar(Vᴴ, (n, minmn, batch_size)) - Ũ = (jobz == 'V') ? U : similar(U, (m, minmn, batch_size)) + # these MUST be "full" sized + Ṽ = 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)) @@ -326,6 +327,7 @@ for (bname, fname, elty, relty) in cuSOLVER.cusolverDnXgesvdjSetTolerance(params[], tol) cuSOLVER.cusolverDnXgesvdjSetMaxSweeps(params[], max_sweeps) dh = cuSOLVER.dense_handle() + resize!(dh.info, batch_size) function bufferSize() out = Ref{Cint}(0) @@ -350,7 +352,8 @@ for (bname, fname, elty, relty) in cuSOLVER.cusolverDnDestroyGesvdjInfo(params[]) if jobz == 'V' - adjoint!(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 From e0f4f8bae24330af610ff1827cb1458c7ca5ad4e Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 14 Sep 2026 13:28:45 +0200 Subject: [PATCH 09/46] Remove redundant dummy one liners --- src/implementations/svd.jl | 7 ------- 1 file changed, 7 deletions(-) diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 18a54a7ba..e5fc2340c 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -241,13 +241,6 @@ for f! in (:gesdd!, :gesvd!, :gesvdj!, :gesvdp!, :gesvdx!, :gesvdr!, :gesdvd!, : @eval $f!(driver::Driver, args...) = throw(ArgumentError("$driver does not provide $($(f!))")) end -# declare these as dummies so the GPU extensions can import them safely -function gesvd_batched! end -function gesdd_batched! end -function gesvdj_batched! end -function gesvdx_batched! end - - """ svd_via_adjoint!(f!, driver, A, S, U, Vᴴ; kwargs...) From a02ce631f2b46f39dba223a39192485e204968c7 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 14 Sep 2026 09:17:48 -0400 Subject: [PATCH 10/46] Switch to batched_svd_f and get rid of Batched versions of algorithms --- .../MatrixAlgebraKitAMDGPUExt.jl | 12 +- .../MatrixAlgebraKitCUDAExt.jl | 13 ++- src/MatrixAlgebraKit.jl | 10 +- src/implementations/svd.jl | 68 ++++++------ src/interface/batched_svd.jl | 66 +++++++++++ src/interface/decompositions.jl | 104 +----------------- test/decompositions/svd.jl | 2 +- test/testsuite/decompositions/svd.jl | 24 ++-- 8 files changed, 144 insertions(+), 155 deletions(-) create mode 100644 src/interface/batched_svd.jl diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index b00034d48..a61ddbf98 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -32,11 +32,21 @@ function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T return DivideAndConquer(; kwargs...) end +function MatrixAlgebraKit.one!(A::StridedROCArray{T, 3}) where {T <: BlasFloat} + length(A) > 0 || return A + zero!(A) + # TODO use mapslices? + for a in eachslice(A, dims = 3) + diagview(a) .= one(eltype(a)) + end + return A +end + 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, :divide_and_conquer, :bisection, :qr_iteration_batched, :jacobi_batched, :divide_and_conquer_batched, :bisection_batched) +MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :divide_and_conquer, :bisection) function gesvd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) m, n = size(A) diff --git a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl index 8e2bc3fc3..a00e46a2b 100644 --- a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl +++ b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl @@ -23,7 +23,7 @@ function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T < return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedCuArray{<:BlasFloat, 3}} - return JacobiBatched(; kwargs...) + return Jacobi(; kwargs...) end function MatrixAlgebraKit.default_eig_algorithm(::Type{T}; kwargs...) where {T <: StridedCuVecOrMat{<:BlasFloat}} return QRIteration(; kwargs...) @@ -32,6 +32,15 @@ function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T return DivideAndConquer(; kwargs...) end +function MatrixAlgebraKit.one!(A::StridedCuArray{T, 3}) where {T <: BlasFloat} + length(A) > 0 || return A + zero!(A) + # TODO use mapslices? + for a in eachslice(A, dims = 3) + diagview(a) .= one(eltype(a)) + end + return A +end for f in (:geqrf!, :ungqr!, :unmqr!) @eval $f(::CUSOLVER, args...) = YACUSOLVER.$f(args...) @@ -39,7 +48,7 @@ end MatrixAlgebraKit.prefers_ungqr(::CUSOLVER) = true -MatrixAlgebraKit.supports_svd_full(::CUSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :svd_polar, :jacobi_batched) +MatrixAlgebraKit.supports_svd_full(::CUSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :svd_polar) function gesvd!(::CUSOLVER, A::StridedCuMatrix, S::StridedCuVector, U::StridedCuMatrix, Vᴴ::StridedCuMatrix; kwargs...) m, n = size(A) diff --git a/src/MatrixAlgebraKit.jl b/src/MatrixAlgebraKit.jl index 59f9cf42e..89e2c68a6 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 @@ -34,7 +36,7 @@ export exponential, exponential! export Householder, Native_HouseholderQR, Native_HouseholderLQ export DivideAndConquer, SafeDivideAndConquer, QRIteration, Bisection, Jacobi, SVDViaPolar -export RobustRepresentations, DivideAndConquerBatched, QRIterationBatched, BisectionBatched, JacobiBatched +export RobustRepresentations export LAPACK_HouseholderQR, LAPACK_HouseholderLQ, LAPACK_Simple, LAPACK_Expert, LAPACK_QRIteration, LAPACK_Bisection, LAPACK_MultipleRelativelyRobustRepresentations, LAPACK_DivideAndConquer, LAPACK_Jacobi, LAPACK_SafeDivideAndConquer @@ -46,10 +48,9 @@ export DefaultAlgorithm export DiagonalAlgorithm export NativeBlocked export CUSOLVER_Simple, CUSOLVER_HouseholderQR, CUSOLVER_QRIteration, CUSOLVER_SVDPolar, - CUSOLVER_Jacobi, CUSOLVER_Randomized, CUSOLVER_DivideAndConquer, CUSOLVER_JacobiBatched + CUSOLVER_Jacobi, CUSOLVER_Randomized, CUSOLVER_DivideAndConquer export ROCSOLVER_HouseholderQR, ROCSOLVER_QRIteration, ROCSOLVER_Jacobi, - ROCSOLVER_DivideAndConquer, ROCSOLVER_Bisection, ROCSOLVER_QRIterationBatched, ROCSOLVER_JacobiBatched, - ROCSOLVER_DivideAndConquerBatched, ROCSOLVER_BisectionBatched + ROCSOLVER_DivideAndConquer, ROCSOLVER_Bisection export notrunc, truncrank, trunctol, truncerror, truncfilter @@ -109,6 +110,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") diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index e5fc2340c..336c4366b 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -1,7 +1,9 @@ # Input # ------ -copy_input(::typeof(svd_full), As::AbstractVector{<:AbstractMatrix}) = map(A -> copy!(similar(A, float(eltype(A))), A), As) -copy_input(::typeof(svd_full), A::AbstractArray{T, 3}) where {T} = copy!(similar(A, float(T)), A) +copy_input(::typeof(batched_svd_full), As::AbstractVector{<:AbstractMatrix}) = map(A -> copy!(similar(A, float(eltype(A))), A), 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) copy_input(::typeof(svd_full), A::AbstractMatrix) = copy!(similar(A, float(eltype(A))), A) copy_input(::typeof(svd_compact), A) = copy_input(svd_full, A) copy_input(::typeof(svd_vals), A) = copy_input(svd_full, A) @@ -45,7 +47,7 @@ function check_input(::typeof(svd_vals!), A::AbstractMatrix, S, ::AbstractAlgori end # batched varieties -function check_input(::typeof(svd_full!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) +function check_input(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) batch_size = length(A) @@ -59,7 +61,7 @@ function check_input(::typeof(svd_full!), A::AbstractVector{<:AbstractMatrix}, U @check_scalar(Vᴴ, first(A)) return nothing end -function check_input(::typeof(svd_compact!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) +function check_input(::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) batch_size = length(A) @@ -74,7 +76,7 @@ function check_input(::typeof(svd_compact!), A::AbstractVector{<:AbstractMatrix} @check_scalar(Vᴴ, first(A)) return nothing end -function check_input(::typeof(svd_vals!), A::AbstractVector{<:AbstractMatrix}, S, ::AbstractAlgorithm) +function check_input(::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, S, ::AbstractAlgorithm) @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) batch_size = length(A) @@ -84,7 +86,7 @@ function check_input(::typeof(svd_vals!), A::AbstractVector{<:AbstractMatrix}, S @check_scalar(S, first(A), real) return nothing end -function check_input(::typeof(svd_full!), A::AbstractArray{T, 3}, USVᴴ, ::AbstractAlgorithm) where {T} +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 @@ -96,7 +98,7 @@ function check_input(::typeof(svd_full!), A::AbstractArray{T, 3}, USVᴴ, ::Abst @check_scalar(Vᴴ, A) return nothing end -function check_input(::typeof(svd_compact!), A::AbstractArray{T, 3}, USVᴴ, ::AbstractAlgorithm) where {T} +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ᴴ @@ -109,7 +111,7 @@ function check_input(::typeof(svd_compact!), A::AbstractArray{T, 3}, USVᴴ, ::A @check_scalar(Vᴴ, A) return nothing end -function check_input(::typeof(svd_vals!), A::AbstractArray{T, 3}, S, ::AbstractAlgorithm) where {T} +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 @@ -169,21 +171,21 @@ function initialize_output(::Union{typeof(svd_trunc!), typeof(svd_trunc_no_error return initialize_output(svd_compact!, A, alg.alg) end # batched versions -function initialize_output(::typeof(svd_full!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) +function initialize_output(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) m, n = size(first(A)) U = similar(first(A), (m, m, length(A))) S = similar(first(A), real(eltype(first(A))), (m, n, length(A))) # TODO: Rectangular diagonal type? Vᴴ = similar(first(A), (n, n, length(A))) return (U, S, Vᴴ) end -function initialize_output(::typeof(svd_full!), A::AbstractArray{T, 3}, ::AbstractAlgorithm) where {T} +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(svd_compact!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) +function initialize_output(::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) minmn = min(m, n) @@ -192,7 +194,7 @@ function initialize_output(::typeof(svd_compact!), A::AbstractVector{<:AbstractM Vᴴ = similar(first(A), (minmn, n, length(A))) return (U, S, Vᴴ) end -function initialize_output(::typeof(svd_compact!), A::AbstractArray{T, 3}, ::AbstractAlgorithm) where {T} +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)) @@ -200,12 +202,12 @@ function initialize_output(::typeof(svd_compact!), A::AbstractArray{T, 3}, ::Abs Vᴴ = similar(A, (minmn, n, batch_size)) return (U, S, Vᴴ) end -function initialize_output(::typeof(svd_vals!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) +function initialize_output(::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) return similar(first(A), real(eltype(first(A))), (min(m, n), length(A))) end -function initialize_output(::typeof(svd_vals!), A::AbstractArray{T, 3}, ::AbstractAlgorithm) where {T} +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 @@ -345,39 +347,39 @@ end # batched varieties for (f, f_lapack!, Alg) in ( - (:divide_and_conquer_batched, :gesdd_batched!, :DivideAndConquerBatched), - (:qr_iteration_batched, :gesvd_batched!, :QRIterationBatched), - (:bisection_batched, :gesvdx_batched!, :BisectionBatched), - (:jacobi_batched, :gesvdj_batched!, :JacobiBatched), + (:divide_and_conquer, :gesdd_batched!, :DivideAndConquer), + (:qr_iteration, :gesvd_batched!, :QRIteration), + (:bisection, :gesvdx_batched!, :Bisection), + (:jacobi, :gesvdj_batched!, :Jacobi), ) - svd_compact_f! = Symbol(:svd_compact_, f, :!) - svd_full_f! = Symbol(:svd_full_, f, :!) - svd_vals_f! = Symbol(:svd_vals_, f, :!) + 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 svd_compact!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) - check_input(svd_compact!, A, USVᴴ, alg) + function batched_svd_compact!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) + check_input(batched_svd_compact!, A, USVᴴ, alg) return $svd_compact_f!(A, USVᴴ...; alg.kwargs...) end - function svd_compact!(A::AbstractArray{T, 3}, USVᴴ, alg::$Alg) where {T} - check_input(svd_compact!, A, USVᴴ, alg) + 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 svd_full!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) - check_input(svd_full!, A, USVᴴ, alg) + function batched_svd_full!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) + check_input(batched_svd_full!, A, USVᴴ, alg) return $svd_full_f!(A, USVᴴ...; alg.kwargs...) end - function svd_full!(A::AbstractArray{T, 3}, USVᴴ, alg::$Alg) where {T} - check_input(svd_full!, A, USVᴴ, alg) + 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 svd_vals!(A::AbstractVector{<:AbstractMatrix}, S, alg::$Alg) - check_input(svd_vals!, A, S, alg) + function batched_svd_vals!(A::AbstractVector{<:AbstractMatrix}, S, alg::$Alg) + check_input(batched_svd_vals!, A, S, alg) return $svd_vals_f!(A, S; alg.kwargs...) end - function svd_vals!(A::AbstractArray{T, 3}, S, alg::$Alg) where {T} - check_input(svd_vals!, A, S, alg) + 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 end 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/interface/decompositions.jl b/src/interface/decompositions.jl index 90999fb3c..1d423564a 100644 --- a/src/interface/decompositions.jl +++ b/src/interface/decompositions.jl @@ -100,17 +100,6 @@ The optional `driver` keyword can be used to choose between different implementa """ @algdef DivideAndConquer -""" - DivideAndConquerBatched(; [driver], fixgauge = default_fixgauge()) - -Algorithm type for computing the *batched* eigenvalue decompositions of a set of Hermitian matrices, -or the singular value decompositions of a set of general matrices using the divide-and-conquer algorithm. - -$_fixgauge_docs -The optional `driver` keyword can be used to choose between different implementations of this algorithm. -""" -@algdef DivideAndConquerBatched - """ SafeDivideAndConquer(; [driver], fixgauge = default_fixgauge()) @@ -173,49 +162,6 @@ The optional `driver` keyword can be used to choose between different implementa """ @algdef Jacobi -""" - QRIterationBatched(; [driver], fixgauge = default_fixgauge(), kwargs...) - -Algorithm type for computing the *batched* eigenvalue, Schur or singular value decompositions of a set of matrices via QR iteration. - -## Keyword arguments - -Various customizations are available, depending on the type of decomposition this algorithm is used for. - -Schur decompositions are not yet supported. - -For non-Hermitian eigenvalue decompositions there is `permute = true` and `scale = true` to control whether -or not to balance the input matrix before starting the QR iterations. - -For the singular value and eigenvalue decompositions, there is residual freedom in the outputs that can be resolved. -$_fixgauge_docs - -In all cases, the optional `driver` keyword can be used to choose between different implementations of this algorithm. -""" -@algdef QRIterationBatched - -""" - BisectionBatched(; [driver], fixgauge = default_fixgauge()) - -Algorithm type for computing the *batched* eigenvalue decompositions of a set of Hermitian matrices -via the bisection algorithm, or the singular value decompositions of a set of general matrices. - -$_fixgauge_docs -The optional `driver` keyword can be used to choose between different implementations of this algorithm. -""" -@algdef BisectionBatched - -""" - JacobiBatched(; [driver], fixgauge = default_fixgauge()) - -Algorithm type for computing the *batched* eigenvalue decompositions of a set of Hermitian matrices, -or the singular value decompositions of a set of general matrices using the Jacobi algorithm. - -$_fixgauge_docs -The optional `driver` keyword can be used to choose between different implementations of this algorithm. -""" -@algdef JacobiBatched - """ RobustRepresentations(; [driver], fixgauge = default_fixgauge()) @@ -451,15 +397,6 @@ $_fixgauge_docs """ @algdef CUSOLVER_Jacobi -""" - CUSOLVER_JacobiBatched(; fixgauge = default_fixgauge()) - -Algorithm type to denote the CUSOLVER driver for computing the *batched singular value decompositions -of a set of general matrices using the Jacobi algorithm. -$_fixgauge_docs -""" -@algdef CUSOLVER_JacobiBatched - """ CUSOLVER_Randomized(; k, p, niters) @@ -548,45 +485,6 @@ $_fixgauge_docs """ @algdef ROCSOLVER_DivideAndConquer -""" - ROCSOLVER_QRIterationBatched(; fixgauge = default_fixgauge()) - -Algorithm type to denote the ROCSOLVER driver for computing the *batched* eigenvalue decompositions of a -set of Hermitian matrices, or the singular value decompositions of a set of general matrices using the -QR Iteration algorithm. -$_fixgauge_docs -""" -@algdef ROCSOLVER_QRIterationBatched - -""" - ROCSOLVER_JacobiBatched(; fixgauge = default_fixgauge()) - -Algorithm type to denote the ROCSOLVER driver for computing the *batched* singular value decompositions of -a set of general matrices using the Jacobi algorithm. -$_fixgauge_docs -""" -@algdef ROCSOLVER_JacobiBatched - -""" - ROCSOLVER_BisectionBatched(; fixgauge = default_fixgauge()) - -Algorithm type to denote the ROCSOLVER driver for computing the *batched* eigenvalue decomposition of a -set of Hermitian matrices, or the singular value decompositions of a set of general matrices using the -Bisection algorithm. -$_fixgauge_docs -""" -@algdef ROCSOLVER_BisectionBatched - -""" - ROCSOLVER_DivideAndConquerBatched(; fixgauge = default_fixgauge()) - -Algorithm type to denote the ROCSOLVER driver for computing the *batched* eigenvalue decomposition of a -set of Hermitian matrices, or the singular value decompositions of a set of general matrices using the -Divide and Conquer algorithm. -$_fixgauge_docs -""" -@algdef ROCSOLVER_DivideAndConquerBatched - # Various consts and unions # ------------------------- @@ -611,6 +509,7 @@ const CUSOLVER_SVDAlgorithm = Union{ CUSOLVER_QRIteration, CUSOLVER_SVDPolar, CUSOLVER_Jacobi, CUSOLVER_Randomized, } const GPU_SVDAlgorithm = Union{CUSOLVER_SVDAlgorithm, ROCSOLVER_SVDAlgorithm} + const LAPACK_EighAlgorithm = Union{ LAPACK_QRIteration, LAPACK_Bisection, @@ -625,6 +524,7 @@ const LAPACK_EigAlgorithm = Union{LAPACK_Simple, LAPACK_Expert} const CUSOLVER_EigAlgorithm = Union{CUSOLVER_Simple} const GPU_EigAlgorithm = Union{GPU_Simple} + # List of available algorithms - for docs and convenience purposes const SVDAlgorithms = Union{ SafeDivideAndConquer, diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index 3e7c1d05e..777fd97b2 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -58,7 +58,7 @@ if CUDA.functional() TestSuite.test_svd_algs(CuMatrix{T}, (m, n), CUDA_SVD_ALGS) TestSuite.test_svd_batched(CuMatrix{T}, (m, n), batch_size) - CUDA_SVD_ALGS = (JacobiBatched(),) + CUDA_SVD_ALGS = (Jacobi(),) TestSuite.test_svd_batched_algs(CuMatrix{T}, (m, n), batch_size, CUDA_SVD_ALGS) end diff --git a/test/testsuite/decompositions/svd.jl b/test/testsuite/decompositions/svd.jl index 1b020de86..f895413d6 100644 --- a/test/testsuite/decompositions/svd.jl +++ b/test/testsuite/decompositions/svd.jl @@ -86,7 +86,7 @@ function test_svd_compact_batched( Ac = deepcopy(Ad) m, n = size(first(As)) minmn = min(m, n) - U, S, Vᴴ = @testinferred svd_compact(Ad) + 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) @@ -98,7 +98,7 @@ function test_svd_compact_batched( end Sc = similar(diagview(S)) - U2, S2, V2ᴴ = @testinferred svd_compact!(Ac, (U, S, Vᴴ)) + 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) @@ -107,7 +107,7 @@ function test_svd_compact_batched( end if test_vals - Sd = @testinferred svd_vals(Ad) + Sd = @testinferred batched_svd_vals(Ad) for (s, sd) in zip(eachslice(S, dims = 2), eachslice(Sd, dims = 2)) @test s ≈ sd end @@ -160,7 +160,7 @@ function test_svd_compact_algs_batched( Ac = deepcopy(Ad) m, n = size(first(As)) minmn = min(m, n) - U, S, Vᴴ = @testinferred svd_compact(Ad; alg) + 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) @@ -171,7 +171,7 @@ function test_svd_compact_algs_batched( @test isposdef(Diagonal(s)) end - U2, S2, V2ᴴ = @testinferred svd_compact!(Ac, (U, S, Vᴴ); alg) + 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) @@ -180,7 +180,7 @@ function test_svd_compact_algs_batched( end if test_vals - Sd = @testinferred svd_vals(Ad; alg) + Sd = @testinferred batched_svd_vals(Ad; alg) for (s, sd) in zip(eachslice(S, dims = 2), eachslice(Sd, dims = 2)) @test s ≈ sd end @@ -233,7 +233,7 @@ function test_svd_full_batched( Ac = deepcopy(Ad) m, n = size(first(As)) minmn = min(m, n) - U, S, Vᴴ = @testinferred svd_full(Ad) + 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) @@ -244,7 +244,7 @@ function test_svd_full_batched( @test all(isposdef, diagview(s)) end - U2, S2, V2ᴴ = @testinferred svd_full!(Ac, (U, S, Vᴴ)) + 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) @@ -253,7 +253,7 @@ function test_svd_full_batched( end Sc = similar(first(As), real(eltype(T)), min(m, n), batch_size) - Sc2 = @testinferred svd_vals!(copy!(Ac, Ad), Sc) + 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 @@ -305,7 +305,7 @@ function test_svd_full_algs_batched( Ac = deepcopy(Ad) m, n = size(first(As)) minmn = min(m, n) - U, S, Vᴴ = @testinferred svd_full(Ad; alg) + 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) @@ -316,7 +316,7 @@ function test_svd_full_algs_batched( @test all(isposdef, diagview(s)) end - U2, S2, V2ᴴ = @testinferred svd_full!(Ac, (U, S, Vᴴ); alg) + 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) @@ -325,7 +325,7 @@ function test_svd_full_algs_batched( end Sc = similar(first(As), real(eltype(T)), min(m, n), batch_size) - Sc2 = @testinferred svd_vals!(copy!(Ac, Ad), Sc; alg) + 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 From e8834e7c53093264251bf32503b1e385fa93e87d Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 14 Sep 2026 14:32:58 -0400 Subject: [PATCH 11/46] Fixup AMD default algos --- ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index a61ddbf98..93bd99fb5 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -23,10 +23,10 @@ function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T < return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCArray{<:BlasFloat, 3}} - return QRIterationBatched(; kwargs...) + return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} - return QRIterationBatched(; kwargs...) + return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}} return DivideAndConquer(; kwargs...) From 721c09ebdba298e7d95f3023be4aa1ecefd74674 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 15 Sep 2026 05:55:39 -0400 Subject: [PATCH 12/46] Move ragged batch handling and cutoffs etc into MAK --- .../MatrixAlgebraKitCUDAExt.jl | 3 + src/MatrixAlgebraKit.jl | 1 + src/implementations/batched_svd.jl | 351 ++++++++++++++++++ src/implementations/svd.jl | 229 +----------- test/testsuite/decompositions/svd.jl | 24 ++ 5 files changed, 390 insertions(+), 218 deletions(-) create mode 100644 src/implementations/batched_svd.jl diff --git a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl index a00e46a2b..bf6a043a7 100644 --- a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl +++ b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl @@ -50,6 +50,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, ::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ᴴ) diff --git a/src/MatrixAlgebraKit.jl b/src/MatrixAlgebraKit.jl index 89e2c68a6..b56548fdd 100644 --- a/src/MatrixAlgebraKit.jl +++ b/src/MatrixAlgebraKit.jl @@ -124,6 +124,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/implementations/batched_svd.jl b/src/implementations/batched_svd.jl new file mode 100644 index 000000000..a7f6ddfd2 --- /dev/null +++ b/src/implementations/batched_svd.jl @@ -0,0 +1,351 @@ +# Inputs +# ------ +copy_input(::typeof(batched_svd_full), As::AbstractVector{<:AbstractMatrix}) = map(A -> copy!(similar(A, float(eltype(A))), A), 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) + +function check_input(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(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) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(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) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(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 +# 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{<:AbstractVector}, 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, Diagonal(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::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 +# ------- +function initialize_output(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) + m, n = size(first(A)) + U = similar(first(A), (m, m, length(A))) + S = similar(first(A), real(eltype(first(A))), (m, n, length(A))) # TODO: Rectangular diagonal type? + Vᴴ = similar(first(A), (n, n, length(A))) + return (U, S, Vᴴ) +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) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(A)) + minmn = min(m, n) + U = similar(first(A), (m, minmn, length(A))) + S = similar(first(A), real(eltype(first(A))), minmn, length(A)) + Vᴴ = similar(first(A), (minmn, n, length(A))) + return (U, S, Vᴴ) +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) + @assert all(==(size(first(A))), size.(A)) + m, n = size(first(A)) + return similar(first(A), real(eltype(first(A))), (min(m, n), length(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 + +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) + 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) + 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) + 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{<:AbstractVector}, 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) + Ub, Sb, Vᴴb = batched_svd_compact!(Ab, initialize_output(batched_svd_compact!, Ab, alg), alg) + for (j, i) in enumerate(inds) + copyto!(Us[i], view(Ub, axes(Us[i])..., j)) + copyto!(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], Diagonal(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) + Sb = batched_svd_vals!(Ab, initialize_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...) + isempty(A) && return one!(U), zero!(S), one!(Vᴴ) + $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) + if fixgauge + for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) + gaugefix!(svd_compact!, u, vᴴ) + end + end + 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!)))`"))) + isempty(A) && return one!(U), zero!(S), one!(Vᴴ) + 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...) + for (s, sd) in zip(eachslice(S, dims = 3), eachslice(Sd, dims = 2)) + diagview(s) .= sd + end + if fixgauge + for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) + gaugefix!(svd_full!, u, vᴴ) + end + end + return U, S, Vᴴ + end + function $svd_vals_f!(driver::Driver, A::AbstractArray{T, 3}, S::AbstractMatrix; fixgauge::Bool = true, kwargs...) where {T} + isempty(A) && return zero!(S) + U, Vᴴ = similar(A, (0, 0, 0)), similar(A, (0, 0, 0)) + $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(A) && return zero!(S) + U, Vᴴ = similar(first(A), (0, 0, 0)), similar(first(A), (0, 0, 0)) + $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) + return S + end + end +end + +# Ragged batches +# -------------- +""" + max_batched_blocksize(alg, T::Type) -> Int + +Largest matrix dimension that the batched driver for `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(::AbstractAlgorithm, ::Type) = typemax(Int) + +# 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 + +# Split a ragged batch into batches the driver can handle: matrices of equal size are +# batched together, and 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) + batches = Tuple{Vector{Int}, Tuple{Int, Int}}[] + rest = Int[] + isempty(A) && return batches, rest + lim = max_batched_blocksize(alg, 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) <= lim && (!needs_tall || m >= n) + push!(batches, (inds, (m, n))) + else + append!(rest, inds) + end + end + # 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) <= lim + push!(batches, (rest, padded)) + rest = Int[] + end + return batches, rest +end + +# Copy `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) + uniform = all(i -> size(A[i]) == (m, n), 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/implementations/svd.jl b/src/implementations/svd.jl index 336c4366b..8f4b7e445 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -1,9 +1,5 @@ # Input # ------ -copy_input(::typeof(batched_svd_full), As::AbstractVector{<:AbstractMatrix}) = map(A -> copy!(similar(A, float(eltype(A))), A), 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) copy_input(::typeof(svd_full), A::AbstractMatrix) = copy!(similar(A, float(eltype(A))), A) copy_input(::typeof(svd_compact), A) = copy_input(svd_full, A) copy_input(::typeof(svd_vals), A) = copy_input(svd_full, A) @@ -46,80 +42,6 @@ function check_input(::typeof(svd_vals!), A::AbstractMatrix, S, ::AbstractAlgori return nothing end -# batched varieties -function check_input(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) - @assert all(==(size(first(A))), size.(A)) - m, n = size(first(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) - @assert all(==(size(first(A))), size.(A)) - m, n = size(first(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) - @assert all(==(size(first(A))), size.(A)) - m, n = size(first(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 - function check_input(::typeof(svd_full!), A::AbstractMatrix, USVᴴ, ::DiagonalAlgorithm) m, n = size(A) @assert m == n && isdiag(A) @@ -170,47 +92,6 @@ end function initialize_output(::Union{typeof(svd_trunc!), typeof(svd_trunc_no_error!)}, A, alg::TruncatedAlgorithm) return initialize_output(svd_compact!, A, alg.alg) end -# batched versions -function initialize_output(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) - m, n = size(first(A)) - U = similar(first(A), (m, m, length(A))) - S = similar(first(A), real(eltype(first(A))), (m, n, length(A))) # TODO: Rectangular diagonal type? - Vᴴ = similar(first(A), (n, n, length(A))) - return (U, S, Vᴴ) -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) - @assert all(==(size(first(A))), size.(A)) - m, n = size(first(A)) - minmn = min(m, n) - U = similar(first(A), (m, minmn, length(A))) - S = similar(first(A), real(eltype(first(A))), minmn, length(A)) - Vᴴ = similar(first(A), (minmn, n, length(A))) - return (U, S, Vᴴ) -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) - @assert all(==(size(first(A))), size.(A)) - m, n = size(first(A)) - return similar(first(A), real(eltype(first(A))), (min(m, n), length(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 function initialize_output(::typeof(svd_full!), A::Diagonal, ::DiagonalAlgorithm) TA = eltype(A) @@ -239,7 +120,7 @@ end # IMPLEMENTATIONS # ========================== -for f! in (:gesdd!, :gesvd!, :gesvdj!, :gesvdp!, :gesvdx!, :gesvdr!, :gesdvd!, :gesdd_batched!, :gesvd_batched!, :gesvdj_batched!, :gesvdx_batched!) +for f! in (:gesdd!, :gesvd!, :gesvdj!, :gesvdp!, :gesvdx!, :gesvdr!, :gesdvd!) @eval $f!(driver::Driver, args...) = throw(ArgumentError("$driver does not provide $($(f!))")) end @@ -345,104 +226,6 @@ for (f, f_lapack!, Alg) in ( end end -# batched varieties -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) - 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) - 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) - 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 - 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...) - isempty(A) && return one!(U), zero!(S), one!(Vᴴ) - $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) - if fixgauge - for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) - gaugefix!(svd_compact!, u, vᴴ) - end - end - 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!)))`"))) - isempty(A) && return one!(U), zero!(S), one!(Vᴴ) - 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...) - for (s, sd) in zip(eachslice(S, dims = 3), eachslice(Sd, dims = 2)) - diagview(s) .= sd - end - if fixgauge - for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) - gaugefix!(svd_full!, u, vᴴ) - end - end - return U, S, Vᴴ - end - function $svd_vals_f!(driver::Driver, A::AbstractArray{T, 3}, S::AbstractMatrix; fixgauge::Bool = true, kwargs...) where {T} - isempty(A) && return zero!(S) - U, Vᴴ = similar(A, (0, 0, 0)), similar(A, (0, 0, 0)) - $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(A) && return zero!(S) - U, Vᴴ = similar(first(A), (0, 0, 0)), similar(first(A), (0, 0, 0)) - $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) - return S - end - end -end - supports_svd_full(::Driver, ::Symbol) = false supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration, :bisection) @@ -466,6 +249,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/test/testsuite/decompositions/svd.jl b/test/testsuite/decompositions/svd.jl index f895413d6..0fdd6fa72 100644 --- a/test/testsuite/decompositions/svd.jl +++ b/test/testsuite/decompositions/svd.jl @@ -185,6 +185,30 @@ function test_svd_compact_algs_batched( @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 = [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) + for (a, u, s, vᴴ) in zip(Ar, U3, S3, V3ᴴ) + @test u * Diagonal(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(s) ≈ collect(s2) + end + end + end end end From dbedd55cda9f18bd9c84df4e658fcf2ad7cb5a84 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 15 Sep 2026 13:59:39 +0200 Subject: [PATCH 13/46] Add support for a batched svd_via_adjoint for algos requiring tall matrices --- .../MatrixAlgebraKitAMDGPUExt.jl | 15 +++++++--- src/implementations/batched_svd.jl | 30 +++++++++++++++++++ 2 files changed, 41 insertions(+), 4 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 93bd99fb5..d06049203 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -66,10 +66,17 @@ function gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid return S, U, Vᴴ end -gesvd_batched!(::ROCSOLVER, As::Vector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = - YArocSOLVER.gesvd_batched!(As, Ss, Us, Vᴴs; kwargs...) -gesvd_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = - YArocSOLVER.gesvd_strided_batched!(As, Ss, Us, Vᴴs; kwargs...) +# rocSOLVER's batched `gesvd` requires m ≥ n, so wide matrices go through the adjoint +function gesvd_batched!(::ROCSOLVER, As::Vector{<: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::Vector{<: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...) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index a7f6ddfd2..fe4d7661f 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -149,6 +149,36 @@ for f! in (:gesdd_batched!, :gesvd_batched!, :gesvdj_batched!, :gesvdx_batched!) @eval $f!(driver::Driver, args...) = throw(ArgumentError("$driver does not provide $($(f!))")) end +# 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) + +""" + 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), From 8e4bcc42f29228321843b2b6583663e7e0a716b2 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 15 Sep 2026 13:41:39 -0400 Subject: [PATCH 14/46] Add batched svd_full too --- src/implementations/batched_svd.jl | 71 ++++++++++++++++++++-------- test/testsuite/decompositions/svd.jl | 43 +++++++++++++++++ 2 files changed, 95 insertions(+), 19 deletions(-) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index fe4d7661f..f3b7e7e84 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -58,6 +58,19 @@ function check_input( 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 @@ -104,12 +117,12 @@ 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) - m, n = size(first(A)) - U = similar(first(A), (m, m, length(A))) - S = similar(first(A), real(eltype(first(A))), (m, n, length(A))) # TODO: Rectangular diagonal type? - Vᴴ = similar(first(A), (n, n, length(A))) - return (U, S, Vᴴ) + 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) @@ -119,13 +132,10 @@ function initialize_output(::typeof(batched_svd_full!), A::AbstractArray{T, 3}, return (U, S, Vᴴ) end function initialize_output(::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) - @assert all(==(size(first(A))), size.(A)) - m, n = size(first(A)) - minmn = min(m, n) - U = similar(first(A), (m, minmn, length(A))) - S = similar(first(A), real(eltype(first(A))), minmn, length(A)) - Vᴴ = similar(first(A), (minmn, n, length(A))) - return (U, S, Vᴴ) + Us = [similar(a, (size(a, 1), minimum(size(a)))) for a in A] + Ss = [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) @@ -136,9 +146,7 @@ function initialize_output(::typeof(batched_svd_compact!), A::AbstractArray{T, 3 return (U, S, Vᴴ) end function initialize_output(::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) - @assert all(==(size(first(A))), size.(A)) - m, n = size(first(A)) - return similar(first(A), real(eltype(first(A))), (min(m, n), length(A))) + 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) @@ -239,6 +247,30 @@ for (f, f_lapack!, Alg) in ( 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) + Ub, Sb, Vᴴb = batched_svd_full!(Ab, initialize_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 ) @@ -332,11 +364,11 @@ max_batched_blocksize(::AbstractAlgorithm, ::Type) = typemax(Int) const BATCHED_SVD_THRESHOLD::Int = 4 # Split a ragged batch into batches the driver can handle: matrices of equal size are -# batched together, and 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. +# 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) +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 @@ -353,6 +385,7 @@ function _ragged_batches(A::AbstractVector{<:AbstractMatrix}, alg::AbstractAlgor 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`). diff --git a/test/testsuite/decompositions/svd.jl b/test/testsuite/decompositions/svd.jl index 0fdd6fa72..cf53878b2 100644 --- a/test/testsuite/decompositions/svd.jl +++ b/test/testsuite/decompositions/svd.jl @@ -201,12 +201,23 @@ function test_svd_compact_algs_batched( @test isisometric(vᴴ; side = :right) end + U4, S4, V4ᴴ = @testinferred batched_svd_compact(Ar; alg) + for (a, u, s, vᴴ) in zip(Ar, U4, S4, V4ᴴ) + @test u * Diagonal(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(s) ≈ collect(s2) end + Sv3 = @testinferred batched_svd_vals(Ar; alg) + for (s, s3) in zip(S3, Sv3) + @test collect(s) ≈ collect(s3) + end end end end @@ -353,6 +364,38 @@ function test_svd_full_algs_batched( 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 From dc48639ce9b8e4b7fee2ce3588f7ef7c2d6df161 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 16 Sep 2026 09:57:19 +0200 Subject: [PATCH 15/46] Add special path for AMDGPU, implement gesvdx batching there, and fix some errors --- .../MatrixAlgebraKitAMDGPUExt.jl | 13 +- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 125 +++++++++++++++++- src/implementations/batched_svd.jl | 62 +++++++-- 3 files changed, 178 insertions(+), 22 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index d06049203..252efecaf 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -48,6 +48,11 @@ end 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, ::Type{<:StridedROCMatrix{<:BlasFloat}}) = true + function gesvd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) m, n = size(A) m >= n && return YArocSOLVER.gesvd!(A, S, U, Vᴴ) @@ -67,7 +72,7 @@ function gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid end # rocSOLVER's batched `gesvd` requires m ≥ n, so wide matrices go through the adjoint -function gesvd_batched!(::ROCSOLVER, As::Vector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} +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...) @@ -78,17 +83,17 @@ function gesvd_batched!(::ROCSOLVER, As::StridedROCArray{T, 3}, Ss::StridedROCMa return MatrixAlgebraKit.batched_svd_via_adjoint!(gesvd_batched!, ROCSOLVER(), As, Ss, Us, Vᴴs; kwargs...) end -gesdd_batched!(::ROCSOLVER, As::Vector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = +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::Vector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = +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::Vector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} = +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...) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 5d43b061a..d4794eb40 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -107,7 +107,7 @@ for (fname, elty, relty) in ) @eval begin function gesvd_batched!( - A::StridedROCVector{<:StridedROCMatrix{$elty}}, + 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)), @@ -166,7 +166,7 @@ for (fname, elty, relty) in E = ROCArray{$relty}(undef, length(A) * strideE) dh = rocBLAS.handle() dev_info = ROCVector{Cint}(undef, length(A)) - pA = map(pointer, A) + pA = ROCVector(map(pointer, A)) rocSOLVER.$fname( dh, jobu, jobvt, m, n, pA, lda, S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, @@ -350,7 +350,7 @@ for (fname, elty, relty) in ) @eval begin function gesdd_batched!( - A::StridedROCVector{<:StridedROCMatrix{$elty}}, + 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)), @@ -406,7 +406,7 @@ for (fname, elty, relty) in dh = rocBLAS.handle() dev_info = ROCVector{Cint}(undef, length(A)) - pA = map(pointer, A) + pA = ROCVector(map(pointer, A)) rocSOLVER.$fname( dh, jobu, jobvt, m, n, pA, lda, S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, @@ -591,7 +591,7 @@ for (fname, elty, relty) in ) @eval begin function gesvdj_batched!( - A::StridedROCVector{<:StridedROCMatrix{$elty}}, + 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)), @@ -652,7 +652,7 @@ for (fname, elty, relty) in dev_n_sweeps = ROCVector{Cint}(undef, length(A)) dh = rocBLAS.handle() - pA = map(pointer, A) + pA = ROCVector(map(pointer, A)) rocSOLVER.$fname( dh, jobu, jobvt, m, n, pA, lda, tol, dev_residual, max_sweeps, dev_n_sweeps, @@ -848,6 +848,119 @@ for (fname, elty, relty) in 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)); + kwargs... + ) + for A_ in A + chkstride1(A_, U, Vᴴ, S) + end + m, n = size(first(A)) + minmn = min(m, n) + batch_count = 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(A, U, Vᴴ, m, n, maxnsv) + length(S) == minmn * batch_count || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(first(A), 2)) + ldu = max(1, stride(U, 2)) + strideU = ldu * size(U, 2) + ldv = max(1, stride(Vᴴ, 2)) + strideV = ldv * n + strideS = minmn + strideF = minmn + + dh = rocBLAS.handle() + nsv = ROCVector{Cint}(undef, batch_count) + ifail = ROCVector{Cint}(undef, minmn * batch_count) + dev_info = ROCVector{Cint}(undef, batch_count) + 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_count + ) + AMDGPU.unsafe_free!(pA) + + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + + 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)); + kwargs... + ) + chkstride1(A, U, Vᴴ, S) + m, n, batch_count = 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(A, U, Vᴴ, m, n, maxnsv) + length(S) == minmn * batch_count || + throw(DimensionMismatch("length mismatch between A and S")) + + lda = max(1, stride(A, 2)) + strideA = stride(A, 3) + ldu = max(1, stride(U, 2)) + strideU = ldu * size(U, 2) + ldv = max(1, stride(Vᴴ, 2)) + strideV = ldv * n + strideS = minmn + strideF = minmn + + dh = rocBLAS.handle() + nsv = ROCVector{Cint}(undef, batch_count) + ifail = ROCVector{Cint}(undef, minmn * batch_count) + dev_info = ROCVector{Cint}(undef, batch_count) + 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_count + ) + + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + + 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/batched_svd.jl b/src/implementations/batched_svd.jl index f3b7e7e84..c1cc0a2bf 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -6,6 +6,7 @@ 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) function check_input(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) + isempty(A) && return nothing @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) batch_size = length(A) @@ -20,6 +21,7 @@ function check_input(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMa return nothing end function check_input(::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) + isempty(A) && return nothing @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) batch_size = length(A) @@ -35,6 +37,7 @@ function check_input(::typeof(batched_svd_compact!), A::AbstractVector{<:Abstrac return nothing end function check_input(::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, S, ::AbstractAlgorithm) + isempty(A) && return nothing @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) batch_size = length(A) @@ -153,6 +156,9 @@ function initialize_output(::typeof(batched_svd_vals!), A::AbstractArray{T, 3}, return similar(A, real(eltype(A)), (min(m, n), batch_size)) end +_isempty_batch(A::AbstractArray{<:Any, 3}) = isempty(A) +_isempty_batch(A::AbstractVector{<:AbstractMatrix}) = all(isempty, A) + for f! in (:gesdd_batched!, :gesvd_batched!, :gesvdj_batched!, :gesvdx_batched!) @eval $f!(driver::Driver, args...) = throw(ArgumentError("$driver does not provide $($(f!))")) end @@ -234,8 +240,8 @@ for (f, f_lapack!, Alg) in ( Us, Ss, Vᴴs = USVᴴ batches, rest = _ragged_batches(A, alg) for (inds, (m, n)) in batches - Ab = _ragged_pack(A, inds, m, n) - Ub, Sb, Vᴴb = batched_svd_compact!(Ab, initialize_output(batched_svd_compact!, Ab, alg), alg) + 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!(Ss[i], view(Sb, axes(Ss[i], 1), j)) @@ -258,8 +264,8 @@ for (f, f_lapack!, Alg) in ( # `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) - Ub, Sb, Vᴴb = batched_svd_full!(Ab, initialize_output(batched_svd_full!, Ab, alg), alg) + 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)) @@ -277,8 +283,8 @@ for (f, f_lapack!, Alg) in ( 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) - Sb = batched_svd_vals!(Ab, initialize_output(batched_svd_vals!, Ab, alg), alg) + 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 @@ -306,7 +312,7 @@ for (f, f_lapack!, Alg) in ( # Implementation @eval begin function $svd_compact_f!(driver::Driver, A, U, S, Vᴴ; fixgauge::Bool = true, kwargs...) - isempty(A) && return one!(U), zero!(S), one!(Vᴴ) + _isempty_batch(A) && return one!(U), zero!(S), one!(Vᴴ) $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) if fixgauge for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) @@ -318,7 +324,7 @@ for (f, f_lapack!, Alg) in ( 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!)))`"))) - isempty(A) && return one!(U), zero!(S), one!(Vᴴ) + _isempty_batch(A) && return one!(U), zero!(S), one!(Vᴴ) zero!(S) m, n, batch_size = size(S) minmn = min(m, n) @@ -335,13 +341,13 @@ for (f, f_lapack!, Alg) in ( return U, S, Vᴴ end function $svd_vals_f!(driver::Driver, A::AbstractArray{T, 3}, S::AbstractMatrix; fixgauge::Bool = true, kwargs...) where {T} - isempty(A) && return zero!(S) + _isempty_batch(A) && return zero!(S) U, Vᴴ = similar(A, (0, 0, 0)), similar(A, (0, 0, 0)) $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(A) && return zero!(S) + _isempty_batch(A) && return zero!(S) U, Vᴴ = similar(first(A), (0, 0, 0)), similar(first(A), (0, 0, 0)) $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) return S @@ -359,6 +365,16 @@ Larger matrices in a ragged batch are decomposed one at a time instead. Unlimite """ max_batched_blocksize(::AbstractAlgorithm, ::Type) = typemax(Int) +""" + supports_pointer_batch(alg, T::Type) -> Bool + +Whether the low-level batched driver version of `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, ::Type) = false + # 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 @@ -399,9 +415,31 @@ function _ragged_batches(A::AbstractVector{<:AbstractMatrix}, alg::AbstractAlgor return batches, rest end -# Copy `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) +# 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) +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 + +# 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) uniform = all(i -> size(A[i]) == (m, n), inds) + uniform && supports_pointer_batch(alg, 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)) From da3adb824eca4ad9fb771d59c2ea78405853a987 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 16 Sep 2026 15:07:55 +0200 Subject: [PATCH 16/46] Some fixes and actually test Bisection on AMD --- .../MatrixAlgebraKitAMDGPUExt.jl | 14 +++++++++++--- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 4 ++-- 2 files changed, 13 insertions(+), 5 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 252efecaf..c05aaee05 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -20,13 +20,21 @@ MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedROCArray{<:BlasF MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} = ROCSOLVER() function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCMatrix{<:BlasFloat}} - return QRIteration(; kwargs...) + return Bisection(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCArray{<:BlasFloat, 3}} - return QRIteration(; kwargs...) + return Bisection(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} - return QRIteration(; kwargs...) + return Bisection(; kwargs...) +end + +for f in (:svd_full!, :batched_svd_full!) + @eval function MatrixAlgebraKit.default_algorithm( + ::typeof(MatrixAlgebraKit.$f), ::Type{T}; kwargs... + ) where {T <: Union{StridedROCMatrix{<:BlasFloat}, StridedROCArray{<:BlasFloat, 3}, AbstractVector{<:StridedROCMatrix{<:BlasFloat}}}} + return Bisection(; kwargs...) + end end function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}} return DivideAndConquer(; kwargs...) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index d4794eb40..51f302877 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -872,7 +872,7 @@ for (fname, elty, relty) in batch_count = 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(A, U, Vᴴ, m, n, maxnsv) + jobu, jobvt = _gesvdx_jobs(U, Vᴴ, m, n, maxnsv) length(S) == minmn * batch_count || throw(DimensionMismatch("length mismatch between A and S")) @@ -927,7 +927,7 @@ for (fname, elty, relty) in 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(A, U, Vᴴ, m, n, maxnsv) + jobu, jobvt = _gesvdx_jobs(U, Vᴴ, m, n, maxnsv) length(S) == minmn * batch_count || throw(DimensionMismatch("length mismatch between A and S")) From d3d4fd4ec8d0f1ca4124f754bd171ad3a6499202 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 17 Sep 2026 07:01:29 +0200 Subject: [PATCH 17/46] A few more fixes for Bisection --- .../MatrixAlgebraKitAMDGPUExt.jl | 14 +-- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 89 ++++++++++++++----- test/decompositions/svd.jl | 1 - 3 files changed, 68 insertions(+), 36 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index c05aaee05..252efecaf 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -20,21 +20,13 @@ MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: StridedROCArray{<:BlasF MatrixAlgebraKit.default_driver(::Type{TA}) where {TA <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} = ROCSOLVER() function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCMatrix{<:BlasFloat}} - return Bisection(; kwargs...) + return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: StridedROCArray{<:BlasFloat, 3}} - return Bisection(; kwargs...) + return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <: AbstractVector{<:StridedROCMatrix{<:BlasFloat}}} - return Bisection(; kwargs...) -end - -for f in (:svd_full!, :batched_svd_full!) - @eval function MatrixAlgebraKit.default_algorithm( - ::typeof(MatrixAlgebraKit.$f), ::Type{T}; kwargs... - ) where {T <: Union{StridedROCMatrix{<:BlasFloat}, StridedROCArray{<:BlasFloat, 3}, AbstractVector{<:StridedROCMatrix{<:BlasFloat}}}} - return Bisection(; kwargs...) - end + return QRIteration(; kwargs...) end function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}} return DivideAndConquer(; kwargs...) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 51f302877..bedbb001a 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -110,7 +110,8 @@ for (fname, elty, relty) in 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)), + 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) @@ -175,8 +176,9 @@ for (fname, elty, relty) in ) AMDGPU.unsafe_free!(pA) AMDGPU.unsafe_free!(E) - - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end return (S, U, Vᴴ) end @@ -195,7 +197,8 @@ for (fname, elty, relty) in 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)), + 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) @@ -258,7 +261,9 @@ for (fname, elty, relty) in ) AMDGPU.unsafe_free!(E) - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end return (S, U, Vᴴ) end @@ -278,7 +283,8 @@ for (fname, elty, relty) in A::StridedROCMatrix{$elty}, S::StridedROCVector{$relty} = similar(A, $relty, min(size(A)...)), U::StridedROCMatrix{$elty} = similar(A, $elty, size(A, 1), min(size(A)...)), - Vᴴ::StridedROCMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)) + Vᴴ::StridedROCMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); + check::Bool = CHECK_LIBRARY_CALLS[], ) chkstride1(A, U, Vᴴ, S) m, n = size(A) @@ -332,9 +338,10 @@ for (fname, elty, relty) in dev_info ) - info = @allowscalar dev_info[1] - rocSOLVER.chkargsok(BlasInt(info)) - + if check + info = @allowscalar dev_info[1] + rocSOLVER.chkargsok(BlasInt(info)) + end return (S, U, Vᴴ) end end @@ -353,7 +360,8 @@ for (fname, elty, relty) in 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)), + 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) @@ -413,7 +421,9 @@ for (fname, elty, relty) in dev_info, length(A) ) AMDGPU.unsafe_free!(pA) - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end return (S, U, Vᴴ) end @@ -432,7 +442,8 @@ for (fname, elty, relty) in 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)), + 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) @@ -489,8 +500,9 @@ for (fname, elty, relty) in S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, dev_info, batch_size ) - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) - + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end return (S, U, Vᴴ) end end @@ -597,6 +609,7 @@ for (fname, elty, relty) in 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) @@ -659,9 +672,9 @@ for (fname, elty, relty) in S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, dev_info, length(A) ) - - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) - + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end AMDGPU.unsafe_free!(pA) AMDGPU.unsafe_free!(dev_residual) AMDGPU.unsafe_free!(dev_n_sweeps) @@ -685,6 +698,7 @@ for (fname, elty, relty) in 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) @@ -746,8 +760,9 @@ for (fname, elty, relty) in dev_info, batch_size ) - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) - + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end AMDGPU.unsafe_free!(dev_residual) AMDGPU.unsafe_free!(dev_n_sweeps) return (S, U, Vᴴ) @@ -793,6 +808,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 ( @@ -837,8 +872,7 @@ 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) @@ -862,6 +896,7 @@ for (fname, elty, relty) in 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 @@ -897,7 +932,10 @@ for (fname, elty, relty) in ) AMDGPU.unsafe_free!(pA) - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + if check + rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + end + _gesvdx_zero_unconverged!(S, nsv) AMDGPU.unsafe_free!(nsv) AMDGPU.unsafe_free!(ifail) @@ -920,6 +958,7 @@ for (fname, elty, relty) in 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) @@ -950,8 +989,10 @@ for (fname, elty, relty) in S, strideS, U, ldu, strideU, Vᴴ, ldv, strideV, ifail, strideF, dev_info, batch_count ) - - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + 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/test/decompositions/svd.jl b/test/decompositions/svd.jl index 777fd97b2..cb12535ff 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -90,7 +90,6 @@ if AMDGPU.functional() 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) - AMD_SVD_ALGS = (QRIterationBatched(), JacobiBatched(), DivideAndConquerBatched(), BisectionBatched()) TestSuite.test_svd_batched_algs(ROCMatrix{T}, (m, n), batch_size, AMD_SVD_ALGS) end From 91b3dc84098e9f332f88990f2b26807b835a629f Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 17 Sep 2026 09:41:54 +0200 Subject: [PATCH 18/46] Support svd_full for Bisection --- .../MatrixAlgebraKitAMDGPUExt.jl | 15 ++++++++--- src/implementations/svd.jl | 25 +++++++++++++++++++ 2 files changed, 37 insertions(+), 3 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 252efecaf..c28ce9589 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -93,13 +93,22 @@ gesvdj_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::Strided 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} = +function 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} = + _complete_svd_basis!(Us, Vᴴs, size(Ss, 1)) + return Ss, Us, Vᴴs +end +function 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...) + _complete_svd_basis!(Us, Vᴴs, size(Ss, 1)) + return Ss, Us, Vᴴs +end -gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) = +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 gesdd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) = YArocSOLVER.gesdd!(A, S, U, Vᴴ; kwargs...) diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 8f4b7e445..ea59163a1 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -249,6 +249,31 @@ function complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int return U, Vᴴ end +# Some methods (e.g. `gesvdx`) only compute the leading `min(m, n)` singular vectors. +# If `U` or `Vᴴ` is square (`(batched_)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 + N = qr_null!(copy(view(U, :, 1:minmn))) + copyto!(view(U, :, (minmn + 1):size(U, 2)), N) + end + if size(Vᴴ, 1) > minmn + V = view(Vᴴ, 1:minmn, :) + N = qr_null!(adjoint!(similar(V, reverse(size(V))), V)) + adjoint!(view(Vᴴ, (minmn + 1):size(Vᴴ, 1), :), N) + end + return U, Vᴴ +end +function _complete_svd_basis!(U::AbstractArray{<:Any, 3}, Vᴴ::AbstractArray{<:Any, 3}, minmn::Int) + (size(U, 2) > minmn || size(Vᴴ, 1) > minmn) || return U, Vᴴ + for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) + _complete_svd_basis!(u, vᴴ, minmn) + end + return U, Vᴴ +end + + """ requires_tall(alg) -> Bool From a2a45f0070aa629c0f9fef20291c55d4accdc46d Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 18 Sep 2026 04:49:31 -0400 Subject: [PATCH 19/46] Remove batched and non-batched bisection for now --- .../MatrixAlgebraKitAMDGPUExt.jl | 18 +------------ src/implementations/svd.jl | 25 ------------------- test/decompositions/svd.jl | 2 +- 3 files changed, 2 insertions(+), 43 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index c28ce9589..8a2ea53d2 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -46,7 +46,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, :divide_and_conquer, :bisection) +MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :divide_and_conquer) # rocSOLVER's `gesvd*_batched` functions take the batch as an array # of device pointers, so a group of equally sized matrices @@ -93,22 +93,6 @@ gesvdj_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::Strided 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...) -function 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...) - _complete_svd_basis!(Us, Vᴴs, size(Ss, 1)) - return Ss, Us, Vᴴs -end -function 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...) - _complete_svd_basis!(Us, Vᴴs, size(Ss, 1)) - return Ss, Us, Vᴴs -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 gesdd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) = YArocSOLVER.gesdd!(A, S, U, Vᴴ; kwargs...) diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index ea59163a1..8f4b7e445 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -249,31 +249,6 @@ function complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int return U, Vᴴ end -# Some methods (e.g. `gesvdx`) only compute the leading `min(m, n)` singular vectors. -# If `U` or `Vᴴ` is square (`(batched_)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 - N = qr_null!(copy(view(U, :, 1:minmn))) - copyto!(view(U, :, (minmn + 1):size(U, 2)), N) - end - if size(Vᴴ, 1) > minmn - V = view(Vᴴ, 1:minmn, :) - N = qr_null!(adjoint!(similar(V, reverse(size(V))), V)) - adjoint!(view(Vᴴ, (minmn + 1):size(Vᴴ, 1), :), N) - end - return U, Vᴴ -end -function _complete_svd_basis!(U::AbstractArray{<:Any, 3}, Vᴴ::AbstractArray{<:Any, 3}, minmn::Int) - (size(U, 2) > minmn || size(Vᴴ, 1) > minmn) || return U, Vᴴ - for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) - _complete_svd_basis!(u, vᴴ, minmn) - end - return U, Vᴴ -end - - """ requires_tall(alg) -> Bool diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index cb12535ff..85223991c 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -87,7 +87,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(), DivideAndConquer(), Bisection()) + AMD_SVD_ALGS = (QRIteration(), Jacobi(), DivideAndConquer()) 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) From 97c4b7b73f1ec37f25ea1009315aa534670fd75e Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 24 Sep 2026 10:02:07 -0400 Subject: [PATCH 20/46] Restore Bisection test --- test/decompositions/svd.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index 85223991c..cb12535ff 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -87,7 +87,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(), DivideAndConquer()) + 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) From 1db311e502ffb6f8fc49b509d0dab52bc8dc3d65 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 24 Sep 2026 10:09:09 -0400 Subject: [PATCH 21/46] Avoid some allocs --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 12 ++++++------ ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 7 +++++-- 2 files changed, 11 insertions(+), 8 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index bedbb001a..1a011537d 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -177,7 +177,7 @@ for (fname, elty, relty) in AMDGPU.unsafe_free!(pA) AMDGPU.unsafe_free!(E) if check - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) end return (S, U, Vᴴ) @@ -262,7 +262,7 @@ for (fname, elty, relty) in AMDGPU.unsafe_free!(E) if check - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) end return (S, U, Vᴴ) @@ -422,7 +422,7 @@ for (fname, elty, relty) in ) AMDGPU.unsafe_free!(pA) if check - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) end return (S, U, Vᴴ) @@ -501,7 +501,7 @@ for (fname, elty, relty) in dev_info, batch_size ) if check - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) end return (S, U, Vᴴ) end @@ -673,7 +673,7 @@ for (fname, elty, relty) in dev_info, length(A) ) if check - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) end AMDGPU.unsafe_free!(pA) AMDGPU.unsafe_free!(dev_residual) @@ -761,7 +761,7 @@ for (fname, elty, relty) in ) if check - rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) + foreach(rocSOLVER.chkargsok ∘ BlasInt, collect(dev_info)) end AMDGPU.unsafe_free!(dev_residual) AMDGPU.unsafe_free!(dev_n_sweeps) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 7c0472618..69c2492a3 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -293,6 +293,7 @@ for (bname, fname, elty, relty) in 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 @@ -346,8 +347,10 @@ for (bname, fname, elty, relty) in ) end - info = collect(dh.info) - cuSOLVER.chkargsok.(BlasInt.(info)) + if check + info = collect(dh.info) + foreach(cuSOLVER.chkargsok ∘ BlasInt, info) + end cuSOLVER.cusolverDnDestroyGesvdjInfo(params[]) From 379538b90c6923842e2f3165c1db3a83b63c7660 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 25 Sep 2026 02:19:31 -0400 Subject: [PATCH 22/46] one to one-liner and more checks for batch_size --- .../MatrixAlgebraKitAMDGPUExt.jl | 10 --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 70 +++++++++++-------- .../MatrixAlgebraKitCUDAExt.jl | 10 --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 7 +- src/implementations/batched_svd.jl | 4 +- 5 files changed, 48 insertions(+), 53 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 8a2ea53d2..22e12383a 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -32,16 +32,6 @@ function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T return DivideAndConquer(; kwargs...) end -function MatrixAlgebraKit.one!(A::StridedROCArray{T, 3}) where {T <: BlasFloat} - length(A) > 0 || return A - zero!(A) - # TODO use mapslices? - for a in eachslice(A, dims = 3) - diagview(a) .= one(eltype(a)) - end - return A -end - for f in (:geqrf!, :ungqr!, :unmqr!) @eval $f(::ROCSOLVER, args...) = YArocSOLVER.$f(args...) end diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 1a011537d..8872a80ed 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -119,6 +119,8 @@ for (fname, elty, relty) in m, n = size(first(A)) (m < n) && throw(ArgumentError("rocSOLVER's gesvd_batched requires m ≥ n")) minmn = min(m, n) + length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) if length(U) == 0 jobu = rocSOLVER.rocblas_svect_none else @@ -153,8 +155,8 @@ for (fname, elty, relty) in throw(DimensionMismatch("invalid row size of Vᴴ")) end end - length(S) == minmn * length(A) || - throw(DimensionMismatch("length mismatch between A and S")) + 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)) @@ -204,6 +206,8 @@ for (fname, elty, relty) in m, n, batch_size = size(A) (m < n) && throw(ArgumentError("rocSOLVER's gesvd_strided_batched requires m ≥ n")) minmn = min(m, n) + 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ᴴ")) if length(U) == 0 jobu = rocSOLVER.rocblas_svect_none else @@ -238,8 +242,8 @@ for (fname, elty, relty) in throw(DimensionMismatch("invalid row size of Vᴴ")) end end - length(S) == minmn * batch_size || - throw(DimensionMismatch("length mismatch between A and S")) + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) strideA = lda * n @@ -368,6 +372,8 @@ for (fname, elty, relty) in end m, n = size(first(A)) minmn = min(m, n) + length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) if length(U) == 0 jobu = rocSOLVER.rocblas_svect_none else @@ -402,8 +408,8 @@ for (fname, elty, relty) in throw(DimensionMismatch("invalid row size of Vᴴ")) end end - length(S) == minmn * length(A) || - throw(DimensionMismatch("length mismatch between A and S")) + 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)) @@ -448,6 +454,8 @@ for (fname, elty, relty) in chkstride1(A, U, Vᴴ, S) m, n, batch_size = size(A) minmn = min(m, n) + 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ᴴ")) if length(U) == 0 jobu = rocSOLVER.rocblas_svect_none else @@ -482,8 +490,8 @@ for (fname, elty, relty) in throw(DimensionMismatch("invalid row size of Vᴴ")) end end - length(S) == minmn * batch_size || - throw(DimensionMismatch("length mismatch between A and S")) + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) strideA = lda * n @@ -616,7 +624,8 @@ for (fname, elty, relty) in end m, n = size(first(A)) minmn = min(m, n) - + length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) + length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) if length(U) == 0 jobu = rocSOLVER.rocblas_svect_none else @@ -651,8 +660,8 @@ for (fname, elty, relty) in throw(DimensionMismatch("invalid row size of Vᴴ")) end end - length(S) == minmn * length(A) || - throw(DimensionMismatch("length mismatch between A and S")) + 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)) @@ -703,7 +712,8 @@ for (fname, elty, relty) in chkstride1(A, U, Vᴴ, S) m, n, batch_size = size(A) minmn = min(m, n) - + 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ᴴ")) if length(U) == 0 jobu = rocSOLVER.rocblas_svect_none else @@ -738,8 +748,8 @@ for (fname, elty, relty) in throw(DimensionMismatch("invalid row size of Vᴴ")) end end - length(S) == minmn * batch_size || - throw(DimensionMismatch("length mismatch between A and S")) + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) strideA = lda * n @@ -904,12 +914,14 @@ for (fname, elty, relty) in end m, n = size(first(A)) minmn = min(m, n) - batch_count = length(A) + batch_size = length(A) + 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ᴴ")) 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 * batch_count || - throw(DimensionMismatch("length mismatch between A and S")) + 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)) @@ -920,15 +932,15 @@ for (fname, elty, relty) in strideF = minmn dh = rocBLAS.handle() - nsv = ROCVector{Cint}(undef, batch_count) - ifail = ROCVector{Cint}(undef, minmn * batch_count) - dev_info = ROCVector{Cint}(undef, batch_count) + nsv = ROCVector{Cint}(undef, batch_size) + ifail = ROCVector{Cint}(undef, minmn * 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_count + ifail, strideF, dev_info, batch_size ) AMDGPU.unsafe_free!(pA) @@ -962,13 +974,15 @@ for (fname, elty, relty) in kwargs... ) chkstride1(A, U, Vᴴ, S) - m, n, batch_count = size(A) + m, n, batch_size = size(A) + 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ᴴ")) 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 * batch_count || - throw(DimensionMismatch("length mismatch between A and S")) + length(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) strideA = stride(A, 3) @@ -980,14 +994,14 @@ for (fname, elty, relty) in strideF = minmn dh = rocBLAS.handle() - nsv = ROCVector{Cint}(undef, batch_count) - ifail = ROCVector{Cint}(undef, minmn * batch_count) - dev_info = ROCVector{Cint}(undef, batch_count) + nsv = ROCVector{Cint}(undef, batch_size) + ifail = ROCVector{Cint}(undef, minmn * 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_count + ifail, strideF, dev_info, batch_size ) if check rocSOLVER.chkargsok.(BlasInt.(collect(dev_info))) diff --git a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl index bf6a043a7..bf2a0f324 100644 --- a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl +++ b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl @@ -32,16 +32,6 @@ function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T return DivideAndConquer(; kwargs...) end -function MatrixAlgebraKit.one!(A::StridedCuArray{T, 3}) where {T <: BlasFloat} - length(A) > 0 || return A - zero!(A) - # TODO use mapslices? - for a in eachslice(A, dims = 3) - diagview(a) .= one(eltype(a)) - end - return A -end - for f in (:geqrf!, :ungqr!, :unmqr!) @eval $f(::CUSOLVER, args...) = YACUSOLVER.$f(args...) end diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 69c2492a3..44dc57920 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -300,7 +300,8 @@ for (bname, fname, elty, relty) in chkstride1(A, U, Vᴴ, S) m, n, batch_size = size(A) minmn = min(m, n) - + 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ᴴ")) if length(U) == 0 && length(Vᴴ) == 0 jobz = 'N' else @@ -313,8 +314,8 @@ for (bname, fname, elty, relty) in throw(DimensionMismatch("invalid column size of U or row size of Vᴴ")) end end - length(S) == minmn * batch_size || - throw(DimensionMismatch("length mismatch between A and S")) + size(S) == (minmn, batch_size) || + throw(DimensionMismatch("size mismatch between A and S")) # these MUST be "full" sized Ṽ = similar(Vᴴ, (n, n, batch_size)) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index c1cc0a2bf..be262da07 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -312,7 +312,7 @@ for (f, f_lapack!, Alg) in ( # Implementation @eval begin function $svd_compact_f!(driver::Driver, A, U, S, Vᴴ; fixgauge::Bool = true, kwargs...) - _isempty_batch(A) && return one!(U), zero!(S), one!(Vᴴ) + _isempty_batch(A) && return foreach(one!, eachslice(U, dims = 3)), zero!(S), foreach(one!, eachslice(Vᴴ, dims = 3)) $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) if fixgauge for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) @@ -324,7 +324,7 @@ for (f, f_lapack!, Alg) in ( 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!)))`"))) - _isempty_batch(A) && return one!(U), zero!(S), one!(Vᴴ) + _isempty_batch(A) && return foreach(one!, eachslice(U, dims = 3)), zero!(S), foreach(one!, eachslice(Vᴴ, dims = 3)) zero!(S) m, n, batch_size = size(S) minmn = min(m, n) From 4f4974e93c4dbdabd127cd95247c2ec3610b729d Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 28 Sep 2026 08:27:22 +0200 Subject: [PATCH 23/46] Apply batched suggestions from code review Co-authored-by: Lukas Devos Co-authored-by: Jutho --- src/implementations/batched_svd.jl | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index be262da07..6eb85b152 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -1,14 +1,14 @@ # Inputs # ------ -copy_input(::typeof(batched_svd_full), As::AbstractVector{<:AbstractMatrix}) = map(A -> copy!(similar(A, float(eltype(A))), A), As) +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) function check_input(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) isempty(A) && return nothing - @assert all(==(size(first(A))), size.(A)) 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 @@ -22,8 +22,8 @@ function check_input(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMa end function check_input(::typeof(batched_svd_compact!), A::AbstractVector{<:AbstractMatrix}, USVᴴ, ::AbstractAlgorithm) isempty(A) && return nothing - @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) + @assert all(==((m, n)) ∘ size, A) batch_size = length(A) minmn = min(m, n) U, S, Vᴴ = USVᴴ @@ -38,8 +38,8 @@ function check_input(::typeof(batched_svd_compact!), A::AbstractVector{<:Abstrac end function check_input(::typeof(batched_svd_vals!), A::AbstractVector{<:AbstractMatrix}, S, ::AbstractAlgorithm) isempty(A) && return nothing - @assert all(==(size(first(A))), size.(A)) m, n = size(first(A)) + @assert all(==((m, n)) ∘ size, A) batch_size = length(A) minmn = min(m, n) @assert S isa AbstractMatrix From 7a7786611de02e8c099594904f5c965bc113ce76 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 28 Sep 2026 06:37:23 -0400 Subject: [PATCH 24/46] Fix sizes of dummy arrays --- src/implementations/batched_svd.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index 6eb85b152..1a695e2fc 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -342,13 +342,13 @@ for (f, f_lapack!, Alg) in ( 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, 0)), similar(A, (0, 0, 0)) + 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, 0)), similar(first(A), (0, 0, 0)) + 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 From fcb5c9cffb497b3e2ad45630131790848f2ea301 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 28 Sep 2026 08:51:11 -0400 Subject: [PATCH 25/46] Don't return foreach --- src/implementations/batched_svd.jl | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index 1a695e2fc..fdf22d77c 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -312,7 +312,12 @@ for (f, f_lapack!, Alg) in ( # Implementation @eval begin function $svd_compact_f!(driver::Driver, A, U, S, Vᴴ; fixgauge::Bool = true, kwargs...) - _isempty_batch(A) && return foreach(one!, eachslice(U, dims = 3)), zero!(S), foreach(one!, eachslice(Vᴴ, dims = 3)) + 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...) if fixgauge for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) @@ -324,7 +329,12 @@ for (f, f_lapack!, Alg) in ( 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!)))`"))) - _isempty_batch(A) && return foreach(one!, eachslice(U, dims = 3)), zero!(S), foreach(one!, eachslice(Vᴴ, dims = 3)) + 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) From ead9a690c04b0aa0c8bd61ead8934e6fb97b3412 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 09:03:16 +0200 Subject: [PATCH 26/46] Missing batched Bisection piping --- ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 22e12383a..f777fc5bf 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -83,6 +83,11 @@ gesvdj_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::Strided 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...) From 09953498270e31a8b09e5673392c01c0fe721be7 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 09:11:32 +0200 Subject: [PATCH 27/46] Refactor out gaugefixing for batches --- src/common/gauge.jl | 14 ++++++++++++++ src/implementations/batched_svd.jl | 12 ++---------- 2 files changed, 16 insertions(+), 10 deletions(-) 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 index fdf22d77c..a6ba37548 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -319,11 +319,7 @@ for (f, f_lapack!, Alg) in ( return U, S, Vᴴ end $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) - if fixgauge - for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) - gaugefix!(svd_compact!, u, vᴴ) - end - end + 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...) @@ -343,11 +339,7 @@ for (f, f_lapack!, Alg) in ( for (s, sd) in zip(eachslice(S, dims = 3), eachslice(Sd, dims = 2)) diagview(s) .= sd end - if fixgauge - for (u, vᴴ) in zip(eachslice(U, dims = 3), eachslice(Vᴴ, dims = 3)) - gaugefix!(svd_full!, u, vᴴ) - end - 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} From bfe9245951689ab94a58cc68304f21c628b19f9b Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 04:11:44 -0400 Subject: [PATCH 28/46] Add a flag for whether the driver supports ragged batches --- src/implementations/batched_svd.jl | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index a6ba37548..4dbc33bff 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -318,6 +318,8 @@ for (f, f_lapack!, Alg) in ( foreach(one!, eachslice(Vᴴ, dims = 3)) return U, S, Vᴴ end + supports_ragged_batches(driver) || + throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) fixgauge && gaugefix!(batched_svd_compact!, U, Vᴴ) return U, S, Vᴴ @@ -331,6 +333,8 @@ for (f, f_lapack!, Alg) in ( foreach(one!, eachslice(Vᴴ, dims = 3)) return U, S, Vᴴ end + supports_ragged_batches(driver) || + throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) zero!(S) m, n, batch_size = size(S) minmn = min(m, n) @@ -344,12 +348,16 @@ for (f, f_lapack!, Alg) in ( 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) + supports_ragged_batches(driver) || + throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) 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) + supports_ragged_batches(driver) || + throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) 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 @@ -377,6 +385,14 @@ to the driver as a vector of pointers, instead of being copied into one contiguo """ supports_pointer_batch(::AbstractAlgorithm, ::Type) = false +""" + supports_ragged_batch(alg, T::Type) -> Bool + +Whether the driver accepts a *ragged* batch of matrices of type `T` which do not have +uniform size. `true` by default. +""" +supports_ragged_batch(::AbstractAlgorithm, ::Type) = true + # 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 From 4d02a3ca429d4061428f203fae77014f6c2a5cc4 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 04:23:21 -0400 Subject: [PATCH 29/46] Dumb typos --- src/implementations/batched_svd.jl | 31 +++++++++++++++--------------- 1 file changed, 15 insertions(+), 16 deletions(-) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index 4dbc33bff..9b9b48c92 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -193,6 +193,15 @@ function batched_svd_via_adjoint!(f!::F, driver::Driver, A, S, U, Vᴴ; kwargs.. return S, U, Vᴴ end +""" + supports_ragged_batch(alg, T::Type) -> Bool + +Whether the driver accepts a *ragged* batch of matrices of type `T` which do not have +uniform size. `true` by default. +""" +supports_ragged_batch(::AbstractAlgorithm, ::Type) = true + + for (f, f_lapack!, Alg) in ( (:divide_and_conquer, :gesdd_batched!, :DivideAndConquer), (:qr_iteration, :gesvd_batched!, :QRIteration), @@ -207,6 +216,8 @@ for (f, f_lapack!, Alg) in ( @eval begin function batched_svd_compact!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) check_input(batched_svd_compact!, A, USVᴴ, alg) + supports_ragged_batch(get(alg.kwargs, :driver, DefaultDriver()), 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} @@ -215,6 +226,8 @@ for (f, f_lapack!, Alg) in ( end function batched_svd_full!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) check_input(batched_svd_full!, A, USVᴴ, alg) + supports_ragged_batch(get(alg.kwargs, :driver, DefaultDriver()), 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} @@ -223,6 +236,8 @@ for (f, f_lapack!, Alg) in ( end function batched_svd_vals!(A::AbstractVector{<:AbstractMatrix}, S, alg::$Alg) check_input(batched_svd_vals!, A, S, alg) + supports_ragged_batch(get(alg.kwargs, :driver, DefaultDriver()), 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} @@ -318,8 +333,6 @@ for (f, f_lapack!, Alg) in ( foreach(one!, eachslice(Vᴴ, dims = 3)) return U, S, Vᴴ end - supports_ragged_batches(driver) || - throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) $f_lapack!(driver, A, S, U, Vᴴ; kwargs...) fixgauge && gaugefix!(batched_svd_compact!, U, Vᴴ) return U, S, Vᴴ @@ -333,8 +346,6 @@ for (f, f_lapack!, Alg) in ( foreach(one!, eachslice(Vᴴ, dims = 3)) return U, S, Vᴴ end - supports_ragged_batches(driver) || - throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) zero!(S) m, n, batch_size = size(S) minmn = min(m, n) @@ -348,16 +359,12 @@ for (f, f_lapack!, Alg) in ( 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) - supports_ragged_batches(driver) || - throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) 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) - supports_ragged_batches(driver) || - throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) 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 @@ -385,14 +392,6 @@ to the driver as a vector of pointers, instead of being copied into one contiguo """ supports_pointer_batch(::AbstractAlgorithm, ::Type) = false -""" - supports_ragged_batch(alg, T::Type) -> Bool - -Whether the driver accepts a *ragged* batch of matrices of type `T` which do not have -uniform size. `true` by default. -""" -supports_ragged_batch(::AbstractAlgorithm, ::Type) = true - # 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 From 1e9646e81b6fe3165d54b905ab93d528e231d5c1 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 05:59:50 -0400 Subject: [PATCH 30/46] Have batched_svd_compact use Diagonal --- src/implementations/batched_svd.jl | 12 ++++++------ test/testsuite/decompositions/svd.jl | 12 +++++++----- 2 files changed, 13 insertions(+), 11 deletions(-) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index 9b9b48c92..03d1fe74c 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -50,14 +50,14 @@ end # 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{<:AbstractVector}, 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, Diagonal(s), vᴴ), alg) + check_input(svd_compact!, a, (u, s, vᴴ), alg) end return nothing end @@ -136,7 +136,7 @@ function initialize_output(::typeof(batched_svd_full!), A::AbstractArray{T, 3}, 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 = [similar(a, real(eltype(a)), 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 @@ -248,7 +248,7 @@ for (f, f_lapack!, Alg) in ( # ragged batches: pack into 3D batches, see `_ragged_batches` function batched_svd_compact!( A::AbstractVector{<:AbstractMatrix}, - USVᴴ::Tuple{AbstractVector{<:AbstractMatrix}, AbstractVector{<:AbstractVector}, AbstractVector{<:AbstractMatrix}}, + USVᴴ::Tuple{AbstractVector{<:AbstractMatrix}, AbstractVector{<:Diagonal}, AbstractVector{<:AbstractMatrix}}, alg::$Alg ) check_input(batched_svd_compact!, A, USVᴴ, alg) @@ -259,12 +259,12 @@ for (f, f_lapack!, Alg) in ( 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!(Ss[i], view(Sb, axes(Ss[i], 1), 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], Diagonal(Ss[i]), Vᴴs[i]), alg) + svd_compact!(A[i], (Us[i], Ss[i], Vᴴs[i]), alg) end return USVᴴ end diff --git a/test/testsuite/decompositions/svd.jl b/test/testsuite/decompositions/svd.jl index cf53878b2..7e867da06 100644 --- a/test/testsuite/decompositions/svd.jl +++ b/test/testsuite/decompositions/svd.jl @@ -192,18 +192,20 @@ function test_svd_compact_algs_batched( @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 = [similar(a, real(eltype(T)), 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 * Diagonal(s) * vᴴ ≈ a + @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 * Diagonal(s) * vᴴ ≈ a + @test u * s * vᴴ ≈ a @test isisometric(u) @test isisometric(vᴴ; side = :right) end @@ -212,11 +214,11 @@ function test_svd_compact_algs_batched( 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(s) ≈ collect(s2) + @test collect(diagview(s)) ≈ collect(s2) end Sv3 = @testinferred batched_svd_vals(Ar; alg) for (s, s3) in zip(S3, Sv3) - @test collect(s) ≈ collect(s3) + @test collect(diagview(s)) ≈ collect(s3) end end end From 25529dac6266b5189c5d4487bb68b97c6870c494 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 13:56:07 +0200 Subject: [PATCH 31/46] Fix dimension check in gesvdx_strided_batched! --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 8872a80ed..c1e293d55 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -981,7 +981,7 @@ for (fname, elty, relty) in 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, batch_size) || + size(S) == (minmn, batch_size) || throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) From 23f487098bfc839b8079a2572435b517eca3df89 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 13:56:20 +0200 Subject: [PATCH 32/46] supports_ragged_batch takes a driver, not an algorithm --- src/implementations/batched_svd.jl | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index 03d1fe74c..9b2492390 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -199,7 +199,7 @@ end Whether the driver accepts a *ragged* batch of matrices of type `T` which do not have uniform size. `true` by default. """ -supports_ragged_batch(::AbstractAlgorithm, ::Type) = true +supports_ragged_batch(::Driver, ::Type) = true for (f, f_lapack!, Alg) in ( @@ -216,7 +216,8 @@ for (f, f_lapack!, Alg) in ( @eval begin function batched_svd_compact!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) check_input(batched_svd_compact!, A, USVᴴ, alg) - supports_ragged_batch(get(alg.kwargs, :driver, DefaultDriver()), eltype(A)) || + driver = get(alg.kwargs, :driver, DefaultDriver()) + supports_ragged_batch(driver, eltype(A)) || throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) return $svd_compact_f!(A, USVᴴ...; alg.kwargs...) end @@ -226,7 +227,8 @@ for (f, f_lapack!, Alg) in ( end function batched_svd_full!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) check_input(batched_svd_full!, A, USVᴴ, alg) - supports_ragged_batch(get(alg.kwargs, :driver, DefaultDriver()), eltype(A)) || + driver = get(alg.kwargs, :driver, DefaultDriver()) + supports_ragged_batch(driver, eltype(A)) || throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) return $svd_full_f!(A, USVᴴ...; alg.kwargs...) end @@ -236,7 +238,8 @@ for (f, f_lapack!, Alg) in ( end function batched_svd_vals!(A::AbstractVector{<:AbstractMatrix}, S, alg::$Alg) check_input(batched_svd_vals!, A, S, alg) - supports_ragged_batch(get(alg.kwargs, :driver, DefaultDriver()), eltype(A)) || + driver = get(alg.kwargs, :driver, DefaultDriver()) + supports_ragged_batch(driver, eltype(A)) || throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) return $svd_vals_f!(A, S; alg.kwargs...) end From 04df1572e2b3f00d1c4a4a29f4c1a025ff4a7008 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 14:28:21 +0200 Subject: [PATCH 33/46] Make sure to mark bisection as supporting svd full here --- ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index f777fc5bf..2896d26aa 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -36,7 +36,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, :divide_and_conquer) +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 From b9ddc59f815d7c0988f74107feb6c41955cec652 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 16:39:11 +0200 Subject: [PATCH 34/46] Move some batching utilities into a common file --- src/MatrixAlgebraKit.jl | 1 + src/common/batches.jl | 99 ++++++++++++++++++++++++++++++ src/implementations/batched_svd.jl | 99 ------------------------------ 3 files changed, 100 insertions(+), 99 deletions(-) create mode 100644 src/common/batches.jl diff --git a/src/MatrixAlgebraKit.jl b/src/MatrixAlgebraKit.jl index b56548fdd..33f858ce1 100644 --- a/src/MatrixAlgebraKit.jl +++ b/src/MatrixAlgebraKit.jl @@ -97,6 +97,7 @@ include("common/view.jl") include("common/regularinv.jl") include("common/matrixproperties.jl") include("common/balancing.jl") +include("common/batches.jl") include("common/utility.jl") include("yalapack.jl") diff --git a/src/common/batches.jl b/src/common/batches.jl new file mode 100644 index 000000000..8b7dffa86 --- /dev/null +++ b/src/common/batches.jl @@ -0,0 +1,99 @@ +_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) + +""" + supports_ragged_batch(alg, T::Type) -> Bool + +Whether the driver accepts a *ragged* batch of matrices of type `T` which do not have +uniform size. `true` by default. +""" +supports_ragged_batch(::Driver, ::Type) = true + +# Ragged batches +# -------------- +""" + max_batched_blocksize(alg, T::Type) -> Int + +Largest matrix dimension that the driver for the batched versoin 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(::AbstractAlgorithm, ::Type) = typemax(Int) + +""" + supports_pointer_batch(alg, T::Type) -> Bool + +Whether the low-level batched driver version of `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, ::Type) = false + +# 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 + batch_size_limit = max_batched_blocksize(alg, 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) + uniform = all(i -> size(A[i]) == (m, n), inds) + uniform && supports_pointer_batch(alg, 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/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index 9b2492390..d68124928 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -156,25 +156,10 @@ function initialize_output(::typeof(batched_svd_vals!), A::AbstractArray{T, 3}, return similar(A, real(eltype(A)), (min(m, n), batch_size)) end -_isempty_batch(A::AbstractArray{<:Any, 3}) = isempty(A) -_isempty_batch(A::AbstractVector{<:AbstractMatrix}) = all(isempty, A) - for f! in (:gesdd_batched!, :gesvd_batched!, :gesvdj_batched!, :gesvdx_batched!) @eval $f!(driver::Driver, args...) = throw(ArgumentError("$driver does not provide $($(f!))")) end -# 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) - """ batched_svd_via_adjoint!(f!, driver, A, S, U, Vᴴ; kwargs...) @@ -193,15 +178,6 @@ function batched_svd_via_adjoint!(f!::F, driver::Driver, A, S, U, Vᴴ; kwargs.. return S, U, Vᴴ end -""" - supports_ragged_batch(alg, T::Type) -> Bool - -Whether the driver accepts a *ragged* batch of matrices of type `T` which do not have -uniform size. `true` by default. -""" -supports_ragged_batch(::Driver, ::Type) = true - - for (f, f_lapack!, Alg) in ( (:divide_and_conquer, :gesdd_batched!, :DivideAndConquer), (:qr_iteration, :gesvd_batched!, :QRIteration), @@ -375,70 +351,10 @@ for (f, f_lapack!, Alg) in ( end end -# Ragged batches -# -------------- -""" - max_batched_blocksize(alg, T::Type) -> Int - -Largest matrix dimension that the batched driver for `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(::AbstractAlgorithm, ::Type) = typemax(Int) - -""" - supports_pointer_batch(alg, T::Type) -> Bool - -Whether the low-level batched driver version of `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, ::Type) = false - # 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 -# 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 - lim = max_batched_blocksize(alg, 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) <= lim && (!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) <= lim - 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) function _packed_output(::typeof(batched_svd_full!), A::AbstractVector{<:AbstractMatrix}, ::AbstractAlgorithm) a = first(A) m, n = size(a) @@ -455,18 +371,3 @@ function _packed_output(::typeof(batched_svd_vals!), A::AbstractVector{<:Abstrac a = first(A) return similar(a, real(eltype(a)), (min(size(a)...), length(A))) end - -# 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) - uniform = all(i -> size(A[i]) == (m, n), inds) - uniform && supports_pointer_batch(alg, 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 From d404f6747197ddd763909f3e72124715a4fb4250 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 16:55:53 +0200 Subject: [PATCH 35/46] Move supports_ragged_batch to Driver definition --- src/algorithms.jl | 8 ++++++++ src/common/batches.jl | 8 -------- src/implementations/batched_svd.jl | 6 +++--- 3 files changed, 11 insertions(+), 11 deletions(-) diff --git a/src/algorithms.jl b/src/algorithms.jl index 65a25bc18..34807bc0b 100644 --- a/src/algorithms.jl +++ b/src/algorithms.jl @@ -237,6 +237,14 @@ default_driver(::Type{TA}) where {TA <: YALAPACK.MaybeBlasVecOrMat} = LAPACK() @inline default_driver(::Type{<:SubArray{T, N, A}}) where {T, N, A} = default_driver(A) @inline default_driver(::Type{<:Base.ReshapedArray{T, N, A}}) where {T, N, A} = default_driver(A) +""" + supports_ragged_batch(f!, driver::Driver, T::Type) -> Bool + +Whether the 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!, driver::Driver, ::Type) = true + # Truncation strategy # ------------------- """ diff --git a/src/common/batches.jl b/src/common/batches.jl index 8b7dffa86..18c80f0dc 100644 --- a/src/common/batches.jl +++ b/src/common/batches.jl @@ -13,14 +13,6 @@ function batched_adjoint(A::AbstractArray{<:Any, 3}) end batched_adjoint(A::AbstractVector{<:AbstractMatrix}) = map(a -> adjoint!(similar(a'), a), A) -""" - supports_ragged_batch(alg, T::Type) -> Bool - -Whether the driver accepts a *ragged* batch of matrices of type `T` which do not have -uniform size. `true` by default. -""" -supports_ragged_batch(::Driver, ::Type) = true - # Ragged batches # -------------- """ diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index d68124928..cb9e5b18d 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -193,7 +193,7 @@ for (f, f_lapack!, Alg) in ( function batched_svd_compact!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) check_input(batched_svd_compact!, A, USVᴴ, alg) driver = get(alg.kwargs, :driver, DefaultDriver()) - supports_ragged_batch(driver, eltype(A)) || + supports_ragged_batch(batched_svd_compact!, driver, eltype(A)) || throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) return $svd_compact_f!(A, USVᴴ...; alg.kwargs...) end @@ -204,7 +204,7 @@ for (f, f_lapack!, Alg) in ( function batched_svd_full!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) check_input(batched_svd_full!, A, USVᴴ, alg) driver = get(alg.kwargs, :driver, DefaultDriver()) - supports_ragged_batch(driver, eltype(A)) || + supports_ragged_batch(batched_svd_full!, driver, eltype(A)) || throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) return $svd_full_f!(A, USVᴴ...; alg.kwargs...) end @@ -215,7 +215,7 @@ for (f, f_lapack!, Alg) in ( function batched_svd_vals!(A::AbstractVector{<:AbstractMatrix}, S, alg::$Alg) check_input(batched_svd_vals!, A, S, alg) driver = get(alg.kwargs, :driver, DefaultDriver()) - supports_ragged_batch(driver, eltype(A)) || + supports_ragged_batch(batched_svd_vals!, driver, eltype(A)) || throw(ArgumentError(LazyString("driver ", driver, " does not suppport ragged (non-uniform) batches"))) return $svd_vals_f!(A, S; alg.kwargs...) end From d9d8c5e0469c198f0fcbfa0e948a54a3a4774b4b Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 29 Sep 2026 17:09:06 +0200 Subject: [PATCH 36/46] Just put batches below algos --- src/MatrixAlgebraKit.jl | 2 +- src/algorithms.jl | 8 -------- src/{common => }/batches.jl | 9 +++++++++ 3 files changed, 10 insertions(+), 9 deletions(-) rename src/{common => }/batches.jl (94%) diff --git a/src/MatrixAlgebraKit.jl b/src/MatrixAlgebraKit.jl index 33f858ce1..78d36bfe5 100644 --- a/src/MatrixAlgebraKit.jl +++ b/src/MatrixAlgebraKit.jl @@ -97,11 +97,11 @@ include("common/view.jl") include("common/regularinv.jl") include("common/matrixproperties.jl") include("common/balancing.jl") -include("common/batches.jl") include("common/utility.jl") include("yalapack.jl") include("algorithms.jl") +include("batches.jl") include("interface/projections.jl") include("interface/decompositions.jl") diff --git a/src/algorithms.jl b/src/algorithms.jl index 34807bc0b..65a25bc18 100644 --- a/src/algorithms.jl +++ b/src/algorithms.jl @@ -237,14 +237,6 @@ default_driver(::Type{TA}) where {TA <: YALAPACK.MaybeBlasVecOrMat} = LAPACK() @inline default_driver(::Type{<:SubArray{T, N, A}}) where {T, N, A} = default_driver(A) @inline default_driver(::Type{<:Base.ReshapedArray{T, N, A}}) where {T, N, A} = default_driver(A) -""" - supports_ragged_batch(f!, driver::Driver, T::Type) -> Bool - -Whether the 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!, driver::Driver, ::Type) = true - # Truncation strategy # ------------------- """ diff --git a/src/common/batches.jl b/src/batches.jl similarity index 94% rename from src/common/batches.jl rename to src/batches.jl index 18c80f0dc..3da7a20a0 100644 --- a/src/common/batches.jl +++ b/src/batches.jl @@ -15,6 +15,15 @@ batched_adjoint(A::AbstractVector{<:AbstractMatrix}) = map(a -> adjoint!(similar # Ragged batches # -------------- + +""" + supports_ragged_batch(f!, driver::Driver, T::Type) -> Bool + +Whether the 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!, driver::Driver, ::Type) = true + """ max_batched_blocksize(alg, T::Type) -> Int From dfada5bd9742e17f10a647af60357a1b0ca3ba6f Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 09:24:11 +0200 Subject: [PATCH 37/46] Put the basis completion in the implementation file --- .../MatrixAlgebraKitAMDGPUExt.jl | 7 ++----- src/implementations/batched_svd.jl | 6 ++++++ src/implementations/svd.jl | 10 +++------- 3 files changed, 11 insertions(+), 12 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 2896d26aa..dd6660d2e 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -9,7 +9,7 @@ using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_ 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 @@ -55,11 +55,8 @@ 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ᴴ -end # 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} diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index cb9e5b18d..b3a198386 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -330,6 +330,12 @@ for (f, f_lapack!, Alg) in ( 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 diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 8f4b7e445..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ᴴ) From bb619b5e55016ab2dfa2d6d197bdb505cac3d205 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 08:56:11 -0400 Subject: [PATCH 38/46] Move batch_size length checks into else --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 24 ++++++++++---------- 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index c1e293d55..c501ca155 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -119,11 +119,10 @@ for (fname, elty, relty) in m, n = size(first(A)) (m < n) && throw(ArgumentError("rocSOLVER's gesvd_batched requires m ≥ n")) minmn = min(m, n) - length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) - length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) 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 @@ -141,6 +140,7 @@ for (fname, elty, relty) in 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 @@ -206,11 +206,10 @@ for (fname, elty, relty) in m, n, batch_size = size(A) (m < n) && throw(ArgumentError("rocSOLVER's gesvd_strided_batched requires m ≥ n")) minmn = min(m, n) - 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ᴴ")) 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 @@ -228,6 +227,7 @@ for (fname, elty, relty) in 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 @@ -372,11 +372,10 @@ for (fname, elty, relty) in end m, n = size(first(A)) minmn = min(m, n) - length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) - length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) 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 @@ -394,6 +393,7 @@ for (fname, elty, relty) in 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 @@ -454,11 +454,10 @@ for (fname, elty, relty) in chkstride1(A, U, Vᴴ, S) m, n, batch_size = size(A) minmn = min(m, n) - 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ᴴ")) 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 @@ -476,6 +475,7 @@ for (fname, elty, relty) in 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 @@ -624,11 +624,10 @@ for (fname, elty, relty) in end m, n = size(first(A)) minmn = min(m, n) - length(A) != size(U, 3) && throw(ArgumentError("batch size mismatch between A and U")) - length(A) != size(Vᴴ, 3) && throw(ArgumentError("batch size mismatch between A and Vᴴ")) 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 @@ -646,6 +645,7 @@ for (fname, elty, relty) in 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 @@ -712,11 +712,10 @@ for (fname, elty, relty) in chkstride1(A, U, Vᴴ, S) m, n, batch_size = size(A) minmn = min(m, n) - 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ᴴ")) 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 @@ -734,6 +733,7 @@ for (fname, elty, relty) in 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 From 3aef61599179ab344a225bbcfdc2c291f83ce0f8 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 15:49:12 +0200 Subject: [PATCH 39/46] Apply batched suggestions from code review Co-authored-by: Jutho --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index c501ca155..939c6a4a5 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -985,12 +985,12 @@ for (fname, elty, relty) in throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) - strideA = stride(A, 3) + strideA = max(1, stride(A, 3)) ldu = max(1, stride(U, 2)) - strideU = ldu * size(U, 2) + strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) - strideV = ldv * n - strideS = minmn + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) strideF = minmn dh = rocBLAS.handle() From 0b7a6ccffc8eb30a1c9142f97d269a0b2f9fa358 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 10:03:50 -0400 Subject: [PATCH 40/46] Strides for safety --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 52 ++++++++++---------- 1 file changed, 26 insertions(+), 26 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 939c6a4a5..73f7ceb4d 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -160,10 +160,10 @@ for (fname, elty, relty) in lda = max(1, stride(first(A), 2)) ldu = max(1, stride(U, 2)) - strideU = ldu * size(U, 2) + strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) - strideV = ldv * n - strideS = minmn + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) strideE = minmn - 1 E = ROCArray{$relty}(undef, length(A) * strideE) @@ -246,12 +246,12 @@ for (fname, elty, relty) in throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) - strideA = lda * n + strideA = max(1, stride(A, 3)) ldu = max(1, stride(U, 2)) - strideU = ldu * size(U, 2) + strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) - strideV = ldv * n - strideS = minmn + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) strideE = minmn - 1 E = ROCArray{$relty}(undef, batch_size * strideE) @@ -413,10 +413,10 @@ for (fname, elty, relty) in lda = max(1, stride(first(A), 2)) ldu = max(1, stride(U, 2)) - strideU = ldu * size(U, 2) + strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) - strideV = ldv * n - strideS = minmn + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) dh = rocBLAS.handle() dev_info = ROCVector{Cint}(undef, length(A)) @@ -494,12 +494,12 @@ for (fname, elty, relty) in throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) - strideA = lda * n + strideA = max(1, stride(A, 3)) ldu = max(1, stride(U, 2)) - strideU = ldu * size(U, 2) + strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) - strideV = ldv * n - strideS = minmn + strideV = max(1, stride(Vᴴ, 3)) + strideS = max(1, stride(S, 2)) dh = rocBLAS.handle() dev_info = ROCVector{Cint}(undef, batch_size) @@ -665,10 +665,10 @@ for (fname, elty, relty) in lda = max(1, stride(first(A), 2)) ldu = max(1, stride(U, 2)) - strideU = ldu * size(U, 2) + strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) - strideV = ldv * n - strideS = minmn + 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)) @@ -752,12 +752,12 @@ for (fname, elty, relty) in throw(DimensionMismatch("size mismatch between A and S")) lda = max(1, stride(A, 2)) - strideA = lda * n + strideA = max(1, stride(A, 3)) ldu = max(1, stride(U, 2)) - strideU = ldu * size(U, 2) + strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) - strideV = ldv * n - strideS = minmn + 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) @@ -925,15 +925,15 @@ for (fname, elty, relty) in lda = max(1, stride(first(A), 2)) ldu = max(1, stride(U, 2)) - strideU = ldu * size(U, 2) + strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) - strideV = ldv * n - strideS = minmn + 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, minmn * batch_size) + ifail = ROCVector{Cint}(undef, strideF * batch_size) dev_info = ROCVector{Cint}(undef, batch_size) pA = ROCVector(map(pointer, A)) rocSOLVER.$fname( @@ -995,7 +995,7 @@ for (fname, elty, relty) in dh = rocBLAS.handle() nsv = ROCVector{Cint}(undef, batch_size) - ifail = ROCVector{Cint}(undef, minmn * 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, From 7cbfec51f2af99e94219a4d351c662701dee5dfa Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 10:36:00 -0400 Subject: [PATCH 41/46] Use both alg and driver in the batch checks --- .../MatrixAlgebraKitAMDGPUExt.jl | 2 +- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 2 +- .../MatrixAlgebraKitCUDAExt.jl | 2 +- src/batches.jl | 26 ++++++++++--------- src/implementations/batched_svd.jl | 6 ++--- 5 files changed, 20 insertions(+), 18 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index dd6660d2e..2d722fe77 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -41,7 +41,7 @@ MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration # 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, ::Type{<:StridedROCMatrix{<:BlasFloat}}) = true +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) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 73f7ceb4d..d2f2691d0 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -251,7 +251,7 @@ for (fname, elty, relty) in strideU = max(1, stride(U, 3)) ldv = max(1, stride(Vᴴ, 2)) strideV = max(1, stride(Vᴴ, 3)) - strideS = max(1, stride(S, 2)) + strideS = max(1, stride(S, 2)) strideE = minmn - 1 E = ROCArray{$relty}(undef, batch_size * strideE) diff --git a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl index bf2a0f324..8f24e3470 100644 --- a/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl +++ b/ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl @@ -41,7 +41,7 @@ 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, ::Type{<:AnyCuArray}) = 32 +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) diff --git a/src/batches.jl b/src/batches.jl index 3da7a20a0..72ef0ddaa 100644 --- a/src/batches.jl +++ b/src/batches.jl @@ -17,31 +17,31 @@ batched_adjoint(A::AbstractVector{<:AbstractMatrix}) = map(a -> adjoint!(similar # -------------- """ - supports_ragged_batch(f!, driver::Driver, T::Type) -> Bool + supports_ragged_batch(f!, alg::AbstractAlgorithm, driver::Driver, T::Type) -> Bool -Whether the driver accepts a *ragged* batch of matrices of type `T` which do not have -uniform size for function `f!`. `true` by default. +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!, driver::Driver, ::Type) = true +supports_ragged_batch(f!, alg::AbstractAlgorithm, driver::Driver, ::Type) = true """ - max_batched_blocksize(alg, T::Type) -> Int + max_batched_blocksize(alg, driver::Driver, T::Type) -> Int -Largest matrix dimension that the driver for the batched versoin of `alg` accepts for arrays +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(::AbstractAlgorithm, ::Type) = typemax(Int) +max_batched_blocksize(alg::AbstractAlgorithm, driver::Driver, ::Type) = typemax(Int) """ - supports_pointer_batch(alg, T::Type) -> Bool + supports_pointer_batch(alg, driver::Driver, T::Type) -> Bool -Whether the low-level batched driver version of `alg` accepts a batch of matrices of type `T` +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, ::Type) = false +supports_pointer_batch(::AbstractAlgorithm, driver::Driver, ::Type) = false # 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. @@ -52,7 +52,8 @@ function _ragged_batches(A::AbstractVector{<:AbstractMatrix}, alg::AbstractAlgor batches = Tuple{Vector{Int}, Tuple{Int, Int}}[] rest = Int[] isempty(A) && return batches, rest - batch_size_limit = max_batched_blocksize(alg, typeof(first(A))) + 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) @@ -86,8 +87,9 @@ _packed_output(f!, A::AbstractArray{<:Any, 3}, alg::AbstractAlgorithm) = initial # 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, typeof(A[first(inds)])) && return view(A, 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)) diff --git a/src/implementations/batched_svd.jl b/src/implementations/batched_svd.jl index b3a198386..35b440f0d 100644 --- a/src/implementations/batched_svd.jl +++ b/src/implementations/batched_svd.jl @@ -193,7 +193,7 @@ for (f, f_lapack!, Alg) in ( function batched_svd_compact!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) check_input(batched_svd_compact!, A, USVᴴ, alg) driver = get(alg.kwargs, :driver, DefaultDriver()) - supports_ragged_batch(batched_svd_compact!, driver, eltype(A)) || + supports_ragged_batch(batched_svd_compact!, 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 @@ -204,7 +204,7 @@ for (f, f_lapack!, Alg) in ( function batched_svd_full!(A::AbstractVector{<:AbstractMatrix}, USVᴴ, alg::$Alg) check_input(batched_svd_full!, A, USVᴴ, alg) driver = get(alg.kwargs, :driver, DefaultDriver()) - supports_ragged_batch(batched_svd_full!, driver, eltype(A)) || + supports_ragged_batch(batched_svd_full!, 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 @@ -215,7 +215,7 @@ for (f, f_lapack!, Alg) in ( function batched_svd_vals!(A::AbstractVector{<:AbstractMatrix}, S, alg::$Alg) check_input(batched_svd_vals!, A, S, alg) driver = get(alg.kwargs, :driver, DefaultDriver()) - supports_ragged_batch(batched_svd_vals!, driver, eltype(A)) || + supports_ragged_batch(batched_svd_vals!, 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 From 382f50ea865410af2d607706f0b219ee5bca165f Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 08:20:16 +0200 Subject: [PATCH 42/46] batch_size checks for batched gesvdx --- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index d2f2691d0..cbe7730db 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -915,11 +915,11 @@ for (fname, elty, relty) in m, n = size(first(A)) minmn = min(m, n) batch_size = length(A) - 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ᴴ")) 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")) @@ -975,12 +975,12 @@ for (fname, elty, relty) in ) chkstride1(A, U, Vᴴ, S) m, n, batch_size = size(A) - 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ᴴ")) 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")) From 818792ad2cfe11659593322daddd481bb9c00461 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 05:43:37 -0400 Subject: [PATCH 43/46] Force driver lookup for support checks --- src/batches.jl | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/src/batches.jl b/src/batches.jl index 72ef0ddaa..ee86c7700 100644 --- a/src/batches.jl +++ b/src/batches.jl @@ -23,6 +23,8 @@ Whether the algorithm `alg` running on `driver` accepts a *ragged* batch of matr 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 @@ -32,6 +34,8 @@ of type `T`. Larger matrices in a ragged batch are decomposed one at a time inst 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 @@ -42,6 +46,8 @@ to the driver as a vector of pointers, instead of being copied into one contiguo `false` by default. """ supports_pointer_batch(::AbstractAlgorithm, driver::Driver, ::Type) = false +supports_pointer_batch(f!, alg::AbstractAlgorithm, ::DefaultDriver, ::Type{TA}) where {TA} = + supports_pointer_batch(f!, 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. From 27df039d2d29ec90852f79a0e23c73b6d5b6009e Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 06:00:26 -0400 Subject: [PATCH 44/46] Dumb typo --- src/batches.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/batches.jl b/src/batches.jl index ee86c7700..b0be7e96a 100644 --- a/src/batches.jl +++ b/src/batches.jl @@ -46,8 +46,8 @@ to the driver as a vector of pointers, instead of being copied into one contiguo `false` by default. """ supports_pointer_batch(::AbstractAlgorithm, driver::Driver, ::Type) = false -supports_pointer_batch(f!, alg::AbstractAlgorithm, ::DefaultDriver, ::Type{TA}) where {TA} = - supports_pointer_batch(f!, alg, default_driver(alg, TA), TA) +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. From e37c01057e0a6c7d25b33c198d8cace136dee6a6 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 06:12:13 -0400 Subject: [PATCH 45/46] Actually make sure matrices that are too big aren't sent to the batched function --- test/decompositions/svd.jl | 4 +++ test/testsuite/decompositions/svd.jl | 38 ++++++++++++++++++++++++++++ 2 files changed, 42 insertions(+) diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index cb12535ff..9d9288ddb 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -61,6 +61,10 @@ if CUDA.functional() 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: for T in BLASFloats, m in (0, 23), n in (0, 17, m, 27) diff --git a/test/testsuite/decompositions/svd.jl b/test/testsuite/decompositions/svd.jl index 7e867da06..b6cfb7d2e 100644 --- a/test/testsuite/decompositions/svd.jl +++ b/test/testsuite/decompositions/svd.jl @@ -401,6 +401,44 @@ function test_svd_full_algs_batched( 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 + 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)), From bb0d96d9730063c4567ee895e23bc7544188786b Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 13:04:29 +0200 Subject: [PATCH 46/46] Move batch_size check for CUSOLVER too --- ext/MatrixAlgebraKitCUDAExt/yacusolver.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl index 44dc57920..193364588 100644 --- a/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl +++ b/ext/MatrixAlgebraKitCUDAExt/yacusolver.jl @@ -300,11 +300,11 @@ for (bname, fname, elty, relty) in chkstride1(A, U, Vᴴ, S) m, n, batch_size = size(A) minmn = min(m, n) - 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ᴴ")) 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"))