From 511973d8ccaf4845fa2356a0a4b85a7d5a4caed8 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 1 Jun 2026 19:16:16 +0200 Subject: [PATCH 1/8] Initial steps to Mooncake support --- Project.toml | 18 ++- ext/PEPSKitMooncakeExt.jl | 61 ++++++++++ test/Project.toml | 1 + test/mooncake/svd_wrapper.jl | 223 +++++++++++++++++++++++++++++++++++ 4 files changed, 302 insertions(+), 1 deletion(-) create mode 100644 ext/PEPSKitMooncakeExt.jl create mode 100644 test/mooncake/svd_wrapper.jl diff --git a/Project.toml b/Project.toml index 68ece85cd..a14dbb5e3 100644 --- a/Project.toml +++ b/Project.toml @@ -8,6 +8,7 @@ projects = ["test", "docs", "benchmark"] [deps] Accessors = "7d9f7c33-5ae7-4f3b-8dc6-eff91059b697" +BlockTensorKit = "5f87ffc2-9cf1-4a46-8172-465d160bd8cd" ChainRulesCore = "d360d2e6-b24c-11e9-a2a3-2a2ae2dbcce4" Compat = "34da2185-b29b-5c13-b0c7-acf172513d20" DocStringExtensions = "ffbed154-4ef7-542d-bbb7-c09d3a79fcae" @@ -29,6 +30,20 @@ TupleTools = "9d95972d-f1c8-5527-a6e0-b4b365fa01f6" VectorInterface = "409d34a3-91d5-4945-b6ec-7529ddf182d8" Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f" +[weakdeps] +Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" + +[extensions] +PEPSKitMooncakeExt = "Mooncake" + +[sources] +BlockTensorKit = {path = "/Users/khyatt/.julia/dev/BlockTensorKit"} +TensorKit = {path = "/Users/khyatt/.julia/dev/TensorKit"} +MPSKit = {url = "https://github.com/quantumkithub/mpskit.jl", rev = "main"} +OptimKit = {url = "https://github.com/kshyatt/optimkit.jl", rev = "patch-1"} +KrylovKit = {url = "https://github.com/kshyatt/krylovkit.jl", rev = "patch-1"} +VectorInterface = {path = "/Users/khyatt/.julia/dev/VectorInterface"} + [compat] Accessors = "0.1" ChainRulesCore = "1.0" @@ -41,6 +56,7 @@ LoggingExtras = "1" MPSKit = "0.13.9" MPSKitModels = "0.4" MatrixAlgebraKit = "0.6.5" +Mooncake = "0.5.27" OhMyThreads = "0.7, 0.8" OptimKit = "0.4" Printf = "1" @@ -49,6 +65,6 @@ Statistics = "1" TensorKit = "0.16.5, 0.17" TensorOperations = "5" TupleTools = "1.6.0" -VectorInterface = "0.4, 0.5, 0.6" +VectorInterface = "0.6" Zygote = "0.6, 0.7" julia = "1.10" diff --git a/ext/PEPSKitMooncakeExt.jl b/ext/PEPSKitMooncakeExt.jl new file mode 100644 index 000000000..24963d654 --- /dev/null +++ b/ext/PEPSKitMooncakeExt.jl @@ -0,0 +1,61 @@ +module PEPSKitMooncakeExt + +using PEPSKit, TensorKit, Mooncake, MatrixAlgebraKit +using PEPSKit: SVDAdjoint +using Mooncake: DefaultCtx, CoDual, Dual, NoRData, primal, rrule!!, arrayify, @is_primitive + +_warn_pullback_truncerror(dϵ::Real; tol = MatrixAlgebraKit.defaulttol(dϵ)) = + abs(dϵ) ≤ tol || @warn "Pullback ignores non-zero tangents for truncation error" + +Mooncake.tangent_type(::Type{<:PEPSKit.SVDAdjoint}) = Mooncake.NoTangent + +@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(svd_trunc), TensorKit.AbstractTensorMap, SVDAdjoint} +function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{SVDAdjoint{F, R}}) where {F, R <: PEPSKit.FullPullback} + # TODO: filter out any decomposition algorithm that doesn't give access to the full spectrum + t, Δt = arrayify(t_dt) + alg = primal(alg_dalg) + # requires access to the full decomposition + U, S, V⁺ = svd_compact!(t, alg.fwd_alg.alg) + (Ũ, S̃, Ṽ⁺), inds = MatrixAlgebraKit.truncate(svd_trunc!, (U, S, V⁺), alg.fwd_alg.trunc) + truncerror = MatrixAlgebraKit.truncation_error(diagview(S), inds) + + gtol = PEPSKit._get_pullback_gauge_tol(alg.rrule_alg.verbosity) + output = (Ũ, S̃, Ṽ⁺, truncerror) + USVᴴtrunc = (Ũ, S̃, Ṽ⁺) + output_codual = CoDual(output, Mooncake.fdata(Mooncake.zero_tangent(output))) + ΔUSVᴴtrunc = last.(arrayify.(USVᴴtrunc, Base.front(Mooncake.tangent(output_codual)))) + function svd_trunc!_full_pullback((_, _, _, dϵ)::Tuple{NoRData, NoRData, NoRData, Real}) + _warn_pullback_truncerror(dϵ) + Δt = MatrixAlgebraKit.svd_pullback!( + Δt, t, (U, S, V⁺), ΔUSVᴴtrunc, inds; + gauge_atol = gtol(ΔUSVᴴtrunc), degeneracy_atol = alg.rrule_alg.degeneracy_atol, + ) + return NoRData(), NoRData(), NoRData() + end + return output_codual, svd_trunc!_full_pullback +end + +function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{SVDAdjoint{F, R}}) where {F, R <: PEPSKit.TruncPullback} + t, Δt = arrayify(t_dt) + alg = primal(alg_dalg) + gtol = PEPSKit._get_pullback_gauge_tol(alg.rrule_alg.verbosity) + output = svd_trunc(t, alg) + + output_codual = CoDual(output, Mooncake.fdata(Mooncake.zero_tangent(output))) + function svd_trunc!_trunc_pullback((_, _, _, dϵ)::Tuple{NoRData, NoRData, NoRData, Real}) + Utrunc, Strunc, Vᴴtrunc, ϵ = Mooncake.primal(output_codual) + dUtrunc_, dStrunc_, dVᴴtrunc_, _ = Mooncake.tangent(output_codual) + _warn_pullback_truncerror(dϵ) + U, dU = arrayify(Utrunc, dUtrunc_) + S, dS = arrayify(Strunc, dStrunc_) + Vᴴ, dVᴴ = arrayify(Vᴴtrunc, dVᴴtrunc_) + MatrixAlgebraKit.svd_trunc_pullback!(Δt, t, (U, S, Vᴴ), (dU, dS, dVᴴ)) + MatrixAlgebraKit.zero!(dU) + MatrixAlgebraKit.zero!(dS) + MatrixAlgebraKit.zero!(dVᴴ) + return NoRData(), NoRData(), NoRData() + end + return output_codual, svd_trunc!_trunc_pullback +end + +end diff --git a/test/Project.toml b/test/Project.toml index 2742a2812..56b69e010 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -9,6 +9,7 @@ LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" MPSKit = "bb1c41ca-d63c-52ed-829e-0820dda26502" MPSKitModels = "ca635005-6f8c-4cd1-b51d-8491250ef2ab" MatrixAlgebraKit = "6c742aac-3347-4629-af66-fc926824e5e4" +Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" OptimKit = "77e91f04-9b3b-57a6-a776-40b61faaebe0" PEPSKit = "52969e89-939e-4361-9b68-9bc7cde4bdeb" ParallelTestRunner = "d3525ed8-44d0-4b2c-a655-542cee43accc" diff --git a/test/mooncake/svd_wrapper.jl b/test/mooncake/svd_wrapper.jl new file mode 100644 index 000000000..af0ac2acb --- /dev/null +++ b/test/mooncake/svd_wrapper.jl @@ -0,0 +1,223 @@ +using Test +using Random +using LinearAlgebra +using TensorKit +using Mooncake +using Accessors +using PEPSKit + +using MatrixAlgebraKit: TruncatedAlgorithm, diagview + +# Gauge-invariant loss function +function lossfun(A, alg, R = randn(space(A)), trunc = notrunc()) + alg = @set alg.fwd_alg = TruncatedAlgorithm(alg.fwd_alg, trunc) + U, S, V, = svd_trunc(A, alg) + return real(dot(R, U * V)) + dot(S, S) # Overlap with random tensor R is gauge-invariant and differentiable, also for m≠n +end + + +dtype = ComplexF64 +m, n = 20, 30 +χ = 12 +trunc = truncspace(ℂ^χ) +rtol = 1.0e-9 +Random.seed!(12345678) +r = randn(dtype, ℂ^m, ℂ^n) +R = randn(space(r)) + +full_alg = SVDAdjoint(; rrule_alg = (; alg = :FullPullback, degeneracy_atol = 1.0e-13)) +trunc_alg = SVDAdjoint(; rrule_alg = (; alg = :TruncPullback, degeneracy_atol = 1.0e-13)) +iter_alg = SVDAdjoint(; fwd_alg = (; alg = :GKL)) + +@testset "Non-truncated SVD" begin + full_lossfun = A -> lossfun(A, full_alg, R) + trunc_lossfun = A -> lossfun(A, trunc_alg, R) + iter_lossfun = A -> lossfun(A, iter_alg, R) + + full_rrule = Mooncake.build_rrule(full_lossfun, r) + trunc_rrule = Mooncake.build_rrule(trunc_lossfun, r) + iter_rrule = Mooncake.build_rrule(iter_lossfun, r) + + l_full, g_full = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, r) + l_trunc, g_trunc = Mooncake.value_and_gradient!!(trunc_rrule, trunc_lossfun, r) + l_iter, g_iter = Mooncake.value_and_gradient!!(iter_rrule, iter_lossfun, r) + + @test l_full ≈ l_trunc ≈ l_iter + @test g_full[2] ≈ g_trunc[2] rtol = rtol + @test g_full[2] ≈ g_iter[2] rtol = rtol + @test g_trunc[2] ≈ g_iter[2] rtol = rtol +end + +@testset "Truncated SVD with χ=$χ" begin + full_lossfun = A -> lossfun(A, full_alg, R, trunc) + trunc_lossfun = A -> lossfun(A, trunc_alg, R, trunc) + iter_lossfun = A -> lossfun(A, iter_alg, R, trunc) + + full_rrule = Mooncake.build_rrule(full_lossfun, r) + trunc_rrule = Mooncake.build_rrule(trunc_lossfun, r) + iter_rrule = Mooncake.build_rrule(iter_lossfun, r) + + l_full, g_full = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, r) + l_trunc, g_trunc = Mooncake.value_and_gradient!!(trunc_rrule, trunc_lossfun, r) + l_iter, g_iter = Mooncake.value_and_gradient!!(iter_rrule, iter_lossfun, r) + + @test l_full ≈ l_trunc ≈ l_iter + @test g_full[2] ≈ g_trunc[2] rtol = rtol + @test g_full[2] ≈ g_iter[2] rtol = rtol + @test g_trunc[2] ≈ g_iter[2] rtol = rtol +end + +@testset "Truncated SVD broadening for $(alg.rrule_alg)" for alg in [full_alg, trunc_alg] + u, s, v, = svd_compact(r) + s.data[1:2:m] .= s.data[2:2:m] # make every singular value two-fold degenerate + r_degen = u * s * v + + no_broadening_no_cutoff_alg = @set full_alg.rrule_alg.degeneracy_atol = 1.0e-30 + small_broadening_alg = @set full_alg.rrule_alg.degeneracy_atol = 1.0e-13 + + full_lossfun = A -> lossfun(A, full_alg, R, trunc) + no_broadening_lossfun = A -> lossfun(A, no_broadening_no_cutoff_alg, R, trunc) + small_broadening_lossfun = A -> lossfun(A, small_broadening_alg, R, trunc) + + full_rrule = Mooncake.build_rrule(full_lossfun, r_degen) + no_broadening_rrule = Mooncake.build_rrule(no_broadening_lossfun, r_degen) + small_broadening_rrule = Mooncake.build_rrule(small_broadening_lossfun, r_degen) + + l_only_cutoff, g_only_cutoff = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, r_degen) # cutoff sets degenerate difference to zero + l_no_broadening_no_cutoff, g_no_broadening_no_cutoff = Mooncake.value_and_gradient!!( # degenerate singular value differences lead to divergent contributions + no_broadening_rrule, no_broadening_lossfun, r_degen, + ) + l_small_broadening, g_small_broadening = Mooncake.value_and_gradient!!( # broadening smoothens divergent contributions + small_broadening_rrule, small_broadening_lossfun, r_degen, + ) + + @test l_only_cutoff ≈ l_no_broadening_no_cutoff ≈ l_small_broadening + @test norm(g_no_broadening_no_cutoff[2] - g_small_broadening[2]) > 1.0e-2 # divergences mess up the gradient + @test g_only_cutoff[2] ≈ g_small_broadening[2] rtol = rtol # cutoff and broadening have similar effect +end + +symm_m, symm_n = 18, 24 +symm_space = Z2Space(0 => symm_m, 1 => symm_n) +symm_trspace = truncspace(Z2Space(0 => symm_m ÷ 2, 1 => symm_n ÷ 3)) +symm_r = randn(dtype, symm_space, symm_space) +symm_R = randn(dtype, space(symm_r)) + +@testset "IterSVD of symmetric tensors" begin + full_lossfun = A -> lossfun(A, full_alg, symm_R) + trunc_lossfun = A -> lossfun(A, trunc_alg, symm_R) + iter_lossfun = A -> lossfun(A, iter_alg, symm_R) + + full_rrule = Mooncake.build_rrule(full_lossfun, symm_r) + trunc_rrule = Mooncake.build_rrule(trunc_lossfun, symm_r) + iter_rrule = Mooncake.build_rrule(iter_lossfun, symm_r) + + l_full, g_full = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, symm_r) + l_trunc, g_trunc = Mooncake.value_and_gradient!!(trunc_rrule, trunc_lossfun, symm_r) + l_iter, g_iter = Mooncake.value_and_gradient!!(iter_rrule, iter_lossfun, symm_r) + + @test l_full ≈ l_trunc ≈ l_iter + @test g_full[2] ≈ g_trunc[2] rtol = rtol + @test g_full[2] ≈ g_iter[2] rtol = rtol + @test g_trunc[2] ≈ g_iter[2] rtol = rtol + + full_lossfun = A -> lossfun(A, full_alg, symm_R, symm_trspace) + trunc_lossfun = A -> lossfun(A, trunc_alg, symm_R, symm_trspace) + iter_lossfun = A -> lossfun(A, iter_alg, symm_R, symm_trspace) + + full_rrule = Mooncake.build_rrule(full_lossfun, symm_r) + trunc_rrule = Mooncake.build_rrule(trunc_lossfun, symm_r) + iter_rrule = Mooncake.build_rrule(iter_lossfun, symm_r) + + l_full_tr, g_full_tr = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, symm_r) + l_trunc_tr, g_trunc_tr = Mooncake.value_and_gradient!!(trunc_rrule, trunc_lossfun, symm_r) + l_iter_tr, g_iter_tr = Mooncake.value_and_gradient!!(iter_rrule, iter_lossfun, symm_r) + @test l_full_tr ≈ l_trunc_tr ≈ l_iter_tr + @test g_full_tr[2] ≈ g_trunc_tr[2] rtol = rtol + @test g_full_tr[2] ≈ g_iter_tr[2] rtol = rtol + @test g_trunc_tr[2] ≈ g_iter_tr[2] rtol = rtol + + iter_alg_fallback = @set iter_alg.fwd_alg.fallback_threshold = 0.4 # Do dense decomposition in one block, sparse one in the other + + fb_lossfun = A -> lossfun(A, iter_alg_fallback, symm_R, symm_trspace) + fb_rrule = Mooncake.build_rrule(fb_lossfun, symm_r) + l_iter_fb, g_iter_fb = Mooncake.value_and_gradient!!(fb_rrule, fb_lossfun, symm_r) + @test l_iter_fb ≈ l_trunc_tr ≈ l_full_tr + @test g_full_tr[2] ≈ g_iter_fb[2] rtol = rtol + @test g_trunc_tr[2] ≈ g_iter_fb[2] rtol = rtol +end +#= +@testset "Truncated symmetric SVD broadening for $(alg.rrule_alg)" for alg in [full_alg, trunc_alg] + u, s, v, = svd_compact(symm_r) + # make every singular value in the 0-sector three-fold degenerate + b0 = diagview(block(s, Z2Irrep(0))) + b0[1:3:symm_m] .= b0[3:3:symm_m] + b0[2:3:symm_m] .= b0[3:3:symm_m] + # make every singular value in the 1-sector two-fold degenerate + b1 = diagview(block(s, Z2Irrep(1))) + b1[1:2:symm_n] .= b1[2:2:symm_n] + symm_r_degen = u * s * v + + no_broadening_no_cutoff_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-30 + small_broadening_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-13 + + l_only_cutoff, g_only_cutoff = Mooncake.value_and_gradient!!( + A -> lossfun(A, alg, symm_R, symm_trspace), symm_r_degen + ) # cutoff sets degenerate difference to zero + l_no_broadening_no_cutoff, g_no_broadening_no_cutoff = Mooncake.value_and_gradient!!( # degenerate singular value differences lead to divergent contributions + A -> lossfun(A, no_broadening_no_cutoff_alg, symm_R, symm_trspace), + symm_r_degen, + ) + l_small_broadening, g_small_broadening = Mooncake.value_and_gradient!!( # broadening smoothens divergent contributions + A -> lossfun(A, small_broadening_alg, symm_R, symm_trspace), + symm_r_degen, + ) + + @test l_only_cutoff ≈ l_no_broadening_no_cutoff ≈ l_small_broadening + @test norm(g_no_broadening_no_cutoff[1] - g_small_broadening[1]) > 1.0e-2 # divergences mess up the gradient + @test g_only_cutoff[1] ≈ g_small_broadening[1] rtol = rtol # cutoff and broadening have similar effect +end +=# +# TODO: Add when IterSVD is implemented for HalfInfiniteEnv +# χbond = 2 +# χenv = 6 +# ctm_alg = CTMRG(; tol=1e-10, verbosity=2, svd_alg=SVDAdjoint()) +# Random.seed!(91283219347) +# H = heisenberg_XYZ(InfiniteSquare()) +# psi = InfinitePEPS(ComplexSpace(2), ComplexSpace(χbond)) +# env = leading_boundary(CTMRGEnv(psi, ComplexSpace(χenv)), psi, ctm_alg); +# hienv = HalfInfiniteEnv( +# env.corners[1], +# env.corners[2], +# env.edges[4], +# env.edges[1], +# env.edges[1], +# env.edges[2], +# psi[1], +# psi[1], +# psi[1], +# psi[1], +# ) +# hienv_dense = hienv() +# env_R = randn(space(hienv)) + +# svd_trunc!(hienv, iter_alg) + +# @testset "IterSVD with HalfInfiniteEnv function handle" begin +# # Equivalence of dense and sparse contractions +# x₀ = PEPSKit.random_start_vector(hienv) +# x′ = hienv(x₀, Val(false)) +# x″ = hienv(x′, Val(true)) +# x‴ = hienv(x″, Val(false)) + +# a = hienv_dense * x₀ +# b = hienv_dense' * a +# c = hienv_dense * b +# @test a ≈ x′ +# @test b ≈ x″ +# @test c ≈ x‴ + +# # l_fullsvd, g_fullsvd = withgradient(A -> lossfun(A, full_alg, env_R), hienv_dense) +# # l_itersvd, g_itersvd = withgradient(A -> lossfun(A, iter_alg, env_R), hienv) +# # @test l_itersvd ≈ l_fullsvd +# # @test g_fullsvd[1] ≈ g_itersvd[1] rtol = rtol +# end From 11083f071d83529b6ae3c00f9a4086684d2b74c0 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 1 Jun 2026 19:27:26 +0200 Subject: [PATCH 2/8] Gotta wait for some tags --- Project.toml | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/Project.toml b/Project.toml index a14dbb5e3..cc8584868 100644 --- a/Project.toml +++ b/Project.toml @@ -37,12 +37,12 @@ Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" PEPSKitMooncakeExt = "Mooncake" [sources] -BlockTensorKit = {path = "/Users/khyatt/.julia/dev/BlockTensorKit"} -TensorKit = {path = "/Users/khyatt/.julia/dev/TensorKit"} +BlockTensorKit = {url = "https://github.com/quantumkithub/BlockTensorKit.jl", rev = "main"} +TensorKit = {url = "https://github.com/quantumkithub/TensorKit.jl", rev = "main"} MPSKit = {url = "https://github.com/quantumkithub/mpskit.jl", rev = "main"} OptimKit = {url = "https://github.com/kshyatt/optimkit.jl", rev = "patch-1"} KrylovKit = {url = "https://github.com/kshyatt/krylovkit.jl", rev = "patch-1"} -VectorInterface = {path = "/Users/khyatt/.julia/dev/VectorInterface"} +VectorInterface = {url = "https://github.com/quantumkithub/vectorinterface.jl", rev = "main"} [compat] Accessors = "0.1" From 34d1b311739e342a5cf5a21620ab6d64a1f914d7 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 1 Jun 2026 20:21:02 +0200 Subject: [PATCH 3/8] Eigh pbs and tests --- ext/PEPSKitMooncakeExt.jl | 53 +++++++++- test/mooncake/eigh_wrapper.jl | 181 ++++++++++++++++++++++++++++++++++ test/mooncake/svd_wrapper.jl | 6 +- 3 files changed, 236 insertions(+), 4 deletions(-) create mode 100644 test/mooncake/eigh_wrapper.jl diff --git a/ext/PEPSKitMooncakeExt.jl b/ext/PEPSKitMooncakeExt.jl index 24963d654..74b528fcf 100644 --- a/ext/PEPSKitMooncakeExt.jl +++ b/ext/PEPSKitMooncakeExt.jl @@ -1,13 +1,14 @@ module PEPSKitMooncakeExt using PEPSKit, TensorKit, Mooncake, MatrixAlgebraKit -using PEPSKit: SVDAdjoint +using PEPSKit: SVDAdjoint, EighAdjoint using Mooncake: DefaultCtx, CoDual, Dual, NoRData, primal, rrule!!, arrayify, @is_primitive _warn_pullback_truncerror(dϵ::Real; tol = MatrixAlgebraKit.defaulttol(dϵ)) = abs(dϵ) ≤ tol || @warn "Pullback ignores non-zero tangents for truncation error" Mooncake.tangent_type(::Type{<:PEPSKit.SVDAdjoint}) = Mooncake.NoTangent +Mooncake.tangent_type(::Type{<:PEPSKit.EighAdjoint}) = Mooncake.NoTangent @is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(svd_trunc), TensorKit.AbstractTensorMap, SVDAdjoint} function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{SVDAdjoint{F, R}}) where {F, R <: PEPSKit.FullPullback} @@ -58,4 +59,54 @@ function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::Co return output_codual, svd_trunc!_trunc_pullback end +@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(eigh_trunc), TensorKit.AbstractTensorMap, EighAdjoint} +function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.eigh_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{EighAdjoint{F, R}}) where {F, R <: PEPSKit.FullPullback} + t, dt = arrayify(t_dt) + alg = primal(alg_dalg) + + D, V = eigh_full!(t; alg.fwd_alg.alg) + (D̃, Ṽ), inds = MatrixAlgebraKit.truncate(eigh_trunc!, (D, V), alg.fwd_alg.trunc) + ϵ = MatrixAlgebraKit.truncation_error(diagview(D), inds) + + DVtrunc = (D̃, Ṽ) + # pack output + DVtrunc_dDVtrunc = Mooncake.zero_fcodual((DVtrunc..., ϵ)) + + # define pullback + dDVtrunc = last.(arrayify.(DVtrunc, Base.front(Mooncake.tangent(DVtrunc_dDVtrunc)))) + + gtol = PEPSKit._get_pullback_gauge_tol(alg.rrule_alg.verbosity) + function eigh_trunc!_full_pullback((_, _, dϵ)::Tuple{NoRData, NoRData, Real}) + _warn_pullback_truncerror(dϵ) + MatrixAlgebraKit.eigh_pullback!(dt, t, (D, V), dDVtrunc, inds; gauge_atol = gtol(dDVtrunc), degeneracy_atol = alg.rrule_alg.degeneracy_atol) + MatrixAlgebraKit.zero!.(dDVtrunc) # since this is allocated in this function this is probably not required + return ntuple(Returns(NoRData()), 3) + end + return DVtrunc_dDVtrunc, eigh_trunc!_full_pullback +end + +function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.eigh_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{EighAdjoint{F, R}}) where {F, R <: PEPSKit.TruncPullback} + t, dt = arrayify(t_dt) + alg = primal(alg_dalg) + + D, V, truncerror = eigh_trunc(t, alg) + gtol = PEPSKit._get_pullback_gauge_tol(alg.rrule_alg.verbosity) + output = (D, V, truncerror) + output_codual = CoDual(output, Mooncake.fdata(Mooncake.zero_tangent(output))) + + gtol = PEPSKit._get_pullback_gauge_tol(alg.rrule_alg.verbosity) + function eigh_trunc!_trunc_pullback((_, _, dϵ)::Tuple{NoRData, NoRData, Real}) + _warn_pullback_truncerror(dϵ) + Dtrunc, Vtrunc, ϵ = Mooncake.primal(output_codual) + dDtrunc_, dVtrunc_, dϵ = Mooncake.tangent(output_codual) + D, dD = arrayify(Dtrunc, dDtrunc_) + V, dV = arrayify(Vtrunc, dVtrunc_) + MatrixAlgebraKit.eigh_trunc_pullback!(dt, t, (D, V), (dD, dV); gauge_atol = gtol((dD, dV)), degeneracy_atol = alg.rrule_alg.degeneracy_atol) + MatrixAlgebraKit.zero!(dD) # since this is allocated in this function this is probably not required + MatrixAlgebraKit.zero!(dV) # since this is allocated in this function this is probably not required + return ntuple(Returns(NoRData()), 3) + end + return output_codual, eigh_trunc!_trunc_pullback +end + end diff --git a/test/mooncake/eigh_wrapper.jl b/test/mooncake/eigh_wrapper.jl new file mode 100644 index 000000000..a05adaf41 --- /dev/null +++ b/test/mooncake/eigh_wrapper.jl @@ -0,0 +1,181 @@ + +using Test +using Random +using LinearAlgebra +using TensorKit +using Mooncake +using Accessors +using PEPSKit + +using MatrixAlgebraKit: TruncatedAlgorithm, diagview + +# Gauge-invariant loss function +function lossfun(A, alg, R = randn(space(A)), trunc = notrunc()) + alg = @set alg.fwd_alg = TruncatedAlgorithm(alg.fwd_alg, trunc) + D, V, = eigh_trunc(project_hermitian(A), alg) + return real(dot(R, V * V')) + dot(D, D) # Overlap with random tensor R is gauge-invariant and differentiable +end + +dtype = ComplexF64 +n = 20 +χ = 10 +trunc = truncspace(ℂ^χ) +rtol = 1.0e-9 +Random.seed!(123456789) +r = randn(dtype, ℂ^n, ℂ^n) +r = 0.5 * (r + r') # make r Hermitian +R = randn(space(r)) +R = 0.5 * (R + R') + +full_alg = EighAdjoint(; fwd_alg = (; alg = :QRIteration), rrule_alg = (; alg = :FullPullback)) +trunc_alg = EighAdjoint(; fwd_alg = (; alg = :QRIteration), rrule_alg = (; alg = :TruncPullback)) +iter_alg = EighAdjoint(; fwd_alg = (; alg = :Lanczos), rrule_alg = (; alg = :TruncPullback)) + +@testset "Non-truncated eigh" begin + full_lossfun = A -> lossfun(A, full_alg, R) + trunc_lossfun = A -> lossfun(A, trunc_alg, R) + iter_lossfun = A -> lossfun(A, iter_alg, R) + + full_rrule = Mooncake.build_rrule(full_lossfun, r) + trunc_rrule = Mooncake.build_rrule(trunc_lossfun, r) + iter_rrule = Mooncake.build_rrule(iter_lossfun, r) + + l_full, g_full = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, r) + l_trunc, g_trunc = Mooncake.value_and_gradient!!(trunc_rrule, trunc_lossfun, r) + l_iter, g_iter = Mooncake.value_and_gradient!!(iter_rrule, iter_lossfun, r) + + @test l_full ≈ l_trunc ≈ l_iter + @test g_full[2] ≈ g_trunc[2] rtol = rtol + @test g_full[2] ≈ g_iter[2] rtol = rtol + @test g_trunc[2] ≈ g_iter[2] rtol = rtol +end + +@testset "Truncated eigh with χ=$χ" begin + full_lossfun = A -> lossfun(A, full_alg, R, trunc) + trunc_lossfun = A -> lossfun(A, trunc_alg, R, trunc) + iter_lossfun = A -> lossfun(A, iter_alg, R, trunc) + + full_rrule = Mooncake.build_rrule(full_lossfun, r) + trunc_rrule = Mooncake.build_rrule(trunc_lossfun, r) + iter_rrule = Mooncake.build_rrule(iter_lossfun, r) + + l_full, g_full = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, r) + l_trunc, g_trunc = Mooncake.value_and_gradient!!(trunc_rrule, trunc_lossfun, r) + l_iter, g_iter = Mooncake.value_and_gradient!!(iter_rrule, iter_lossfun, r) + + @test l_full ≈ l_trunc ≈ l_iter + @test g_full[2] ≈ g_trunc[2] rtol = rtol + @test g_full[2] ≈ g_iter[2] rtol = rtol + @test g_trunc[2] ≈ g_iter[2] rtol = rtol +end + +@testset "Truncated eigh broadening for $(alg.rrule_alg)" for alg in [full_alg, trunc_alg] + d, v = eigh_full(r) + d.data[1:2:n] .= d.data[2:2:n] # make every eigenvalue two-fold degenerate + r_degen = v * d * v' + + no_broadening_no_cutoff_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-30 + small_broadening_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-13 + + only_lossfun = A -> lossfun(A, alg, R, trunc) + no_broadening_lossfun = A -> lossfun(A, no_broadening_no_cutoff_alg, R, trunc) + small_broadening_lossfun = A -> lossfun(A, small_broadening_alg, R, trunc) + + only_rrule = Mooncake.build_rrule(only_lossfun, r_degen) + no_broadening_rrule = Mooncake.build_rrule(no_broadening_lossfun, r_degen) + small_broadening_rrule = Mooncake.build_rrule(small_broadening_lossfun, r_degen) + + l_only_cutoff, g_only_cutoff = Mooncake.value_and_gradient!!(only_rrule, only_lossfun, r_degen) # cutoff sets degenerate difference to zero + l_no_broadening_no_cutoff, g_no_broadening_no_cutoff = Mooncake.value_and_gradient!!( # degenerate singular value differences lead to divergent contributions + no_broadening_rrule, no_broadening_lossfun, r_degen, + ) + l_small_broadening, g_small_broadening = Mooncake.value_and_gradient!!( # broadening smoothens divergent contributions + small_broadening_rrule, small_broadening_lossfun, r_degen, + ) + + @test l_only_cutoff ≈ l_no_broadening_no_cutoff ≈ l_small_broadening + @test norm(g_no_broadening_no_cutoff[2] - g_small_broadening[2]) > 1.0e-2 # divergences mess up the gradient + @test g_only_cutoff[2] ≈ g_small_broadening[2] rtol = rtol # cutoff and broadening have similar effect +end + +symm_m, symm_n = 18, 24 +symm_space = Z2Space(0 => symm_m, 1 => symm_n) +symm_trspace = truncspace(Z2Space(0 => symm_m ÷ 2, 1 => symm_n ÷ 3)) +symm_r = randn(dtype, symm_space, symm_space) +symm_r = 0.5 * (symm_r + symm_r') +symm_R = randn(dtype, space(symm_r)) +symm_R = 0.5 * (symm_R + symm_R') + +@testset "IterEig of symmetric tensors" begin + full_lossfun = A -> lossfun(A, full_alg, symm_R) + trunc_lossfun = A -> lossfun(A, trunc_alg, symm_R) + iter_lossfun = A -> lossfun(A, iter_alg, symm_R) + + full_rrule = Mooncake.build_rrule(full_lossfun, symm_r) + trunc_rrule = Mooncake.build_rrule(trunc_lossfun, symm_r) + iter_rrule = Mooncake.build_rrule(iter_lossfun, symm_r) + + l_full, g_full = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, symm_r) + l_trunc, g_trunc = Mooncake.value_and_gradient!!(trunc_rrule, trunc_lossfun, symm_r) + l_iter, g_iter = Mooncake.value_and_gradient!!(iter_rrule, iter_lossfun, symm_r) + + @test l_full ≈ l_trunc ≈ l_iter + @test g_full[2] ≈ g_trunc[2] rtol = rtol + @test g_full[2] ≈ g_iter[2] rtol = rtol + @test g_trunc[2] ≈ g_iter[2] rtol = rtol + + full_lossfun = A -> lossfun(A, full_alg, symm_R, symm_trspace) + trunc_lossfun = A -> lossfun(A, trunc_alg, symm_R, symm_trspace) + iter_lossfun = A -> lossfun(A, iter_alg, symm_R, symm_trspace) + + full_rrule = Mooncake.build_rrule(full_lossfun, symm_r) + trunc_rrule = Mooncake.build_rrule(trunc_lossfun, symm_r) + iter_rrule = Mooncake.build_rrule(iter_lossfun, symm_r) + + l_full_tr, g_full_tr = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, symm_r) + l_trunc_tr, g_trunc_tr = Mooncake.value_and_gradient!!(trunc_rrule, trunc_lossfun, symm_r) + l_iter_tr, g_iter_tr = Mooncake.value_and_gradient!!(iter_rrule, iter_lossfun, symm_r) + @test l_full_tr ≈ l_trunc_tr ≈ l_iter_tr + @test g_full_tr[2] ≈ g_trunc_tr[2] rtol = rtol + @test g_full_tr[2] ≈ g_iter_tr[2] rtol = rtol + @test g_trunc_tr[2] ≈ g_iter_tr[2] rtol = rtol + + iter_alg_fallback = @set iter_alg.fwd_alg.fallback_threshold = 0.4 # Do dense decomposition in one block, sparse one in the other + fb_lossfun = A -> lossfun(A, iter_alg_fallback, symm_R, symm_trspace) + fb_rrule = Mooncake.build_rrule(fb_lossfun, symm_r) + l_iter_fb, g_iter_fb = Mooncake.value_and_gradient!!(fb_rrule, fb_lossfun, symm_r) + @test l_iter_fb ≈ l_trunc_tr ≈ l_full_tr + @test g_full_tr[2] ≈ g_iter_fb[2] rtol = rtol + @test g_trunc_tr[2] ≈ g_iter_fb[2] rtol = rtol +end +#= +@testset "Truncated symmetric eigh broadening for $(alg.rrule_alg)" for alg in [full_alg, trunc_alg] + d, v = eigh_full(symm_r) + # make every singular value in the 0-sector three-fold degenerate + b0 = diagview(block(d, Z2Irrep(0))) + b0[1:3:symm_m] .= b0[3:3:symm_m] + b0[2:3:symm_m] .= b0[3:3:symm_m] + # make every singular value in the 1-sector two-fold degenerate + b1 = diagview(block(d, Z2Irrep(1))) + b1[1:2:symm_n] .= b1[2:2:symm_n] + symm_r_degen = v * d * v' + + no_broadening_no_cutoff_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-30 + small_broadening_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-13 + + l_only_cutoff, g_only_cutoff = withgradient( + A -> lossfun(A, alg, symm_R, symm_trspace), symm_r_degen + ) # cutoff sets degenerate difference to zero + l_no_broadening_no_cutoff, g_no_broadening_no_cutoff = withgradient( # degenerate singular value differences lead to divergent contributions + A -> lossfun(A, no_broadening_no_cutoff_alg, symm_R, symm_trspace), + symm_r_degen, + ) + l_small_broadening, g_small_broadening = withgradient( # broadening smoothens divergent contributions + A -> lossfun(A, small_broadening_alg, symm_R, symm_trspace), + symm_r_degen, + ) + + @test l_only_cutoff ≈ l_no_broadening_no_cutoff ≈ l_small_broadening + @test norm(g_no_broadening_no_cutoff[1] - g_small_broadening[1]) > 1.0e-2 # divergences mess up the gradient + @test g_only_cutoff[1] ≈ g_small_broadening[1] rtol = rtol # cutoff and broadening have similar effect +end=# diff --git a/test/mooncake/svd_wrapper.jl b/test/mooncake/svd_wrapper.jl index af0ac2acb..0a549ac39 100644 --- a/test/mooncake/svd_wrapper.jl +++ b/test/mooncake/svd_wrapper.jl @@ -75,15 +75,15 @@ end no_broadening_no_cutoff_alg = @set full_alg.rrule_alg.degeneracy_atol = 1.0e-30 small_broadening_alg = @set full_alg.rrule_alg.degeneracy_atol = 1.0e-13 - full_lossfun = A -> lossfun(A, full_alg, R, trunc) + only_lossfun = A -> lossfun(A, alg, R, trunc) no_broadening_lossfun = A -> lossfun(A, no_broadening_no_cutoff_alg, R, trunc) small_broadening_lossfun = A -> lossfun(A, small_broadening_alg, R, trunc) - full_rrule = Mooncake.build_rrule(full_lossfun, r_degen) + only_rrule = Mooncake.build_rrule(only_lossfun, r_degen) no_broadening_rrule = Mooncake.build_rrule(no_broadening_lossfun, r_degen) small_broadening_rrule = Mooncake.build_rrule(small_broadening_lossfun, r_degen) - l_only_cutoff, g_only_cutoff = Mooncake.value_and_gradient!!(full_rrule, full_lossfun, r_degen) # cutoff sets degenerate difference to zero + l_only_cutoff, g_only_cutoff = Mooncake.value_and_gradient!!(only_rrule, only_lossfun, r_degen) # cutoff sets degenerate difference to zero l_no_broadening_no_cutoff, g_no_broadening_no_cutoff = Mooncake.value_and_gradient!!( # degenerate singular value differences lead to divergent contributions no_broadening_rrule, no_broadening_lossfun, r_degen, ) From 44b253e877af16dc793003b9afecd06899ea613f Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 1 Jun 2026 20:28:59 +0200 Subject: [PATCH 4/8] Enable some more tests --- test/mooncake/svd_wrapper.jl | 26 +++++++++++++++----------- 1 file changed, 15 insertions(+), 11 deletions(-) diff --git a/test/mooncake/svd_wrapper.jl b/test/mooncake/svd_wrapper.jl index 0a549ac39..2c010ca28 100644 --- a/test/mooncake/svd_wrapper.jl +++ b/test/mooncake/svd_wrapper.jl @@ -145,7 +145,7 @@ symm_R = randn(dtype, space(symm_r)) @test g_full_tr[2] ≈ g_iter_fb[2] rtol = rtol @test g_trunc_tr[2] ≈ g_iter_fb[2] rtol = rtol end -#= + @testset "Truncated symmetric SVD broadening for $(alg.rrule_alg)" for alg in [full_alg, trunc_alg] u, s, v, = svd_compact(symm_r) # make every singular value in the 0-sector three-fold degenerate @@ -159,24 +159,28 @@ end no_broadening_no_cutoff_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-30 small_broadening_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-13 + + only_lossfun = A -> lossfun(A, alg, symm_R, symm_trspace) + no_broadening_lossfun = A -> lossfun(A, no_broadening_no_cutoff_alg, symm_R, symm_trspace) + small_broadening_lossfun = A -> lossfun(A, small_broadening_alg, symm_R, symm_trspace) + + only_rrule = Mooncake.build_rrule(only_lossfun, symm_r_degen) + no_broadening_rrule = Mooncake.build_rrule(no_broadening_lossfun, symm_r_degen) + small_broadening_rrule = Mooncake.build_rrule(small_broadening_lossfun, symm_r_degen) - l_only_cutoff, g_only_cutoff = Mooncake.value_and_gradient!!( - A -> lossfun(A, alg, symm_R, symm_trspace), symm_r_degen - ) # cutoff sets degenerate difference to zero + l_only_cutoff, g_only_cutoff = Mooncake.value_and_gradient!!(only_rrule, only_lossfun, symm_r_degen) # cutoff sets degenerate difference to zero l_no_broadening_no_cutoff, g_no_broadening_no_cutoff = Mooncake.value_and_gradient!!( # degenerate singular value differences lead to divergent contributions - A -> lossfun(A, no_broadening_no_cutoff_alg, symm_R, symm_trspace), - symm_r_degen, + no_broadening_rrule, no_broadening_lossfun, symm_r_degen, ) l_small_broadening, g_small_broadening = Mooncake.value_and_gradient!!( # broadening smoothens divergent contributions - A -> lossfun(A, small_broadening_alg, symm_R, symm_trspace), - symm_r_degen, + small_broadening_rrule, small_broadening_lossfun, symm_r_degen, ) @test l_only_cutoff ≈ l_no_broadening_no_cutoff ≈ l_small_broadening - @test norm(g_no_broadening_no_cutoff[1] - g_small_broadening[1]) > 1.0e-2 # divergences mess up the gradient - @test g_only_cutoff[1] ≈ g_small_broadening[1] rtol = rtol # cutoff and broadening have similar effect + #@test norm(g_no_broadening_no_cutoff[2] - g_small_broadening[2]) > 1.0e-2 # divergences mess up the gradient + @test g_only_cutoff[2] ≈ g_small_broadening[2] rtol = rtol # cutoff and broadening have similar effect end -=# + # TODO: Add when IterSVD is implemented for HalfInfiniteEnv # χbond = 2 # χenv = 6 From e442d8fd4c098b91b6c39bfa0725106911d12397 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Tue, 2 Jun 2026 12:43:43 +0200 Subject: [PATCH 5/8] Add support for QRAdjoint --- ext/PEPSKitMooncakeExt.jl | 22 +++++++++++++++++++++- 1 file changed, 21 insertions(+), 1 deletion(-) diff --git a/ext/PEPSKitMooncakeExt.jl b/ext/PEPSKitMooncakeExt.jl index 74b528fcf..9a1aab366 100644 --- a/ext/PEPSKitMooncakeExt.jl +++ b/ext/PEPSKitMooncakeExt.jl @@ -1,7 +1,7 @@ module PEPSKitMooncakeExt using PEPSKit, TensorKit, Mooncake, MatrixAlgebraKit -using PEPSKit: SVDAdjoint, EighAdjoint +using PEPSKit: SVDAdjoint, EighAdjoint, QRAdjoint using Mooncake: DefaultCtx, CoDual, Dual, NoRData, primal, rrule!!, arrayify, @is_primitive _warn_pullback_truncerror(dϵ::Real; tol = MatrixAlgebraKit.defaulttol(dϵ)) = @@ -9,6 +9,7 @@ _warn_pullback_truncerror(dϵ::Real; tol = MatrixAlgebraKit.defaulttol(dϵ)) = Mooncake.tangent_type(::Type{<:PEPSKit.SVDAdjoint}) = Mooncake.NoTangent Mooncake.tangent_type(::Type{<:PEPSKit.EighAdjoint}) = Mooncake.NoTangent +Mooncake.tangent_type(::Type{<:PEPSKit.QRAdjoint}) = Mooncake.NoTangent @is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(svd_trunc), TensorKit.AbstractTensorMap, SVDAdjoint} function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{SVDAdjoint{F, R}}) where {F, R <: PEPSKit.FullPullback} @@ -109,4 +110,23 @@ function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.eigh_trunc)}, t_dt::C return output_codual, eigh_trunc!_trunc_pullback end +@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(left_orth), TensorKit.AbstractTensorMap, QRAdjoint} +function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.left_orth)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{QRAdjoint}) + t, dt = arrayify(t_dt) + alg = primal(alg_dalg) + + QR = left_orth(t, alg) + gtol = PEPSKit._get_pullback_gauge_tol(alg.rrule_alg.verbosity) + + output_codual = Mooncake.zero_fcodual(QR) + dQ_, dR_ = Mooncake.tangent(output_codual) + Q, dQ = arrayify(Q, dQ_) + R, dR = arrayify(R, dR_) + function left_orth_pullback(::NoRData) + MatrixAlgebraKit.qr_pullback!(dt, t, QR, (dQ, dR); gauge_atol = gtol(dQR)) + return ntuple(Returns(NoRData()), 3) + end + return output_codual, left_orth_pullback +end + end From f3a8b7cd788f9ccae02016c8c8fb64c57d00ed98 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 10 Jun 2026 15:33:29 +0200 Subject: [PATCH 6/8] Some more Mooncake support --- ext/PEPSKitMooncakeExt.jl | 28 +++++++++++++++++++++------- 1 file changed, 21 insertions(+), 7 deletions(-) diff --git a/ext/PEPSKitMooncakeExt.jl b/ext/PEPSKitMooncakeExt.jl index 9a1aab366..c06084400 100644 --- a/ext/PEPSKitMooncakeExt.jl +++ b/ext/PEPSKitMooncakeExt.jl @@ -1,8 +1,15 @@ module PEPSKitMooncakeExt -using PEPSKit, TensorKit, Mooncake, MatrixAlgebraKit -using PEPSKit: SVDAdjoint, EighAdjoint, QRAdjoint -using Mooncake: DefaultCtx, CoDual, Dual, NoRData, primal, rrule!!, arrayify, @is_primitive +using PEPSKit, MPSKit, TensorKit, Mooncake, MatrixAlgebraKit +using PEPSKit: SVDAdjoint, EighAdjoint, QRAdjoint, CTMRGAlgorithm, FixedPointGradient, sdiag_pow +import PEPSKit: real_inner +using Mooncake: DefaultCtx, CoDual, Dual, NoRData, primal, tangent, rrule!!, arrayify, @is_primitive + +function Mooncake.arrayify(ψ::PEPSKit.InfinitePEPS{T}, dψ) where {T} + Δψmat = map((a, da) -> Mooncake.arrayify(a, da)[2], ψ.A, dψ.fields.A) + Δψ = PEPSKit.InfinitePEPS{T}(Δψmat) + return ψ, Δψ +end _warn_pullback_truncerror(dϵ::Real; tol = MatrixAlgebraKit.defaulttol(dϵ)) = abs(dϵ) ≤ tol || @warn "Pullback ignores non-zero tangents for truncation error" @@ -10,6 +17,13 @@ _warn_pullback_truncerror(dϵ::Real; tol = MatrixAlgebraKit.defaulttol(dϵ)) = Mooncake.tangent_type(::Type{<:PEPSKit.SVDAdjoint}) = Mooncake.NoTangent Mooncake.tangent_type(::Type{<:PEPSKit.EighAdjoint}) = Mooncake.NoTangent Mooncake.tangent_type(::Type{<:PEPSKit.QRAdjoint}) = Mooncake.NoTangent +Mooncake.tangent_type(::Type{<:PEPSKit.CTMRGAlgorithm}) = Mooncake.NoTangent +Mooncake.tangent_type(::Type{<:PEPSKit.FixedPointGradient}) = Mooncake.NoTangent + +Mooncake.@zero_derivative Mooncake.DefaultCtx Tuple{typeof(PEPSKit.eachcoordinate), Any, Any} +Mooncake.@zero_derivative Mooncake.DefaultCtx Tuple{typeof(PEPSKit._next_coordinate), Int, Int} +Mooncake.@zero_derivative Mooncake.DefaultCtx Tuple{typeof(PEPSKit._set_decomposition_truncation), Any, Any} +Mooncake.@zero_derivative Mooncake.DefaultCtx Tuple{typeof(PEPSKit.CTMRGEnv), Union{PEPSKit.InfinitePartitionFunction, PEPSKit.InfinitePEPS}, Vararg} @is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(svd_trunc), TensorKit.AbstractTensorMap, SVDAdjoint} function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{SVDAdjoint{F, R}}) where {F, R <: PEPSKit.FullPullback} @@ -32,7 +46,7 @@ function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::Co Δt, t, (U, S, V⁺), ΔUSVᴴtrunc, inds; gauge_atol = gtol(ΔUSVᴴtrunc), degeneracy_atol = alg.rrule_alg.degeneracy_atol, ) - return NoRData(), NoRData(), NoRData() + return NoRData(), NoRData(), NoRData(), zero(dϵ) end return output_codual, svd_trunc!_full_pullback end @@ -64,11 +78,11 @@ end function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.eigh_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{EighAdjoint{F, R}}) where {F, R <: PEPSKit.FullPullback} t, dt = arrayify(t_dt) alg = primal(alg_dalg) - + D, V = eigh_full!(t; alg.fwd_alg.alg) (D̃, Ṽ), inds = MatrixAlgebraKit.truncate(eigh_trunc!, (D, V), alg.fwd_alg.trunc) ϵ = MatrixAlgebraKit.truncation_error(diagview(D), inds) - + DVtrunc = (D̃, Ṽ) # pack output DVtrunc_dDVtrunc = Mooncake.zero_fcodual((DVtrunc..., ϵ)) @@ -89,7 +103,7 @@ end function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.eigh_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{EighAdjoint{F, R}}) where {F, R <: PEPSKit.TruncPullback} t, dt = arrayify(t_dt) alg = primal(alg_dalg) - + D, V, truncerror = eigh_trunc(t, alg) gtol = PEPSKit._get_pullback_gauge_tol(alg.rrule_alg.verbosity) output = (D, V, truncerror) From 80afa290e86318c727c263f38617f1ed1c0b00ae Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Wed, 10 Jun 2026 15:36:39 +0200 Subject: [PATCH 7/8] Cleanup Project.toml --- Project.toml | 15 +-------------- 1 file changed, 1 insertion(+), 14 deletions(-) diff --git a/Project.toml b/Project.toml index cc8584868..47efcf211 100644 --- a/Project.toml +++ b/Project.toml @@ -8,7 +8,6 @@ projects = ["test", "docs", "benchmark"] [deps] Accessors = "7d9f7c33-5ae7-4f3b-8dc6-eff91059b697" -BlockTensorKit = "5f87ffc2-9cf1-4a46-8172-465d160bd8cd" ChainRulesCore = "d360d2e6-b24c-11e9-a2a3-2a2ae2dbcce4" Compat = "34da2185-b29b-5c13-b0c7-acf172513d20" DocStringExtensions = "ffbed154-4ef7-542d-bbb7-c09d3a79fcae" @@ -19,6 +18,7 @@ LoggingExtras = "e6f89c97-d47a-5376-807f-9c37f3926c36" MPSKit = "bb1c41ca-d63c-52ed-829e-0820dda26502" MPSKitModels = "ca635005-6f8c-4cd1-b51d-8491250ef2ab" MatrixAlgebraKit = "6c742aac-3347-4629-af66-fc926824e5e4" +Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" OhMyThreads = "67456a42-1dca-4109-a031-0a68de7e3ad5" OptimKit = "77e91f04-9b3b-57a6-a776-40b61faaebe0" Printf = "de0858da-6303-5e67-8744-51eddeeeb8d7" @@ -28,22 +28,10 @@ TensorKit = "07d1fe3e-3e46-537d-9eac-e9e13d0d4cec" TensorOperations = "6aa20fa7-93e2-5fca-9bc0-fbd0db3c71a2" TupleTools = "9d95972d-f1c8-5527-a6e0-b4b365fa01f6" VectorInterface = "409d34a3-91d5-4945-b6ec-7529ddf182d8" -Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f" - -[weakdeps] -Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" [extensions] PEPSKitMooncakeExt = "Mooncake" -[sources] -BlockTensorKit = {url = "https://github.com/quantumkithub/BlockTensorKit.jl", rev = "main"} -TensorKit = {url = "https://github.com/quantumkithub/TensorKit.jl", rev = "main"} -MPSKit = {url = "https://github.com/quantumkithub/mpskit.jl", rev = "main"} -OptimKit = {url = "https://github.com/kshyatt/optimkit.jl", rev = "patch-1"} -KrylovKit = {url = "https://github.com/kshyatt/krylovkit.jl", rev = "patch-1"} -VectorInterface = {url = "https://github.com/quantumkithub/vectorinterface.jl", rev = "main"} - [compat] Accessors = "0.1" ChainRulesCore = "1.0" @@ -66,5 +54,4 @@ TensorKit = "0.16.5, 0.17" TensorOperations = "5" TupleTools = "1.6.0" VectorInterface = "0.6" -Zygote = "0.6, 0.7" julia = "1.10" From 99efe722b24cd8da2f23f27f1ac3e7eb5b17fa58 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 15 Jun 2026 10:34:28 +0200 Subject: [PATCH 8/8] Incremental stuff --- Project.toml | 11 +- ext/PEPSKitMooncakeExt.jl | 147 ++++++++++++++++-- src/PEPSKit.jl | 2 +- src/algorithms/ctmrg/c4v.jl | 8 +- src/algorithms/ctmrg/projectors.jl | 16 +- src/algorithms/ctmrg/sequential.jl | 4 +- src/algorithms/ctmrg/simultaneous.jl | 40 ++--- .../fixed_point_differentiation.jl | 9 +- src/algorithms/toolbox.jl | 10 +- src/environments/ctmrg_environments.jl | 18 +-- src/utility/diffable_threads.jl | 6 +- src/utility/indexing.jl | 2 + src/utility/util.jl | 17 -- test/Project.toml | 1 - test/ctmrg/jacobian_real_linear.jl | 1 - test/ctmrg/pepo.jl | 62 ++++++-- test/mooncake/eigh_wrapper.jl | 5 +- test/mooncake/svd_wrapper.jl | 6 +- 18 files changed, 251 insertions(+), 114 deletions(-) diff --git a/Project.toml b/Project.toml index 47efcf211..babc0ae4a 100644 --- a/Project.toml +++ b/Project.toml @@ -18,7 +18,6 @@ LoggingExtras = "e6f89c97-d47a-5376-807f-9c37f3926c36" MPSKit = "bb1c41ca-d63c-52ed-829e-0820dda26502" MPSKitModels = "ca635005-6f8c-4cd1-b51d-8491250ef2ab" MatrixAlgebraKit = "6c742aac-3347-4629-af66-fc926824e5e4" -Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" OhMyThreads = "67456a42-1dca-4109-a031-0a68de7e3ad5" OptimKit = "77e91f04-9b3b-57a6-a776-40b61faaebe0" Printf = "de0858da-6303-5e67-8744-51eddeeeb8d7" @@ -29,9 +28,15 @@ TensorOperations = "6aa20fa7-93e2-5fca-9bc0-fbd0db3c71a2" TupleTools = "9d95972d-f1c8-5527-a6e0-b4b365fa01f6" VectorInterface = "409d34a3-91d5-4945-b6ec-7529ddf182d8" +[weakdeps] +Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" + [extensions] PEPSKitMooncakeExt = "Mooncake" +[sources] +TensorKit = {url = "https://github.com/QuantumKitHub/TensorKit.jl", rev = "main"} + [compat] Accessors = "0.1" ChainRulesCore = "1.0" @@ -46,11 +51,11 @@ MPSKitModels = "0.4" MatrixAlgebraKit = "0.6.5" Mooncake = "0.5.27" OhMyThreads = "0.7, 0.8" -OptimKit = "0.4" +OptimKit = "0.5" Printf = "1" Random = "1" Statistics = "1" -TensorKit = "0.16.5, 0.17" +TensorKit = "0.17" TensorOperations = "5" TupleTools = "1.6.0" VectorInterface = "0.6" diff --git a/ext/PEPSKitMooncakeExt.jl b/ext/PEPSKitMooncakeExt.jl index c06084400..1e82cbc06 100644 --- a/ext/PEPSKitMooncakeExt.jl +++ b/ext/PEPSKitMooncakeExt.jl @@ -1,9 +1,9 @@ module PEPSKitMooncakeExt using PEPSKit, MPSKit, TensorKit, Mooncake, MatrixAlgebraKit -using PEPSKit: SVDAdjoint, EighAdjoint, QRAdjoint, CTMRGAlgorithm, FixedPointGradient, sdiag_pow +using PEPSKit: SVDAdjoint, EighAdjoint, QRAdjoint, CTMRGAlgorithm, FixedPointGradient, sdiag_pow, eachcoordinate import PEPSKit: real_inner -using Mooncake: DefaultCtx, CoDual, Dual, NoRData, primal, tangent, rrule!!, arrayify, @is_primitive +using Mooncake: DefaultCtx, MinimalCtx, CoDual, Dual, NoRData, primal, tangent, rrule!!, arrayify, @is_primitive function Mooncake.arrayify(ψ::PEPSKit.InfinitePEPS{T}, dψ) where {T} Δψmat = map((a, da) -> Mooncake.arrayify(a, da)[2], ψ.A, dψ.fields.A) @@ -20,12 +20,13 @@ Mooncake.tangent_type(::Type{<:PEPSKit.QRAdjoint}) = Mooncake.NoTangent Mooncake.tangent_type(::Type{<:PEPSKit.CTMRGAlgorithm}) = Mooncake.NoTangent Mooncake.tangent_type(::Type{<:PEPSKit.FixedPointGradient}) = Mooncake.NoTangent -Mooncake.@zero_derivative Mooncake.DefaultCtx Tuple{typeof(PEPSKit.eachcoordinate), Any, Any} -Mooncake.@zero_derivative Mooncake.DefaultCtx Tuple{typeof(PEPSKit._next_coordinate), Int, Int} -Mooncake.@zero_derivative Mooncake.DefaultCtx Tuple{typeof(PEPSKit._set_decomposition_truncation), Any, Any} -Mooncake.@zero_derivative Mooncake.DefaultCtx Tuple{typeof(PEPSKit.CTMRGEnv), Union{PEPSKit.InfinitePartitionFunction, PEPSKit.InfinitePEPS}, Vararg} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(PEPSKit.eachcoordinate), Any} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(PEPSKit.eachcoordinate), Any, Any} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(PEPSKit._next_coordinate), Int, Int} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(PEPSKit._set_decomposition_truncation), Any, Any} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(PEPSKit.CTMRGEnv), Union{PEPSKit.InfinitePartitionFunction, PEPSKit.InfinitePEPS}, Vararg} -@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(svd_trunc), TensorKit.AbstractTensorMap, SVDAdjoint} +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(svd_trunc), TensorKit.AbstractTensorMap, SVDAdjoint} function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{SVDAdjoint{F, R}}) where {F, R <: PEPSKit.FullPullback} # TODO: filter out any decomposition algorithm that doesn't give access to the full spectrum t, Δt = arrayify(t_dt) @@ -74,7 +75,7 @@ function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.svd_trunc)}, t_dt::Co return output_codual, svd_trunc!_trunc_pullback end -@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(eigh_trunc), TensorKit.AbstractTensorMap, EighAdjoint} +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(eigh_trunc), TensorKit.AbstractTensorMap, EighAdjoint} function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.eigh_trunc)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{EighAdjoint{F, R}}) where {F, R <: PEPSKit.FullPullback} t, dt = arrayify(t_dt) alg = primal(alg_dalg) @@ -124,7 +125,7 @@ function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.eigh_trunc)}, t_dt::C return output_codual, eigh_trunc!_trunc_pullback end -@is_primitive Mooncake.DefaultCtx Mooncake.ReverseMode Tuple{typeof(left_orth), TensorKit.AbstractTensorMap, QRAdjoint} +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(left_orth), TensorKit.AbstractTensorMap, QRAdjoint} function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.left_orth)}, t_dt::CoDual{<:TensorKit.AbstractTensorMap}, alg_dalg::CoDual{QRAdjoint}) t, dt = arrayify(t_dt) alg = primal(alg_dalg) @@ -143,4 +144,132 @@ function Mooncake.rrule!!(::CoDual{typeof(MatrixAlgebraKit.left_orth)}, t_dt::Co return output_codual, left_orth_pullback end +PEPSKit.real_inner(_, η₁::Mooncake.Tangent, η₂::Mooncake.Tangent) = Mooncake._dot(η₁, η₂) + +# 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 +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(Core.kwcall), NamedTuple, typeof(PEPSKit.dtmap), Any, AbstractArray} +function Mooncake.rrule!!(::CoDual{typeof(Core.kwcall)}, kw::CoDual{<:NamedTuple{(:scheduler,), Tuple{R}}}, ::CoDual{typeof(PEPSKit.dtmap)}, f_df::CoDual, A_dA::CoDual{<:AbstractArray}) where {R} + scheduler = get(Mooncake.primal(kw), :scheduler, PEPSKit.Defaults.scheduler[]) + f = Mooncake.primal(f_df) + A, ΔA = Mooncake.arrayify(A_dA) + el_rrules = tmap(A; scheduler) do a + cache = Mooncake.prepare_pullback_cache(f, a) + return Mooncake.value_and_pullback!!(cache, f, a) + end + y = map(first, el_rrules) + y_dy = Mooncake.zero_fcodual(y) + Δys = Mooncake.arrayify(y_dy)[2] + function dtmap_pullback(::NoRData) + backevals = tmap(el_rrules, Δys; scheduler) do el_rrule, Δy + last(el_rrule)(Δy) + end + ΔA .= map(last, backevals) + return ntuple(Returns(NoRData()), 5) + end + return y_dy, dtmap_pullback +end + +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(Core.kwcall), NamedTuple, typeof(PEPSKit.dtmap!!), Any, AbstractArray, AbstractArray} +function Mooncake.rrule!!(::CoDual{typeof(Core.kwcall)}, kw::CoDual, ::CoDual{typeof(PEPSKit.dtmap!!)}, f_df::CoDual, C_dC::CoDual{<:AbstractArray}, A_dA::CoDual{<:AbstractArray}) + C, dtmap_pullback = rrule(config, dtmap, f, A; kwargs...) + function dtmap!!_pullback(dy) + dtmap, df, dA = dtmap_pullback(dy) + return dtmap, df, NoTangent, dA + end + return C_dC, dtmap!!_pullback +end + +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(Core.kwcall), NamedTuple, typeof(PEPSKit.sdiag_pow), AbstractTensorMap, Real} +function Mooncake.rrule!!(::CoDual{typeof(Core.kwcall)}, kw::CoDual{<:NamedTuple{(:tol,), Tuple{<:Real}}}, ::CoDual{typeof(PEPSKit.sdiag_pow)}, s_ds::CoDual{<:AbstractTensorMap}, p_dp::CoDual{<:Real}) + s, Δs = arrayify(s_ds) + tol = get(primal(kw), :tol, eps(real(TensorKit.scalartype(s)))^(3 / 4)) + tol *= norm(s, Inf) + pow = primal(p_df) + spow = sdiag_pow(s, pow; tol) + spow_minus1_conj = scale!(sdiag_pow(s', pow - 1; tol), pow) + spow_dspow = Mooncake.zero_fcodual(spow) + spow, Δspow = arrayify(spow_dspow) + function sdiag_pow_pullback(::NoRData) + PEPSKit._elementwise_mult(Δs, spow_minus1_conj) + return NoRData(), NoRData(), NoRData(), NoRData(), zero(pow) + end + return spow_dspow, sdiag_pow_pullback +end + +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(PEPSKit.sdiag_pow), AbstractTensorMap, Real} +function Mooncake.rrule!!(::CoDual{typeof(PEPSKit.sdiag_pow)}, s_ds::CoDual{<:AbstractTensorMap}, p_dp::CoDual{<:Real}) + s, Δs = arrayify(s_ds) + tol = eps(real(TensorKit.scalartype(s)))^(3 / 4) + tol *= norm(s, Inf) + pow = primal(p_dp) + spow = sdiag_pow(s, pow; tol) + spow_minus1_conj = scale!(sdiag_pow(s', pow - 1; tol), pow) + spow_dspow = Mooncake.zero_fcodual(spow) + spow, Δspow = arrayify(spow_dspow) + function sdiag_pow_pullback(::NoRData) + PEPSKit._elementwise_mult(Δs, spow_minus1_conj) + return NoRData(), NoRData(), zero(pow) + end + return spow_dspow, sdiag_pow_pullback +end + +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(PEPSKit.CTMRGEnv), Any, Any} +function Mooncake.rrule!!(::CoDual{typeof(PEPSKit.CTMRGEnv)}, c_dc::CoDual{Array{C, 3}}, e_de::CoDual{Array{T, 3}}) where {C, T} + corners, dcorners = arrayify(c_dc) + edges, dedges = arrayify(e_de) + env = CTMRGEnv(corners, edges) + denv = CTMRGEnv(dcorners, dedges) + ctmrgenv_pullback(::NoRData) = NoRData(), env.corners, env.edges + return Mooncake.CoDual(env, denv), ctmrgenv_pullback +end + +Mooncake.tangent_type(::Type{NamedTuple{(:converged, :convergence_error, :contraction_metrics), Tuple{Bool, Float64, NamedTuple{(:truncation_error,), Tuple{Float64}}}}}) = Mooncake.NoTangent +Mooncake.tangent_type(::Type{NamedTuple{(:alg_rrule,), Tuple{NamedTuple{(:solver_alg,), Tuple{NamedTuple{(:orth, :krylovdim, :maxiter, :tol, :eager, :verbosity), Tuple{O, Int, Int, Float64, Bool, Int}}}}}}}) where {O} = Mooncake.NoTangent + +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(Core.kwcall), NamedTuple, typeof(PEPSKit.hook_pullback), Any, Vararg{Any}} +function Mooncake.rrule!!(::CoDual{typeof(Core.kwcall)}, kw::CoDual{<:NamedTuple{(:alg_rrule,), Tuple{R}}}, hpb::CoDual{typeof(PEPSKit.hook_pullback)}, f_df::CoDual, args_dargs::CoDual...) where {R} + alg_rrule = Mooncake.primal(kw)[:alg_rrule] + y, f_pullback = PEPSKit._rrule(alg_rrule, Mooncake.primal(f_df), Mooncake.primal.(args_dargs)...) + hook_pullback_pullback(Δ) = (NoRData(), f_pullback(Δ)...) + return y, hook_pullback_pullback +end + +# compute the CTMRG gradient through fixed-point differentiation +@is_primitive Mooncake.MinimalCtx Mooncake.ReverseMode Tuple{typeof(MPSKit.leading_boundary), Any, Any, CTMRGAlgorithm} +function Mooncake.rrule!!(::CoDual{typeof(MPSKit.leading_boundary)}, envinit_denvinit::CoDual, state_dstate::CoDual, alg_dalg::CoDual{<:CTMRGAlgorithm}) + alg = Mooncake.primal(alg_dalg) + state = Mooncake.primal(state_dstate) + envinit = Mooncake.primal(envinit_denvinit) + #PEPSKit._check_algorithm_combination(alg, gradmode) + + env, = MPSKit.leading_boundary(envinit, state, alg) + + # prepare iterating function corresponding to a single gauge-fixed CTMRG iteration + alg_fixed = PEPSKit._set_fixed_truncation(alg) # fix spaces during differentiation + alg_gauge = PEPSKit._scrambling_env_gauge(alg) # select appropriate gauge-fixing algorithm + env_conv, info = PEPSKit.ctmrg_iteration(InfiniteSquareNetwork(state), env, alg_fixed) + signs, corner_phases, edge_phases = PEPSKit.compute_gauge_fix_gauge(env_conv, env, alg_gauge) + # prepare its pullback + #sig = Tuple{typeof(gauge_fixed_iteration), typeof(state), typeof(env), typeof(alg_fixed), typeof(signs), typeof(corner_phases), typeof(edge_phases)} + #rule = Mooncake.build_rrule(gauge_fixed_iteration, state, env, alg_fixed, signs, corner_phases, edge_phases) + #_, env_vjp = Mooncake.value_and_gradient!!(rule, gauge_fixed_iteration, state, env, alg_fixed, signs, corner_phases, edge_phases) + #out, env_vjp = Mooncake.rrule!!(CoDual(gauge_fixed_iteration, Mooncake.NoFData()), Mooncake.zero_fcodual(state), Mooncake.zero_fcodual(env)) + cache = Mooncake.prepare_pullback_cache(PEPSKit.gauge_fixed_iteration, state, env, alg_fixed, signs, corner_phases, edge_phases) + _, env_vjp = Mooncake.value_and_pullback!!(cache, env_conv, PEPSKit.gauge_fixed_iteration, state, env, alg_fixed, signs, corner_phases, edge_phases) + # split off state and environment parts + ∂f∂A(x)::typeof(state) = env_vjp(x)[2] + ∂f∂x(x)::typeof(env) = env_vjp(x)[3] + + output_doutput = Mooncake.zero_fcodual((env, info)) + denv = Mooncake.tangent(output_doutput)[1] + env, Δenv = Mooncake.arrayify(env, denv) + function leading_boundary_fixed_pullback(::NoRData) + # evaluate the geometric sum + ∂F∂env = PEPSKit.fixedpoint_gradient(Δenv, ∂f∂x, ∂f∂A, Δenv, gradmode.solver_alg) + return ntuple(Returns(NoRData()), 4) + end + return (env, invo), leading_boundary_fixed_pullback +end + end diff --git a/src/PEPSKit.jl b/src/PEPSKit.jl index a6e199c54..311b71107 100644 --- a/src/PEPSKit.jl +++ b/src/PEPSKit.jl @@ -23,7 +23,7 @@ using KrylovKit using KrylovKit: Lanczos, BlockLanczos using TensorOperations, OptimKit -using ChainRulesCore, Zygote +using ChainRulesCore using LoggingExtras import TupleTools diff --git a/src/algorithms/ctmrg/c4v.jl b/src/algorithms/ctmrg/c4v.jl index e986e8967..4e5198d56 100644 --- a/src/algorithms/ctmrg/c4v.jl +++ b/src/algorithms/ctmrg/c4v.jl @@ -202,11 +202,9 @@ function c4v_projector!(enlarged_corner, alg::C4vEighProjector) D, V, truncation_error = eigh_trunc!(enlarged_corner, eigh_alg) # Check for degenerate eigenvalues - Zygote.isderiving() && ignore_derivatives() do - if alg.verbosity > 0 && is_degenerate_spectrum(D) - vals = TensorKit.SectorDict(c => diag(b) for (c, b) in blocks(D)) - @warn("degenerate eigenvalues detected: ", vals) - end + if alg.verbosity > 0 && is_degenerate_spectrum(D) + vals = TensorKit.SectorDict(c => diag(b) for (c, b) in blocks(D)) + @warn("degenerate eigenvalues detected: ", vals) end return D / norm(D), V, (; D, V, truncation_error) diff --git a/src/algorithms/ctmrg/projectors.jl b/src/algorithms/ctmrg/projectors.jl index bd81041b1..cf1a4dffc 100644 --- a/src/algorithms/ctmrg/projectors.jl +++ b/src/algorithms/ctmrg/projectors.jl @@ -183,11 +183,9 @@ function compute_projector(enlarged_corners, alg::HalfInfiniteProjector) truncation_error = truncation_error / norm(S) # normalize truncation error # Check for degenerate singular values - Zygote.isderiving() && ignore_derivatives() do - if alg.verbosity > 0 && is_degenerate_spectrum(S) - svals = TensorKit.SectorDict(c => diag(b) for (c, b) in blocks(S)) - @warn("degenerate singular values detected: ", svals) - end + if alg.verbosity > 0 && is_degenerate_spectrum(S) + svals = TensorKit.SectorDict(c => diag(b) for (c, b) in blocks(S)) + @warn("degenerate singular values detected: ", svals) end P_left, P_right = contract_projectors(U, S, V, enlarged_corners...) @@ -206,11 +204,9 @@ function compute_projector(enlarged_corners, alg::FullInfiniteProjector) truncation_error = truncation_error / norm(S) # normalize truncation error # Check for degenerate singular values - Zygote.isderiving() && ignore_derivatives() do - if alg.verbosity > 0 && is_degenerate_spectrum(S) - svals = TensorKit.SectorDict(c => diag(b) for (c, b) in blocks(S)) - @warn("degenerate singular values detected: ", svals) - end + if alg.verbosity > 0 && is_degenerate_spectrum(S) + svals = TensorKit.SectorDict(c => diag(b) for (c, b) in blocks(S)) + @warn("degenerate singular values detected: ", svals) end P_left, P_right = contract_projectors(U, S, V, halfinf_left, halfinf_right) diff --git a/src/algorithms/ctmrg/sequential.jl b/src/algorithms/ctmrg/sequential.jl index d68c41640..e546e9346 100644 --- a/src/algorithms/ctmrg/sequential.jl +++ b/src/algorithms/ctmrg/sequential.jl @@ -124,8 +124,8 @@ end Renormalize one column of the CTMRG environment. """ function renormalize_sequentially(col::Int, projectors, network, env) - corners = Zygote.Buffer(env.corners) - edges = Zygote.Buffer(env.edges) + corners = env.corners + edges = env.edges for (dir, r, c) in eachcoordinate(network, 1:4) (c == col && dir in [SOUTHWEST, NORTHWEST]) && continue diff --git a/src/algorithms/ctmrg/simultaneous.jl b/src/algorithms/ctmrg/simultaneous.jl index ecaf52c7b..a72423bc6 100644 --- a/src/algorithms/ctmrg/simultaneous.jl +++ b/src/algorithms/ctmrg/simultaneous.jl @@ -37,16 +37,10 @@ end CTMRG_SYMBOLS[:SimultaneousCTMRG] = SimultaneousCTMRG function ctmrg_iteration(network, env::CTMRGEnv, alg::SimultaneousCTMRG) - coordinates = eachcoordinate(network, 1:4) - T_corners = Base.promote_op( - TensorMap ∘ EnlargedCorner, typeof(network), typeof(env), eltype(coordinates) - ) - enlarged_corners′ = similar(coordinates, T_corners) - enlarged_corners::typeof(enlarged_corners′) = - dtmap!!(enlarged_corners′, eachcoordinate(network, 1:4)) do idx - return TensorMap(EnlargedCorner(network, env, idx)) - end # expand environment + coords = eachcoordinate(network, 1:4) + enlarged_corners = [TensorMap(EnlargedCorner(network, env, idx)) for idx in coords] # expand environment projectors, info = simultaneous_projectors(enlarged_corners, env, alg.projector_alg) # compute projectors on all coordinates + # problem is here! env′ = renormalize_simultaneously(enlarged_corners, projectors, network, env) # renormalize enlarged corners info = (; contraction_metrics = (; info.truncation_error), @@ -79,13 +73,7 @@ function simultaneous_projectors( enlarged_corners::Array{E, 3}, env::CTMRGEnv, alg::ProjectorAlgorithm ) where {E} coordinates = eachcoordinate(env, 1:4) - T_dst = Base.promote_op( - simultaneous_projectors, - NTuple{3, Int}, typeof(enlarged_corners), typeof(env), typeof(alg), - ) - proj_and_info′ = similar(coordinates, T_dst) - proj_and_info::typeof(proj_and_info′) = - dtmap!!(proj_and_info′, coordinates) do coordinate + proj_and_info = map(coordinates) do coordinate return simultaneous_projectors(coordinate, enlarged_corners, env, alg) end return _split_proj_and_info(proj_and_info) @@ -93,8 +81,8 @@ end function simultaneous_projectors( coordinate, enlarged_corners::Array{E, 3}, env, alg::HalfInfiniteProjector ) where {E} - coordinate′ = _next_coordinate(coordinate, size(env)[2:3]...) - trunc = truncation_strategy(alg, env.edges[coordinate[1], coordinate′[2:3]...]) + coordinate′ = _next_coordinate(coordinate, size(env, 2), size(env, 3)) + trunc = truncation_strategy(alg, env.edges[coordinate[1], coordinate′[2], coordinate′[3]]) alg′ = _set_decomposition_truncation(alg, trunc) ec = (enlarged_corners[coordinate...], enlarged_corners[coordinate′...]) return compute_projector(ec, alg′) @@ -126,25 +114,26 @@ function renormalize_simultaneously(enlarged_corners, projectors, network, env) P_left, P_right = projectors coordinates = eachcoordinate(env, 1:4) T_CE = Tuple{cornertype(env), edgetype(env)} - corners_edges′ = similar(coordinates, T_CE) - corners_edges::typeof(corners_edges′) = - dtmap!!(corners_edges′, coordinates) do (dir, r, c) - if dir == NORTH + corners_edges = similar(coordinates, T_CE) + #dtmap!!(corners_edges, coordinates) do coord + map!(corners_edges, coordinates) do coord + direction, r, c = coord + if direction == NORTH corner = renormalize_northwest_corner( (r, c), enlarged_corners, P_left, P_right ) edge = renormalize_north_edge((r, c), env, P_left, P_right, network) - elseif dir == EAST + elseif direction == EAST corner = renormalize_northeast_corner( (r, c), enlarged_corners, P_left, P_right ) edge = renormalize_east_edge((r, c), env, P_left, P_right, network) - elseif dir == SOUTH + elseif direction == SOUTH corner = renormalize_southeast_corner( (r, c), enlarged_corners, P_left, P_right ) edge = renormalize_south_edge((r, c), env, P_left, P_right, network) - elseif dir == WEST + elseif direction == WEST corner = renormalize_southwest_corner( (r, c), enlarged_corners, P_left, P_right ) @@ -152,6 +141,5 @@ function renormalize_simultaneously(enlarged_corners, projectors, network, env) end return corner / norm(corner), edge / norm(edge) end - return CTMRGEnv(map(first, corners_edges), map(last, corners_edges)) end diff --git a/src/algorithms/optimization/fixed_point_differentiation.jl b/src/algorithms/optimization/fixed_point_differentiation.jl index 2c2a54991..89e5cff80 100644 --- a/src/algorithms/optimization/fixed_point_differentiation.jl +++ b/src/algorithms/optimization/fixed_point_differentiation.jl @@ -210,6 +210,11 @@ function _set_fixed_truncation(alg::CTMRGAlgorithm) return alg_fixed end +function gauge_fixed_iteration(A, x, alg_fixed, signs, corner_phases, edge_phases) + x′ = ctmrg_iteration(InfiniteSquareNetwork(A), x, alg_fixed)[1] + return fix_phases(x′, signs, corner_phases, edge_phases) +end + # compute the CTMRG gradient through fixed-point differentiation function _rrule( gradmode::FixedPointGradient, @@ -228,12 +233,12 @@ function _rrule( alg_gauge = _scrambling_env_gauge(alg) # select appropriate gauge-fixing algorithm env_conv, info = ctmrg_iteration(InfiniteSquareNetwork(state), env, alg_fixed) signs, corner_phases, edge_phases = compute_gauge_fix_gauge(env_conv, env, alg_gauge) - function gauge_fixed_iteration(A, x) + #=function gauge_fixed_iteration(A, x) return fix_phases( ctmrg_iteration(InfiniteSquareNetwork(A), x, alg_fixed)[1], signs, corner_phases, edge_phases, ) - end + end=# # prepare its pullback _, env_vjp = rrule_via_ad(config, gauge_fixed_iteration, state, env) # split off state and environment parts diff --git a/src/algorithms/toolbox.jl b/src/algorithms/toolbox.jl index c532d3e1d..8baf99119 100644 --- a/src/algorithms/toolbox.jl +++ b/src/algorithms/toolbox.jl @@ -81,9 +81,13 @@ Return the value (per unit cell) of a given contractible network contracted usin CTMRG environment. """ function network_value(network::InfiniteSquareNetwork, env::CTMRGEnv) - return prod(Iterators.product(axes(network)...)) do (r, c) - return _contract_site((r, c), network, env) * _contract_corners((r, c), env) / - _contract_vertical_edges((r, c), env) / _contract_horizontal_edges((r, c), env) + ax_prod = collect(Iterators.product(axes(network)...)) + return prod(ax_prod) do (r, c) + site_val = _contract_site((r, c), network, env) + corner_val = _contract_corners((r, c), env) + vertical_val = _contract_vertical_edges((r, c), env) + horizontal_val = _contract_horizontal_edges((r, c), env) + return site_val * corner_val / vertical_val / horizontal_val end end network_value(state, env::CTMRGEnv) = network_value(InfiniteSquareNetwork(state), env) diff --git a/src/environments/ctmrg_environments.jl b/src/environments/ctmrg_environments.jl index cabc0e764..8a485d904 100644 --- a/src/environments/ctmrg_environments.jl +++ b/src/environments/ctmrg_environments.jl @@ -332,10 +332,8 @@ end # Rotate corners & edges counter-clockwise function Base.rotl90(env::CTMRGEnv{C, T}) where {C, T} # Initialize rotated corners & edges with rotated sizes - corners′ = Zygote.Buffer( - Array{C, 3}(undef, 4, size(env.corners, 3), size(env.corners, 2)) - ) - edges′ = Zygote.Buffer(Array{T, 3}(undef, 4, size(env.edges, 3), size(env.edges, 2))) + corners′ = Array{C, 3}(undef, 4, size(env.corners, 3), size(env.corners, 2)) + edges′ = Array{T, 3}(undef, 4, size(env.edges, 3), size(env.edges, 2)) for dir in 1:4 dir2 = _prev(dir, 4) corners′[dir2, :, :] = rotl90(env.corners[dir, :, :]) @@ -347,10 +345,8 @@ end # Rotate corners & edges clockwise function Base.rotr90(env::CTMRGEnv{C, T}) where {C, T} # Initialize rotated corners & edges with rotated sizes - corners′ = Zygote.Buffer( - Array{C, 3}(undef, 4, size(env.corners, 3), size(env.corners, 2)) - ) - edges′ = Zygote.Buffer(Array{T, 3}(undef, 4, size(env.edges, 3), size(env.edges, 2))) + corners′ = Array{C, 3}(undef, 4, size(env.corners, 3), size(env.corners, 2)) + edges′ = Array{T, 3}(undef, 4, size(env.edges, 3), size(env.edges, 2)) for dir in 1:4 dir2 = _next(dir, 4) corners′[dir2, :, :] = rotr90(env.corners[dir, :, :]) @@ -362,10 +358,8 @@ end # Rotate corners & edges by 180 degrees function Base.rot180(env::CTMRGEnv{C, T}) where {C, T} # Initialize rotated corners & edges with rotated sizes - corners′ = Zygote.Buffer( - Array{C, 3}(undef, 4, size(env.corners, 2), size(env.corners, 3)) - ) - edges′ = Zygote.Buffer(Array{T, 3}(undef, 4, size(env.edges, 2), size(env.edges, 3))) + corners′ = Array{C, 3}(undef, 4, size(env.corners, 2), size(env.corners, 3)) + edges′ = Array{T, 3}(undef, 4, size(env.edges, 2), size(env.edges, 3)) for dir in 1:4 dir2 = _next(_next(dir, 4), 4) corners′[dir2, :, :] = rot180(env.corners[dir, :, :]) diff --git a/src/utility/diffable_threads.jl b/src/utility/diffable_threads.jl index e6b6270e8..4c70851a0 100644 --- a/src/utility/diffable_threads.jl +++ b/src/utility/diffable_threads.jl @@ -62,11 +62,7 @@ macro fwdthreads(ex) @assert ex.head === :for "@fwdthreads expects a for loop:\n$ex" diffable_ex = quote - if Zygote.isderiving() - $ex - else - Threads.@threads $ex - end + Threads.@threads $ex end return esc(diffable_ex) diff --git a/src/utility/indexing.jl b/src/utility/indexing.jl index cc5fa1249..88f9ac6be 100644 --- a/src/utility/indexing.jl +++ b/src/utility/indexing.jl @@ -71,6 +71,8 @@ function _next_coordinate((dir, row, col), rowsize, colsize) return (_next(dir, 4), row, _prev(col, colsize)) elseif dir == 4 return (_next(dir, 4), _prev(row, rowsize), col) + else + error(lazy"invalid dir $dir") end end function _prev_coordinate((dir, row, col), rowsize, colsize) diff --git a/src/utility/util.jl b/src/utility/util.jl index 2c58f8801..932d31276 100644 --- a/src/utility/util.jl +++ b/src/utility/util.jl @@ -168,23 +168,6 @@ function ChainRulesCore.rrule(::typeof(rotr90), a::AbstractMatrix) return rotr90(a), rotr90_pullback end -# TODO: link to Zygote.showgrad once they update documenter.jl -""" - @showtypeofgrad(x) - -Macro utility to show to type of the gradient that is about to accumulate for `x`. - -See also `Zygote.@showgrad`. -""" -macro showtypeofgrad(x) - return :( - Zygote.hook($(esc(x))) do x̄ - println($"∂($x) = ", repr(typeof(x̄))) - x̄ - end - ) -end - """ Randomly take the dual of `ElementarySpace`s in `Vs` with propability `p` """ diff --git a/test/Project.toml b/test/Project.toml index 56b69e010..7a50e87f4 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -19,7 +19,6 @@ TensorKit = "07d1fe3e-3e46-537d-9eac-e9e13d0d4cec" Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" TestExtras = "5ed8adda-3752-4e41-b88a-e8b09835ee3a" VectorInterface = "409d34a3-91d5-4945-b6ec-7529ddf182d8" -Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f" [sources] PEPSKit = {path = ".."} diff --git a/test/ctmrg/jacobian_real_linear.jl b/test/ctmrg/jacobian_real_linear.jl index 9738aaa16..6808d9238 100644 --- a/test/ctmrg/jacobian_real_linear.jl +++ b/test/ctmrg/jacobian_real_linear.jl @@ -1,7 +1,6 @@ using Test using Random using Accessors -using Zygote using TensorKit, KrylovKit, PEPSKit using PEPSKit: ctmrg_iteration, compute_gauge_fix_gauge, fix_phases, ScramblingEnvGauge diff --git a/test/ctmrg/pepo.jl b/test/ctmrg/pepo.jl index e2e0ec05a..4692ea610 100644 --- a/test/ctmrg/pepo.jl +++ b/test/ctmrg/pepo.jl @@ -1,11 +1,25 @@ using Test using Random using LinearAlgebra -using PEPSKit +using PEPSKit, MPSKit using TensorKit using KrylovKit using OptimKit -using Zygote +using Mooncake +using MatrixAlgebraKit +using PEPSKit: LoggingExtras + +const MCExt = Base.get_extension(PEPSKit, :PEPSKitMooncakeExt) +@assert !isnothing(MCExt) + +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(Core.current_scope)} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(time)} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(Base.CoreLogging.with_logstate), Any, Any} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(Base.CoreLogging._invoked_min_enabled_level), Any} +Mooncake.@zero_derivative Mooncake.MinimalCtx Tuple{typeof(PEPSKit.LoggingExtras.withlevel), Any, Int} + +Mooncake.tangent_type(::Type{<:Base.HashArrayMappedTries.HAMT}) = Mooncake.NoTangent +Mooncake.tangent_type(::Type{Base.HashArrayMappedTries.Leaf}) = Mooncake.NoTangent ## Setup @@ -80,6 +94,12 @@ projector_algs = [:HalfInfiniteProjector, :FullInfiniteProjector] end end +#=f1=Mooncake.zero_fcodual +rule_tester=Mooncake.build_rrule(Complex,1.0,1.0) +rule_tester(f1(Complex),f1(1.0),f1(1.0))=# + +#@show Mooncake.is_primitive(Mooncake.DefaultCtx, Mooncake.ReverseMode, Tuple{Complex, Float64, Float64}, Base.get_world_counter()) +#@show Base.which(Mooncake.rrule!!, (Complex,Float64, Float64)) @testset "Fixed-point computation for 3D classical ising model" begin Random.seed!(81812781144) @@ -104,6 +124,23 @@ end env2_0 = CTMRGEnv(InfiniteSquareNetwork(psi0), χenv) env3_0 = CTMRGEnv(InfiniteSquareNetwork(psi0, T), χenv) + function energ_test(ψ, env2, env3) + n2 = InfiniteSquareNetwork(ψ) + env2′, info = PEPSKit.hook_pullback( + leading_boundary, env2, n2, ctm_alg; alg_rrule = gradient_alg + ) + n3 = InfiniteSquareNetwork(ψ, T) + env3′, info = PEPSKit.hook_pullback( + leading_boundary, env3, n3, ctm_alg; alg_rrule = gradient_alg + ) + PEPSKit.update!(env2, env2′) + PEPSKit.update!(env3, env3′) + λ3 = network_value(n3, env3) + λ2 = network_value(n2, env2) + return -log(abs(λ3 / λ2)) + end + @show energ_test(psi0, env2_0, env3_0) + # optimize free energy per site (psi_final, env2_final, env3_final), f, = optimize( (psi0, env2_0, env3_0), @@ -112,24 +149,27 @@ end retract = pepo_retract, (transport!) = (pepo_transport!), ) do (psi, env2, env3) - E, gs = withgradient(psi) do ψ + function energ(ψ) n2 = InfiniteSquareNetwork(ψ) - env2′, info = PEPSKit.hook_pullback( + #=env2′, info = PEPSKit.hook_pullback( leading_boundary, env2, n2, ctm_alg; alg_rrule = gradient_alg ) n3 = InfiniteSquareNetwork(ψ, T) env3′, info = PEPSKit.hook_pullback( leading_boundary, env3, n3, ctm_alg; alg_rrule = gradient_alg - ) - PEPSKit.ignore_derivatives() do - PEPSKit.update!(env2, env2′) - PEPSKit.update!(env3, env3′) - end + )=# + env2′, info = PEPSKit.leading_boundary(env2, n2, ctm_alg) + n3 = InfiniteSquareNetwork(ψ, T) + env3′, info = PEPSKit.leading_boundary(env3, n3, ctm_alg) + PEPSKit.update!(env2, env2′) + PEPSKit.update!(env3, env3′) λ3 = network_value(n3, env3) λ2 = network_value(n2, env2) - return -log(real(λ3 / λ2)) + return -log(abs(λ3 / λ2)) end - g = only(gs) + cache = prepare_gradient_cache(energ, psi) + E, gs = value_and_gradient!!(cache, energ, psi) + _, g = Mooncake.arrayify(psi, gs[2]) return E, g end diff --git a/test/mooncake/eigh_wrapper.jl b/test/mooncake/eigh_wrapper.jl index a05adaf41..9017c0a12 100644 --- a/test/mooncake/eigh_wrapper.jl +++ b/test/mooncake/eigh_wrapper.jl @@ -1,9 +1,8 @@ - using Test using Random using LinearAlgebra using TensorKit -using Mooncake +using Mooncake using Accessors using PEPSKit @@ -35,7 +34,7 @@ iter_alg = EighAdjoint(; fwd_alg = (; alg = :Lanczos), rrule_alg = (; alg = :Tru full_lossfun = A -> lossfun(A, full_alg, R) trunc_lossfun = A -> lossfun(A, trunc_alg, R) iter_lossfun = A -> lossfun(A, iter_alg, R) - + full_rrule = Mooncake.build_rrule(full_lossfun, r) trunc_rrule = Mooncake.build_rrule(trunc_lossfun, r) iter_rrule = Mooncake.build_rrule(iter_lossfun, r) diff --git a/test/mooncake/svd_wrapper.jl b/test/mooncake/svd_wrapper.jl index 2c010ca28..a99c8d489 100644 --- a/test/mooncake/svd_wrapper.jl +++ b/test/mooncake/svd_wrapper.jl @@ -119,7 +119,7 @@ symm_R = randn(dtype, space(symm_r)) @test g_full[2] ≈ g_trunc[2] rtol = rtol @test g_full[2] ≈ g_iter[2] rtol = rtol @test g_trunc[2] ≈ g_iter[2] rtol = rtol - + full_lossfun = A -> lossfun(A, full_alg, symm_R, symm_trspace) trunc_lossfun = A -> lossfun(A, trunc_alg, symm_R, symm_trspace) iter_lossfun = A -> lossfun(A, iter_alg, symm_R, symm_trspace) @@ -137,7 +137,7 @@ symm_R = randn(dtype, space(symm_r)) @test g_trunc_tr[2] ≈ g_iter_tr[2] rtol = rtol iter_alg_fallback = @set iter_alg.fwd_alg.fallback_threshold = 0.4 # Do dense decomposition in one block, sparse one in the other - + fb_lossfun = A -> lossfun(A, iter_alg_fallback, symm_R, symm_trspace) fb_rrule = Mooncake.build_rrule(fb_lossfun, symm_r) l_iter_fb, g_iter_fb = Mooncake.value_and_gradient!!(fb_rrule, fb_lossfun, symm_r) @@ -159,7 +159,7 @@ end no_broadening_no_cutoff_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-30 small_broadening_alg = @set alg.rrule_alg.degeneracy_atol = 1.0e-13 - + only_lossfun = A -> lossfun(A, alg, symm_R, symm_trspace) no_broadening_lossfun = A -> lossfun(A, no_broadening_no_cutoff_alg, symm_R, symm_trspace) small_broadening_lossfun = A -> lossfun(A, small_broadening_alg, symm_R, symm_trspace)