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
69 changes: 68 additions & 1 deletion test/testsuite/ad_utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,57 @@ test in-place Hermitian eigendecomposition rules via Mooncake's non-primitive AD
"""
eigh!_wrapper(f!, A, alg) = (F = f!(project_hermitian!(A), alg); MatrixAlgebraKit.zero!(A); F)

"""
eig_vals_wrapper(f, A, alg)

Wrapper that sorts the eigenvalues returned by `f(A, alg)` by modulus and then by imaginary part.
LAPACK's ordering of the eigenvalues can change discontinuously under small perturbations of
`A`, which breaks finite-difference checks. The ordering imposed here is smooth for matrices
built with `make_eig_matrix`, whose eigenvalues have distinct moduli up to conjugate pairs.
"""
eig_vals_wrapper(f, A, alg) = sort_eigvals(f(A, alg))

"""
eig_vals!_wrapper(f!, A, alg)

In-place variant of [`eig_vals_wrapper`](@ref), which zeros `A` after calling `f!`.
"""
eig_vals!_wrapper(f!, A, alg) = sort_eigvals(call_and_zero!(f!, A, alg))

"""
eig_trunc_wrapper(f, A, alg)

Wrapper that sorts the eigenpairs returned by `f(A, alg)` (`(D, V)` or `(D, V, ϵ)`) by
modulus and then by imaginary part of the eigenvalues. Truncation by value keeps the retained
eigenpairs in LAPACK's ordering, which can change discontinuously under small perturbations
of `A` and breaks finite-difference checks; see [`eig_vals_wrapper`](@ref).
"""
eig_trunc_wrapper(f, A, alg) = sort_eig_trunc(f(A, alg))

"""
eig_trunc!_wrapper(f!, A, alg)

In-place variant of [`eig_trunc_wrapper`](@ref), which zeros `A` after calling `f!`.
"""
eig_trunc!_wrapper(f!, A, alg) = sort_eig_trunc(call_and_zero!(f!, A, alg))

# sortperm is used here because Mooncake CAN differentiate that on CUDA,
# but CANNOT differentiate sort
eigvals_sortperm(D) = sortperm(collect(D); by = λ -> (abs(λ), imag(λ)))
sort_eigvals(D) = D[eigvals_sortperm(D)]

# reorder `(D, V, rest...)` (or its tangent) by the permutation `p` of the eigenvalues
permute_eigpairs(DV, p) = (Diagonal(diagview(DV[1])[p]), permute_columns(DV[2], p), Base.tail(Base.tail(DV))...)

# equivalent to `V[:, p]`, but written with linear indexing because Mooncake CAN differentiate
# that on CUDA, but CANNOT differentiate `V[:, p]`
function permute_columns(V, p)
m = size(V, 1)
lin = vec((1:m) .+ m .* (p' .- 1))
return reshape(vec(V)[lin], m, length(p))
end
sort_eig_trunc(DV) = permute_eigpairs(DV, eigvals_sortperm(diagview(DV[1])))

"""
qr_gauge_invariant_wrapper(f, A, alg, r)

Expand Down Expand Up @@ -149,11 +200,27 @@ function stabilize_eigvals!(D::AbstractVector)
n = maximum(p)
# rescale eigenvalues so that they lie on distinct radii in the complex plane
# that are chosen randomly in non-overlapping intervals [10 * k/n, 10 * (k+0.5)/n)] for k=1,...,n
radii = 10 .* ((1:n) .+ rand(real(eltype(D)), n) ./ 2) ./ n
radii = 10 .* ((1:n) .+ rand(rng, real(eltype(D)), n) ./ 2) ./ n
hD = sign.(collect(D)) .* radii[p]
copyto!(D, hD)
return D
end
"""
midgap_tol(vals)

Return a truncation tolerance halfway across the widest gap between consecutive values of
`abs.(vals)`, restricted to the middle half so that truncation keeps a nontrivial subset.
This keeps the number of retained values fixed under the perturbations used by
finite-difference checks.
"""
function midgap_tol(vals)
s = sort!(collect(abs.(vals)))
Comment thread
lkdvos marked this conversation as resolved.
n = length(s)
gaps = (max(1, n ÷ 4)):(min(n - 1, (3n) ÷ 4))
_, i = findmax(i -> s[i + 1] - s[i], gaps)
return (s[gaps[i]] + s[gaps[i] + 1]) / 2
end

function make_eig_matrix(T, sz)
A = instantiate_matrix(T, sz)
D, V = eig_full(A)
Expand Down
6 changes: 3 additions & 3 deletions test/testsuite/chainrules.jl
Original file line number Diff line number Diff line change
Expand Up @@ -442,7 +442,7 @@ function test_chainrules_eigh(
@test isequal(ΔDVtrunc, ΔDVtrunc_copy)
end
D, ΔD = ad_eigh_vals_setup(A / 2)
truncalg = TruncatedAlgorithm(alg, trunctol(; atol = maximum(abs, D) / 2))
truncalg = TruncatedAlgorithm(alg, trunctol(; atol = midgap_tol(eigh_vals(A))))
DV, DVtrunc, ΔDV, ΔDVtrunc = ad_eigh_trunc_setup(A, truncalg)
ind = MatrixAlgebraKit.findtruncated(diagview(DV[1]), truncalg.trunc)
ot = (ΔDVtrunc..., zero(real(T)))
Expand Down Expand Up @@ -601,7 +601,7 @@ function test_chainrules_svd(
@test isequal(ΔUSVᴴtrunc, ΔUSVᴴtrunc_copy)
end
S, ΔS = ad_svd_vals_setup(A)
truncalg = TruncatedAlgorithm(alg, trunctol(atol = S[1, 1] / 2))
truncalg = TruncatedAlgorithm(alg, trunctol(atol = midgap_tol(S)))
USVᴴ, _, ΔUSVᴴ, ΔUSVᴴtrunc = ad_svd_trunc_setup(A, truncalg)
ot = (ΔUSVᴴtrunc..., zero(real(T)))
ot_copy = deepcopy(ot)
Expand All @@ -625,7 +625,7 @@ function test_chainrules_svd(
dA1 = MatrixAlgebraKit.svd_pullback!(zero(A), A, USVᴴ, ΔUSVᴴtrunc, ind)
dA2 = MatrixAlgebraKit.svd_trunc_pullback!(zero(A), A, (Utrunc, Strunc, Vᴴtrunc), ΔUSVᴴtrunc)
@test isapprox(dA1, dA2; atol = atol, rtol = rtol)
trunc = trunctol(; atol = S[1, 1] / 2)
trunc = truncalg.trunc
ind = MatrixAlgebraKit.findtruncated(diagview(S), trunc)
ot = (ΔUSVᴴtrunc..., zero(real(T)))
ot_copy = deepcopy(ot)
Expand Down
28 changes: 15 additions & 13 deletions test/testsuite/enzyme/eig.jl
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ Test the Enzyme foward- and reverse-mode AD rule for `eig_full` and its in-place
"""
function test_enzyme_eig_full(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "eig_full: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -40,17 +40,17 @@ Test the Enzyme forward- and reverse-mode AD rule for `eig_vals` and its in-plac
"""
function test_enzyme_eig_vals(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "eig_vals: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
A = make_eig_matrix(T, sz)
alg = MatrixAlgebraKit.select_algorithm(eig_vals, A)
D, ΔD = ad_eig_vals_setup(A)
test_reverse(eig_vals, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm)
test_reverse(call_and_zero!, RT, (eig_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm)
test_forward(eig_vals, RT, (A, TA), (alg, Const); atol, rtol, fdm)
test_forward(call_and_zero!, RT, (eig_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm)
test_reverse(eig_vals_wrapper, RT, (eig_vals, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm)
test_reverse(eig_vals!_wrapper, RT, (eig_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm)
test_forward(eig_vals_wrapper, RT, (eig_vals, Const), (A, TA), (alg, Const); atol, rtol, fdm)
test_forward(eig_vals!_wrapper, RT, (eig_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm)
end
end

Expand All @@ -62,7 +62,7 @@ in-place variants, over a range of truncation ranks and a tolerance-based trunca
"""
function test_enzyme_eig_trunc(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "eig_trunc reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -83,14 +83,16 @@ function test_enzyme_eig_trunc(
@testset "trunctol" begin
A = make_eig_matrix(T, sz)
D = eig_vals(A)
trunc = trunctol(atol = maximum(abs, D) / 2; by = abs)
trunc = trunctol(atol = midgap_tol(D); by = abs)
truncalg = TruncatedAlgorithm(alg, trunc)
DV, _, ΔDV, ΔDVtrunc = ad_eig_trunc_setup(A, truncalg)
test_reverse(eig_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm)
test_reverse(call_and_zero!, RT, (eig_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm)
DV, DVtrunc, ΔDV, ΔDVtrunc = ad_eig_trunc_setup(A, truncalg)
# trunctol keeps LAPACK's eigenvalue ordering, so sort the outputs (and the tangent)
ΔDVtrunc = permute_eigpairs(ΔDVtrunc, eigvals_sortperm(diagview(DVtrunc[1])))
test_reverse(eig_trunc_wrapper, RT, (eig_trunc_no_error, Const), (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm)
test_reverse(eig_trunc!_wrapper, RT, (eig_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm)
# use max range here to try to dodge issues when the gap between eigenvalues is close to the FD perturbation
test_forward(eig_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, fdm = EnzymeTestUtils.FiniteDifferences.central_fdm(5, 1, max_range = 1.0e-3))
test_forward(call_and_zero!, RT, (eig_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, fdm = EnzymeTestUtils.FiniteDifferences.central_fdm(5, 1, max_range = 1.0e-3))
test_forward(eig_trunc_wrapper, RT, (eig_trunc_no_error, Const), (A, TA), (truncalg, Const); atol, rtol, fdm = EnzymeTestUtils.FiniteDifferences.central_fdm(5, 1, max_range = 1.0e-3))
test_forward(eig_trunc!_wrapper, RT, (eig_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, fdm = EnzymeTestUtils.FiniteDifferences.central_fdm(5, 1, max_range = 1.0e-3))
end
end
end
10 changes: 5 additions & 5 deletions test/testsuite/enzyme/eigh.jl
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `eigh_full` and its in-pla
"""
function test_enzyme_eigh_full(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "eigh_full: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -40,7 +40,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `eigh_vals` and its in-pla
"""
function test_enzyme_eigh_vals(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "eigh_vals: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -62,7 +62,7 @@ in-place variants, over a range of truncation ranks.
"""
function test_enzyme_eigh_trunc(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "eigh_trunc reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -82,8 +82,8 @@ function test_enzyme_eigh_trunc(
end
@testset "trunctol" begin
A = make_eigh_matrix(T, sz)
D = eigh_vals(A / 2, alg)
trunc = trunctol(; atol = maximum(abs, D) / 2)
D = eigh_vals(A, alg)
trunc = trunctol(; atol = midgap_tol(D))
truncalg = TruncatedAlgorithm(alg, trunc)
DV, _, ΔDV, ΔDVtrunc = ad_eigh_trunc_setup(A, truncalg)
test_reverse(eigh_wrapper, RT, (eigh_trunc_no_error, Const), (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm)
Expand Down
8 changes: 4 additions & 4 deletions test/testsuite/enzyme/lq.jl
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ end

function test_enzyme_lq_compact(
T::Type, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "lq_compact: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -31,7 +31,7 @@ end

function test_enzyme_lq_compact_rank_deficient(
T::Type, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "lq_compact rank deficient A: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -52,7 +52,7 @@ end

function test_enzyme_lq_full(
T::Type, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "lq_full reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -70,7 +70,7 @@ end

function test_enzyme_lq_null(
T::Type, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "lq_null reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand Down
8 changes: 4 additions & 4 deletions test/testsuite/enzyme/orthnull.jl
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@ algorithms, and their in-place variants.
"""
function test_enzyme_left_orth(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "left_orth reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand Down Expand Up @@ -61,7 +61,7 @@ algorithms, and their in-place variants.
"""
function test_enzyme_right_orth(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "right_orth reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand Down Expand Up @@ -99,7 +99,7 @@ in-place variant.
"""
function test_enzyme_left_null(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "left_null: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -124,7 +124,7 @@ in-place variant.
"""
function test_enzyme_right_null(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "right_null: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand Down
4 changes: 2 additions & 2 deletions test/testsuite/enzyme/polar.jl
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ Only runs for tall or square matrices (`m >= n`).
"""
function test_enzyme_left_polar(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T)
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T)
)
return @testset "left_polar: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
A = instantiate_matrix(T, sz)
Expand All @@ -44,7 +44,7 @@ Only runs for wide or square matrices (`m <= n`).
"""
function test_enzyme_right_polar(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T)
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T)
)
return @testset "right_polar: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
A = instantiate_matrix(T, sz)
Expand Down
4 changes: 2 additions & 2 deletions test/testsuite/enzyme/projections.jl
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `project_hermitian` and it
"""
function test_enzyme_project_hermitian(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "project_hermitian: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -42,7 +42,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `project_antihermitian` an
"""
function test_enzyme_project_antihermitian(
T, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "project_antihermitian: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand Down
8 changes: 4 additions & 4 deletions test/testsuite/enzyme/qr.jl
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ end

function test_enzyme_qr_compact(
T::Type, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "qr_compact reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -31,7 +31,7 @@ end

function test_enzyme_qr_compact_rank_deficient(
T::Type, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "qr_compact rank deficient A reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -52,7 +52,7 @@ end

function test_enzyme_qr_full(
T::Type, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "qr_full reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -70,7 +70,7 @@ end

function test_enzyme_qr_null(
T::Type, sz;
rng = Random.default_rng(), atol::Real = 0, rtol::Real = precision(T),
rng = TestSuite.rng, atol::Real = 0, rtol::Real = precision(T),
fdm = enzyme_fdm(T)
)
return @testset "qr_null reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand Down
Loading
Loading