From dea46d6ae53ac04cf9d32aa02f0563f4d27bc065 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 11:56:36 +0200 Subject: [PATCH] Fix a bug in the eig pushfwd and enable more tests --- src/pushforwards/eig.jl | 1 + test/testsuite/enzyme/eig.jl | 18 ++++++------------ test/testsuite/enzyme/eigh.jl | 14 +++++--------- 3 files changed, 12 insertions(+), 21 deletions(-) diff --git a/src/pushforwards/eig.jl b/src/pushforwards/eig.jl index 46f3de1e8..5862c7354 100644 --- a/src/pushforwards/eig.jl +++ b/src/pushforwards/eig.jl @@ -12,6 +12,7 @@ function eig_pushforward!( if !iszerotangent(ΔV) ∂K .*= inv_safe.(transpose(diagview(D)) .- diagview(D), degeneracy_atol) mul!(ΔV, V, ∂K) + ΔV .-= V .* real.(sum(conj.(V) .* ΔV; dims = 1)) if eltype(V) <: Complex # fix gauge for `gaugefix!` compatibility _, I = findmax(abs, V; dims = 1) infinitesimal_phases = imag.(ΔV[I] ./ V[I]) diff --git a/test/testsuite/enzyme/eig.jl b/test/testsuite/enzyme/eig.jl index b260e557b..1f11949eb 100644 --- a/test/testsuite/enzyme/eig.jl +++ b/test/testsuite/enzyme/eig.jl @@ -27,12 +27,9 @@ function test_enzyme_eig_full( alg = MatrixAlgebraKit.select_algorithm(eig_full, A) DV, ΔDV = ad_eig_full_setup(A) test_reverse(eig_full, RT, (A, TA), (alg, Const); atol, rtol, output_tangent = ΔDV, fdm) - test_reverse(call_and_zero!, RT, (eig_full!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔDV, fdm) - if eltype(T) <: Real && T <: Diagonal - A = make_eig_matrix(T, sz) - test_forward(eig_full, RT, (A, TA), (alg, Const); atol, rtol, fdm) - test_forward(call_and_zero!, RT, (eig_full!, Const), (A, TA), (alg, Const); atol, rtol, fdm) - end + test_reverse(call_and_zero!, RT, (eig_full!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔDV, fdm) + test_forward(eig_full, RT, (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (eig_full!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end end @@ -51,12 +48,9 @@ function test_enzyme_eig_vals( 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), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm) - if eltype(T) <: Real - A = make_eig_matrix(T, sz) - test_forward(eig_vals, RT, (A, TA), (alg, Const); atol, rtol, fdm) - test_forward(call_and_zero!, RT, (eig_vals!, Const), (A, TA), (alg, Const); atol, rtol, fdm) - end + 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) end end diff --git a/test/testsuite/enzyme/eigh.jl b/test/testsuite/enzyme/eigh.jl index 83c81296d..4e42d7b3c 100644 --- a/test/testsuite/enzyme/eigh.jl +++ b/test/testsuite/enzyme/eigh.jl @@ -27,12 +27,9 @@ function test_enzyme_eigh_full( alg = MatrixAlgebraKit.select_algorithm(eigh_full, A) DV, ΔDV = ad_eigh_full_setup(A) test_reverse(eigh_wrapper, RT, (eigh_full, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔDV, fdm) - test_reverse(eigh!_wrapper, RT, (eigh_full!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔDV, fdm) - if eltype(T) <: Real - A = make_eigh_matrix(T, sz) - test_forward(eigh_wrapper, RT, (eigh_full, Const), (A, TA), (alg, Const); atol, rtol, fdm) - test_forward(eigh!_wrapper, RT, (eigh_full!, Const), (A, TA), (alg, Const); atol, rtol, fdm) - end + test_reverse(eigh!_wrapper, RT, (eigh_full!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔDV, fdm) + test_forward(eigh_wrapper, RT, (eigh_full, Const), (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(eigh!_wrapper, RT, (eigh_full!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end end @@ -51,10 +48,9 @@ function test_enzyme_eigh_vals( alg = MatrixAlgebraKit.select_algorithm(eigh_vals, A) D, ΔD = ad_eigh_vals_setup(A) test_reverse(eigh_wrapper, RT, (eigh_vals, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm) - test_reverse(eigh!_wrapper, RT, (eigh_vals!, Const), (A, TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm) - A = make_eigh_matrix(T, sz) + test_reverse(eigh!_wrapper, RT, (eigh_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, output_tangent = ΔD, fdm) test_forward(eigh_wrapper, RT, (eigh_vals, Const), (A, TA), (alg, Const); atol, rtol, fdm) - test_forward(eigh!_wrapper, RT, (eigh_vals!, Const), (A, TA), (alg, Const); atol, rtol, fdm) + test_forward(eigh!_wrapper, RT, (eigh_vals!, Const), (copy(A), TA), (alg, Const); atol, rtol, fdm) end end