Skip to content

Commit 0bd5c7f

Browse files
committed
Also implement Schur decomposition
1 parent 58fc98a commit 0bd5c7f

3 files changed

Lines changed: 112 additions & 49 deletions

File tree

ext/MatrixAlgebraKitGenericSchurExt.jl

Lines changed: 27 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@ module MatrixAlgebraKitGenericSchurExt
22

33
using MatrixAlgebraKit
44
using MatrixAlgebraKit: check_input, GS
5-
import MatrixAlgebraKit: geev!
5+
import MatrixAlgebraKit: geev!, gees!, eig_full!, eig_vals!, schur_full!, schur_vals!
66
using LinearAlgebra: Diagonal, sorteig!
77
using GenericSchur
88

@@ -14,37 +14,43 @@ end
1414

1515
MatrixAlgebraKit.default_driver(::Type{<:Simple}, ::Type{TA}) where {TA <: StridedMatrix{<:GSFloat}} = GS()
1616

17+
supports_schur(::GS, f::Symbol) = f === :simple
18+
1719
function geev!(::GS, A::AbstractMatrix, Dd::AbstractVector, V::AbstractMatrix; kwargs...)
1820
D, Vmat = GenericSchur.eigen!(A)
1921
copyto!(Dd, D)
2022
length(V) > 0 && copyto!(V, Vmat)
2123
return Dd, V
2224
end
2325

26+
function gees!(::GS, A::AbstractMatrix, Z::AbstractMatrix, vals::AbstractVector)
27+
S = GenericSchur.gschur(A)
28+
copyto!(A, S.T)
29+
if length(Z) > 0
30+
copyto!(Z, S.Z)
31+
copyto!(vals, S.values)
32+
else
33+
copyto!(vals, sorteig!(S.values))
34+
end
35+
return A, Z, vals
36+
end
37+
2438
Base.@deprecate(
25-
MatrixAlgebraKit.eig_full!(A, DV, alg::GS_QRIteration),
26-
MatrixAlgebraKit.eig_full!(A, DV, Simple(; driver = GS(), alg.kwargs...))
39+
eig_full!(A, DV, alg::GS_QRIteration),
40+
eig_full!(A, DV, Simple(; driver = GS(), alg.kwargs...))
2741
)
2842
Base.@deprecate(
29-
MatrixAlgebraKit.eig_vals!(A, D, alg::GS_QRIteration),
30-
MatrixAlgebraKit.eig_vals!(A, D, Simple(; driver = GS(), alg.kwargs...))
43+
eig_vals!(A, D, alg::GS_QRIteration),
44+
eig_vals!(A, D, Simple(; driver = GS(), alg.kwargs...))
3145
)
3246

33-
function MatrixAlgebraKit.schur_full!(A::AbstractMatrix, TZv, alg::GS_QRIteration)
34-
check_input(schur_full!, A, TZv, alg)
35-
T, Z, vals = TZv
36-
S = GenericSchur.gschur(A)
37-
copyto!(T, S.T)
38-
copyto!(Z, S.Z)
39-
copyto!(vals, S.values)
40-
return T, Z, vals
41-
end
42-
43-
function MatrixAlgebraKit.schur_vals!(A::AbstractMatrix, vals, alg::GS_QRIteration)
44-
check_input(schur_vals!, A, vals, alg)
45-
S = GenericSchur.gschur(A)
46-
copyto!(vals, sorteig!(S.values))
47-
return vals
48-
end
47+
Base.@deprecate(
48+
schur_full!(A, TZv, alg::GS_QRIteration),
49+
schur_full!(A, TZv, Simple(; driver = GS(), alg.kwargs...))
50+
)
51+
Base.@deprecate(
52+
schur_vals!(A, vals, alg::GS_QRIteration),
53+
schur_vals!(A, vals, Simple(; driver = GS(), alg.kwargs...))
54+
)
4955

5056
end

src/implementations/schur.jl

Lines changed: 84 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -49,35 +49,92 @@ for f! in (:schur_full!, :schur_vals!)
4949
end
5050
end
5151

52-
# Implementation
53-
# --------------
54-
function schur_full!(A::AbstractMatrix, TZv, alg::LAPACK_EigAlgorithm)
55-
check_input(schur_full!, A, TZv, alg)
56-
T, Z, vals = TZv
57-
if alg isa LAPACK_Simple
58-
isempty(alg.kwargs) ||
59-
throw(ArgumentError("LAPACK_Simple Schur (gees) does not accept any keyword arguments"))
60-
YALAPACK.gees!(A, Z, vals)
61-
else # alg isa LAPACK_Expert
62-
isempty(alg.kwargs) ||
63-
throw(ArgumentError("LAPACK_Expert Schur (geesx) does not accept any keyword arguments"))
64-
YALAPACK.geesx!(A, Z, vals)
52+
# ==========================
53+
# IMPLEMENTATIONS
54+
# ==========================
55+
56+
for f! in (:gees!, :geesx!)
57+
@eval $f!(driver::Driver, args...) = throw(ArgumentError("$driver does not provide $f!"))
58+
end
59+
60+
# LAPACK implementations
61+
for f! in (:gees!, :geesx!)
62+
@eval $f!(::LAPACK, args...; kwargs...) = YALAPACK.$f!(args...; kwargs...)
63+
end
64+
65+
supports_schur(::Driver, ::Symbol) = false
66+
supports_schur(::LAPACK, f::Symbol) = f in (:simple, :expert)
67+
68+
for (f, f_lapack!, Alg) in (
69+
(:simple, :gees!, :Simple),
70+
(:expert, :geesx!, :Expert),
71+
)
72+
f_schur_full! = Symbol(f, :_schur_full!)
73+
f_schur_vals! = Symbol(f, :_schur_vals!)
74+
75+
# MatrixAlgebraKit wrappers
76+
@eval begin
77+
function schur_full!(A::AbstractMatrix, TZv, alg::$Alg)
78+
check_input(schur_full!, A, TZv, alg)
79+
T, Z, vals = TZv
80+
$f_schur_full!(A, T, Z, vals; alg.kwargs...)
81+
return T, Z, vals
82+
end
83+
function schur_vals!(A::AbstractMatrix, vals, alg::$Alg)
84+
check_input(schur_vals!, A, vals, alg)
85+
Z = similar(A, eltype(A), (size(A, 1), 0))
86+
$f_schur_vals!(A, Z, vals; alg.kwargs...)
87+
return vals
88+
end
89+
end
90+
91+
# driver dispatch
92+
@eval begin
93+
@inline $f_schur_full!(A, T, Z, vals; driver::Driver = DefaultDriver(), kwargs...) =
94+
$f_schur_full!(driver, A, T, Z, vals; kwargs...)
95+
@inline $f_schur_vals!(A, Z, vals; driver::Driver = DefaultDriver(), kwargs...) =
96+
$f_schur_vals!(driver, A, Z, vals; kwargs...)
97+
98+
@inline $f_schur_full!(::DefaultDriver, A, T, Z, vals; kwargs...) =
99+
$f_schur_full!(default_driver($Alg, A), A, T, Z, vals; kwargs...)
100+
@inline $f_schur_vals!(::DefaultDriver, A, Z, vals; kwargs...) =
101+
$f_schur_vals!(default_driver($Alg, A), A, Z, vals; kwargs...)
102+
end
103+
104+
# Implementation
105+
@eval begin
106+
function $f_schur_full!(driver::Driver, A, T, Z, vals; kwargs...)
107+
supports_schur(driver, $(QuoteNode(f))) ||
108+
throw(ArgumentError(LazyString("driver ", driver, " does not provide `$($(QuoteNode(f_lapack!)))`")))
109+
isempty(kwargs) ||
110+
throw(ArgumentError(LazyString("invalid keyword arguments for ", driver, " schur")))
111+
$f_lapack!(driver, A, Z, vals)
112+
T === A || copy!(T, A)
113+
return T, Z, vals
114+
end
115+
function $f_schur_vals!(driver::Driver, A, Z, vals; kwargs...)
116+
supports_schur(driver, $(QuoteNode(f))) ||
117+
throw(ArgumentError(LazyString("driver ", driver, " does not provide `$($(QuoteNode(f_lapack!)))`")))
118+
isempty(kwargs) ||
119+
throw(ArgumentError(LazyString("invalid keyword arguments for ", driver, " schur")))
120+
$f_lapack!(driver, A, Z, vals)
121+
return vals
122+
end
65123
end
66-
T === A || copy!(T, A)
67-
return T, Z, vals
68124
end
69125

70-
function schur_vals!(A::AbstractMatrix, vals, alg::LAPACK_EigAlgorithm)
71-
check_input(schur_vals!, A, vals, alg)
72-
Z = similar(A, eltype(A), (size(A, 1), 0))
73-
if alg isa LAPACK_Simple
74-
isempty(alg.kwargs) ||
75-
throw(ArgumentError("LAPACK_Simple (gees) does not accept any keyword arguments"))
76-
YALAPACK.gees!(A, Z, vals)
77-
else # alg isa LAPACK_Expert
78-
isempty(alg.kwargs) ||
79-
throw(ArgumentError("LAPACK_Expert (geesx) does not accept any keyword arguments"))
80-
YALAPACK.geesx!(A, Z, vals)
126+
# Deprecations
127+
# ------------
128+
for algtype in (:Simple, :Expert)
129+
lapack_algtype = Symbol(:LAPACK_, algtype)
130+
@eval begin
131+
Base.@deprecate(
132+
schur_full!(A, TZv, alg::$lapack_algtype),
133+
schur_full!(A, TZv, $algtype(; driver = LAPACK(), alg.kwargs...))
134+
)
135+
Base.@deprecate(
136+
schur_vals!(A, vals, alg::$lapack_algtype),
137+
schur_vals!(A, vals, $algtype(; driver = LAPACK(), alg.kwargs...))
138+
)
81139
end
82-
return vals
83140
end

test/schur.jl

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@ for T in (BLASFloats..., GenericFloats...)
2828
if !is_buildkite
2929
TestSuite.test_schur(T, (m, m))
3030
if T BLASFloats
31-
LAPACK_SCHUR_ALGS = (LAPACK_Simple(), LAPACK_Expert())
31+
LAPACK_SCHUR_ALGS = (Simple(), Expert())
3232
TestSuite.test_schur_algs(T, (m, m), LAPACK_SCHUR_ALGS)
3333
end
3434
#AT = Diagonal{T, Vector{T}}

0 commit comments

Comments
 (0)