Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 65 additions & 9 deletions ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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...)
Expand Down Expand Up @@ -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}
Expand All @@ -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...) =
Expand Down
21 changes: 19 additions & 2 deletions test/decompositions/svd.jl
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
using MatrixAlgebraKit
using LinearAlgebra: Diagonal
using Test
using LinearAlgebra: Diagonal, isposdef
using CUDA, AMDGPU

if @isdefined(fast_tests) && fast_tests
Expand All @@ -12,6 +13,7 @@ end

@isdefined(TestSuite) || include("../testsuite/TestSuite.jl")
using .TestSuite
using .TestSuite: testargs_summary

is_buildkite = get(ENV, "BUILDKITE", "false") == "true"

Expand Down Expand Up @@ -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())
Expand All @@ -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)
Expand Down
4 changes: 2 additions & 2 deletions test/testsuite/decompositions/svd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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]
Expand Down
Loading