diff --git a/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl b/ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl index 2d722fe77..7e2a44fd0 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 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 function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}} return DivideAndConquer(; kwargs...) @@ -55,8 +55,14 @@ function gesvdj!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid return MatrixAlgebraKit.svd_via_adjoint!(gesvdj!, ROCSOLVER(), A, S, U, Vᴴ; kwargs...) end -gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) = - YArocSOLVER.gesvdx!(A, S, U, Vᴴ; kwargs...) +function gesvdx!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) + if min(size(A)...) == 1 + _gesvdx_rank1!(A, S, U, Vᴴ; kwargs...) + else + YArocSOLVER.gesvdx!(A, S, U, Vᴴ; kwargs...) + end + 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} @@ -80,14 +86,64 @@ 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...) +function gesvdx_batched!(::ROCSOLVER, As::AbstractVector{<:StridedROCMatrix}, Ss::StridedROCMatrix, Us::StridedROCArray{T, 3}, Vᴴs::StridedROCArray{T, 3}; kwargs...) where {T <: BlasFloat} + if min(size(first(As))...) == 1 + _gesvdx_rank1_batched!(As, Ss, Us, Vᴴs; kwargs...) + else + YArocSOLVER.gesvdx_batched!(As, Ss, Us, Vᴴs; kwargs...) + end + 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} + if min(size(As, 1), size(As, 2)) == 1 + _gesvdx_rank1_batched!(As, Ss, Us, Vᴴs; kwargs...) + else + YArocSOLVER.gesvdx_strided_batched!(As, Ss, Us, Vᴴs; kwargs...) + end + return Ss, Us, Vᴴs +end gesdd!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) = YArocSOLVER.gesdd!(A, S, U, Vᴴ; kwargs...) +function _gesvdx_rank1!(A::StridedROCMatrix, S::StridedROCVector, U::StridedROCMatrix, Vᴴ::StridedROCMatrix; kwargs...) + select_kwargs = Base.structdiff(NamedTuple(kwargs), (; irange = nothing, vl = nothing, vu = nothing)) + nrmA = norm(A) + fill!(S, nrmA) + if !isempty(U) && !isempty(Vᴴ) + u, v = view(U, :, 1), view(Vᴴ, 1, :) + x, e = size(A, 1) == 1 ? (v, u) : (u, v) + if iszero(nrmA) + zero!(x) + fill!(view(x, 1:1), one(eltype(x))) + else + copyto!(x, vec(A)) + x ./= nrmA + end + fill!(e, one(eltype(e))) + end + _gesvdx_apply_range!(S, select_kwargs) + return S, U, Vᴴ +end + +function _gesvdx_rank1_batched!(As, Ss::StridedROCMatrix, Us, Vᴴs; kwargs...) + select_kwargs = Base.structdiff(NamedTuple(kwargs), (; irange = nothing, vl = nothing, vu = nothing)) + gesvdj_batched!(ROCSOLVER(), As, Ss, Us, Vᴴs; select_kwargs...) + _gesvdx_apply_range!(Ss, select_kwargs) + return Ss, Us, Vᴴs +end + +function _gesvdx_apply_range!(S, select::NamedTuple) + if haskey(select, :irange) + 1 in convert(UnitRange{Int}, select.irange) || zero!(S) + elseif haskey(select, :vl) || haskey(select, :vu) + vl = convert(eltype(S), get(select, :vl, -Inf)) + vu = convert(eltype(S), get(select, :vu, Inf)) + S .= ifelse.((vl .<= S) .& (S .< vu), S, zero(eltype(S))) + end + return S +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/test/decompositions/svd.jl b/test/decompositions/svd.jl index 9d9288ddb..a0faf4343 100644 --- a/test/decompositions/svd.jl +++ b/test/decompositions/svd.jl @@ -1,5 +1,6 @@ using MatrixAlgebraKit -using LinearAlgebra: Diagonal +using Test +using LinearAlgebra: Diagonal, isposdef using CUDA, AMDGPU if @isdefined(fast_tests) && fast_tests @@ -12,6 +13,7 @@ end @isdefined(TestSuite) || include("../testsuite/TestSuite.jl") using .TestSuite +using .TestSuite: testargs_summary is_buildkite = get(ENV, "BUILDKITE", "false") == "true" @@ -88,7 +90,7 @@ end # ------------ if AMDGPU.functional() # ROCSOLVER algorithms: - for T in BLASFloats, m in (0, 23), n in (0, 17, m, 27) + for T in BLASFloats, m in (0, 1, 23), n in (0, 1, 17, m, 27) TestSuite.seed_rng!(123) TestSuite.test_svd(ROCMatrix{T}, (m, n)) AMD_SVD_ALGS = (QRIteration(), Jacobi(), DivideAndConquer(), Bisection()) @@ -97,6 +99,21 @@ if AMDGPU.functional() TestSuite.test_svd_batched_algs(ROCMatrix{T}, (m, n), batch_size, AMD_SVD_ALGS) end + @testset "Bisection with min(m, n) == 1 $(testargs_summary(T, sz))" for T in BLASFloats, + sz in ((1, 1), (2, 1), (5, 1), (1, 2), (1, 5)) + + TestSuite.seed_rng!(123) + for _ in 1:16 + A = TestSuite.instantiate_matrix(ROCMatrix{T}, sz) + U, S, Vᴴ = svd_compact(A; alg = Bisection()) + @test U * S * Vᴴ ≈ A + @test isisometric(U) + @test isisometric(Vᴴ; side = :right) + @test isposdef(S) + @test Array(S) ≈ Diagonal(svd_vals(Array(A))) + end + end + # Diagonal: for T in BLASFloats, m in (0, 23) TestSuite.seed_rng!(123) diff --git a/test/testsuite/decompositions/svd.jl b/test/testsuite/decompositions/svd.jl index a0fddda61..914f28054 100644 --- a/test/testsuite/decompositions/svd.jl +++ b/test/testsuite/decompositions/svd.jl @@ -468,7 +468,7 @@ function test_svd_trunc( S₀ = collect(svd_vals(A)) r = minmn - 2 - if m > 0 && n > 0 + if m > 0 && n > 0 && r >= 0 U1, S1, V1ᴴ, ϵ1 = @testinferred svd_trunc(A; trunc = truncrank(r)) @test length(diagview(S1)) == r @test collect(diagview(S1)) ≈ S₀[1:r] @@ -573,7 +573,7 @@ function test_svd_trunc_algs( S₀ = collect(svd_vals(A)) r = minmn - 2 - if m > 0 && n > 0 + if m > 0 && n > 0 && r >= 0 U1, S1, V1ᴴ, ϵ1 = @testinferred svd_trunc(A; trunc = truncrank(r), alg) @test length(diagview(S1)) == r @test collect(diagview(S1)) ≈ S₀[1:r]