Skip to content

Commit f8787ea

Browse files
committed
Also implement Schur decomposition
1 parent c1954b1 commit f8787ea

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
@@ -38,35 +38,92 @@ function initialize_output(::typeof(schur_vals!), A::AbstractMatrix, ::AbstractA
3838
return vals
3939
end
4040

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

59-
function schur_vals!(A::AbstractMatrix, vals, alg::LAPACK_EigAlgorithm)
60-
check_input(schur_vals!, A, vals, alg)
61-
Z = similar(A, eltype(A), (size(A, 1), 0))
62-
if alg isa LAPACK_Simple
63-
isempty(alg.kwargs) ||
64-
throw(ArgumentError("LAPACK_Simple (gees) does not accept any keyword arguments"))
65-
YALAPACK.gees!(A, Z, vals)
66-
else # alg isa LAPACK_Expert
67-
isempty(alg.kwargs) ||
68-
throw(ArgumentError("LAPACK_Expert (geesx) does not accept any keyword arguments"))
69-
YALAPACK.geesx!(A, Z, vals)
115+
# Deprecations
116+
# ------------
117+
for algtype in (:Simple, :Expert)
118+
lapack_algtype = Symbol(:LAPACK_, algtype)
119+
@eval begin
120+
Base.@deprecate(
121+
schur_full!(A, TZv, alg::$lapack_algtype),
122+
schur_full!(A, TZv, $algtype(; driver = LAPACK(), alg.kwargs...))
123+
)
124+
Base.@deprecate(
125+
schur_vals!(A, vals, alg::$lapack_algtype),
126+
schur_vals!(A, vals, $algtype(; driver = LAPACK(), alg.kwargs...))
127+
)
70128
end
71-
return vals
72129
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)