diff --git a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl index 944833237..009ebee81 100644 --- a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl +++ b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl @@ -2,7 +2,7 @@ module MatrixAlgebraKitEnzymeExt using MatrixAlgebraKit using MatrixAlgebraKit: copy_input, initialize_output, zero!, has_equal_storage -using MatrixAlgebraKit: diagview, inv_safe, truncate +using MatrixAlgebraKit: diagview, inv_safe, truncate, TruncationStrategy using MatrixAlgebraKit: qr_pullback!, lq_pullback! using MatrixAlgebraKit: qr_pushforward!, lq_pushforward! using MatrixAlgebraKit: qr_null_pullback!, lq_null_pullback! @@ -19,7 +19,7 @@ using Enzyme.EnzymeCore: EnzymeRules using LinearAlgebra @inline EnzymeRules.inactive_type(::Type{Alg}) where {Alg <: MatrixAlgebraKit.AbstractAlgorithm} = true -@inline EnzymeRules.inactive_type(::Type{TS}) where {TS <: MatrixAlgebraKit.TruncationStrategy} = true +@inline EnzymeRules.inactive_type(::Type{TS}) where {TS <: TruncationStrategy} = true @inline EnzymeRules.inactive(::typeof(MatrixAlgebraKit.select_algorithm), func::F, A::AbstractMatrix, alg::Alg) where {F, Alg} = true @inline EnzymeRules.inactive(::typeof(MatrixAlgebraKit.default_algorithm), func::F, A::AbstractMatrix) where {F} = true @inline EnzymeRules.inactive(::typeof(MatrixAlgebraKit.check_input), func::F, A::AbstractMatrix, alg::Alg) where {F, Alg} = true @@ -552,4 +552,56 @@ function EnzymeRules.forward( end end +function EnzymeRules.forward( + config::EnzymeRules.FwdConfigWidth{1}, + func::Const{typeof(truncate)}, + ::Type{RT}, + f::Const{<:Union{typeof(eigh_trunc!), typeof(eig_trunc!)}}, + DV::Annotation, + strategy::Annotation{<:TruncationStrategy}, + ) where {RT} + D, V = DV.val + ind = MatrixAlgebraKit.findtruncated(diagview(D), strategy.val) + Dtrunc = Diagonal(diagview(D)[ind]) + Vtrunc = V[:, ind] + dDtrunc = isa(DV, Const) ? nothing : Diagonal(diagview(DV.dval[1])[ind]) + dVtrunc = isa(DV, Const) ? nothing : DV.dval[2][:, ind] + if EnzymeRules.needs_primal(config) && EnzymeRules.needs_shadow(config) + return Duplicated(((Dtrunc, Vtrunc), ind), ((dDtrunc, dVtrunc), make_zero(ind))) + elseif EnzymeRules.needs_primal(config) + return ((Dtrunc, Vtrunc), ind) + elseif EnzymeRules.needs_shadow(config) + return ((dDtrunc, dVtrunc), make_zero(ind)) + else + return nothing + end +end + +function EnzymeRules.forward( + config::EnzymeRules.FwdConfigWidth{1}, + func::Const{typeof(truncate)}, + ::Type{RT}, + f::Const{typeof(svd_trunc!)}, + USVᴴ::Annotation, + strategy::Annotation{<:TruncationStrategy}, + ) where {RT} + U, S, Vᴴ = USVᴴ.val + ind = MatrixAlgebraKit.findtruncated_svd(diagview(S), strategy.val) + Utrunc = U[:, ind] + Strunc = Diagonal(diagview(S)[ind]) + Vᴴtrunc = Vᴴ[ind, :] + dUtrunc = isa(USVᴴ, Const) ? nothing : USVᴴ.dval[1][:, ind] + dStrunc = isa(USVᴴ, Const) ? nothing : Diagonal(diagview(USVᴴ.dval[2])[ind]) + dVᴴtrunc = isa(USVᴴ, Const) ? nothing : USVᴴ.dval[3][ind, :] + if EnzymeRules.needs_primal(config) && EnzymeRules.needs_shadow(config) + return Duplicated(((Utrunc, Strunc, Vᴴtrunc), ind), ((dUtrunc, dStrunc, dVᴴtrunc), make_zero(ind))) + elseif EnzymeRules.needs_primal(config) + return ((Utrunc, Strunc, Vᴴtrunc), ind) + elseif EnzymeRules.needs_shadow(config) + return ((dUtrunc, dStrunc, dVᴴtrunc), make_zero(ind)) + else + return nothing + end +end + end diff --git a/test/testsuite/enzyme/eig.jl b/test/testsuite/enzyme/eig.jl index 1f11949eb..4eaf878f3 100644 --- a/test/testsuite/enzyme/eig.jl +++ b/test/testsuite/enzyme/eig.jl @@ -76,7 +76,9 @@ function test_enzyme_eig_trunc( A = make_eig_matrix(T, sz) 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), (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) + test_forward(eig_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (eig_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, fdm) end @testset "trunctol" begin A = make_eig_matrix(T, sz) @@ -85,7 +87,9 @@ function test_enzyme_eig_trunc( 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), (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) + 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)) end end end diff --git a/test/testsuite/enzyme/eigh.jl b/test/testsuite/enzyme/eigh.jl index 4e42d7b3c..8143678ec 100644 --- a/test/testsuite/enzyme/eigh.jl +++ b/test/testsuite/enzyme/eigh.jl @@ -76,7 +76,9 @@ function test_enzyme_eigh_trunc( A = make_eigh_matrix(T, sz) 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) - test_reverse(eigh!_wrapper, RT, (eigh_trunc_no_error!, Const), (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm) + test_reverse(eigh!_wrapper, RT, (eigh_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm) + test_forward(eigh_wrapper, RT, (eigh_trunc_no_error, Const), (A, TA), (truncalg, Const); atol, rtol, fdm) + test_forward(eigh!_wrapper, RT, (eigh_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, fdm) end @testset "trunctol" begin A = make_eigh_matrix(T, sz) @@ -85,7 +87,9 @@ function test_enzyme_eigh_trunc( 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) - test_reverse(eigh!_wrapper, RT, (eigh_trunc_no_error!, Const), (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm) + test_reverse(eigh!_wrapper, RT, (eigh_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔDVtrunc, fdm) + test_forward(eigh_wrapper, RT, (eigh_trunc_no_error, Const), (A, TA), (truncalg, Const); atol, rtol, fdm = EnzymeTestUtils.FiniteDifferences.central_fdm(5, 1, max_range = 1.0e-3)) + test_forward(eigh!_wrapper, RT, (eigh_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 diff --git a/test/testsuite/enzyme/svd.jl b/test/testsuite/enzyme/svd.jl index 07a35d6d4..cddd1b8ee 100644 --- a/test/testsuite/enzyme/svd.jl +++ b/test/testsuite/enzyme/svd.jl @@ -94,16 +94,20 @@ function test_enzyme_svd_trunc( trunc = truncrank(r) truncalg = TruncatedAlgorithm(alg, trunc) USVᴴ, _, ΔUSVᴴ, ΔUSVᴴtrunc = ad_svd_trunc_setup(A, truncalg) - test_reverse(svd_trunc_no_error, RT, (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm) + test_reverse(svd_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm) test_reverse(call_and_zero!, RT, (svd_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm) + test_forward(svd_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, fdm) + test_forward(call_and_zero!, RT, (svd_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, fdm) end @testset "trunctol" begin S = svd_vals(A, alg) trunc = trunctol(atol = maximum(S) / 2) truncalg = TruncatedAlgorithm(alg, trunc) USVᴴ, _, ΔUSVᴴ, ΔUSVᴴtrunc = ad_svd_trunc_setup(A, truncalg) - test_reverse(svd_trunc_no_error, RT, (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm) + test_reverse(svd_trunc_no_error, RT, (A, TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm) test_reverse(call_and_zero!, RT, (svd_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm) + test_forward(svd_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, (svd_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