Skip to content
Merged
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
12 changes: 9 additions & 3 deletions ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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...) =
Expand Down
90 changes: 90 additions & 0 deletions ext/MatrixAlgebraKitAMDGPUExt/yarocsolver.jl
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,96 @@ 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_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

# 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 = srange == rocSOLVER.rocblas_srange_index ? iu - il + 1 : minmn
jobu, jobvt = _gesvdx_jobs(U, Vᴴ, m, n, maxnsv)
Comment thread
kshyatt marked this conversation as resolved.
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))
# 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)
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),
Expand Down
30 changes: 28 additions & 2 deletions src/implementations/svd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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ᴴ)
Expand Down Expand Up @@ -221,7 +227,27 @@ 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
# `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
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
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

function svd_trunc_no_error!(A, USVᴴ, alg::TruncatedAlgorithm)
U, S, Vᴴ = svd_compact!(A, USVᴴ, alg.alg)
Expand Down
14 changes: 10 additions & 4 deletions src/yalapack.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2235,17 +2235,23 @@ 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) ||
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
jobvt = 'N'
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) ||
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 ||
Expand Down
4 changes: 2 additions & 2 deletions test/decompositions/svd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down
Loading