Skip to content
Open
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
39 changes: 38 additions & 1 deletion test/testsuite/ad_utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,27 @@ 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))

# sortperm is used here because Mooncake CAN differentiate that on CUDA,
# but CANNOT differentiate sort
sort_eigvals(D) = D[sortperm(collect(D); by = λ -> (abs(λ), imag(λ)))]
Comment thread
kshyatt marked this conversation as resolved.

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

Expand Down Expand Up @@ -149,11 +170,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)))

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

why do we need the collect, shouldn't abs already give a vector?

Suggested change
s = sort!(collect(abs.(vals)))
s = sort!(abs.(vals))

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Not if vals lives on the GPU, right?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

ah, I missed the GPU case :) doesn't sort! also work on the GPU though?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

yeah it does but I was skittish of the logic later on so I just did everything on CPU instead 😭

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
16 changes: 8 additions & 8 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,7 +83,7 @@ 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)
Expand Down
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
10 changes: 5 additions & 5 deletions test/testsuite/enzyme/svd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `svd_compact` and its in-p
"""
function test_enzyme_svd_compact(
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 "svd_compact: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -37,7 +37,7 @@ gauge-dependent extra columns of `U` and rows of `Vᴴ` are zeroed out in the co
"""
function test_enzyme_svd_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 "svd_full: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -60,7 +60,7 @@ Test the Enzyme forward- and reverse-mode AD rule for `svd_vals` and its in-plac
"""
function test_enzyme_svd_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 "svd_vals: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -82,7 +82,7 @@ in-place variants, over a range of truncation ranks and a tolerance-based trunca
"""
function test_enzyme_svd_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 "svd_trunc reverse: RT $RT, TA $TA" for RT in (Duplicated,), TA in (Duplicated,)
Expand All @@ -101,7 +101,7 @@ function test_enzyme_svd_trunc(
end
@testset "trunctol" begin
S = svd_vals(A, alg)
trunc = trunctol(atol = maximum(S) / 2)
trunc = trunctol(atol = midgap_tol(S))
truncalg = TruncatedAlgorithm(alg, trunc)
USVᴴ, _, ΔUSVᴴ, ΔUSVᴴtrunc = ad_svd_trunc_setup(A, truncalg)
test_reverse(svd_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm)
Expand Down
Loading
Loading