From 3658cb3ad02ba29887fb4b1eb29725881a136742 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 25 Jun 2026 21:40:48 +0200 Subject: [PATCH 1/9] Add Enzyme rules for factorizations --- .github/workflows/CI.yml | 2 +- Project.toml | 2 +- ext/TensorKitEnzymeExt/TensorKitEnzymeExt.jl | 1 + ext/TensorKitEnzymeExt/factorizations.jl | 69 ++++++++++ ext/TensorKitEnzymeExt/utility.jl | 1 + src/factorizations/diagonal.jl | 1 + src/factorizations/pullbacks.jl | 41 ++++++ test/Project.toml | 1 + test/enzyme-factorizations/factorizations.jl | 137 +++++++++++++++++++ 9 files changed, 253 insertions(+), 2 deletions(-) create mode 100644 ext/TensorKitEnzymeExt/factorizations.jl create mode 100644 test/enzyme-factorizations/factorizations.jl diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index 3e620460d..8dad11b0a 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -22,6 +22,6 @@ jobs: with: fast: "${{ github.event.pull_request.draft == true }}" exclude: '["cuda", "amd"]' - timeout-minutes: 120 + timeout-minutes: 150 secrets: CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} diff --git a/Project.toml b/Project.toml index 49cb54457..5e279ce87 100644 --- a/Project.toml +++ b/Project.toml @@ -65,7 +65,7 @@ Preferences = "1" Printf = "1" Random = "1" ScopedValues = "1.3.0" -Strided = "2.6.1" +Strided = "2.6.3" TensorKitSectors = "0.3.7" TensorOperations = "5.5.2, 5.6" TimerOutputs = "1" diff --git a/ext/TensorKitEnzymeExt/TensorKitEnzymeExt.jl b/ext/TensorKitEnzymeExt/TensorKitEnzymeExt.jl index cadae496a..2095b3c52 100644 --- a/ext/TensorKitEnzymeExt/TensorKitEnzymeExt.jl +++ b/ext/TensorKitEnzymeExt/TensorKitEnzymeExt.jl @@ -15,5 +15,6 @@ include("utility.jl") include("linalg.jl") include("indexmanipulations.jl") include("tensoroperations.jl") +include("factorizations.jl") end diff --git a/ext/TensorKitEnzymeExt/factorizations.jl b/ext/TensorKitEnzymeExt/factorizations.jl new file mode 100644 index 000000000..4e6b6e962 --- /dev/null +++ b/ext/TensorKitEnzymeExt/factorizations.jl @@ -0,0 +1,69 @@ +# need these due to Enzyme choking on blocks + +for f in (:project_hermitian, :project_antihermitian) + f! = Symbol(f, :!) + @eval begin + function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + func::Const{typeof($f!)}, + ::Type{RT}, + A::Annotation{<:AbstractTensorMap}, + arg::Annotation{<:AbstractTensorMap}, + alg::Const, + ) where {RT} + $f!(A.val, arg.val, alg.val) + primal = EnzymeRules.needs_primal(config) ? arg.val : nothing + shadow = EnzymeRules.needs_shadow(config) ? arg.dval : nothing + cache = nothing + return EnzymeRules.AugmentedReturn(primal, shadow, cache) + end + function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + func::Const{typeof($f!)}, + ::Type{RT}, + cache, + A::Annotation{<:AbstractTensorMap}, + arg::Annotation{<:AbstractTensorMap}, + alg::Const, + ) where {RT} + if !isa(A, Const) && !isa(arg, Const) + $f!(arg.dval, arg.dval, alg.val) + if A.dval !== arg.dval + A.dval .+= arg.dval + make_zero!(arg.dval) + end + end + return (nothing, nothing, nothing) + end + function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + func::Const{typeof($f)}, + ::Type{RT}, + A::Annotation{<:AbstractTensorMap}, + alg::Const, + ) where {RT} + ret = $f(A.val, alg.val) + dret = make_zero(ret) + primal = EnzymeRules.needs_primal(config) ? ret : nothing + shadow = EnzymeRules.needs_shadow(config) ? dret : nothing + cache = dret + return EnzymeRules.AugmentedReturn(primal, shadow, cache) + end + function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + func::Const{typeof($f)}, + ::Type{RT}, + cache, + A::Annotation{<:AbstractTensorMap}, + alg::Const, + ) where {RT} + dret = cache + if !isa(A, Const) + $f!(dret, dret, alg.val) + add!(A.dval, dret) + end + make_zero!(dret) + return (nothing, nothing) + end + end +end diff --git a/ext/TensorKitEnzymeExt/utility.jl b/ext/TensorKitEnzymeExt/utility.jl index 35ce295fd..910423f4f 100644 --- a/ext/TensorKitEnzymeExt/utility.jl +++ b/ext/TensorKitEnzymeExt/utility.jl @@ -148,6 +148,7 @@ end @inline EnzymeRules.inactive(::typeof(TensorKit.insertleftunit), ::HomSpace, ::Any) = nothing @inline EnzymeRules.inactive(::typeof(TensorKit.insertrightunit), ::HomSpace, ::Any) = nothing @inline EnzymeRules.inactive(::typeof(TensorKit.removeunit), ::HomSpace, ::Any) = nothing +@inline EnzymeRules.inactive(::typeof(TensorKit.infimum), ::Any, ::Any) = nothing @inline EnzymeRules.inactive(::typeof(TensorKit.sectorstructure), ::Any) = nothing @inline EnzymeRules.inactive(::typeof(TensorKit.degeneracystructure), ::Any) = nothing @inline EnzymeRules.inactive(::typeof(TensorKit.select), s::HomSpace, i::Index2Tuple) = nothing diff --git a/src/factorizations/diagonal.jl b/src/factorizations/diagonal.jl index dae550ea1..49261f24e 100644 --- a/src/factorizations/diagonal.jl +++ b/src/factorizations/diagonal.jl @@ -3,6 +3,7 @@ _repack_diagonal(d::DiagonalTensorMap) = Diagonal(d.data) _repack_diagonal(d::SectorVector) = Diagonal(parent(d)) +MAK.diagonal(t::SectorVector) = DiagonalTensorMap(t) MAK.diagview(t::DiagonalTensorMap) = SectorVector(t.data, TensorKit.diagonalblockstructure(space(t))) for f in ( diff --git a/src/factorizations/pullbacks.jl b/src/factorizations/pullbacks.jl index a74acf8df..1b395bba5 100644 --- a/src/factorizations/pullbacks.jl +++ b/src/factorizations/pullbacks.jl @@ -11,6 +11,16 @@ for pullback! in ( end return Δt end + @eval function MAK.$pullback!( + Δt::AbstractTensorMap, ::Nothing, F, ΔF; kwargs... + ) + foreachblock(Δt) do c, (Δb,) + Fc = block.(F, Ref(c)) + ΔFc = block.(ΔF, Ref(c)) + return MAK.$pullback!(Δb, nothing, Fc, ΔFc; kwargs...) + end + return Δt + end end for pullback! in (:qr_null_pullback!, :lq_null_pullback!) @eval function MAK.$pullback!( @@ -41,6 +51,33 @@ for pullback! in (:svd_pullback!, :eig_pullback!, :eigh_pullback!) end return Δt end + @eval function MAK.$pullback!( + Δt::AbstractTensorMap, ::Nothing, F, ΔF, inds; kwargs... + ) + foreachblock(Δt) do c, (Δb,) + haskey(inds, c) || return nothing + ind = inds[c] + Fc = block.(F, Ref(c)) + ΔFc = map(ΔFc -> isnothing(ΔFc) ? nothing : block(ΔFc, c), ΔF) + return MAK.$pullback!(Δb, nothing, Fc, ΔFc, ind; kwargs...) + end + return Δt + end + @eval function MAK.$pullback!( + Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF, ::Colon; kwargs... + ) + return MAK.$pullback!(Δt, t, F, ΔF, _notrunc_ind(t); kwargs...) + end + @eval function MAK.$pullback!( + Δt::AbstractTensorMap, ::Nothing, F, ΔF; kwargs... + ) + return MAK.$pullback!(Δt, nothing, F, ΔF, _notrunc_ind(Δt); kwargs...) + end + @eval function MAK.$pullback!( + Δt::AbstractTensorMap, ::Nothing, F, ΔF, ::Colon; kwargs... + ) + return MAK.$pullback!(Δt, nothing, F, ΔF, _notrunc_ind(Δt); kwargs...) + end end for pullback_trunc! in (:svd_trunc_pullback!, :eig_trunc_pullback!, :eigh_trunc_pullback!) @@ -97,3 +134,7 @@ function MAK.remove_svd_gauge_dependence!( end return ΔU, ΔVᴴ end + +MAK.has_equal_storage(A::AbstractTensorMap, B::AbstractTensorMap) = A === B +MAK.has_equal_storage(A::AbstractTensorMap, B::SectorVector) = false +MAK.has_equal_storage(A::SectorVector, B::AbstractTensorMap) = false diff --git a/test/Project.toml b/test/Project.toml index ce51b70a3..a11773a35 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -32,6 +32,7 @@ cuTENSOR = "011b41b2-24ef-40a8-b3eb-fa098493e9e1" [sources] TensorKit = {path = ".."} +Enzyme = {url = "https://github.com/EnzymeAD/Enzyme.jl", rev = "ksh/nested-only-entry"} [compat] Aqua = "0.6, 0.7, 0.8" diff --git a/test/enzyme-factorizations/factorizations.jl b/test/enzyme-factorizations/factorizations.jl new file mode 100644 index 000000000..97ba2779b --- /dev/null +++ b/test/enzyme-factorizations/factorizations.jl @@ -0,0 +1,137 @@ +using Test, TestExtras +using TensorKit +using TensorOperations +using MatrixAlgebraKit +using MatrixAlgebraKit: remove_svd_gauge_dependence! +using MatrixAlgebraKit: remove_eig_gauge_dependence! +using MatrixAlgebraKit: remove_eigh_gauge_dependence! +using MatrixAlgebraKit: remove_lq_gauge_dependence!, remove_lq_null_gauge_dependence! +using MatrixAlgebraKit: remove_qr_gauge_dependence!, remove_qr_null_gauge_dependence! +using Enzyme, EnzymeTestUtils +using Random + +is_ci = get(ENV, "CI", "false") == "true" + +spacelist = ad_spacelist(fast_tests) +eltypes = (Float64, ComplexF64) + +@timedtestset "Enzyme - Factorizations: $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2]), randn(T, V[1] ⊗ V[2] ← (V[3] ⊗ V[4] ⊗ V[5])')) + atol = default_tol(T) + rtol = default_tol(T) + + @testset "SVD" begin + if !is_ci + S = svd_vals(t) + EnzymeTestUtils.test_reverse(svd_vals, Duplicated, (t, Duplicated); atol, rtol) + + USVᴴ = svd_full(t) + ΔUSVᴴ = EnzymeTestUtils.rand_tangent(USVᴴ) + remove_svd_gauge_dependence!(ΔUSVᴴ[1], ΔUSVᴴ[3], USVᴴ...) + EnzymeTestUtils.test_reverse(svd_full, Duplicated, (t, Duplicated); output_tangent = ΔUSVᴴ, atol, rtol) + + USVᴴ = svd_compact(t) + ΔUSVᴴ = EnzymeTestUtils.rand_tangent(USVᴴ) + remove_svd_gauge_dependence!(ΔUSVᴴ[1], ΔUSVᴴ[3], USVᴴ...) + EnzymeTestUtils.test_reverse(svd_compact, Duplicated, (t, Duplicated); output_tangent = ΔUSVᴴ, atol, rtol) + end + + V_trunc = spacetype(t)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) + trunc = truncspace(V_trunc) + alg = MatrixAlgebraKit.select_algorithm(svd_trunc_no_error, t, nothing; trunc) + USVᴴtrunc = svd_trunc_no_error(t, alg) + ΔUSVᴴtrunc = EnzymeTestUtils.rand_tangent(USVᴴtrunc) + remove_svd_gauge_dependence!(ΔUSVᴴtrunc[1], ΔUSVᴴtrunc[3], USVᴴtrunc...) + EnzymeTestUtils.test_reverse(svd_trunc_no_error, Duplicated, (t, Duplicated), (alg, Const); output_tangent = ΔUSVᴴtrunc, atol, rtol) + end + + @testset "LQ" begin + EnzymeTestUtils.test_reverse(lq_compact, Duplicated, (t, Duplicated); atol, rtol) + + if !is_ci + # lq_full/lq_null requires being careful with gauges + LQ = lq_full(t) + ΔLQ = EnzymeTestUtils.rand_tangent(LQ) + remove_lq_gauge_dependence!(ΔLQ..., t, LQ...) + EnzymeTestUtils.test_reverse(lq_full, Duplicated, (t, Duplicated); output_tangent = ΔLQ, atol, rtol) + + Nᴴ = lq_null(t) + Q = lq_compact(t)[2] + ΔNᴴ = EnzymeTestUtils.rand_tangent(Nᴴ) + remove_lq_null_gauge_dependence!(ΔNᴴ, Q, Nᴴ) + EnzymeTestUtils.test_reverse(lq_null, Duplicated, (t, Duplicated); output_tangent = ΔNᴴ, atol, rtol) + end + end + + @testset "QR" begin + EnzymeTestUtils.test_reverse(qr_compact, Duplicated, (t, Duplicated); atol, rtol) + + if !is_ci + # qr_full/qr_null requires being careful with gauges + QR = qr_full(t) + ΔQR = EnzymeTestUtils.rand_tangent(QR) + remove_qr_gauge_dependence!(ΔQR..., t, QR...) + EnzymeTestUtils.test_reverse(qr_full, Duplicated, (t, Duplicated); output_tangent = ΔQR, atol, rtol) + + N = qr_null(t) + Q = qr_compact(t)[1] + ΔN = EnzymeTestUtils.rand_tangent(N) + remove_qr_null_gauge_dependence!(ΔN, t, N) + EnzymeTestUtils.test_reverse(qr_null, Duplicated, (t, Duplicated); atol, rtol, output_tangent = ΔN) + end + end +end + +@timedtestset "Enzyme - Factorizations (EIGH/EIG): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ← V[1]), rand(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2])) + atol = default_tol(T) + rtol = default_tol(T) + + @testset "EIG" begin + if !is_ci + DV = eig_full(t) + ΔDV = EnzymeTestUtils.rand_tangent(DV) + remove_eig_gauge_dependence!(ΔDV[2], DV...) + EnzymeTestUtils.test_reverse(eig_full, Duplicated, (t, Duplicated); output_tangent = ΔDV, atol, rtol) + + D = eig_vals(t) + EnzymeTestUtils.test_reverse(eig_vals, Duplicated, (t, Duplicated); atol, rtol) + end + + V_trunc = spacetype(t)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) + trunc = truncspace(V_trunc) + alg = MatrixAlgebraKit.select_algorithm(eig_trunc_no_error, t, nothing; trunc) + DVtrunc = eig_trunc_no_error(t, alg) + ΔDVtrunc = EnzymeTestUtils.rand_tangent(DVtrunc) + remove_eig_gauge_dependence!(ΔDVtrunc[2], DVtrunc...) + EnzymeTestUtils.test_reverse(eig_trunc_no_error, Duplicated, (t, Duplicated), (alg, Const); output_tangent = ΔDVtrunc, atol, rtol) + end + + @testset "EIGH" begin + th = project_hermitian(t) + if !is_ci + DV = eigh_full(th) + ΔDV = EnzymeTestUtils.rand_tangent(DV) + remove_eigh_gauge_dependence!(ΔDV[2], DV...) + proj_eigh_full(t) = eigh_full(project_hermitian(t)) + EnzymeTestUtils.test_reverse(proj_eigh_full, Duplicated, (th, Duplicated); output_tangent = ΔDV, atol, rtol) + + D = eigh_vals(th) + EnzymeTestUtils.test_reverse(eigh_vals ∘ project_hermitian, Duplicated, (th, Duplicated); atol, rtol) + end + + V_trunc = spacetype(th)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) + trunc = truncspace(V_trunc) + alg = MatrixAlgebraKit.select_algorithm(eigh_trunc_no_error, th, nothing; trunc) + DVtrunc = eigh_trunc_no_error(th, alg) + ΔDVtrunc = EnzymeTestUtils.rand_tangent(DVtrunc) + remove_eigh_gauge_dependence!(ΔDVtrunc[2], DVtrunc...) + proj_eigh(t, alg) = eigh_trunc_no_error(project_hermitian(t), alg) + EnzymeTestUtils.test_reverse(proj_eigh, Duplicated, (th, Duplicated), (alg, Const); output_tangent = ΔDVtrunc, atol, rtol) + end + + @testset "Projections" begin + EnzymeTestUtils.test_reverse(project_hermitian, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_reverse(project_antihermitian, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_reverse(project_hermitian!, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_reverse(project_antihermitian!, Duplicated, (t, Duplicated); atol, rtol) + end +end From 27a3604e01eca6afe75d10b0f38d6bf481027379 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 14 Sep 2026 14:55:05 +0200 Subject: [PATCH 2/9] Remove stale sources entry --- test/Project.toml | 1 - 1 file changed, 1 deletion(-) diff --git a/test/Project.toml b/test/Project.toml index a11773a35..ce51b70a3 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -32,7 +32,6 @@ cuTENSOR = "011b41b2-24ef-40a8-b3eb-fa098493e9e1" [sources] TensorKit = {path = ".."} -Enzyme = {url = "https://github.com/EnzymeAD/Enzyme.jl", rev = "ksh/nested-only-entry"} [compat] Aqua = "0.6, 0.7, 0.8" From ffac8a17babb54a77fc8e6eaf38fd0f4431b7521 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 20:04:22 +0200 Subject: [PATCH 3/9] Forward tests too --- test/enzyme-factorizations/factorizations.jl | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/test/enzyme-factorizations/factorizations.jl b/test/enzyme-factorizations/factorizations.jl index 97ba2779b..69b4f1352 100644 --- a/test/enzyme-factorizations/factorizations.jl +++ b/test/enzyme-factorizations/factorizations.jl @@ -23,11 +23,13 @@ eltypes = (Float64, ComplexF64) if !is_ci S = svd_vals(t) EnzymeTestUtils.test_reverse(svd_vals, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(svd_vals, Duplicated, (t, Duplicated); atol, rtol) USVᴴ = svd_full(t) ΔUSVᴴ = EnzymeTestUtils.rand_tangent(USVᴴ) remove_svd_gauge_dependence!(ΔUSVᴴ[1], ΔUSVᴴ[3], USVᴴ...) EnzymeTestUtils.test_reverse(svd_full, Duplicated, (t, Duplicated); output_tangent = ΔUSVᴴ, atol, rtol) + EnzymeTestUtils.test_forward(svd_full, Duplicated, (t, Duplicated); atol, rtol) USVᴴ = svd_compact(t) ΔUSVᴴ = EnzymeTestUtils.rand_tangent(USVᴴ) @@ -46,6 +48,7 @@ eltypes = (Float64, ComplexF64) @testset "LQ" begin EnzymeTestUtils.test_reverse(lq_compact, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(lq_compact, Duplicated, (t, Duplicated); atol, rtol) if !is_ci # lq_full/lq_null requires being careful with gauges @@ -53,17 +56,20 @@ eltypes = (Float64, ComplexF64) ΔLQ = EnzymeTestUtils.rand_tangent(LQ) remove_lq_gauge_dependence!(ΔLQ..., t, LQ...) EnzymeTestUtils.test_reverse(lq_full, Duplicated, (t, Duplicated); output_tangent = ΔLQ, atol, rtol) + EnzymeTestUtils.test_forward(lq_full, Duplicated, (t, Duplicated); atol, rtol) Nᴴ = lq_null(t) Q = lq_compact(t)[2] ΔNᴴ = EnzymeTestUtils.rand_tangent(Nᴴ) remove_lq_null_gauge_dependence!(ΔNᴴ, Q, Nᴴ) EnzymeTestUtils.test_reverse(lq_null, Duplicated, (t, Duplicated); output_tangent = ΔNᴴ, atol, rtol) + EnzymeTestUtils.test_forward(lq_null, Duplicated, (t, Duplicated); atol, rtol) end end @testset "QR" begin EnzymeTestUtils.test_reverse(qr_compact, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(qr_compact, Duplicated, (t, Duplicated); atol, rtol) if !is_ci # qr_full/qr_null requires being careful with gauges @@ -71,12 +77,14 @@ eltypes = (Float64, ComplexF64) ΔQR = EnzymeTestUtils.rand_tangent(QR) remove_qr_gauge_dependence!(ΔQR..., t, QR...) EnzymeTestUtils.test_reverse(qr_full, Duplicated, (t, Duplicated); output_tangent = ΔQR, atol, rtol) + EnzymeTestUtils.test_forward(qr_full, Duplicated, (t, Duplicated); atol, rtol) N = qr_null(t) Q = qr_compact(t)[1] ΔN = EnzymeTestUtils.rand_tangent(N) remove_qr_null_gauge_dependence!(ΔN, t, N) EnzymeTestUtils.test_reverse(qr_null, Duplicated, (t, Duplicated); atol, rtol, output_tangent = ΔN) + EnzymeTestUtils.test_forward(qr_null, Duplicated, (t, Duplicated); atol, rtol) end end end @@ -91,9 +99,11 @@ end ΔDV = EnzymeTestUtils.rand_tangent(DV) remove_eig_gauge_dependence!(ΔDV[2], DV...) EnzymeTestUtils.test_reverse(eig_full, Duplicated, (t, Duplicated); output_tangent = ΔDV, atol, rtol) + EnzymeTestUtils.test_forward(eig_full, Duplicated, (t, Duplicated); atol, rtol) D = eig_vals(t) EnzymeTestUtils.test_reverse(eig_vals, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(eig_vals, Duplicated, (t, Duplicated); atol, rtol) end V_trunc = spacetype(t)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) @@ -113,9 +123,11 @@ end remove_eigh_gauge_dependence!(ΔDV[2], DV...) proj_eigh_full(t) = eigh_full(project_hermitian(t)) EnzymeTestUtils.test_reverse(proj_eigh_full, Duplicated, (th, Duplicated); output_tangent = ΔDV, atol, rtol) + EnzymeTestUtils.test_forward(proj_eigh_full, Duplicated, (th, Duplicated); atol, rtol) D = eigh_vals(th) EnzymeTestUtils.test_reverse(eigh_vals ∘ project_hermitian, Duplicated, (th, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(eigh_vals ∘ project_hermitian, Duplicated, (th, Duplicated); atol, rtol) end V_trunc = spacetype(th)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) @@ -133,5 +145,9 @@ end EnzymeTestUtils.test_reverse(project_antihermitian, Duplicated, (t, Duplicated); atol, rtol) EnzymeTestUtils.test_reverse(project_hermitian!, Duplicated, (t, Duplicated); atol, rtol) EnzymeTestUtils.test_reverse(project_antihermitian!, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(project_hermitian, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(project_antihermitian, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(project_hermitian!, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(project_antihermitian!, Duplicated, (t, Duplicated); atol, rtol) end end From 6a911a7e7f4930abcd8b10cf58898affa2e04d2c Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 30 Sep 2026 20:10:43 +0200 Subject: [PATCH 4/9] Use MAK main --- Project.toml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/Project.toml b/Project.toml index 5e279ce87..6c26a8359 100644 --- a/Project.toml +++ b/Project.toml @@ -72,3 +72,6 @@ TimerOutputs = "1" TupleTools = "1.5" VectorInterface = "0.6" julia = "1.10" + +[sources] +MatrixAlgebraKit = {url = "https://github.com/QuantumKitHub/MatrixAlgebraKit.jl/", rev = "main"} From 371153e6b88996b29e25e1fba6e0ce49e56e2678 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 09:48:39 +0200 Subject: [PATCH 5/9] Fixes for the vals methods --- ext/TensorKitEnzymeTestUtilsExt.jl | 27 +++++++++++++++++++++++++++ src/factorizations/pullbacks.jl | 15 +++++++++++++++ 2 files changed, 42 insertions(+) diff --git a/ext/TensorKitEnzymeTestUtilsExt.jl b/ext/TensorKitEnzymeTestUtilsExt.jl index 4a1f393b1..eefe6f929 100644 --- a/ext/TensorKitEnzymeTestUtilsExt.jl +++ b/ext/TensorKitEnzymeTestUtilsExt.jl @@ -43,6 +43,33 @@ function EnzymeTestUtils.to_vec(t::TensorKit.DiagonalTensorMap, seen_vecs::Enzym parent_vec, parent_t = to_vec(TensorMap(t), seen_vecs) return parent_vec, TensorKit.DiagonalTensorMap ∘ parent_t end +function EnzymeTestUtils.to_vec(v::TensorKit.SectorVector, seen_vecs::EnzymeTestUtils.AliasDict) + has_seen = haskey(seen_vecs, v) + is_const = Enzyme.Compiler.guaranteed_const(Core.Typeof(v)) + if has_seen || is_const + v_vec = Float32[] + else + vec_of_vecs = [b * TensorKit.sqrtdim(c) for (c, b) in pairs(v)] + v_vec, back = to_vec(vec_of_vecs) + seen_vecs[v] = v_vec + end + function SectorVector_from_vec(v_vec_new::AbstractVector, seen_xs::EnzymeTestUtils.AliasDict) + if xor(has_seen, haskey(seen_xs, v)) + throw(ErrorException("Arrays must be reconstructed in the same order as they are vectorized.")) + end + has_seen && return seen_xs[v] + is_const && return v + + v_new = similar(v) + vvec_of_vecs = back(v_vec_new) + for (i, (c, b)) in enumerate(pairs(v_new)) + scale!(b, vvec_of_vecs[i], TensorKit.invsqrtdim(c)) + end + seen_xs[v] = v_new + return v_new + end + return v_vec, SectorVector_from_vec +end # generate random tangents for testing function EnzymeTestUtils.rand_tangent(rng, t::TensorMap) diff --git a/src/factorizations/pullbacks.jl b/src/factorizations/pullbacks.jl index 1b395bba5..013150883 100644 --- a/src/factorizations/pullbacks.jl +++ b/src/factorizations/pullbacks.jl @@ -36,6 +36,21 @@ for pullback! in (:qr_null_pullback!, :lq_null_pullback!) end _notrunc_ind(t) = SectorDict(c => Colon() for c in blocksectors(t)) +for pullback! in (:eig_vals_pullback!, :eigh_vals_pullback!) + @eval function MAK.$pullback!( + Δt::AbstractTensorMap, ::Nothing, DV::Tuple{Diagonal, <:AbstractTensorMap}, ΔD, inds; + kwargs... + ) + return MAK.$pullback!(Δt, nothing, (MAK.diagonal(parent(DV[1])), DV[2]), ΔD, inds; kwargs...) + end +end +function MAK.svd_vals_pullback!( + Δt::AbstractTensorMap, ::Nothing, USVᴴ::Tuple{<:AbstractTensorMap, Diagonal, <:AbstractTensorMap}, ΔS, ind; + kwargs... + ) + return MAK.svd_vals_pullback!(Δt, nothing, (USVᴴ[1], MAK.diagonal(parent(USVᴴ[2])), USVᴴ[3]), ΔS, ind; kwargs...) +end + for pullback! in (:svd_pullback!, :eig_pullback!, :eigh_pullback!) @eval function MAK.$pullback!( Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF, inds = _notrunc_ind(t); From 5f0ee013f345270c4414d73c7f44ffece893ccef Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 14:48:24 +0200 Subject: [PATCH 6/9] Cleanup pbs and split test files --- src/factorizations/pullbacks.jl | 29 +--- test/enzyme-factorizations/eig.jl | 37 +++++ test/enzyme-factorizations/eigh.jl | 37 +++++ test/enzyme-factorizations/factorizations.jl | 153 ------------------- test/enzyme-factorizations/lq.jl | 36 +++++ test/enzyme-factorizations/projections.jl | 24 +++ test/enzyme-factorizations/qr.jl | 36 +++++ test/enzyme-factorizations/svd.jl | 41 +++++ 8 files changed, 215 insertions(+), 178 deletions(-) create mode 100644 test/enzyme-factorizations/eig.jl create mode 100644 test/enzyme-factorizations/eigh.jl delete mode 100644 test/enzyme-factorizations/factorizations.jl create mode 100644 test/enzyme-factorizations/lq.jl create mode 100644 test/enzyme-factorizations/projections.jl create mode 100644 test/enzyme-factorizations/qr.jl create mode 100644 test/enzyme-factorizations/svd.jl diff --git a/src/factorizations/pullbacks.jl b/src/factorizations/pullbacks.jl index 013150883..fd2551cdc 100644 --- a/src/factorizations/pullbacks.jl +++ b/src/factorizations/pullbacks.jl @@ -51,48 +51,27 @@ function MAK.svd_vals_pullback!( return MAK.svd_vals_pullback!(Δt, nothing, (USVᴴ[1], MAK.diagonal(parent(USVᴴ[2])), USVᴴ[3]), ΔS, ind; kwargs...) end +nothing_or_block(x, c) = isnothing(x) ? x : block(x, c) for pullback! in (:svd_pullback!, :eig_pullback!, :eigh_pullback!) @eval function MAK.$pullback!( - Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF, inds = _notrunc_ind(t); + Δt::AbstractTensorMap, t, F, ΔF, inds = _notrunc_ind(Δt); kwargs... ) foreachblock(Δt, t) do c, (Δb, b) ind = get(inds, c, nothing) isnothing(ind) && return nothing - Fc = block.(F, Ref(c)) - ΔFc = block.(ΔF, Ref(c)) + Fc = nothing_or_block.(F, Ref(c)) + ΔFc = nothing_or_block.(ΔF, Ref(c)) MAK.$pullback!(Δb, b, Fc, ΔFc, ind; kwargs...) return nothing end return Δt end - @eval function MAK.$pullback!( - Δt::AbstractTensorMap, ::Nothing, F, ΔF, inds; kwargs... - ) - foreachblock(Δt) do c, (Δb,) - haskey(inds, c) || return nothing - ind = inds[c] - Fc = block.(F, Ref(c)) - ΔFc = map(ΔFc -> isnothing(ΔFc) ? nothing : block(ΔFc, c), ΔF) - return MAK.$pullback!(Δb, nothing, Fc, ΔFc, ind; kwargs...) - end - return Δt - end @eval function MAK.$pullback!( Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF, ::Colon; kwargs... ) return MAK.$pullback!(Δt, t, F, ΔF, _notrunc_ind(t); kwargs...) end - @eval function MAK.$pullback!( - Δt::AbstractTensorMap, ::Nothing, F, ΔF; kwargs... - ) - return MAK.$pullback!(Δt, nothing, F, ΔF, _notrunc_ind(Δt); kwargs...) - end - @eval function MAK.$pullback!( - Δt::AbstractTensorMap, ::Nothing, F, ΔF, ::Colon; kwargs... - ) - return MAK.$pullback!(Δt, nothing, F, ΔF, _notrunc_ind(Δt); kwargs...) - end end for pullback_trunc! in (:svd_trunc_pullback!, :eig_trunc_pullback!, :eigh_trunc_pullback!) diff --git a/test/enzyme-factorizations/eig.jl b/test/enzyme-factorizations/eig.jl new file mode 100644 index 000000000..e78222925 --- /dev/null +++ b/test/enzyme-factorizations/eig.jl @@ -0,0 +1,37 @@ +using Test, TestExtras +using TensorKit +using TensorOperations +using MatrixAlgebraKit +using MatrixAlgebraKit: remove_eig_gauge_dependence! +using Enzyme, EnzymeTestUtils +using Random + +is_ci = get(ENV, "CI", "false") == "true" + +spacelist = ad_spacelist(fast_tests) +eltypes = (Float64, ComplexF64) + +@timedtestset "Enzyme - Factorizations (EIG): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ← V[1]), rand(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2])) + atol = default_tol(T) + rtol = default_tol(T) + + if !is_ci + DV = eig_full(t) + ΔDV = EnzymeTestUtils.rand_tangent(DV) + remove_eig_gauge_dependence!(ΔDV[2], DV...) + EnzymeTestUtils.test_reverse(eig_full, Duplicated, (t, Duplicated); output_tangent = ΔDV, atol, rtol) + EnzymeTestUtils.test_forward(eig_full, Duplicated, (t, Duplicated); atol, rtol) + + D = eig_vals(t) + EnzymeTestUtils.test_reverse(eig_vals, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(eig_vals, Duplicated, (t, Duplicated); atol, rtol) + end + + V_trunc = spacetype(t)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) + trunc = truncspace(V_trunc) + alg = MatrixAlgebraKit.select_algorithm(eig_trunc_no_error, t, nothing; trunc) + DVtrunc = eig_trunc_no_error(t, alg) + ΔDVtrunc = EnzymeTestUtils.rand_tangent(DVtrunc) + remove_eig_gauge_dependence!(ΔDVtrunc[2], DVtrunc...) + EnzymeTestUtils.test_reverse(eig_trunc_no_error, Duplicated, (t, Duplicated), (alg, Const); output_tangent = ΔDVtrunc, atol, rtol) +end diff --git a/test/enzyme-factorizations/eigh.jl b/test/enzyme-factorizations/eigh.jl new file mode 100644 index 000000000..9ec94aed8 --- /dev/null +++ b/test/enzyme-factorizations/eigh.jl @@ -0,0 +1,37 @@ +using Test, TestExtras +using TensorKit +using TensorOperations +using MatrixAlgebraKit +using MatrixAlgebraKit: remove_eigh_gauge_dependence! +using Enzyme, EnzymeTestUtils +using Random + +is_ci = get(ENV, "CI", "false") == "true" + +spacelist = ad_spacelist(fast_tests) +eltypes = (Float64, ComplexF64) + +@timedtestset "Enzyme - Factorizations (EIGH): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ← V[1]), rand(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2])) + th = project_hermitian(t) + if !is_ci + DV = eigh_full(th) + ΔDV = EnzymeTestUtils.rand_tangent(DV) + remove_eigh_gauge_dependence!(ΔDV[2], DV...) + proj_eigh_full(t) = eigh_full(project_hermitian(t)) + EnzymeTestUtils.test_reverse(proj_eigh_full, Duplicated, (th, Duplicated); output_tangent = ΔDV, atol, rtol) + EnzymeTestUtils.test_forward(proj_eigh_full, Duplicated, (th, Duplicated); atol, rtol) + + D = eigh_vals(th) + EnzymeTestUtils.test_reverse(eigh_vals ∘ project_hermitian, Duplicated, (th, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(eigh_vals ∘ project_hermitian, Duplicated, (th, Duplicated); atol, rtol) + end + + V_trunc = spacetype(th)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) + trunc = truncspace(V_trunc) + alg = MatrixAlgebraKit.select_algorithm(eigh_trunc_no_error, th, nothing; trunc) + DVtrunc = eigh_trunc_no_error(th, alg) + ΔDVtrunc = EnzymeTestUtils.rand_tangent(DVtrunc) + remove_eigh_gauge_dependence!(ΔDVtrunc[2], DVtrunc...) + proj_eigh(t, alg) = eigh_trunc_no_error(project_hermitian(t), alg) + EnzymeTestUtils.test_reverse(proj_eigh, Duplicated, (th, Duplicated), (alg, Const); output_tangent = ΔDVtrunc, atol, rtol) +end diff --git a/test/enzyme-factorizations/factorizations.jl b/test/enzyme-factorizations/factorizations.jl deleted file mode 100644 index 69b4f1352..000000000 --- a/test/enzyme-factorizations/factorizations.jl +++ /dev/null @@ -1,153 +0,0 @@ -using Test, TestExtras -using TensorKit -using TensorOperations -using MatrixAlgebraKit -using MatrixAlgebraKit: remove_svd_gauge_dependence! -using MatrixAlgebraKit: remove_eig_gauge_dependence! -using MatrixAlgebraKit: remove_eigh_gauge_dependence! -using MatrixAlgebraKit: remove_lq_gauge_dependence!, remove_lq_null_gauge_dependence! -using MatrixAlgebraKit: remove_qr_gauge_dependence!, remove_qr_null_gauge_dependence! -using Enzyme, EnzymeTestUtils -using Random - -is_ci = get(ENV, "CI", "false") == "true" - -spacelist = ad_spacelist(fast_tests) -eltypes = (Float64, ComplexF64) - -@timedtestset "Enzyme - Factorizations: $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2]), randn(T, V[1] ⊗ V[2] ← (V[3] ⊗ V[4] ⊗ V[5])')) - atol = default_tol(T) - rtol = default_tol(T) - - @testset "SVD" begin - if !is_ci - S = svd_vals(t) - EnzymeTestUtils.test_reverse(svd_vals, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(svd_vals, Duplicated, (t, Duplicated); atol, rtol) - - USVᴴ = svd_full(t) - ΔUSVᴴ = EnzymeTestUtils.rand_tangent(USVᴴ) - remove_svd_gauge_dependence!(ΔUSVᴴ[1], ΔUSVᴴ[3], USVᴴ...) - EnzymeTestUtils.test_reverse(svd_full, Duplicated, (t, Duplicated); output_tangent = ΔUSVᴴ, atol, rtol) - EnzymeTestUtils.test_forward(svd_full, Duplicated, (t, Duplicated); atol, rtol) - - USVᴴ = svd_compact(t) - ΔUSVᴴ = EnzymeTestUtils.rand_tangent(USVᴴ) - remove_svd_gauge_dependence!(ΔUSVᴴ[1], ΔUSVᴴ[3], USVᴴ...) - EnzymeTestUtils.test_reverse(svd_compact, Duplicated, (t, Duplicated); output_tangent = ΔUSVᴴ, atol, rtol) - end - - V_trunc = spacetype(t)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) - trunc = truncspace(V_trunc) - alg = MatrixAlgebraKit.select_algorithm(svd_trunc_no_error, t, nothing; trunc) - USVᴴtrunc = svd_trunc_no_error(t, alg) - ΔUSVᴴtrunc = EnzymeTestUtils.rand_tangent(USVᴴtrunc) - remove_svd_gauge_dependence!(ΔUSVᴴtrunc[1], ΔUSVᴴtrunc[3], USVᴴtrunc...) - EnzymeTestUtils.test_reverse(svd_trunc_no_error, Duplicated, (t, Duplicated), (alg, Const); output_tangent = ΔUSVᴴtrunc, atol, rtol) - end - - @testset "LQ" begin - EnzymeTestUtils.test_reverse(lq_compact, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(lq_compact, Duplicated, (t, Duplicated); atol, rtol) - - if !is_ci - # lq_full/lq_null requires being careful with gauges - LQ = lq_full(t) - ΔLQ = EnzymeTestUtils.rand_tangent(LQ) - remove_lq_gauge_dependence!(ΔLQ..., t, LQ...) - EnzymeTestUtils.test_reverse(lq_full, Duplicated, (t, Duplicated); output_tangent = ΔLQ, atol, rtol) - EnzymeTestUtils.test_forward(lq_full, Duplicated, (t, Duplicated); atol, rtol) - - Nᴴ = lq_null(t) - Q = lq_compact(t)[2] - ΔNᴴ = EnzymeTestUtils.rand_tangent(Nᴴ) - remove_lq_null_gauge_dependence!(ΔNᴴ, Q, Nᴴ) - EnzymeTestUtils.test_reverse(lq_null, Duplicated, (t, Duplicated); output_tangent = ΔNᴴ, atol, rtol) - EnzymeTestUtils.test_forward(lq_null, Duplicated, (t, Duplicated); atol, rtol) - end - end - - @testset "QR" begin - EnzymeTestUtils.test_reverse(qr_compact, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(qr_compact, Duplicated, (t, Duplicated); atol, rtol) - - if !is_ci - # qr_full/qr_null requires being careful with gauges - QR = qr_full(t) - ΔQR = EnzymeTestUtils.rand_tangent(QR) - remove_qr_gauge_dependence!(ΔQR..., t, QR...) - EnzymeTestUtils.test_reverse(qr_full, Duplicated, (t, Duplicated); output_tangent = ΔQR, atol, rtol) - EnzymeTestUtils.test_forward(qr_full, Duplicated, (t, Duplicated); atol, rtol) - - N = qr_null(t) - Q = qr_compact(t)[1] - ΔN = EnzymeTestUtils.rand_tangent(N) - remove_qr_null_gauge_dependence!(ΔN, t, N) - EnzymeTestUtils.test_reverse(qr_null, Duplicated, (t, Duplicated); atol, rtol, output_tangent = ΔN) - EnzymeTestUtils.test_forward(qr_null, Duplicated, (t, Duplicated); atol, rtol) - end - end -end - -@timedtestset "Enzyme - Factorizations (EIGH/EIG): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ← V[1]), rand(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2])) - atol = default_tol(T) - rtol = default_tol(T) - - @testset "EIG" begin - if !is_ci - DV = eig_full(t) - ΔDV = EnzymeTestUtils.rand_tangent(DV) - remove_eig_gauge_dependence!(ΔDV[2], DV...) - EnzymeTestUtils.test_reverse(eig_full, Duplicated, (t, Duplicated); output_tangent = ΔDV, atol, rtol) - EnzymeTestUtils.test_forward(eig_full, Duplicated, (t, Duplicated); atol, rtol) - - D = eig_vals(t) - EnzymeTestUtils.test_reverse(eig_vals, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(eig_vals, Duplicated, (t, Duplicated); atol, rtol) - end - - V_trunc = spacetype(t)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) - trunc = truncspace(V_trunc) - alg = MatrixAlgebraKit.select_algorithm(eig_trunc_no_error, t, nothing; trunc) - DVtrunc = eig_trunc_no_error(t, alg) - ΔDVtrunc = EnzymeTestUtils.rand_tangent(DVtrunc) - remove_eig_gauge_dependence!(ΔDVtrunc[2], DVtrunc...) - EnzymeTestUtils.test_reverse(eig_trunc_no_error, Duplicated, (t, Duplicated), (alg, Const); output_tangent = ΔDVtrunc, atol, rtol) - end - - @testset "EIGH" begin - th = project_hermitian(t) - if !is_ci - DV = eigh_full(th) - ΔDV = EnzymeTestUtils.rand_tangent(DV) - remove_eigh_gauge_dependence!(ΔDV[2], DV...) - proj_eigh_full(t) = eigh_full(project_hermitian(t)) - EnzymeTestUtils.test_reverse(proj_eigh_full, Duplicated, (th, Duplicated); output_tangent = ΔDV, atol, rtol) - EnzymeTestUtils.test_forward(proj_eigh_full, Duplicated, (th, Duplicated); atol, rtol) - - D = eigh_vals(th) - EnzymeTestUtils.test_reverse(eigh_vals ∘ project_hermitian, Duplicated, (th, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(eigh_vals ∘ project_hermitian, Duplicated, (th, Duplicated); atol, rtol) - end - - V_trunc = spacetype(th)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) - trunc = truncspace(V_trunc) - alg = MatrixAlgebraKit.select_algorithm(eigh_trunc_no_error, th, nothing; trunc) - DVtrunc = eigh_trunc_no_error(th, alg) - ΔDVtrunc = EnzymeTestUtils.rand_tangent(DVtrunc) - remove_eigh_gauge_dependence!(ΔDVtrunc[2], DVtrunc...) - proj_eigh(t, alg) = eigh_trunc_no_error(project_hermitian(t), alg) - EnzymeTestUtils.test_reverse(proj_eigh, Duplicated, (th, Duplicated), (alg, Const); output_tangent = ΔDVtrunc, atol, rtol) - end - - @testset "Projections" begin - EnzymeTestUtils.test_reverse(project_hermitian, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_reverse(project_antihermitian, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_reverse(project_hermitian!, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_reverse(project_antihermitian!, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(project_hermitian, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(project_antihermitian, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(project_hermitian!, Duplicated, (t, Duplicated); atol, rtol) - EnzymeTestUtils.test_forward(project_antihermitian!, Duplicated, (t, Duplicated); atol, rtol) - end -end diff --git a/test/enzyme-factorizations/lq.jl b/test/enzyme-factorizations/lq.jl new file mode 100644 index 000000000..fd5d650fd --- /dev/null +++ b/test/enzyme-factorizations/lq.jl @@ -0,0 +1,36 @@ +using Test, TestExtras +using TensorKit +using TensorOperations +using MatrixAlgebraKit +using MatrixAlgebraKit: remove_lq_gauge_dependence!, remove_lq_null_gauge_dependence! +using Enzyme, EnzymeTestUtils +using Random + +is_ci = get(ENV, "CI", "false") == "true" + +spacelist = ad_spacelist(fast_tests) +eltypes = (Float64, ComplexF64) + +@timedtestset "Enzyme - Factorizations (LQ): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2]), randn(T, V[1] ⊗ V[2] ← (V[3] ⊗ V[4] ⊗ V[5])')) + atol = default_tol(T) + rtol = default_tol(T) + + EnzymeTestUtils.test_reverse(lq_compact, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(lq_compact, Duplicated, (t, Duplicated); atol, rtol) + + if !is_ci + # lq_full/lq_null requires being careful with gauges + LQ = lq_full(t) + ΔLQ = EnzymeTestUtils.rand_tangent(LQ) + remove_lq_gauge_dependence!(ΔLQ..., t, LQ...) + EnzymeTestUtils.test_reverse(lq_full, Duplicated, (t, Duplicated); output_tangent = ΔLQ, atol, rtol) + EnzymeTestUtils.test_forward(lq_full, Duplicated, (t, Duplicated); atol, rtol) + + Nᴴ = lq_null(t) + Q = lq_compact(t)[2] + ΔNᴴ = EnzymeTestUtils.rand_tangent(Nᴴ) + remove_lq_null_gauge_dependence!(ΔNᴴ, Q, Nᴴ) + EnzymeTestUtils.test_reverse(lq_null, Duplicated, (t, Duplicated); output_tangent = ΔNᴴ, atol, rtol) + EnzymeTestUtils.test_forward(lq_null, Duplicated, (t, Duplicated); atol, rtol) + end +end diff --git a/test/enzyme-factorizations/projections.jl b/test/enzyme-factorizations/projections.jl new file mode 100644 index 000000000..fe3eadf6e --- /dev/null +++ b/test/enzyme-factorizations/projections.jl @@ -0,0 +1,24 @@ +using Test, TestExtras +using TensorKit +using TensorOperations +using MatrixAlgebraKit +using Enzyme, EnzymeTestUtils +using Random + +is_ci = get(ENV, "CI", "false") == "true" + +spacelist = ad_spacelist(fast_tests) +eltypes = (Float64, ComplexF64) + +@timedtestset "Enzyme - Factorizations (Projections): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ← V[1]), rand(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2])) + atol = default_tol(T) + rtol = default_tol(T) + EnzymeTestUtils.test_reverse(project_hermitian, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_reverse(project_antihermitian, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_reverse(project_hermitian!, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_reverse(project_antihermitian!, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(project_hermitian, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(project_antihermitian, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(project_hermitian!, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(project_antihermitian!, Duplicated, (t, Duplicated); atol, rtol) +end diff --git a/test/enzyme-factorizations/qr.jl b/test/enzyme-factorizations/qr.jl new file mode 100644 index 000000000..fb3cf7725 --- /dev/null +++ b/test/enzyme-factorizations/qr.jl @@ -0,0 +1,36 @@ +using Test, TestExtras +using TensorKit +using TensorOperations +using MatrixAlgebraKit +using MatrixAlgebraKit: remove_qr_gauge_dependence!, remove_qr_null_gauge_dependence! +using Enzyme, EnzymeTestUtils +using Random + +is_ci = get(ENV, "CI", "false") == "true" + +spacelist = ad_spacelist(fast_tests) +eltypes = (Float64, ComplexF64) + +@timedtestset "Enzyme - Factorizations (QR): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2]), randn(T, V[1] ⊗ V[2] ← (V[3] ⊗ V[4] ⊗ V[5])')) + atol = default_tol(T) + rtol = default_tol(T) + + EnzymeTestUtils.test_reverse(qr_compact, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(qr_compact, Duplicated, (t, Duplicated); atol, rtol) + + if !is_ci + # qr_full/qr_null requires being careful with gauges + QR = qr_full(t) + ΔQR = EnzymeTestUtils.rand_tangent(QR) + remove_qr_gauge_dependence!(ΔQR..., t, QR...) + EnzymeTestUtils.test_reverse(qr_full, Duplicated, (t, Duplicated); output_tangent = ΔQR, atol, rtol) + EnzymeTestUtils.test_forward(qr_full, Duplicated, (t, Duplicated); atol, rtol) + + N = qr_null(t) + Q = qr_compact(t)[1] + ΔN = EnzymeTestUtils.rand_tangent(N) + remove_qr_null_gauge_dependence!(ΔN, t, N) + EnzymeTestUtils.test_reverse(qr_null, Duplicated, (t, Duplicated); atol, rtol, output_tangent = ΔN) + EnzymeTestUtils.test_forward(qr_null, Duplicated, (t, Duplicated); atol, rtol) + end +end diff --git a/test/enzyme-factorizations/svd.jl b/test/enzyme-factorizations/svd.jl new file mode 100644 index 000000000..482e93b73 --- /dev/null +++ b/test/enzyme-factorizations/svd.jl @@ -0,0 +1,41 @@ +using Test, TestExtras +using TensorKit +using TensorOperations +using MatrixAlgebraKit +using MatrixAlgebraKit: remove_svd_gauge_dependence! +using Enzyme, EnzymeTestUtils +using Random + +is_ci = get(ENV, "CI", "false") == "true" + +spacelist = ad_spacelist(fast_tests) +eltypes = (Float64, ComplexF64) + +@timedtestset "Enzyme - Factorizations (SVD): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2]), randn(T, V[1] ⊗ V[2] ← (V[3] ⊗ V[4] ⊗ V[5])')) + atol = default_tol(T) + rtol = default_tol(T) + if !is_ci + S = svd_vals(t) + EnzymeTestUtils.test_reverse(svd_vals, Duplicated, (t, Duplicated); atol, rtol) + EnzymeTestUtils.test_forward(svd_vals, Duplicated, (t, Duplicated); atol, rtol) + + USVᴴ = svd_full(t) + ΔUSVᴴ = EnzymeTestUtils.rand_tangent(USVᴴ) + remove_svd_gauge_dependence!(ΔUSVᴴ[1], ΔUSVᴴ[3], USVᴴ...) + EnzymeTestUtils.test_reverse(svd_full, Duplicated, (t, Duplicated); output_tangent = ΔUSVᴴ, atol, rtol) + EnzymeTestUtils.test_forward(svd_full, Duplicated, (t, Duplicated); atol, rtol) + + USVᴴ = svd_compact(t) + ΔUSVᴴ = EnzymeTestUtils.rand_tangent(USVᴴ) + remove_svd_gauge_dependence!(ΔUSVᴴ[1], ΔUSVᴴ[3], USVᴴ...) + EnzymeTestUtils.test_reverse(svd_compact, Duplicated, (t, Duplicated); output_tangent = ΔUSVᴴ, atol, rtol) + end + + V_trunc = spacetype(t)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t)) + trunc = truncspace(V_trunc) + alg = MatrixAlgebraKit.select_algorithm(svd_trunc_no_error, t, nothing; trunc) + USVᴴtrunc = svd_trunc_no_error(t, alg) + ΔUSVᴴtrunc = EnzymeTestUtils.rand_tangent(USVᴴtrunc) + remove_svd_gauge_dependence!(ΔUSVᴴtrunc[1], ΔUSVᴴtrunc[3], USVᴴtrunc...) + EnzymeTestUtils.test_reverse(svd_trunc_no_error, Duplicated, (t, Duplicated), (alg, Const); output_tangent = ΔUSVᴴtrunc, atol, rtol) +end From 277993c52d977e655e3da5063e407eaefe3388d9 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Thu, 1 Oct 2026 18:33:54 +0200 Subject: [PATCH 7/9] Need the pushforward plumbing too --- src/factorizations/factorizations.jl | 1 + src/factorizations/pushforwards.jl | 73 ++++++++++++++++++++++++++++ 2 files changed, 74 insertions(+) create mode 100644 src/factorizations/pushforwards.jl diff --git a/src/factorizations/factorizations.jl b/src/factorizations/factorizations.jl index e49cc8b29..7bada491c 100644 --- a/src/factorizations/factorizations.jl +++ b/src/factorizations/factorizations.jl @@ -32,6 +32,7 @@ include("truncation.jl") include("adjoint.jl") include("diagonal.jl") include("pullbacks.jl") +include("pushforwards.jl") TensorKit.one!(A::AbstractMatrix) = MatrixAlgebraKit.one!(A) diff --git a/src/factorizations/pushforwards.jl b/src/factorizations/pushforwards.jl new file mode 100644 index 000000000..031ee6d37 --- /dev/null +++ b/src/factorizations/pushforwards.jl @@ -0,0 +1,73 @@ +for pushforward! in ( + :qr_pushforward!, :lq_pushforward!, :left_polar_pushforward!, :right_polar_pushforward!, + ) + @eval function MAK.$pushforward!( + Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF; kwargs... + ) + foreachblock(Δt, t) do c, (Δb, b) + Fc = block.(F, Ref(c)) + ΔFc = block.(ΔF, Ref(c)) + return MAK.$pushforward!(Δb, b, Fc, ΔFc; kwargs...) + end + return Δt + end + @eval function MAK.$pushforward!( + Δt::AbstractTensorMap, ::Nothing, F, ΔF; kwargs... + ) + foreachblock(Δt) do c, (Δb,) + Fc = block.(F, Ref(c)) + ΔFc = block.(ΔF, Ref(c)) + return MAK.$pushforward!(Δb, nothing, Fc, ΔFc; kwargs...) + end + return Δt + end +end +for pushforward! in (:qr_null_pushforward!, :lq_null_pushforward!) + @eval function MAK.$pushforward!( + Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF; kwargs... + ) + foreachblock(Δt, t) do c, (Δb, b) + Fc = block(F, c) + ΔFc = block(ΔF, c) + return MAK.$pushforward!(Δb, b, Fc, ΔFc; kwargs...) + end + return Δt + end +end + +for pushforward! in (:eig_vals_pushforward!, :eigh_vals_pushforward!) + @eval function MAK.$pushforward!( + Δt::AbstractTensorMap, ::Nothing, DV::Tuple{Diagonal, <:AbstractTensorMap}, ΔD, inds; + kwargs... + ) + return MAK.$pushforward!(Δt, nothing, (MAK.diagonal(parent(DV[1])), DV[2]), ΔD, inds; kwargs...) + end +end +function MAK.svd_vals_pushforward!( + Δt::AbstractTensorMap, ::Nothing, USVᴴ::Tuple{<:AbstractTensorMap, Diagonal, <:AbstractTensorMap}, ΔS, ind; + kwargs... + ) + return MAK.svd_vals_pushforward!(Δt, nothing, (USVᴴ[1], MAK.diagonal(parent(USVᴴ[2])), USVᴴ[3]), ΔS, ind; kwargs...) +end + +for pushforward! in (:svd_pushforward!, :eig_pushforward!, :eigh_pushforward!) + @eval function MAK.$pushforward!( + Δt::AbstractTensorMap, t, F, ΔF, inds = _notrunc_ind(Δt); + kwargs... + ) + foreachblock(Δt, t) do c, (Δb, b) + ind = get(inds, c, nothing) + isnothing(ind) && return nothing + Fc = nothing_or_block.(F, Ref(c)) + ΔFc = nothing_or_block.(ΔF, Ref(c)) + MAK.$pushforward!(Δb, b, Fc, ΔFc, ind; kwargs...) + return nothing + end + return Δt + end + @eval function MAK.$pushforward!( + Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF, ::Colon; kwargs... + ) + return MAK.$pushforward!(Δt, t, F, ΔF, _notrunc_ind(t); kwargs...) + end +end From 147ab9ec7e36a0553a8335ec9dc44151fd4a2e16 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 2 Oct 2026 10:45:39 +0200 Subject: [PATCH 8/9] Touchups --- src/factorizations/pullbacks.jl | 9 ++++++--- src/factorizations/pushforwards.jl | 11 ++--------- test/enzyme-factorizations/eigh.jl | 2 ++ 3 files changed, 10 insertions(+), 12 deletions(-) diff --git a/src/factorizations/pullbacks.jl b/src/factorizations/pullbacks.jl index fd2551cdc..7fd3e790a 100644 --- a/src/factorizations/pullbacks.jl +++ b/src/factorizations/pullbacks.jl @@ -52,12 +52,15 @@ function MAK.svd_vals_pullback!( end nothing_or_block(x, c) = isnothing(x) ? x : block(x, c) +nothing_or_block(x::Diagonal, c) = block(MAK.diagonal(parent(x)), c) +nothing_or_foreachblock(f, Δt, t) = isnothing(t) ? foreachblock(f, Δt) : foreachblock(f, Δt, t) for pullback! in (:svd_pullback!, :eig_pullback!, :eigh_pullback!) @eval function MAK.$pullback!( Δt::AbstractTensorMap, t, F, ΔF, inds = _notrunc_ind(Δt); kwargs... ) - foreachblock(Δt, t) do c, (Δb, b) + nothing_or_foreachblock(Δt, t) do c, Δbb + Δb, b = length(Δbb) == 1 ? (only(Δbb), nothing) : Δbb ind = get(inds, c, nothing) isnothing(ind) && return nothing Fc = nothing_or_block.(F, Ref(c)) @@ -68,9 +71,9 @@ for pullback! in (:svd_pullback!, :eig_pullback!, :eigh_pullback!) return Δt end @eval function MAK.$pullback!( - Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF, ::Colon; kwargs... + Δt::AbstractTensorMap, t, F, ΔF, ::Colon; kwargs... ) - return MAK.$pullback!(Δt, t, F, ΔF, _notrunc_ind(t); kwargs...) + return MAK.$pullback!(Δt, t, F, ΔF, _notrunc_ind(Δt); kwargs...) end end diff --git a/src/factorizations/pushforwards.jl b/src/factorizations/pushforwards.jl index 031ee6d37..5558f2978 100644 --- a/src/factorizations/pushforwards.jl +++ b/src/factorizations/pushforwards.jl @@ -52,22 +52,15 @@ end for pushforward! in (:svd_pushforward!, :eig_pushforward!, :eigh_pushforward!) @eval function MAK.$pushforward!( - Δt::AbstractTensorMap, t, F, ΔF, inds = _notrunc_ind(Δt); + Δt::AbstractTensorMap, t, F, ΔF; kwargs... ) foreachblock(Δt, t) do c, (Δb, b) - ind = get(inds, c, nothing) - isnothing(ind) && return nothing Fc = nothing_or_block.(F, Ref(c)) ΔFc = nothing_or_block.(ΔF, Ref(c)) - MAK.$pushforward!(Δb, b, Fc, ΔFc, ind; kwargs...) + MAK.$pushforward!(Δb, b, Fc, ΔFc; kwargs...) return nothing end return Δt end - @eval function MAK.$pushforward!( - Δt::AbstractTensorMap, t::AbstractTensorMap, F, ΔF, ::Colon; kwargs... - ) - return MAK.$pushforward!(Δt, t, F, ΔF, _notrunc_ind(t); kwargs...) - end end diff --git a/test/enzyme-factorizations/eigh.jl b/test/enzyme-factorizations/eigh.jl index 9ec94aed8..ad6fdb82f 100644 --- a/test/enzyme-factorizations/eigh.jl +++ b/test/enzyme-factorizations/eigh.jl @@ -12,6 +12,8 @@ spacelist = ad_spacelist(fast_tests) eltypes = (Float64, ComplexF64) @timedtestset "Enzyme - Factorizations (EIGH): $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes, t in (randn(T, V[1] ← V[1]), rand(T, V[1] ⊗ V[2] ← V[1] ⊗ V[2])) + atol = default_tol(T) + rtol = default_tol(T) th = project_hermitian(t) if !is_ci DV = eigh_full(th) From 1ff395353100b5dc819e4e748c30ed4fffa91226 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Fri, 2 Oct 2026 10:56:58 +0200 Subject: [PATCH 9/9] Formatter --- src/factorizations/pullbacks.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/factorizations/pullbacks.jl b/src/factorizations/pullbacks.jl index 7fd3e790a..c42231d4e 100644 --- a/src/factorizations/pullbacks.jl +++ b/src/factorizations/pullbacks.jl @@ -53,7 +53,7 @@ end nothing_or_block(x, c) = isnothing(x) ? x : block(x, c) nothing_or_block(x::Diagonal, c) = block(MAK.diagonal(parent(x)), c) -nothing_or_foreachblock(f, Δt, t) = isnothing(t) ? foreachblock(f, Δt) : foreachblock(f, Δt, t) +nothing_or_foreachblock(f, Δt, t) = isnothing(t) ? foreachblock(f, Δt) : foreachblock(f, Δt, t) for pullback! in (:svd_pullback!, :eig_pullback!, :eigh_pullback!) @eval function MAK.$pullback!( Δt::AbstractTensorMap, t, F, ΔF, inds = _notrunc_ind(Δt);