From cea6c9fb68ee5acb083a137eb513a9d5f241e23f Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 17 Sep 2026 16:31:50 +0200 Subject: [PATCH 1/8] Support Bisection for ROCSOLVER --- .../MatrixAlgebraKitAMDGPUExt.jl | 12 +- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 103 ++++++++++++++++++ src/implementations/svd.jl | 17 +++ test/decompositions/svd.jl | 2 +- 4 files changed, 130 insertions(+), 4 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 0bdb10497..b5f56576b 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -6,9 +6,9 @@ using MatrixAlgebraKit: one!, zero!, uppertriangular!, lowertriangular! using MatrixAlgebraKit: diagview, sign_safe using MatrixAlgebraKit: ROCSOLVER, LQViaTransposedQR, TruncationStrategy, NoTruncation, TruncationByValue, AbstractAlgorithm using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eigh_algorithm -import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdj! +import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdx!, gesvdj! import MatrixAlgebraKit: heevj!, heevd!, heev!, heevx! -import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback! +import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback!, _complete_svd_basis! using AMDGPU using LinearAlgebra using LinearAlgebra: BlasFloat @@ -28,7 +28,7 @@ for f in (:geqrf!, :ungqr!, :unmqr!) @eval $f(::ROCSOLVER, args...) = YArocSOLVER.$f(args...) end -MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi) +MatrixAlgebraKit.supports_svd_full(::ROCSOLVER, f::Symbol) = f in (:qr_iteration, :jacobi, :bisection) function gesvd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) m, n = size(A) @@ -42,6 +42,12 @@ function gesvdj!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid return MatrixAlgebraKit.svd_via_adjoint!(gesvdj!, ROCSOLVER(), A, S, U, Vᴴ; kwargs...) end +function gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) + YArocSOLVER.gesvdx!(A, S, U, Vᴴ; kwargs...) + _complete_svd_basis!(U, Vᴴ, length(S)) + return S, U, Vᴴ +end + heevj!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = YArocSOLVER.heevj!(A, Dd, V; kwargs...) heevd!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) = diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index e0c5f084d..7af6e8f19 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -175,6 +175,109 @@ for (fname, elty, relty) in end end +# `gesvdx` computes all singular values, those in the half-open interval `[vl, vu)`, or those +# with an index in `irange`. Unlike the other algorithms, it never forms the full `U` and `Vᴴ`, +# so the only job modes are `singular` and `none`. +function _gesvdx_range(::Type{T}, kwargs) where {T <: Real} + if haskey(kwargs, :irange) + irange = convert(UnitRange{Int}, kwargs[:irange]) + return rocSOLVER.rocblas_srange_index, zero(T), zero(T), first(irange), last(irange) + elseif haskey(kwargs, :vl) || haskey(kwargs, :vu) + vl = convert(T, get(kwargs, :vl, -Inf)) + vu = convert(T, get(kwargs, :vu, +Inf)) + return rocSOLVER.rocblas_srange_value, vl, vu, 0, 0 + else + return rocSOLVER.rocblas_srange_all, zero(T), zero(T), 0, 0 + end +end + +function _gesvdx_maxnsv(srange, il::Integer, iu::Integer, minmn::Integer) + return srange == rocSOLVER.rocblas_srange_index ? iu - il + 1 : minmn +end + +function _gesvdx_jobs(U, Vᴴ, m::Integer, n::Integer, maxnsv::Integer) + if length(U) == 0 + jobu = rocSOLVER.rocblas_svect_none + else + size(U, 1) == m || + throw(DimensionMismatch("row size mismatch between A ($m) and U ($(size(U, 1)))")) + size(U, 2) >= maxnsv || + throw(DimensionMismatch("invalid column size of U")) + jobu = rocSOLVER.rocblas_svect_singular + end + if length(Vᴴ) == 0 + jobvt = rocSOLVER.rocblas_svect_none + else + size(Vᴴ, 2) == n || + throw(DimensionMismatch("column size mismatch between A ($n) and Vᴴ ($(size(Vᴴ, 2)))")) + size(Vᴴ, 1) >= maxnsv || + throw(DimensionMismatch("invalid row size of Vᴴ")) + jobvt = rocSOLVER.rocblas_svect_singular + end + return jobu, jobvt +end + +""" + _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 + +# Wrapper for SVD via Bisection +for (fname, elty, relty) in + ( + (:rocsolver_sgesvdx, :Float32, :Float32), + (:rocsolver_dgesvdx, :Float64, :Float64), + (:rocsolver_cgesvdx, :ComplexF32, :Float32), + (:rocsolver_zgesvdx, :ComplexF64, :Float64), + ) + @eval begin + function gesvdx!( + A::StridedROCMatrix{$elty}, + S::StridedROCVector{$relty} = similar(A, $relty, min(size(A)...)), + U::StridedROCMatrix{$elty} = similar(A, $elty, size(A, 1), min(size(A)...)), + Vᴴ::StridedROCMatrix{$elty} = similar(A, $elty, min(size(A)...), size(A, 2)); + kwargs... + ) + chkstride1(A, U, Vᴴ, S) + m, n = size(A) + minmn = min(m, n) + srange, vl, vu, il, iu = _gesvdx_range($relty, kwargs) + maxnsv = _gesvdx_maxnsv(srange, il, iu, minmn) + jobu, jobvt = _gesvdx_jobs(U, Vᴴ, m, n, maxnsv) + length(S) == minmn || + throw(DimensionMismatch("length mismatch between A ($minmn) and S ($(length(S)))")) + + lda = max(1, stride(A, 2)) + ldu = max(1, stride(U, 2)) + ldv = max(1, stride(Vᴴ, 2)) + ifail = ROCVector{Cint}(undef, minmn) + nsv = ROCVector{Cint}(undef, 1) + dh = rocBLAS.handle() + dev_info = ROCVector{Cint}(undef, 1) + rocSOLVER.$fname( + dh, jobu, jobvt, srange, m, n, + A, lda, vl, vu, il, iu, nsv, + S, U, ldu, Vᴴ, ldv, ifail, + dev_info + ) + info = @allowscalar dev_info[1] + rocSOLVER.chkargsok(BlasInt(info)) + _gesvdx_zero_unconverged!(S, nsv) + + AMDGPU.unsafe_free!(nsv) + AMDGPU.unsafe_free!(ifail) + AMDGPU.unsafe_free!(dev_info) + return (S, U, Vᴴ) + end + end +end + # for (jname, bname, fname, elty, relty) in # ((:sygvd!, :rocsolverDnSsygvd_bufferSize, :rocsolverDnSsygvd, :Float32, :Float32), # (:sygvd!, :rocsolverDnDsygvd_bufferSize, :rocsolverDnDsygvd, :Float64, :Float64), diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 3d20e96d4..ce7d5563c 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -223,6 +223,23 @@ end supports_svd_full(::Driver, ::Symbol) = false supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration) +# Some methods (e.g. `gesvdx`) only compute the leading `min(m, n)` singular vectors. +# If `U` or `Vᴴ` is square (`svd_full!`), the remaining columns (row) of +# `U` (`Vᴴ`) need to be filled with an orthonormal basis for the complement of the +# computed singular vectors. +function _complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int) + if size(U, 2) > minmn + 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 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/decompositions/svd.jl b/test/decompositions/svd.jl index c69ed3a0e..cc1506a72 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -81,7 +81,7 @@ if AMDGPU.functional() for T in BLASFloats, m in (0, 23), n in (0, 17, m, 27) TestSuite.seed_rng!(123) TestSuite.test_svd(ROCMatrix{T}, (m, n)) - AMD_SVD_ALGS = (QRIteration(), Jacobi()) + AMD_SVD_ALGS = (QRIteration(), Jacobi(), Bisection()) TestSuite.test_svd_algs(ROCMatrix{T}, (m, n), AMD_SVD_ALGS) end From 09ddeb2721cfca4fe042c44f95e9fd453c01f930 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 17 Sep 2026 16:40:46 +0200 Subject: [PATCH 2/8] Also test Bisection for LAPACK --- src/implementations/svd.jl | 10 ++++++++-- test/decompositions/svd.jl | 2 +- 2 files changed, 9 insertions(+), 3 deletions(-) diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index ce7d5563c..4c9714655 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -142,10 +142,16 @@ function svd_via_adjoint!(f!::F, driver::Driver, A, S, U, Vᴴ; kwargs...) where end # LAPACK -for f! in (:gesdd!, :gesvd!, :gesvdx!, :gesdvd!) +for f! in (:gesdd!, :gesvd!, :gesdvd!) @eval $f!(::LAPACK, args...; kwargs...) = YALAPACK.$f!(args...; kwargs...) end +function gesvdx!(::LAPACK, A, S, U, Vᴴ; kwargs...) + YALAPACK.gesvdx!(A, S, U, Vᴴ; kwargs...) + _complete_svd_basis!(U, Vᴴ, length(S)) + return S, U, Vᴴ +end + function gesvdj!(::LAPACK, A, S, U, Vᴴ; kwargs...) m, n = size(A) m >= n && return YALAPACK.gesvdj!(A, S, U, Vᴴ) @@ -221,7 +227,7 @@ for (f, f_lapack!, Alg) in ( end supports_svd_full(::Driver, ::Symbol) = false -supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration) +supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide_and_conquer, :qr_iteration, :bisection) # Some methods (e.g. `gesvdx`) only compute the leading `min(m, n)` singular vectors. # If `U` or `Vᴴ` is square (`svd_full!`), the remaining columns (row) of diff --git a/test/decompositions/svd.jl b/test/decompositions/svd.jl index cc1506a72..5387252a3 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -21,7 +21,7 @@ if !is_buildkite # LAPACK algorithms: for T in BLASFloats, m in (0, 54), n in (0, 37, m, 63) TestSuite.seed_rng!(123) - LAPACK_SVD_ALGS = (QRIteration(), DivideAndConquer(), SafeDivideAndConquer(; fixgauge = true)) + LAPACK_SVD_ALGS = (QRIteration(), DivideAndConquer(), SafeDivideAndConquer(; fixgauge = true), Bisection()) TestSuite.test_svd(T, (m, n)) TestSuite.test_svd_algs(T, (m, n), LAPACK_SVD_ALGS) @static if VERSION > v"1.11-" # Jacobi broken on 1.10 From f71e50e8e83cf83273de24491de0368ded836b02 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 18 Sep 2026 09:11:50 +0200 Subject: [PATCH 3/8] Respond to comments --- .../MatrixAlgebraKitAMDGPUExt.jl | 4 ++-- ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl | 21 ++++--------------- src/implementations/svd.jl | 6 +++--- 3 files changed, 9 insertions(+), 22 deletions(-) diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index b5f56576b..4a80ae0fd 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl @@ -8,7 +8,7 @@ using MatrixAlgebraKit: ROCSOLVER, LQViaTransposedQR, TruncationStrategy, NoTrun using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eigh_algorithm import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdx!, gesvdj! import MatrixAlgebraKit: heevj!, heevd!, heev!, heevx! -import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback!, _complete_svd_basis! +import MatrixAlgebraKit: _sylvester, svd_rank, svd_pullback!, complete_svd_basis! using AMDGPU using LinearAlgebra using LinearAlgebra: BlasFloat @@ -44,7 +44,7 @@ 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)) + complete_svd_basis!(U, Vᴴ, length(S)) return S, U, Vᴴ end diff --git a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl index 7af6e8f19..28e206660 100644 --- a/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl +++ b/ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl @@ -191,10 +191,6 @@ function _gesvdx_range(::Type{T}, kwargs) where {T <: Real} end end -function _gesvdx_maxnsv(srange, il::Integer, iu::Integer, minmn::Integer) - return srange == rocSOLVER.rocblas_srange_index ? iu - il + 1 : minmn -end - function _gesvdx_jobs(U, Vᴴ, m::Integer, n::Integer, maxnsv::Integer) if length(U) == 0 jobu = rocSOLVER.rocblas_svect_none @@ -217,17 +213,6 @@ 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 - # Wrapper for SVD via Bisection for (fname, elty, relty) in ( @@ -248,7 +233,7 @@ for (fname, elty, relty) in m, n = size(A) minmn = min(m, n) srange, vl, vu, il, iu = _gesvdx_range($relty, kwargs) - maxnsv = _gesvdx_maxnsv(srange, il, iu, minmn) + maxnsv = srange == rocSOLVER.rocblas_srange_index ? iu - il + 1 : minmn jobu, jobvt = _gesvdx_jobs(U, Vᴴ, m, n, maxnsv) length(S) == minmn || throw(DimensionMismatch("length mismatch between A ($minmn) and S ($(length(S)))")) @@ -268,7 +253,9 @@ for (fname, elty, relty) in ) info = @allowscalar dev_info[1] rocSOLVER.chkargsok(BlasInt(info)) - _gesvdx_zero_unconverged!(S, nsv) + # Zero the entries of `S` that `gesvdx` did not write. + nv = @allowscalar Int(nsv[1]) + nv < length(S) && fill!(view(S, (nv + 1):length(S)), zero(eltype(S))) AMDGPU.unsafe_free!(nsv) AMDGPU.unsafe_free!(ifail) diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 4c9714655..71190c1fc 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -233,14 +233,14 @@ supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide # If `U` or `Vᴴ` is square (`svd_full!`), the remaining columns (row) of # `U` (`Vᴴ`) need to be filled with an orthonormal basis for the complement of the # computed singular vectors. -function _complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int) +function complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int) if size(U, 2) > minmn - N = qr_null!(copy(view(U, :, 1:minmn))) + N = qr_null(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)) + N = lq_null!(similar(V, reverse(size(V))), V) adjoint!(view(Vᴴ, (minmn + 1):size(Vᴴ, 1), :), N) end return U, Vᴴ From 9a10cc9b2f991f773ddbff877da7531998b79066 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 18 Sep 2026 10:31:27 +0200 Subject: [PATCH 4/8] Typo --- src/implementations/svd.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 71190c1fc..7fb6d8590 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -148,7 +148,7 @@ end function gesvdx!(::LAPACK, A, S, U, Vᴴ; kwargs...) YALAPACK.gesvdx!(A, S, U, Vᴴ; kwargs...) - _complete_svd_basis!(U, Vᴴ, length(S)) + complete_svd_basis!(U, Vᴴ, length(S)) return S, U, Vᴴ end From ae6b93071c4001af5f993b44035e5661b45245c8 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 18 Sep 2026 11:48:14 +0200 Subject: [PATCH 5/8] Actually fix lq_null --- src/implementations/svd.jl | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index 7fb6d8590..e583b8512 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -239,9 +239,8 @@ function complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int copyto!(view(U, :, (minmn + 1):size(U, 2)), N) end if size(Vᴴ, 1) > minmn - V = view(Vᴴ, 1:minmn, :) - N = lq_null!(similar(V, reverse(size(V))), V) - adjoint!(view(Vᴴ, (minmn + 1):size(Vᴴ, 1), :), N) + N = lq_null(view(Vᴴ, 1:minmn, :)) + copy!(view(Vᴴ, (minmn + 1):size(Vᴴ, 1), :), N) end return U, Vᴴ end From a100c7b13785b5adfdbb7bf48f133d7b45b5dab3 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 22 Sep 2026 14:25:37 +0200 Subject: [PATCH 6/8] Make sure U (V) doesn't have too many columns (rows) --- src/yalapack.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/yalapack.jl b/src/yalapack.jl index 57e974bda..685458518 100644 --- a/src/yalapack.jl +++ b/src/yalapack.jl @@ -2235,7 +2235,7 @@ for (gesvd, gesdd, gesvdx, gejsv, gesvj, elty, relty) in else size(U, 1) == m || throw(DimensionMismatch("row size mismatch between A ($m) and U ($(size(U, 1)))")) - size(U, 2) >= (range == 'I' ? iu - il + 1 : minmn) || + (size(U, 2) >= (range == 'I' ? iu - il + 1 : minmn) && size(U, 2) <= m) || throw(DimensionMismatch("invalid column size of U")) jobu = 'V' end @@ -2244,7 +2244,7 @@ for (gesvd, gesdd, gesvdx, gejsv, gesvj, elty, relty) in else size(Vᴴ, 2) == n || throw(DimensionMismatch("column size mismatch between A ($n) and Vᴴ ($(size(Vᴴ, 2)))")) - size(Vᴴ, 1) >= (range == 'I' ? iu - il + 1 : minmn) || + (size(Vᴴ, 1) >= (range == 'I' ? iu - il + 1 : minmn) && size(Vᴴ, 1) <= n) || throw(DimensionMismatch("invalid row size of Vᴴ")) jobvt = 'V' end From 0c24f41ada84ea913e56750030d4b0903575f9da Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 22 Sep 2026 09:32:23 -0400 Subject: [PATCH 7/8] Use inplace nulls --- src/implementations/svd.jl | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/src/implementations/svd.jl b/src/implementations/svd.jl index e583b8512..080019e71 100644 --- a/src/implementations/svd.jl +++ b/src/implementations/svd.jl @@ -235,12 +235,16 @@ supports_svd_full(::LAPACK, f::Symbol) = f in (:safe_divide_and_conquer, :divide # computed singular vectors. function complete_svd_basis!(U::AbstractMatrix, Vᴴ::AbstractMatrix, minmn::Int) if size(U, 2) > minmn - N = qr_null(view(U, :, 1:minmn)) - copyto!(view(U, :, (minmn + 1):size(U, 2)), N) + Uc = copy_input(qr_null, view(U, :, 1:minmn)) + N = view(U, :, (minmn + 1):size(U, 2)) + N′ = qr_null!(Uc, N) + N′ === N || copyto!(N, N′) end if size(Vᴴ, 1) > minmn - N = lq_null(view(Vᴴ, 1:minmn, :)) - copy!(view(Vᴴ, (minmn + 1):size(Vᴴ, 1), :), N) + Vc = copy_input(lq_null, view(Vᴴ, 1:minmn, :)) + Nᴴ = view(Vᴴ, (minmn + 1):size(Vᴴ, 1), :) + Nᴴ′ = lq_null!(Vc, Nᴴ) + Nᴴ′ === Nᴴ || copyto!(Nᴴ, Nᴴ′) end return U, Vᴴ end From 5e5b697e2f3cccf2320aa2742b8328e997ea031d Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 22 Sep 2026 15:47:44 +0200 Subject: [PATCH 8/8] Try to fix the size check --- src/yalapack.jl | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/src/yalapack.jl b/src/yalapack.jl index 685458518..f6ca7a507 100644 --- a/src/yalapack.jl +++ b/src/yalapack.jl @@ -2235,8 +2235,11 @@ for (gesvd, gesdd, gesvdx, gejsv, gesvj, elty, relty) in else size(U, 1) == m || throw(DimensionMismatch("row size mismatch between A ($m) and U ($(size(U, 1)))")) - (size(U, 2) >= (range == 'I' ? iu - il + 1 : minmn) && size(U, 2) <= m) || - throw(DimensionMismatch("invalid column size of U")) + if range == 'I' + (size(U, 2) >= iu - il + 1 && size(U, 2) <= m) || throw(DimensionMismatch("invalid column size of U")) + else + (size(U, 2) == minmn || size(U, 2) == m) || throw(DimensionMismatch("invalid column size of U")) + end jobu = 'V' end if length(Vᴴ) == 0 @@ -2244,8 +2247,11 @@ for (gesvd, gesdd, gesvdx, gejsv, gesvj, elty, relty) in else size(Vᴴ, 2) == n || throw(DimensionMismatch("column size mismatch between A ($n) and Vᴴ ($(size(Vᴴ, 2)))")) - (size(Vᴴ, 1) >= (range == 'I' ? iu - il + 1 : minmn) && size(Vᴴ, 1) <= n) || - throw(DimensionMismatch("invalid row size of Vᴴ")) + if range == 'I' + (size(Vᴴ, 1) >= iu - il + 1 && size(Vᴴ, 1) <= n) || throw(DimensionMismatch("invalid row size of Vᴴ")) + else + (size(Vᴴ, 1) == minmn || size(Vᴴ, 1) == n) || throw(DimensionMismatch("invalid row size of Vᴴ")) + end jobvt = 'V' end length(S) == minmn ||