Skip to content

Commit c1954b1

Browse files
committed
also implement extensions
1 parent c90c392 commit c1954b1

5 files changed

Lines changed: 55 additions & 37 deletions

File tree

ext/MatrixAlgebraKitAMDGPUExt/MatrixAlgebraKitAMDGPUExt.jl

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,8 @@ using MatrixAlgebraKit: diagview, sign_safe
77
using MatrixAlgebraKit: ROCSOLVER, LQViaTransposedQR, TruncationStrategy, NoTruncation, TruncationByValue, AbstractAlgorithm
88
using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eigh_algorithm
99
import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdj!
10-
import MatrixAlgebraKit: _gpu_heevj!, _gpu_heevd!, _gpu_heev!, _gpu_heevx!, _sylvester, svd_rank
10+
import MatrixAlgebraKit: heevj!, heevd!, heev!, heevx!
11+
import MatrixAlgebraKit: _sylvester, svd_rank
1112
using AMDGPU
1213
using LinearAlgebra
1314
using LinearAlgebra: BlasFloat
@@ -20,7 +21,7 @@ function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <
2021
return QRIteration(; kwargs...)
2122
end
2223
function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedROCVecOrMat{<:BlasFloat}}
23-
return ROCSOLVER_DivideAndConquer(; kwargs...)
24+
return DivideAndConquer(; kwargs...)
2425
end
2526

2627
for f in (:geqrf!, :ungqr!, :unmqr!)
@@ -42,13 +43,13 @@ function gesvdj!(::ROCSOLVER, A::StridedROCMatrix, S::StridedROCVector, U::Strid
4243
return MatrixAlgebraKit.svd_via_adjoint!(gesvdj!, ROCSOLVER(), A, S, U, Vᴴ; kwargs...)
4344
end
4445

45-
_gpu_heevj!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
46+
heevj!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
4647
YArocSOLVER.heevj!(A, Dd, V; kwargs...)
47-
_gpu_heevd!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
48+
heevd!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
4849
YArocSOLVER.heevd!(A, Dd, V; kwargs...)
49-
_gpu_heev!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
50+
heev!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
5051
YArocSOLVER.heev!(A, Dd, V; kwargs...)
51-
_gpu_heevx!(A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
52+
heevx!(::ROCSOLVER, A::StridedROCMatrix, Dd::StridedROCVector, V::StridedROCMatrix; kwargs...) =
5253
YArocSOLVER.heevx!(A, Dd, V; kwargs...)
5354

5455
function MatrixAlgebraKit.findtruncated_svd(values::StridedROCVector, strategy::TruncationByValue)

ext/MatrixAlgebraKitCUDAExt/MatrixAlgebraKitCUDAExt.jl

Lines changed: 9 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,9 @@ using MatrixAlgebraKit: one!, zero!, uppertriangular!, lowertriangular!
66
using MatrixAlgebraKit: diagview, sign_safe
77
using MatrixAlgebraKit: CUSOLVER, LQViaTransposedQR, TruncationByValue, AbstractAlgorithm
88
using MatrixAlgebraKit: default_qr_algorithm, default_lq_algorithm, default_svd_algorithm, default_eig_algorithm, default_eigh_algorithm
9-
import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdp!, gesvdr!, gesvdj!, _gpu_geev!
10-
import MatrixAlgebraKit: _gpu_heevj!, _gpu_heevd!, _gpu_Xgesvdr!, _sylvester, svd_rank
9+
import MatrixAlgebraKit: geqrf!, ungqr!, unmqr!, gesvd!, gesvdp!, gesvdr!, gesvdj!
10+
import MatrixAlgebraKit: heevj!, heevd!, geev!
11+
import MatrixAlgebraKit: _gpu_Xgesvdr!, _sylvester, svd_rank
1112
using CUDA, CUDA.CUBLAS
1213
using CUDA: i32
1314
using LinearAlgebra
@@ -21,10 +22,10 @@ function MatrixAlgebraKit.default_svd_algorithm(::Type{T}; kwargs...) where {T <
2122
return QRIteration(; kwargs...)
2223
end
2324
function MatrixAlgebraKit.default_eig_algorithm(::Type{T}; kwargs...) where {T <: StridedCuVecOrMat{<:BlasFloat}}
24-
return CUSOLVER_Simple(; kwargs...)
25+
return Simple(; kwargs...)
2526
end
2627
function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: StridedCuVecOrMat{<:BlasFloat}}
27-
return CUSOLVER_DivideAndConquer(; kwargs...)
28+
return DivideAndConquer(; kwargs...)
2829
end
2930

3031

@@ -53,12 +54,12 @@ gesvdp!(::CUSOLVER, A::StridedCuMatrix, S::StridedCuVector, U::StridedCuMatrix,
5354
_gpu_Xgesvdr!(A::StridedCuMatrix, S::StridedCuVector, U::StridedCuMatrix, Vᴴ::StridedCuMatrix; kwargs...) =
5455
YACUSOLVER.gesvdr!(A, S, U, Vᴴ; kwargs...)
5556

56-
_gpu_geev!(A::StridedCuMatrix, D::StridedCuVector, V::StridedCuMatrix) =
57-
YACUSOLVER.Xgeev!(A, D, V)
57+
geev!(::CUSOLVER, A::StridedCuMatrix, Dd::StridedCuVector, V::StridedCuMatrix) =
58+
YACUSOLVER.Xgeev!(A, Dd, V)
5859

59-
_gpu_heevj!(A::StridedCuMatrix, Dd::StridedCuVector, V::StridedCuMatrix; kwargs...) =
60+
heevj!(::CUSOLVER, A::StridedCuMatrix, Dd::StridedCuVector, V::StridedCuMatrix; kwargs...) =
6061
YACUSOLVER.heevj!(A, Dd, V; kwargs...)
61-
_gpu_heevd!(A::StridedCuMatrix, Dd::StridedCuVector, V::StridedCuMatrix; kwargs...) =
62+
heevd!(::CUSOLVER, A::StridedCuMatrix, Dd::StridedCuVector, V::StridedCuMatrix; kwargs...) =
6263
YACUSOLVER.heevd!(A, Dd, V; kwargs...)
6364

6465
function MatrixAlgebraKit.findtruncated_svd(values::StridedCuVector, strategy::TruncationByValue)

ext/MatrixAlgebraKitGenericLinearAlgebraExt.jl

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@ module MatrixAlgebraKitGenericLinearAlgebraExt
33
using MatrixAlgebraKit
44
using MatrixAlgebraKit: sign_safe, check_input, diagview, gaugefix!, one!, zero!, default_fixgauge
55
using MatrixAlgebraKit: GLA
6-
import MatrixAlgebraKit: gesvd!
6+
import MatrixAlgebraKit: gesvd!, heev!
77
using GenericLinearAlgebra: svd!, svdvals!, eigen!, eigvals!, Hermitian, qr!
88
using LinearAlgebra: I, Diagonal, lmul!
99

@@ -33,19 +33,19 @@ function gesvd!(::GLA, A::AbstractMatrix, S::AbstractVector, U::AbstractMatrix,
3333
end
3434

3535
function MatrixAlgebraKit.default_eigh_algorithm(::Type{T}; kwargs...) where {T <: GlaStridedVecOrMatrix}
36-
return GLA_QRIteration(; kwargs...)
37-
end
38-
39-
MatrixAlgebraKit.initialize_output(::typeof(eigh_full!), A::AbstractMatrix, ::GLA_QRIteration) = (nothing, nothing)
40-
MatrixAlgebraKit.initialize_output(::typeof(eigh_vals!), A::AbstractMatrix, ::GLA_QRIteration) = nothing
41-
42-
function MatrixAlgebraKit.eigh_full!(A::AbstractMatrix, DV, ::GLA_QRIteration)
43-
eigval, eigvec = eigen!(Hermitian(A); sortby = real)
44-
return Diagonal(eigval::AbstractVector{real(eltype(A))}), eigvec::AbstractMatrix{eltype(A)}
36+
return QRIteration(; kwargs...)
4537
end
4638

47-
function MatrixAlgebraKit.eigh_vals!(A::AbstractMatrix, D, ::GLA_QRIteration)
48-
return eigvals!(Hermitian(A); sortby = real)
39+
function heev!(::GLA, A::AbstractMatrix, Dd::AbstractVector, V::AbstractMatrix; kwargs...)
40+
if length(V) > 0
41+
eigval, eigvec = eigen!(Hermitian(A); sortby = real)
42+
copyto!(Dd, eigval)
43+
copyto!(V, eigvec)
44+
else
45+
eigval = eigvals!(Hermitian(A); sortby = real)
46+
copyto!(Dd, eigval)
47+
end
48+
return Dd, V
4949
end
5050

5151
function MatrixAlgebraKit.householder_qr!(

ext/MatrixAlgebraKitGenericSchurExt.jl

Lines changed: 20 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,25 +1,34 @@
11
module MatrixAlgebraKitGenericSchurExt
22

33
using MatrixAlgebraKit
4-
using MatrixAlgebraKit: check_input
4+
using MatrixAlgebraKit: check_input, GS
5+
import MatrixAlgebraKit: geev!
56
using LinearAlgebra: Diagonal, sorteig!
67
using GenericSchur
78

8-
function MatrixAlgebraKit.default_eig_algorithm(::Type{T}; kwargs...) where {T <: StridedMatrix{<:Union{Float16, ComplexF16, BigFloat, Complex{BigFloat}}}}
9-
return GS_QRIteration(; kwargs...)
9+
const GSFloat = Union{Float16, ComplexF16, BigFloat, Complex{BigFloat}}
10+
11+
function MatrixAlgebraKit.default_eig_algorithm(::Type{T}; kwargs...) where {T <: StridedMatrix{<:GSFloat}}
12+
return Simple(; kwargs...)
1013
end
1114

12-
MatrixAlgebraKit.initialize_output(::typeof(eig_full!), A::AbstractMatrix, ::GS_QRIteration) = (nothing, nothing)
13-
MatrixAlgebraKit.initialize_output(::typeof(eig_vals!), A::AbstractMatrix, ::GS_QRIteration) = nothing
15+
MatrixAlgebraKit.default_driver(::Type{<:Simple}, ::Type{TA}) where {TA <: StridedMatrix{<:GSFloat}} = GS()
1416

15-
function MatrixAlgebraKit.eig_full!(A::AbstractMatrix, DV, ::GS_QRIteration)
16-
D, V = GenericSchur.eigen!(A)
17-
return Diagonal(D), V
17+
function geev!(::GS, A::AbstractMatrix, Dd::AbstractVector, V::AbstractMatrix; kwargs...)
18+
D, Vmat = GenericSchur.eigen!(A)
19+
copyto!(Dd, D)
20+
length(V) > 0 && copyto!(V, Vmat)
21+
return Dd, V
1822
end
1923

20-
function MatrixAlgebraKit.eig_vals!(A::AbstractMatrix, D, ::GS_QRIteration)
21-
return GenericSchur.eigvals!(A)
22-
end
24+
Base.@deprecate(
25+
MatrixAlgebraKit.eig_full!(A, DV, alg::GS_QRIteration),
26+
MatrixAlgebraKit.eig_full!(A, DV, Simple(; driver = GS(), alg.kwargs...))
27+
)
28+
Base.@deprecate(
29+
MatrixAlgebraKit.eig_vals!(A, D, alg::GS_QRIteration),
30+
MatrixAlgebraKit.eig_vals!(A, D, Simple(; driver = GS(), alg.kwargs...))
31+
)
2332

2433
function MatrixAlgebraKit.schur_full!(A::AbstractMatrix, TZv, alg::GS_QRIteration)
2534
check_input(schur_full!, A, TZv, alg)

src/algorithms.jl

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -212,6 +212,13 @@ Driver to select a native implementation in MatrixAlgebraKit as the implementati
212212
"""
213213
struct Native <: Driver end
214214

215+
"""
216+
GS <: Driver
217+
218+
Driver to select GenericSchur.jl as the implementation strategy.
219+
"""
220+
struct GS <: Driver end
221+
215222
# In order to avoid amibiguities, this method is implemented in a tiered way
216223
# default_driver(alg, A) -> default_driver(typeof(alg), typeof(A))
217224
# default_driver(Talg, TA) -> default_driver(TA)

0 commit comments

Comments
 (0)