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