From 3d23599bae3802b380ecfab8ea1a287288bf5bce Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 11:22:52 +0200 Subject: [PATCH 1/4] Forward rule for truncate and tests --- .../MatrixAlgebraKitEnzymeExt.jl | 56 ++++++++++++++++++- test/testsuite/enzyme/eig.jl | 8 ++- test/testsuite/enzyme/eigh.jl | 8 ++- test/testsuite/enzyme/svd.jl | 8 ++- 4 files changed, 72 insertions(+), 8 deletions(-) diff --git a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl index 944833237..dedbd7d1c 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::Const{<: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::Const{<: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 From 8a71c5af9f9cb734ca30b7097e62bf6e88fd186a Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 18:22:33 +0200 Subject: [PATCH 2/4] Loosen annotation for 1.10 --- ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl index dedbd7d1c..009ebee81 100644 --- a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl +++ b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl @@ -558,7 +558,7 @@ function EnzymeRules.forward( ::Type{RT}, f::Const{<:Union{typeof(eigh_trunc!), typeof(eig_trunc!)}}, DV::Annotation, - strategy::Const{<:TruncationStrategy}, + strategy::Annotation{<:TruncationStrategy}, ) where {RT} D, V = DV.val ind = MatrixAlgebraKit.findtruncated(diagview(D), strategy.val) @@ -583,7 +583,7 @@ function EnzymeRules.forward( ::Type{RT}, f::Const{typeof(svd_trunc!)}, USVᴴ::Annotation, - strategy::Const{<:TruncationStrategy}, + strategy::Annotation{<:TruncationStrategy}, ) where {RT} U, S, Vᴴ = USVᴴ.val ind = MatrixAlgebraKit.findtruncated_svd(diagview(S), strategy.val) From d3bf893739cf55f96016441ec6fa6cf2d61a6187 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 2 Oct 2026 16:04:48 +0200 Subject: [PATCH 3/4] Reuse truncate Co-authored-by: Lukas Devos --- .../MatrixAlgebraKitEnzymeExt.jl | 11 ++--------- 1 file changed, 2 insertions(+), 9 deletions(-) diff --git a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl index 009ebee81..bdbf82347 100644 --- a/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl +++ b/ext/MatrixAlgebraKitEnzymeExt/MatrixAlgebraKitEnzymeExt.jl @@ -560,10 +560,7 @@ function EnzymeRules.forward( 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] + (Dtrunc, Vtrunc), ind = MatrixAlgebraKit.truncate(f.val, DV.val, strategy.val) 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) @@ -585,11 +582,7 @@ function EnzymeRules.forward( 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, :] + (Utrunc, Strunc, Vᴴtrunc), ind = MatrixAlgebraKit.truncate(f.val, USVᴴ.val, strategy.val) 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, :] From 20f4eb1e07e59849e4490e4095e10ad398cf398b Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 2 Oct 2026 16:06:54 +0200 Subject: [PATCH 4/4] Explanatory comments about max_range --- test/testsuite/enzyme/eig.jl | 1 + test/testsuite/enzyme/eigh.jl | 1 + test/testsuite/enzyme/svd.jl | 1 + 3 files changed, 3 insertions(+) diff --git a/test/testsuite/enzyme/eig.jl b/test/testsuite/enzyme/eig.jl index 4eaf878f3..a3e7107ee 100644 --- a/test/testsuite/enzyme/eig.jl +++ b/test/testsuite/enzyme/eig.jl @@ -88,6 +88,7 @@ function test_enzyme_eig_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) + # 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)) end diff --git a/test/testsuite/enzyme/eigh.jl b/test/testsuite/enzyme/eigh.jl index 8143678ec..a6a73eee8 100644 --- a/test/testsuite/enzyme/eigh.jl +++ b/test/testsuite/enzyme/eigh.jl @@ -88,6 +88,7 @@ function test_enzyme_eigh_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), (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(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 diff --git a/test/testsuite/enzyme/svd.jl b/test/testsuite/enzyme/svd.jl index cddd1b8ee..b5309b5a9 100644 --- a/test/testsuite/enzyme/svd.jl +++ b/test/testsuite/enzyme/svd.jl @@ -106,6 +106,7 @@ function test_enzyme_svd_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) test_reverse(call_and_zero!, RT, (svd_trunc_no_error!, Const), (copy(A), TA), (truncalg, Const); atol, rtol, output_tangent = ΔUSVᴴtrunc, fdm) + # use max range here to try to dodge issues when the gap between eigenvalues is close to the FD perturbation 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