Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/CI.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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 }}
5 changes: 4 additions & 1 deletion Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -65,10 +65,13 @@ 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"
TupleTools = "1.5"
VectorInterface = "0.6"
julia = "1.10"

[sources]
MatrixAlgebraKit = {url = "https://github.com/QuantumKitHub/MatrixAlgebraKit.jl/", rev = "main"}
1 change: 1 addition & 0 deletions ext/TensorKitEnzymeExt/TensorKitEnzymeExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -15,5 +15,6 @@ include("utility.jl")
include("linalg.jl")
include("indexmanipulations.jl")
include("tensoroperations.jl")
include("factorizations.jl")

end
69 changes: 69 additions & 0 deletions ext/TensorKitEnzymeExt/factorizations.jl
Original file line number Diff line number Diff line change
@@ -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
1 change: 1 addition & 0 deletions ext/TensorKitEnzymeExt/utility.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
27 changes: 27 additions & 0 deletions ext/TensorKitEnzymeTestUtilsExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
1 change: 1 addition & 0 deletions src/factorizations/diagonal.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down
1 change: 1 addition & 0 deletions src/factorizations/factorizations.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
46 changes: 42 additions & 4 deletions src/factorizations/pullbacks.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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!(
Expand All @@ -26,21 +36,45 @@ 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

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::AbstractTensorMap, F, ΔF, inds = _notrunc_ind(t);
Δ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 = 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, t, F, ΔF, ::Colon; kwargs...
)
return MAK.$pullback!(Δt, t, F, ΔF, _notrunc_ind(Δt); kwargs...)
end
end

for pullback_trunc! in (:svd_trunc_pullback!, :eig_trunc_pullback!, :eigh_trunc_pullback!)
Expand Down Expand Up @@ -97,3 +131,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
66 changes: 66 additions & 0 deletions src/factorizations/pushforwards.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
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;
kwargs...
)
foreachblock(Δt, t) do c, (Δb, b)
Fc = nothing_or_block.(F, Ref(c))
ΔFc = nothing_or_block.(ΔF, Ref(c))
MAK.$pushforward!(Δb, b, Fc, ΔFc; kwargs...)
return nothing
end
return Δt
end
end
37 changes: 37 additions & 0 deletions test/enzyme-factorizations/eig.jl
Original file line number Diff line number Diff line change
@@ -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
39 changes: 39 additions & 0 deletions test/enzyme-factorizations/eigh.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
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]))
atol = default_tol(T)
rtol = default_tol(T)
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
Loading
Loading