diff --git a/Project.toml b/Project.toml index 89027f73f..52b46f718 100644 --- a/Project.toml +++ b/Project.toml @@ -30,11 +30,18 @@ TupleTools = "9d95972d-f1c8-5527-a6e0-b4b365fa01f6" VectorInterface = "409d34a3-91d5-4945-b6ec-7529ddf182d8" Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f" +[weakdeps] +Enzyme = "7da242da-08ed-463a-9acd-ee780be4f1d9" + +[extensions] +PEPSKitEnzymeExt = "Enzyme" + [compat] Accessors = "0.1" ChainRulesCore = "1.0" Compat = "3.46, 4.2" DocStringExtensions = "0.9.3" +Enzyme = "0.13.208" FiniteDifferences = "0.12" KrylovKit = "0.9.5, 0.10" LinearAlgebra = "1" @@ -51,6 +58,10 @@ TensorKit = "0.16.5, 0.17" TensorKitTensors = "0.3.1" TensorOperations = "5" TupleTools = "1.6.0" -VectorInterface = "0.4, 0.5, 0.6" +VectorInterface = "0.4, 0.5, 0.6, 0.7" Zygote = "0.6, 0.7" julia = "1.10" + +[sources] +TensorKit = {url = "https://github.com/QuantumKitHub/TensorKit.jl", rev = "main"} +MPSKit = {url = "https://github.com/QuantumKitHub/MPSKit.jl", rev = "main"} diff --git a/ext/PEPSKitEnzymeExt.jl b/ext/PEPSKitEnzymeExt.jl new file mode 100644 index 000000000..59b70e8ec --- /dev/null +++ b/ext/PEPSKitEnzymeExt.jl @@ -0,0 +1,402 @@ +module PEPSKitEnzymeExt + +using PEPSKit, MPSKit, TensorKit, MatrixAlgebraKit +using PEPSKit: SVDAdjoint, EighAdjoint, QRAdjoint, CTMRGAlgorithm, FixedPointGradient, sdiag_pow, dtmap, dtmap!! +using PEPSKit: InfiniteSquareNetwork, InfinitePEPS, InfinitePEPO, _stack_tuples +using PEPSKit: unitcell, ket, bra, pepo +using PEPSKit: _periodic_getindex_dispatch +using PEPSKit: _split_corners_edges +using TensorKit: AbstractTensorMap +using ChainRulesCore: ignore_derivatives +using VectorInterface: add!, One +import PEPSKit: real_inner +using Enzyme +using Enzyme.EnzymeCore: EnzymeRules + +@inline EnzymeRules.inactive_type(::Type{<:SVDAdjoint}) = true +@inline EnzymeRules.inactive_type(::Type{<:QRAdjoint}) = true +@inline EnzymeRules.inactive_type(::Type{<:EighAdjoint}) = true +@inline EnzymeRules.inactive_type(::Type{<:CTMRGAlgorithm}) = true +@inline EnzymeRules.inactive_type(::Type{<:PEPSKit.GradientAlgorithm}) = true + +@inline EnzymeRules.inactive(::typeof(PEPSKit.checklattice), args...) = nothing +@inline EnzymeRules.inactive(::typeof(ignore_derivatives), args...) = nothing +@inline EnzymeRules.inactive(::typeof(PEPSKit.eachcoordinate), args...) = nothing + +function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + func::Const{typeof(MatrixAlgebraKit.svd_trunc_no_error)}, + ::Type{RT}, + t::Annotation, + alg::Const{<:SVDAdjoint{F, R}} + ) where {RT, F, R <: PEPSKit.FullPullback} + # requires access to the full decomposition + U, S, V⁺ = svd_compact(t.val, alg.val.fwd_alg.alg) + (Ũ, S̃, Ṽ⁺), inds = MatrixAlgebraKit.truncate(svd_trunc!, (U, S, V⁺), alg.val.fwd_alg.trunc) + truncerror = MatrixAlgebraKit.truncation_error(diagview(S), inds) + + output = (Ũ, S̃, Ṽ⁺, truncerror) + USVᴴtrunc = (Ũ, S̃, Ṽ⁺) + primal = EnzymeRules.needs_primal(config) ? USVᴴtrunc : nothing + dret = if EnzymeRules.needs_shadow(config) + (zero(USVᴴtrunc[1]), zero(USVᴴtrunc[2]), zero(USVᴴtrunc[3])) + else + nothing + end + return EnzymeRules.AugmentedReturn(primal, dret, (dret, (U, S, V⁺), inds)) +end + +function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + func::Const{typeof(MatrixAlgebraKit.svd_trunc_no_error)}, + ::Type{RT}, + cache, + t::Annotation, + alg::Const{<:SVDAdjoint{F, R}} + ) where {RT, F, R <: PEPSKit.FullPullback} + dUSVᴴtrunc, USV⁺, ind = cache + U, S, V⁺ = USV⁺ + gtol = PEPSKit._get_pullback_gauge_tol(alg.val.rrule_alg.verbosity) + if !isa(t, Const) + MatrixAlgebraKit.svd_pullback!( + t.dval, t.val, (U, S, V⁺), dUSVᴴtrunc, ind; + gauge_atol = gtol(dUSVᴴtrunc), degeneracy_atol = alg.val.rrule_alg.degeneracy_atol, + ) + end + return ntuple(Returns(nothing), 2) +end + +""" +Shared implementation of the CTMRG fixed-point gradient. + +Both the bare `leading_boundary` rule and the `hook_pullback` rule below need the +same augmented-forward work; they differ only in where the solver algorithm comes +from. Keeping one implementation means the two cannot drift apart. +""" +function _leading_boundary_augmented(config, envinit::Annotation, state::Annotation, alg::Const) + env, info = MPSKit.leading_boundary(envinit.val, state.val, alg.val) + alg_fixed = PEPSKit._set_fixed_truncation(alg.val) + alg_gauge = PEPSKit._scrambling_env_gauge(alg.val) + env_conv, _ = PEPSKit.ctmrg_iteration(InfiniteSquareNetwork(state.val), env, alg_fixed) + shadow = EnzymeRules.needs_shadow(config) ? Enzyme.make_zero((env, info)) : nothing + denv = isnothing(shadow) ? nothing : shadow[1] + primal = EnzymeRules.needs_primal(config) ? (env, info) : nothing + signs, corner_phases, edge_phases = PEPSKit.compute_gauge_fix_gauge( + env_conv, env, alg_gauge, + ) + cache = (env, denv, alg_fixed, signs, corner_phases, edge_phases) + return primal, shadow, cache +end + +function _leading_boundary_reverse!(config, cache, state::Annotation, solver_alg) + env, denv, alg_fixed, signs, corner_phases, edge_phases = cache + function gauge_fixed_iteration(A, x) + x′ = PEPSKit.ctmrg_iteration(InfiniteSquareNetwork(A), x, alg_fixed)[1] + return PEPSKit.fix_phases(x′, signs, corner_phases, edge_phases) + end + inner_mode = Enzyme.set_runtime_activity(ReverseSplitWithPrimal, config) + fwd, rev = Enzyme.autodiff_thunk(inner_mode, Const{typeof(gauge_fixed_iteration)}, Duplicated, typeof(state), Duplicated{typeof(env)}) + # NOTE: the vjp MUST NOT touch the caller's shadows. `denv` is the incoming + # cotangent and is simultaneously `∂E∂x` for the fixed-point solve, and + # `state.dval` already holds the gradient contributions accumulated by the + # rest of the reverse sweep. Zeroing either silently destroys the whole gradient. + # Instead we allocate fresh shadows per evaluation. The Krylov solver also retains + # the returned vectors, so they MUST NOT alias a buffer we reuse. + function vjp(Δ) + dstate = Enzyme.make_zero(state.val) + denv_scratch = Enzyme.make_zero(env) + state_dup = Duplicated(state.val, dstate) + env_dup = Duplicated(env, denv_scratch) + + # Enzyme's split-mode tape is single-use, and this vjp is called once per + # Krylov iteration, so build a fresh one for each evaluation. + tape, _, out_shadow = fwd(Const(gauge_fixed_iteration), state_dup, env_dup) + add!(out_shadow, Δ, One(), One()) + rev(Const(gauge_fixed_iteration), state_dup, env_dup, tape) + return dstate, denv_scratch + end + ∂f∂A(x)::typeof(state.val) = vjp(x)[1] + ∂f∂x(x)::typeof(env) = vjp(x)[2] + ∂A = PEPSKit.fixedpoint_gradient(denv, ∂f∂x, ∂f∂A, denv, solver_alg) + if !isa(state, Const) + add!(state.dval, ∂A, One(), One()) + end + return nothing +end + +""" +The solver algorithm `hook_pullback` was asked for. + +`alg_rrule = nothing` selects naive AD through the CTMRG iterations in the +ChainRules path; there is no Enzyme equivalent yet, so it falls back to the +default fixed-point solver, which is what this rule did unconditionally before. +""" +@inline function _fixedpoint_solver_alg(kw::NamedTuple) + gradmode = get(kw, :alg_rrule, nothing) + isnothing(gradmode) && return PEPSKit.FixedPointGradient().solver_alg + return gradmode.solver_alg +end + +const _LeadingBoundary = typeof(MPSKit.leading_boundary) + +function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(Core.kwcall)}, + ::Type{RT}, + kw::Const{<:NamedTuple}, + ::Const{typeof(PEPSKit.hook_pullback)}, + ::Const{_LeadingBoundary}, + envinit::Annotation, + state::Annotation, + alg::Const{<:CTMRGAlgorithm}, + ) where {RT} + primal, shadow, cache = _leading_boundary_augmented(config, envinit, state, alg) + return EnzymeRules.AugmentedReturn(primal, shadow, cache) +end + +function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(Core.kwcall)}, + ::Type{RT}, + cache, + kw::Const{<:NamedTuple}, + ::Const{typeof(PEPSKit.hook_pullback)}, + ::Const{_LeadingBoundary}, + envinit::Annotation, + state::Annotation, + alg::Const{<:CTMRGAlgorithm}, + ) where {RT} + _leading_boundary_reverse!(config, cache, state, _fixedpoint_solver_alg(kw.val)) + return ntuple(Returns(nothing), 6) +end + +function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(MPSKit.leading_boundary)}, + ::Type{RT}, + envinit::Annotation, + state::Annotation, + alg::Const{<:CTMRGAlgorithm} + ) where {RT} + primal, shadow, cache = _leading_boundary_augmented(config, envinit, state, alg) + return EnzymeRules.AugmentedReturn(primal, shadow, cache) +end + +function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(MPSKit.leading_boundary)}, + ::Type{RT}, + cache, + envinit::Annotation, + state::Annotation, + alg::Const{<:CTMRGAlgorithm} + ) where {RT} + _leading_boundary_reverse!(config, cache, state, PEPSKit.FixedPointGradient().solver_alg) + return ntuple(Returns(nothing), 3) +end + +@inline _dtmap_elem(src::Const, i) = Const(src.val[i]) +@inline _dtmap_elem(src::Annotation, i) = Duplicated(src.val[i], src.dval[i]) + +for pb in (:svd_pullback!, :eig_pullback!, :eigh_pullback!) + @eval function MatrixAlgebraKit.$pb( + Δt::AbstractTensorMap, ::Nothing, F, ΔF, + inds = TensorKit.SectorDict(c => Colon() for c in TensorKit.blocksectors(Δt)); + kwargs... + ) + for (c, Δb) in TensorKit.blocks(Δt) + haskey(inds, c) || continue + Fc = TensorKit.block.(F, Ref(c)) + ΔFc = TensorKit.block.(ΔF, Ref(c)) + MatrixAlgebraKit.$pb(Δb, nothing, Fc, ΔFc, inds[c]; kwargs...) + end + return Δt + end +end + +function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(_split_corners_edges)}, + ::Type{RT}, + ce::Annotation{<:AbstractArray}, + ) where {RT} + primal_val = (map(first, ce.val), map(last, ce.val)) + primal = EnzymeRules.needs_primal(config) ? primal_val : nothing + shadow = if EnzymeRules.needs_shadow(config) && !isa(ce, Const) + (map(first, ce.dval), map(last, ce.dval)) + else + nothing + end + return EnzymeRules.AugmentedReturn(primal, shadow, nothing) +end + +function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(_split_corners_edges)}, + ::Type{RT}, + cache, + ce::Annotation{<:AbstractArray}, + ) where {RT} + return (nothing,) +end + +""" +How Enzyme should carry the mapped function's return value. + +The CTMRG routines map to tensors, which travel as `Duplicated`. `expectation_value` +maps to scalars, and a scalar is immutable: asking for `Duplicated` there hands back a +`Base.RefValue` that the destination array cannot store. Scalars must be `Active`, +which also changes the reverse call -- the cotangent is passed in rather than +accumulated into a shadow. +""" +@inline _dtmap_retann(::Type{T}) where {T} = Duplicated{T} +@inline _dtmap_retann(::Type{T}) where {T <: Number} = Active{T} + +function _dtmap_augmented!(config, f::FA, dst, src) where {FA <: Annotation} + ET = eltype(src.val) + DT = eltype(dst.val) + SA = src isa Const ? Const{ET} : Duplicated{ET} + # Propagate the caller's runtime-activity setting into the nested thunk. + # Without this the inner differentiation runs with static activity while the + # outer one does not, and derivative contributions are silently dropped. + mode = Enzyme.set_runtime_activity(ReverseSplitWithPrimal, config) + fwd, rev = Enzyme.autodiff_thunk(mode, FA, _dtmap_retann(DT), SA) + + inds = collect(eachindex(src.val)) + tapes = Vector{Any}(undef, length(inds)) + elems = Vector{Any}(undef, length(inds)) + for (k, i) in enumerate(inds) + arg = _dtmap_elem(src, i) + tape, primal, shadow = fwd(f, arg) + dst.val[i] = primal + # an `Active` return has no shadow to store; its cotangent is seeded in reverse + if !isa(dst, Const) && !(DT <: Number) + dst.dval[i] = shadow + end + tapes[k] = tape + elems[k] = arg + end + return (rev, inds, tapes, elems) +end + +function _dtmap_reverse!(f::FA, dst, cache) where {FA <: Annotation} + rev, inds, tapes, elems = cache + DT = eltype(dst.val) + for k in eachindex(tapes) + if DT <: Number + seed = isa(dst, Const) ? zero(DT) : dst.dval[inds[k]] + rev(f, elems[k], seed, tapes[k]) + else + rev(f, elems[k], tapes[k]) + end + end + return nothing +end + +#= turn this off until the needsReRouting fix is merged at Enzyme +function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(dtmap!!)}, + ::Type{RT}, + f::FA, + dst::Annotation{<:AbstractArray}, + src::Annotation{<:AbstractArray}, + ) where {RT, FA <: Annotation} + cache = _dtmap_augmented!(config, f, dst, src) + primal = EnzymeRules.needs_primal(config) ? dst.val : nothing + shadow = if EnzymeRules.needs_shadow(config) && !isa(dst, Const) + dst.dval + else + nothing + end + return EnzymeRules.AugmentedReturn(primal, shadow, cache) +end + +function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(dtmap!!)}, + ::Type{RT}, + cache, + f::FA, + dst::Annotation{<:AbstractArray}, + src::Annotation{<:AbstractArray}, + ) where {RT, FA <: Annotation} + _dtmap_reverse!(f, dst, cache) + return (nothing, nothing, nothing) +end + +=# + +function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{Type{InfiniteSquareNetwork}}, + ::Type{RT}, + top::Annotation{<:InfinitePEPS}, + mid::Annotation{<:InfinitePEPO}, + bot::Annotation{<:InfinitePEPS}, + ) where {RT} + netw = InfiniteSquareNetwork(top.val, mid.val, bot.val) + primal = EnzymeRules.needs_primal(config) ? netw : nothing + shadow = EnzymeRules.needs_shadow(config) ? Enzyme.make_zero(netw) : nothing + return EnzymeRules.augmented_rule_return_type(config, RT)(primal, shadow, shadow) +end + +function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{Type{InfiniteSquareNetwork}}, + ::Type{RT}, + cache, + top::Annotation{<:InfinitePEPS}, + mid::Annotation{<:InfinitePEPO}, + bot::Annotation{<:InfinitePEPS}, + ) where {RT} + Δnetwork = cache + aliased = !isa(top, Const) && !isa(bot, Const) && (top.dval === bot.dval) + w = aliased ? 0.5 : 1.0 + !isa(top, Const) && add!(top.dval, InfinitePEPS(map(ket, unitcell(Δnetwork))), w, One()) + !isa(bot, Const) && add!(bot.dval, InfinitePEPS(map(bra, unitcell(Δnetwork))), w, One()) + !isa(mid, Const) && add!(mid.dval, InfinitePEPO(_stack_tuples(map(pepo, unitcell(Δnetwork)))), One(), One()) + return (nothing, nothing, nothing) +end + + +const _PeriodicElt = Union{AbstractTensorMap, Tuple{Vararg{AbstractTensorMap}}} + +function EnzymeRules.augmented_primal( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(_periodic_getindex_dispatch)}, + ::Type{RT}, + A::Annotation, + data::Annotation{<:AbstractArray{<:_PeriodicElt}}, + J::Annotation, + ) where {RT} + primal = _periodic_getindex_dispatch(A.val, data.val, J.val) + shadow = if EnzymeRules.needs_shadow(config) + # `data.dval === data.val` means Enzyme shared the container between + # primal and shadow because it considers it inactive. Accumulating into + # it would write the primal, so hand back a throwaway zero instead. + if isa(data, Const) || data.dval === data.val + Enzyme.make_zero(primal) + else + _periodic_getindex_dispatch(A.val, data.dval, J.val) + end + else + nothing + end + p = EnzymeRules.needs_primal(config) ? primal : nothing + return EnzymeRules.AugmentedReturn(p, shadow, nothing) +end + +function EnzymeRules.reverse( + config::EnzymeRules.RevConfigWidth{1}, + ::Const{typeof(_periodic_getindex_dispatch)}, + ::Type{RT}, + cache, + A::Annotation, + data::Annotation{<:AbstractArray{<:_PeriodicElt}}, + J::Annotation, + ) where {RT} + return (nothing, nothing, nothing) +end + +end diff --git a/src/algorithms/contractions/ctmrg/projector.jl b/src/algorithms/contractions/ctmrg/projector.jl index f8791d77f..56a4b5070 100644 --- a/src/algorithms/contractions/ctmrg/projector.jl +++ b/src/algorithms/contractions/ctmrg/projector.jl @@ -239,7 +239,7 @@ Right projector: ``` """ function contract_projectors(U, S, V, Q, Q_next) - isqS = sdiag_pow(S, -0.5) + isqS = TensorMap(sdiag_pow(S, -0.5)) P_left = Q_next * V' * isqS # use * to respect fermionic case P_right = isqS * U' * Q return P_left, P_right diff --git a/src/algorithms/ctmrg/ctmrg.jl b/src/algorithms/ctmrg/ctmrg.jl index f8cc6e53f..7f68717a2 100644 --- a/src/algorithms/ctmrg/ctmrg.jl +++ b/src/algorithms/ctmrg/ctmrg.jl @@ -117,7 +117,8 @@ function leading_boundary( CS, TS = ignore_derivatives() do return convergence_spectra(env₀, alg) end - η = one(real(scalartype(network))) + Tη = real(scalartype(network)) + η::Tη = one(Tη) ctmrg_loginit!(log, η, network, env₀) local info_iter converged = false @@ -217,7 +218,7 @@ Spectra of the tensors that [`convergence_tensors`](@ref) selects. """ function convergence_spectra(env::CTMRGEnv, alg) corners, edges = convergence_tensors(env, alg) - return map(C -> corner_spectrum(C, alg), corners), map(T -> edge_spectrum(T, alg), edges) + return stablemap(C -> corner_spectrum(C, alg), corners), stablemap(T -> edge_spectrum(T, alg), edges) end """ @@ -246,10 +247,10 @@ This determined either from the previous corner and edge spectra `CS_old` and `TS_old`, or alternatively, directly from the old environment. """ function calc_convergence(env, CS_old, TS_old) - CS_new = map(svd_vals, env.corners) + CS_new = stablemap(svd_vals, env.corners) ΔCS = maximum(splat(_singular_value_distance), zip(CS_old, CS_new)) - TS_new = map(svd_vals, env.edges) + TS_new = stablemap(svd_vals, env.edges) ΔTS = maximum(splat(_singular_value_distance), zip(TS_old, TS_new)) @debug "maxᵢ|Cⁿ⁺¹ - Cⁿ|ᵢ = $ΔCS maxᵢ|Tⁿ⁺¹ - Tⁿ|ᵢ = $ΔTS" @@ -257,8 +258,8 @@ function calc_convergence(env, CS_old, TS_old) return max(ΔCS, ΔTS), CS_new, TS_new end function calc_convergence(env_new::CTMRGEnv, env_old::CTMRGEnv) - CS_old = map(svd_vals, env_old.corners) - TS_old = map(svd_vals, env_old.edges) + CS_old = stablemap(svd_vals, env_old.corners) + TS_old = stablemap(svd_vals, env_old.edges) return calc_convergence(env_new, CS_old, TS_old) end @non_differentiable calc_convergence(args...) diff --git a/src/algorithms/ctmrg/gaugefix.jl b/src/algorithms/ctmrg/gaugefix.jl index 6ff44e554..4e071155a 100644 --- a/src/algorithms/ctmrg/gaugefix.jl +++ b/src/algorithms/ctmrg/gaugefix.jl @@ -37,7 +37,7 @@ function compute_gauge_fix_gauge( envfinal::CTMRGEnv{C, T}, envprev::CTMRGEnv{C, T}, alg::G ) where {C, T, G <: Union{ScramblingEnvGauge, ScramblingEnvGaugeC4v}} # Check if spaces in envprev and envfinal are the same - same_spaces = map(eachcoordinate(envfinal, 1:4)) do (dir, r, c) + same_spaces = stablemap(eachcoordinate(envfinal, 1:4)) do (dir, r, c) space(envfinal.edges[dir, r, c]) == space(envprev.edges[dir, r, c]) && space(envfinal.corners[dir, r, c]) == space(envprev.corners[dir, r, c]) end @@ -64,7 +64,7 @@ function compute_relative_phases( envfinal::CTMRGEnv{C, T}, envprev::CTMRGEnv{C, T}, ::ScramblingEnvGauge ) where {C, T} - signs = map(eachcoordinate(envfinal, 1:4)) do (dir, r, c) + signs = stablemap(eachcoordinate(envfinal, 1:4)) do (dir, r, c) # Gather edge tensors and pretend they're InfiniteMPSs if dir == NORTH Tsprev = circshift(envprev.edges[dir, r, :], 1 - c) @@ -81,7 +81,7 @@ function compute_relative_phases( end # Random MPS of same bond dimension - M = map(Tsfinal) do t + M = stablemap(Tsfinal) do t randn(scalartype(t), codomain(t) ← domain(t)) end @@ -166,7 +166,7 @@ end # Explicit fixing of relative phases (doing this compactly in a loop is annoying) function fix_relative_phases(envfinal::CTMRGEnv, signs) - corners_fixed = map(eachcoordinate(envfinal, 1:4)) do (dir, r, c) + corners_fixed = stablemap(eachcoordinate(envfinal, 1:4)) do (dir, r, c) Cf = if dir == NORTHWEST fix_gauge_northwest_corner((r, c), envfinal, signs) elseif dir == NORTHEAST @@ -179,7 +179,7 @@ function fix_relative_phases(envfinal::CTMRGEnv, signs) return Cf end - edges_fixed = map(eachcoordinate(envfinal, 1:4)) do (dir, r, c) + edges_fixed = stablemap(eachcoordinate(envfinal, 1:4)) do (dir, r, c) Ef = if dir == NORTHWEST fix_gauge_north_edge((r, c), envfinal, signs) elseif dir == NORTHEAST @@ -197,7 +197,7 @@ end function fix_relative_phases( U::Array{Ut, 3}, V::Array{Vt, 3}, signs ) where {Ut <: AbstractTensorMap, Vt <: AbstractTensorMap} - U_fixed = map(eachindex(IndexCartesian(), U)) do I + U_fixed = stablemap(eachindex(IndexCartesian(), U)) do I dir, r, c = Tuple(I) Uf = if dir == NORTHWEST fix_gauge_north_left_vecs((r, c), U, signs) @@ -211,7 +211,7 @@ function fix_relative_phases( return Uf end - V_fixed = map(eachindex(IndexCartesian(), V)) do I + V_fixed = stablemap(eachindex(IndexCartesian(), V)) do I dir, r, c = Tuple(I) Vf = if dir == NORTHWEST fix_gauge_north_right_vecs((r, c), V, signs) diff --git a/src/algorithms/ctmrg/simultaneous.jl b/src/algorithms/ctmrg/simultaneous.jl index ecaf52c7b..3187ca8e9 100644 --- a/src/algorithms/ctmrg/simultaneous.jl +++ b/src/algorithms/ctmrg/simultaneous.jl @@ -58,12 +58,12 @@ end # Work-around to stop Zygote from choking on first execution (sometimes) # Split up map returning projectors and info into separate arrays function _split_proj_and_info(proj_and_info) - P_left = map(x -> x[1][1], proj_and_info) - P_right = map(x -> x[1][2], proj_and_info) + P_left = stablemap(x -> x[1][1], proj_and_info) + P_right = stablemap(x -> x[1][2], proj_and_info) truncation_error = maximum(x -> x[2].truncation_error, proj_and_info) - U = map(x -> x[2].U, proj_and_info) - S = map(x -> x[2].S, proj_and_info) - V = map(x -> x[2].V, proj_and_info) + U = stablemap(x -> x[2].U, proj_and_info) + S = stablemap(x -> x[2].S, proj_and_info) + V = stablemap(x -> x[2].V, proj_and_info) info = (; truncation_error, U, S, V) return (P_left, P_right), info end @@ -117,6 +117,13 @@ function simultaneous_projectors( return compute_projector(ec, alg′) end +""" +Split an array of `(corner, edge)` tuples into separate corner and edge arrays. +""" +@noinline function _split_corners_edges(corners_edges) + return stablemap(first, corners_edges), stablemap(last, corners_edges) +end + """ $(SIGNATURES) @@ -153,5 +160,5 @@ function renormalize_simultaneously(enlarged_corners, projectors, network, env) return corner / norm(corner), edge / norm(edge) end - return CTMRGEnv(map(first, corners_edges), map(last, corners_edges)) + return CTMRGEnv(_split_corners_edges(corners_edges)...) end diff --git a/src/algorithms/expectation_value/expectation_value.jl b/src/algorithms/expectation_value/expectation_value.jl index d2ea498f6..a8891b8c6 100644 --- a/src/algorithms/expectation_value/expectation_value.jl +++ b/src/algorithms/expectation_value/expectation_value.jl @@ -14,18 +14,42 @@ function MPSKit.expectation_value( bra::S, O::LocalOperator, ket::S, env ) where {S <: InfiniteState} checklattice(bra, O, ket) + # protect against needing runtime activity + T = ignore_derivatives(() -> promote_type(scalartype(O), scalartype(bra), scalartype(ket), scalartype(env))) + total = zero(T) + for (inds, operator) in O.terms + # the copy here is necessary to avoid RT activity + # until dtmap!! can be used + total += local_expectation_value(inds, bra, copy(operator), ket, env) + end + return total +end +# TEMPORARY - block use of dtmap until https://github.com/EnzymeAD/Enzyme/pull/3152 is merged +#=function MPSKit.expectation_value( + bra::S, O::LocalOperator, ket::S, env + ) where {S <: InfiniteState} + checklattice(bra, O, ket) term_vals = dtmap(collect(O.terms)) do (inds, operator) # OhMyThreads can't iterate over O.terms directly return local_expectation_value(inds, bra, operator, ket, env) end return sum(term_vals) -end +end=# MPSKit.expectation_value(peps::InfinitePEPS, O::LocalOperator, env) = expectation_value(peps, O, peps, env) function MPSKit.expectation_value(state::InfinitePEPO, O::LocalOperator, env) checklattice(state, O) - term_vals = dtmap(collect(O.terms)) do (inds, operator) # OhMyThreads can't iterate over O.terms directly + #=term_vals = dtmap(collect(O.terms)) do (inds, operator) # OhMyThreads can't iterate over O.terms directly return local_expectation_value(inds, state, operator, env) end - return sum(term_vals) + return sum(term_vals)=# + # protect against needing runtime activity + T = ignore_derivatives(() -> promote_type(scalartype(O), scalartype(state), scalartype(env))) + total = zero(T) + for (inds, operator) in O.terms + # the copy here is necessary to avoid RT activity + # until dtmap!! can be used + total += local_expectation_value(inds, state, copy(operator), env) + end + return total end diff --git a/src/networks/infinitesquarenetwork.jl b/src/networks/infinitesquarenetwork.jl index fe62377b3..fbfc22164 100644 --- a/src/networks/infinitesquarenetwork.jl +++ b/src/networks/infinitesquarenetwork.jl @@ -69,6 +69,19 @@ end function VectorInterface.zerovector(A::InfiniteSquareNetwork) return InfiniteSquareNetwork(zerovector(unitcell(A))) end +function VI.add!( + A₁::InfiniteSquareNetwork, A₂::InfiniteSquareNetwork, α::Number, β::Number + ) + _add!(O1, O2) = _add_localsandwich!(O1, O2, α, β) + foreach(_add!, unitcell(A₁), unitcell(A₂)) + return A₁ +end +VI.add!(A₁::InfiniteSquareNetwork, A₂::InfiniteSquareNetwork) = VI.add!(A₁, A₂, One(), One()) +function VI.add!!( + A₁::InfiniteSquareNetwork, A₂::InfiniteSquareNetwork, α::Number, β::Number + ) + return add!(A₁, A₂, α, β) +end ## Math (for Zygote accumulation) diff --git a/src/networks/local_sandwich.jl b/src/networks/local_sandwich.jl index d29f54ea1..04a01878d 100644 --- a/src/networks/local_sandwich.jl +++ b/src/networks/local_sandwich.jl @@ -34,6 +34,8 @@ end # generic local interface _add_localsandwich(O1, O2) = O1 .+ O2 +# in-place: the sandwich tuple itself is immutable, but its tensors are not +_add_localsandwich!(O1, O2, α, β) = (foreach((x, y) -> add!(x, y, α, β), O1, O2); O1) _subtract_localsandwich(O1, O2) = O1 .- O2 _mul_localsandwich(α::Number, O) = α .* O _isapprox_localsandwich(O1, O2; kwargs...) = all(isapprox.(O1, O2; kwargs...)) @@ -47,6 +49,7 @@ _rot180_localsandwich(O::PFTensor) = rot180(O) # specialized local math interface _add_localsandwich(O1::PFTensor, O2::PFTensor) = O1 + O2 +_add_localsandwich!(O1::PFTensor, O2::PFTensor, α, β) = add!(O1, O2, α, β) _subtract_localsandwich(O1::PFTensor, O2::PFTensor) = O1 - O2 _mul_localsandwich(α::Number, O::PFTensor) = α * O _isapprox_localsandwich(O1::PFTensor, O2::PFTensor; kwargs...) = isapprox(O1, O2; kwargs...) diff --git a/src/operators/localoperator.jl b/src/operators/localoperator.jl index c35e5ddb9..a7c8747e8 100644 --- a/src/operators/localoperator.jl +++ b/src/operators/localoperator.jl @@ -47,8 +47,23 @@ function LocalOperator{O}(lattice, terms) where {O} return operator end -# Default to Any for eltype: needs to be abstract anyways so not that much to gain -LocalOperator(lattice, terms) = LocalOperator{Any}(lattice, terms) +""" +Narrowest element type that holds every term. +""" +function _term_eltype(terms) + T = Union{} + for (_, term) in terms + T = typejoin(T, typeof(term)) + end + return T === Union{} ? Any : T +end + +# `terms` may be a generator (`real`, `imag` and `*` build one), so collect before the +# two passes over it. +function LocalOperator(lattice, terms) + collected = collect(terms) + return LocalOperator{_term_eltype(collected)}(lattice, collected) +end LocalOperator(lattice, terms::Pair...) = LocalOperator(lattice, terms) # TODO: add terms beyond AbstractTensorMap # e.g. tensor product of 1-site operators, MPOs diff --git a/src/utility/diffable_threads.jl b/src/utility/diffable_threads.jl index e6b6270e8..2cc7c165a 100644 --- a/src/utility/diffable_threads.jl +++ b/src/utility/diffable_threads.jl @@ -10,6 +10,16 @@ dtmap(args...; scheduler = Defaults.scheduler[]) = tmap(args...; scheduler) dtmap!!(args...; scheduler = Defaults.scheduler[]) = tmap!(args...; scheduler) +# make this serial for now so that Enzyme can differentiate it. Once +# https://github.com/EnzymeAD/Enzyme/pull/3152 is merged and the new +# JLL is built, restore the old version +function dtmap!!(f, dst::AbstractArray, src::AbstractArray; scheduler = Defaults.scheduler[]) + for (i, a) in zip(eachindex(dst), src) + dst[i] = f(a) + end + return dst +end + # Follows the `map` rrule from ChainRules.jl but specified for the case of one AbstractArray that is being mapped # https://github.com/JuliaDiff/ChainRules.jl/blob/e245d50a1ae56ce46fc8c1f0fe9b925964f1146e/src/rulesets/Base/base.jl#L243 function ChainRulesCore.rrule( diff --git a/src/utility/eigh.jl b/src/utility/eigh.jl index 0db3bb7ef..98efecacd 100644 --- a/src/utility/eigh.jl +++ b/src/utility/eigh.jl @@ -18,7 +18,7 @@ Construct a `FullPullback` algorithm struct from the following keyword arguments * `verbosity::Int=0` : Suppresses all output if `≤0`, prints gauge dependency warnings if `1`, and always prints gauge dependency if `≥2`. """ @kwdef struct FullPullback - degeneracy_atol::Real = Defaults.rrule_degeneracy_atol + degeneracy_atol::Float64 = Defaults.rrule_degeneracy_atol verbosity::Int = 0 end @@ -41,7 +41,7 @@ Construct a `TruncPullback` algorithm struct from the following keyword argument * `verbosity::Int=0` : Suppresses all output if `≤0`, prints gauge dependency warnings if `1`, and always prints gauge dependency if `≥2`. """ @kwdef struct TruncPullback - degeneracy_atol::Real = Defaults.rrule_degeneracy_atol + degeneracy_atol::Float64 = Defaults.rrule_degeneracy_atol verbosity::Int = 0 end diff --git a/src/utility/util.jl b/src/utility/util.jl index 603aec675..f42de7f8e 100644 --- a/src/utility/util.jl +++ b/src/utility/util.jl @@ -150,3 +150,33 @@ function _permute_to_last(axes::NTuple{N, Int}, ax::Int) where {N} new_axes = (ntuple(i -> axes[biperm[1][i]], N - 1)..., ax) return new_axes, biperm end + +""" + stablemap(f, A) + +Type-stable replacement for `map(f, A)` on the CTMRG differentiation path. + +`Base.map` builds a `Base.Generator` whose element type is unknown, so `_collect` +falls back to a type-widening loop over untyped storage. Enzyme cannot statically +prove the element type through that and bails out with an `EnzymeNoTypeError`. +Inferring the element type up front and filling a concretely typed destination +keeps the same semantics without the widening machinery. +""" +@inline function stablemap(f::F, A) where {F} + T = Base.promote_op(f, eltype(A)) + if !isconcretetype(T) # inference failed: fall back to Base + return map(f, A) + end + dst = similar(A, T) + @inbounds for (i, a) in zip(eachindex(dst), A) + dst[i] = f(a) + end + return dst +end + +# the fill loop above mutates `dst`, which Zygote cannot differentiate +function ChainRulesCore.rrule( + config::RuleConfig{>:HasReverseMode}, ::typeof(stablemap), f, A + ) + return rrule_via_ad(config, map, f, A) +end diff --git a/test/Project.toml b/test/Project.toml index 86bdd9467..512740053 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -5,6 +5,8 @@ Accessors = "7d9f7c33-5ae7-4f3b-8dc6-eff91059b697" Adapt = "79e6a3ab-5dfb-504d-930d-738a2a938a0e" ChainRulesCore = "d360d2e6-b24c-11e9-a2a3-2a2ae2dbcce4" ChainRulesTestUtils = "cdddcdb0-9152-4a09-a978-84456f9df70a" +Enzyme = "7da242da-08ed-463a-9acd-ee780be4f1d9" +EnzymeTestUtils = "12d8515a-0907-448a-8884-5fe00fdf1c5a" KrylovKit = "0b1a1467-8014-51b9-945f-bf0ae24f4b77" LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" MPSKit = "bb1c41ca-d63c-52ed-829e-0820dda26502" @@ -28,6 +30,7 @@ PEPSKit = {path = ".."} [compat] Adapt = "4" ChainRulesTestUtils = "1.13" +EnzymeTestUtils = "0.2.8" ParallelTestRunner = "2.6.0" QuadGK = "2.11.1" Test = "1" diff --git a/test/ctmrg/enz_pepo.jl b/test/ctmrg/enz_pepo.jl new file mode 100644 index 000000000..10df0e883 --- /dev/null +++ b/test/ctmrg/enz_pepo.jl @@ -0,0 +1,12 @@ +using Test +using PEPSKit + +@isdefined(TestSuite) || include("../testsuite/TestSuite.jl") +using .TestSuite + +is_buildkite = get(ENV, "BUILDKITE", "false") == "true" + +if !is_buildkite + TestSuite.enzyme_ctmrg_pepo_runthroughs(Vector) + TestSuite.enzyme_ctmrg_pepo_fixed_point(Vector) +end diff --git a/test/enzyme_gradients/enz_ctmrg_gradients.jl b/test/enzyme_gradients/enz_ctmrg_gradients.jl new file mode 100644 index 000000000..c928d380d --- /dev/null +++ b/test/enzyme_gradients/enz_ctmrg_gradients.jl @@ -0,0 +1,11 @@ +using PEPSKit + +@isdefined(TestSuite) || include("../testsuite/TestSuite.jl") +using .TestSuite + +is_buildkite = get(ENV, "BUILDKITE", "false") == "true" + +if !is_buildkite + TestSuite.enzyme_gradients_asymmetric(Vector) + TestSuite.enzyme_gradients_asymmetric_276(Vector) +end diff --git a/test/testsuite/TestSuite.jl b/test/testsuite/TestSuite.jl index 018a68a82..518cf193a 100644 --- a/test/testsuite/TestSuite.jl +++ b/test/testsuite/TestSuite.jl @@ -139,6 +139,12 @@ module CTMRGPEPO end using .CTMRGPEPO +module CTMRGEnzymePEPO + include("ctmrg/enz_pepo.jl") + export enzyme_ctmrg_pepo_runthroughs, enzyme_ctmrg_pepo_fixed_point +end +using .CTMRGEnzymePEPO + module CTMRGSUWeight include("ctmrg/suweight.jl") export ctmrg_suweight @@ -197,6 +203,12 @@ module GradientsCTMRG end using .GradientsCTMRG +module EnzymeGradientsCTMRG + include("enzyme_gradients/enz_ctmrg_gradients.jl") + export enzyme_gradients_asymmetric, enzyme_gradients_asymmetric_276 +end +using .EnzymeGradientsCTMRG + # Time evolution # -------------- module TimeEvolClusterProjectors diff --git a/test/testsuite/ctmrg/enz_pepo.jl b/test/testsuite/ctmrg/enz_pepo.jl new file mode 100644 index 000000000..475f5b6d8 --- /dev/null +++ b/test/testsuite/ctmrg/enz_pepo.jl @@ -0,0 +1,154 @@ +using Test +using Random +using LinearAlgebra +using PEPSKit +using TensorKit +using KrylovKit +using OptimKit +using Enzyme +using Adapt + +# Enzyme names basic blocks after Julia types (e.g. `zeroType.`); the tape types +# here exceed LLVM's default 1024-char cap on non-global value names, which its +# textual IR parser rejects when Enzyme round-trips the module in `check_ir!`. +Enzyme.LLVM.clopts("--non-global-value-max-name-size=1048576") + +function three_dimensional_classical_ising(AT; beta, J = 1.0) + K = beta * J + # Boltzmann weights + t = ComplexF64[exp(K) exp(-K); exp(-K) exp(K)] + r = eigen(t) + q = r.vectors * sqrt(LinearAlgebra.Diagonal(r.values)) * r.vectors + + # local partition function tensor + O = zeros(2, 2, 2, 2, 2, 2) + O[1, 1, 1, 1, 1, 1] = 1 + O[2, 2, 2, 2, 2, 2] = 1 + @tensor o[-1 -2; -3 -4 -5 -6] := + O[1 2; 3 4 5 6] * q[-1; 1] * q[-2; 2] * q[-3; 3] * q[-4; 4] * q[-5; 5] * q[-6; 6] + + # magnetization tensor + M = copy(O) + M[2, 2, 2, 2, 2, 2] *= -1 + @tensor m[-1 -2; -3 -4 -5 -6] := + M[1 2; 3 4 5 6] * q[-1; 1] * q[-2; 2] * q[-3; 3] * q[-4; 4] * q[-5; 5] * q[-6; 6] + + # bond interaction tensor and energy-per-site tensor + e = ComplexF64[-J J; J -J] .* q + @tensor e_x[-1 -2; -3 -4 -5 -6] := + O[1 2; 3 4 5 6] * q[-1; 1] * q[-2; 2] * q[-3; 3] * e[-4; 4] * q[-5; 5] * q[-6; 6] + @tensor e_y[-1 -2; -3 -4 -5 -6] := + O[1 2; 3 4 5 6] * q[-1; 1] * q[-2; 2] * e[-3; 3] * q[-4; 4] * q[-5; 5] * q[-6; 6] + @tensor e_z[-1 -2; -3 -4 -5 -6] := + O[1 2; 3 4 5 6] * e[-1; 1] * q[-2; 2] * q[-3; 3] * q[-4; 4] * q[-5; 5] * q[-6; 6] + e = e_x + e_y + e_z + + # fixed tensor map space for all three + TMS = ℂ^2 ⊗ (ℂ^2)' ← ℂ^2 ⊗ ℂ^2 ⊗ (ℂ^2)' ⊗ (ℂ^2)' + + return adapt(AT, TensorMap(o, TMS)), adapt(AT, TensorMap(m, TMS)), adapt(AT, TensorMap(e, TMS)) +end + +## Test + +# initialize +beta = 0.2391 # slightly lower temperature than βc ≈ 0.2216544 +χpeps = ℂ^2 +χenv = ℂ^12 + +# cover all different flavors +ctm_styles = [:SequentialCTMRG, :SimultaneousCTMRG] +projector_algs = [:HalfInfiniteProjector, :FullInfiniteProjector] + +function enzyme_ctmrg_pepo_runthroughs(AT) + return @testset "PEPO CTMRG runthroughs for unitcell=$(unitcell) ($AT)" for unitcell in + [(1, 1, 1), (1, 1, 2)] + Random.seed!(81812781144) + + O, M, E = three_dimensional_classical_ising(AT; beta) + # contract + T = adapt(AT, InfinitePEPO(O; unitcell = unitcell)) + psi0 = initializePEPS(T, χpeps) + n = InfiniteSquareNetwork(psi0, T) + env0 = CTMRGEnv(n, χenv) + + @test spacetype(typeof(T)) === ComplexSpace + @test spacetype(T) === ComplexSpace + @test sectortype(typeof(T)) === Trivial + @test sectortype(T) === Trivial + + @testset "PEPO CTMRG contraction using $alg with $projector_alg" for ( + alg, projector_alg, + ) in Iterators.product(ctm_styles, projector_algs) + env, = leading_boundary(env0, n; alg, maxiter = 150, projector_alg) + end + end +end + +function enzyme_ctmrg_pepo_fixed_point(AT) + return @testset "Fixed-point computation for 3D classical ising model ($AT)" begin + Random.seed!(81812781144) + + # prep + ctm_alg = SimultaneousCTMRG(; maxiter = 150, tol = 1.0e-8, verbosity = 2) + gradient_alg = FixedPointGradient(; + solver_alg = KrylovKit.Arnoldi(; maxiter = 30, tol = 1.0e-6, eager = true), + ) + opt_alg = LBFGS(32; maxiter = 50, gradtol = 1.0e-5, verbosity = 3) + function pepo_retract(x, η, α) + x´_partial, ξ = PEPSKit.peps_retract(x[1:2], η, α) + x´ = (x´_partial..., deepcopy(x[3])) + return x´, ξ + end + function pepo_transport!(ξ, x, η, α, x´) + return PEPSKit.peps_transport!(ξ, x[1:2], η, α, x´[1:2]) + end + + O, M, E = three_dimensional_classical_ising(AT; beta) + # contract + T = adapt(AT, InfinitePEPO(O; unitcell = (1, 1, 1))) + psi0 = initializePEPS(T, χpeps) + env2_0 = CTMRGEnv(InfiniteSquareNetwork(psi0), χenv) + env3_0 = CTMRGEnv(InfiniteSquareNetwork(psi0, T), χenv) + + # optimize free energy per site + (psi_final, env2_final, env3_final), f, = optimize( + (psi0, env2_0, env3_0), + opt_alg; + inner = PEPSKit.real_inner, + retract = pepo_retract, + (transport!) = (pepo_transport!), + ) do (psi, env2, env3) + function energ(ψ) + n2 = InfiniteSquareNetwork(ψ) + env2′, info = leading_boundary(env2, n2, ctm_alg) + n3 = InfiniteSquareNetwork(ψ, T)::InfiniteSquareNetwork{ + Tuple{ + TensorMap{ComplexF64, ComplexSpace, 1, 4, Vector{ComplexF64}}, + TensorMap{ComplexF64, ComplexSpace, 1, 4, Vector{ComplexF64}}, + TensorMap{ComplexF64, ComplexSpace, 2, 4, Vector{ComplexF64}}, + }, + } + env3′, info = leading_boundary(env3, n3, ctm_alg) + PEPSKit.ignore_derivatives() do + PEPSKit.update!(env2, env2′) + PEPSKit.update!(env3, env3′) + end + λ3 = network_value(n3, env3′) + λ2 = network_value(n2, env2′) + return -log(real(λ3 / λ2)) + end + dpsi = zerovector(psi) + _, E = Enzyme.autodiff(ReverseWithPrimal, Const(energ), Active, Duplicated(psi, dpsi)) + return E, dpsi + end + + # check energy + n3_final = InfiniteSquareNetwork(psi_final, T) + m = PEPSKit.contract_local_tensor((1, 1, 1), M, n3_final, env3_final) + nrm3 = PEPSKit._contract_site((1, 1), n3_final, env3_final) + + # compare to Monte-Carlo result from https://www.worldscientific.com/doi/abs/10.1142/S0129183101002383 + @test abs(m / nrm3) ≈ 0.667162 rtol = 1.0e-2 + end +end diff --git a/test/testsuite/enzyme_gradients/enz_ctmrg_gradients.jl b/test/testsuite/enzyme_gradients/enz_ctmrg_gradients.jl new file mode 100644 index 000000000..9c6e3e543 --- /dev/null +++ b/test/testsuite/enzyme_gradients/enz_ctmrg_gradients.jl @@ -0,0 +1,183 @@ +using Test +using Random +using PEPSKit +using TensorKit +using Enzyme +using OptimKit +using KrylovKit +using Adapt +# Enzyme names basic blocks after Julia types (e.g. `zeroType.`); the tape types +# here exceed LLVM's default 1024-char cap on non-global value names, which its +# textual IR parser rejects when Enzyme round-trips the module in `check_ir!`. +Enzyme.LLVM.clopts("--non-global-value-max-name-size=1048576") + +## Test models, gradmodes and CTMRG algorithm +# ------------------------------------------- +χbond = 2 +χenv = 6 +Pspaces = [ComplexSpace(2), Vect[FermionParity](0 => 1, 1 => 1)] +Vspaces = [ComplexSpace(χbond), Vect[FermionParity](0 => χbond / 2, 1 => χbond / 2)] +Espaces = [ComplexSpace(χenv), Vect[FermionParity](0 => χenv / 2, 1 => χenv / 2)] +models = [heisenberg_XYZ(InfiniteSquare()), pwave_superconductor(InfiniteSquare())] +names = ["Heisenberg", "p-wave superconductor"] + +gradtol = 1.0e-4 +ctmrg_verbosity = 0 +ctmrg_algs = [[:SequentialCTMRG, :SimultaneousCTMRG], [:SequentialCTMRG, :SimultaneousCTMRG]] +projector_algs = [[:HalfInfiniteProjector, :FullInfiniteProjector], [:HalfInfiniteProjector, :FullInfiniteProjector]] +svd_rrule_algs = [[:FullPullback, :TruncPullback, :Arnoldi], [:FullPullback, :Arnoldi]] +gradient_algs = [[nothing, :FixedPointGradient], [:FixedPointGradient]] +# the solver only affects the linear solve inside the fixed-point rule, which is +# already covered by the Zygote tests; one solver suffices to exercise Enzyme here +gradient_solver_algs = [[:GMRES], [:GMRES]] +steps = -0.01:0.005:0.01 + +# don't check naive AD gradients for all algorithm combinations, since it's slow +naive_gradient_combinations = [ + (:SimultaneousCTMRG, :HalfInfiniteProjector, :FullPullback), + (:SimultaneousCTMRG, :FullInfiniteProjector, :FullPullback), + #(:SequentialCTMRG, :HalfInfiniteProjector, :FullPullback), +] +naive_gradient_done = Set() + +# fixed-point differentiation is incompatible with sequential CTMRG +function _check_disallowed_combination( + ctmrg_alg, projector_alg, decomposition_rrule_alg, gradient_alg + ) + ctmrg_alg == :SequentialCTMRG && !isnothing(gradient_alg) && return true + return false +end + + +## Tests +# ------ +function enzyme_gradients_asymmetric(AT) + naive_gradient_done = Set() + return @testset "Enzyme AD CTMRG energy gradients for $(names[i]) model ($AT)" verbose = true for i in + eachindex( + models + ) + Pspace = Pspaces[i] + Vspace = Vspaces[i] + Espace = Espaces[i] + calgs = ctmrg_algs[i] + palgs = projector_algs[i] + salgs = svd_rrule_algs[i] + galgs = gradient_algs[i] + gsalgs = gradient_solver_algs[i] + @testset "ctmrg_alg=:$ctmrg_alg, projector_alg=:$projector_alg, svd_rrule_alg=:$svd_rrule_alg, gradient_alg=(; alg = :$gradient_alg, solver_alg = (; alg = :$gradient_solver_alg))" for ( + ctmrg_alg, projector_alg, svd_rrule_alg, gradient_alg, gradient_solver_alg, + ) in Iterators.product( + calgs, palgs, salgs, galgs, gsalgs + ) + + # only run GMRES for the implicit gradient, and skip distinction between decomposition rrule algs + if gradient_alg == :ImplicitGradient + gradient_solver_alg == :GMRES || continue + svd_rrule_alg == first(salgs) || continue + end + + # check for allowed algorithm combinations when testing naive gradient + if isnothing(gradient_alg) + combo = (ctmrg_alg, projector_alg, svd_rrule_alg) + combo in naive_gradient_combinations || continue + combo in naive_gradient_done && continue + push!(naive_gradient_done, combo) + gradient_solver_alg = nothing # unused in naive gradient, so set to nothing to avoid confusion + end + + + # filter disallowed algorithm combinations + if _check_disallowed_combination( + ctmrg_alg, projector_alg, svd_rrule_alg, gradient_alg + ) + # but verify that its use would throw an error + @test_throws ArgumentError PEPSOptimize(; + boundary_alg = (; alg = ctmrg_alg, projector_alg, decomposition_alg = (; rrule_alg = (; alg = svd_rrule_alg))), + gradient_alg = (; alg = gradient_alg, solver_alg = (; alg = gradient_solver_alg, tol = gradtol)), + ) + continue + end + + @info "optimtest of ctmrg_alg=:$ctmrg_alg, projector_alg=:$projector_alg, svd_rrule_alg=:$svd_rrule_alg and gradient_alg=(; alg = :$gradient_alg, solver_alg = (; alg = :$gradient_solver_alg)) on $(names[i])" + Random.seed!(42039482030) + dir = adapt(AT, InfinitePEPS(Pspace, Vspace)) + psi = adapt(AT, InfinitePEPS(Pspace, Vspace)) + # instantiate to avoid having to type this twice... + contrete_ctmrg_alg = PEPSKit.CTMRGAlgorithm(; + alg = ctmrg_alg, + verbosity = ctmrg_verbosity, + projector_alg = projector_alg, + decomposition_alg = SVDAdjoint(; rrule_alg = (; alg = svd_rrule_alg)), + ) + # instantiate because hook_pullback doesn't go through the keyword selector... + concrete_gradient_alg = if isnothing(gradient_alg) + nothing # TODO: add this to the PEPSKit.GradientAlgorithm selector? + else + PEPSKit.GradientAlgorithm(; + alg = gradient_alg, solver_alg = (; alg = gradient_solver_alg, tol = gradtol) + ) + end + env, = leading_boundary(CTMRGEnv(psi, Espace), psi, contrete_ctmrg_alg) + model = adapt(AT, models[i]) + alphas, fs, dfs1, dfs2 = OptimKit.optimtest( + (psi, env), + dir; + alpha = steps, + retract = PEPSKit.peps_retract, + inner = PEPSKit.real_inner, + ) do (peps, env) + function energ(psi) + env2, info = PEPSKit.hook_pullback( + leading_boundary, env, psi, contrete_ctmrg_alg; + alg_rrule = concrete_gradient_alg, + ) + return cost_function(psi, env2, model) + end + dpeps = zerovector(peps) + _, E = Enzyme.autodiff( + set_runtime_activity(ReverseWithPrimal), Const(energ), Active, + Duplicated(peps, dpeps), + ) + return E, dpeps + end + @test dfs1 ≈ dfs2 atol = 1.0e-2 + end + end +end + +function enzyme_gradients_asymmetric_276(AT) + ## Regression test for gradient accuracy (https://github.com/QuantumKitHub/PEPSKit.jl/pull/276) + return @testset "Enzyme AD CTMRG energy gradient accuracy regression test (#276) ($AT)" begin + Random.seed!(1234) + + boundary_alg = PEPSKit.CTMRGAlgorithm(; tol = 1.0e-10) + gradient_alg = PEPSKit.GradientAlgorithm(; tol = 5.0e-8) + + function fg((peps, env)) + function energ(ψ) + env2, = leading_boundary(env, ψ, boundary_alg) + return cost_function(ψ, env2, H) + end + E, gs = Enzyme.autodiff(ReverseWithPrimal, Const(energ), Active, Duplicated(peps, zerovector(peps))) + return E, only(gs) + end + + # initialize randomly + H = adapt(AT, heisenberg_XYZ(InfiniteSquare(1, 1))) + peps = PEPSKit.peps_normalize(adapt(AT, InfinitePEPS(randn, ComplexF64, physicalspace(H)[1], ComplexSpace(3)))) + env0 = CTMRGEnv(randn, ComplexF64, peps, ComplexSpace(20)) + + # test gradient against finite-difference + Δx = 1.0e-5 + _, _, dfs1, dfs2 = OptimKit.optimtest( + fg, (peps, env0); + alpha = LinRange(-Δx, Δx, 2), + retract = PEPSKit.peps_retract, + inner = PEPSKit.real_inner, + ) + + # verify high gradient accuracy for small finite-difference step size + @test dfs1 ≈ dfs2 rtol = 1.0e-2 * Δx + end +end diff --git a/test/utility/enz_diff_maps.jl b/test/utility/enz_diff_maps.jl new file mode 100644 index 000000000..b5c156a76 --- /dev/null +++ b/test/utility/enz_diff_maps.jl @@ -0,0 +1,7 @@ +using Enzyme, EnzymeTestUtils +using PEPSKit: dtmap + +# Can the rrule of dtmap be made inferable? (if check_inferred=true, tests error at the moment) +@testset "Differentiable tmap" begin + test_reverse(dtmap, Duplicated, Const(x -> x^3), (randn(5, 5), Duplicated)) +end