Skip to content
This repository was archived by the owner on Jan 20, 2025. It is now read-only.
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 13 additions & 16 deletions src/ITensorNetworks/boundary_mps.jl
Original file line number Diff line number Diff line change
Expand Up @@ -339,6 +339,17 @@ function boundary_mps(tn::NamedTuple)
)
end

function boundary_mps(tn::Matrix{ITensor}; cutoff, maxdim)
#TODO
# Contract in every direction
combiner_gauge = combiners(linkinds, tn)
tnᶜ = insert_gauge(tn, combiner_gauge)
boundary_mpsᶜ = contract_approx(tnᶜ; maxdim=maxdim, cutoff=cutoff)
tn_cacheᶜ = contraction_cache(tnᶜ, boundary_mpsᶜ)
tn_cache = insert_gauge(tn_cacheᶜ, combiner_gauge)
return boundary_mps(tn_cache)
end

# Return a network that when contracted equals the
# squared norm of the input network.
# TODO: rename norm2_network
Expand All @@ -359,27 +370,13 @@ function sqnorm_approx(ψ::Matrix{ITensor}; center, cutoff, maxdim)
ψ′ = addtags(linkinds, ψ, "ket")
# TODO: implement contract(commoninds, ψ′, ψᴴ)
tn = ψ′ .* ψᴴ

# Contract in every direction
combiner_gauge = combiners(linkinds, tn)
tnᶜ = insert_gauge(tn, combiner_gauge)
boundary_mpsᶜ = contract_approx(tnᶜ; maxdim=maxdim, cutoff=cutoff)

tn_cacheᶜ = contraction_cache(tnᶜ, boundary_mpsᶜ)
tn_cache = insert_gauge(tn_cacheᶜ, combiner_gauge)
_boundary_mps = boundary_mps(tn_cache)

#
# Insert projectors horizontally (to measure e.g. properties
# in a row of the network)
#

tn_projected = insert_projectors(tn, _boundary_mps; center=center)
bmps = boundary_mps(tn; cutoff=cutoff, maxdim=maxdim)
tn_projected = insert_projectors(tn, bmps; center=center)
tn_split, Pl, Pr = tn_projected

ψᴴ_split = split_network(ψᴴ)
ψ′_split = split_network(ψ′)

Pl_flat = reduce(vcat, Pl)
Pr_flat = reduce(vcat, Pr)
return mapreduce(vec, vcat, (ψᴴ_split, ψ′_split, Pl_flat, Pr_flat))
Expand Down
8 changes: 6 additions & 2 deletions src/ITensorNetworks/itensor_network.jl
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,10 @@ function ITensors.addtags(::typeof(linkinds), tn, args...)
return mapinds(x -> addtags(x, args...), linkinds, tn)
end

function ITensors.removetags(::typeof(linkinds), tn, args...)
return mapinds(x -> removetags(x, args...), linkinds, tn)
end

# Compute the sets of combiners that combine the link indices
# of the tensor network so that neighboring tensors only
# share a single larger index.
Expand Down Expand Up @@ -198,8 +202,8 @@ function split_links(H::Union{MPS,MPO}; split_tags=("" => ""), split_plevs=(0 =>
for bond in keys(l)
n1, n2 = bond
lₙ = l[bond]
left_l_n = setprime(addtags(lₙ, left_tags), left_plev)
right_l_n = setprime(addtags(lₙ, right_tags), right_plev)
left_l_n = prime(addtags(lₙ, left_tags), left_plev)
right_l_n = prime(addtags(lₙ, right_tags), right_plev)
Hsplit[n1] = replaceinds(Hsplit[n1], lₙ => left_l_n)
Hsplit[n2] = replaceinds(Hsplit[n2], lₙ => right_l_n)
end
Expand Down
36 changes: 36 additions & 0 deletions src/ITensorNetworks/peps.jl
Original file line number Diff line number Diff line change
Expand Up @@ -81,10 +81,36 @@ broadcast_inner(A::PEPS, B::PEPS) = mapreduce(v -> v[], +, A.data .* B.data)

ITensors.prime(P::PEPS, n::Integer=1) = PEPS(map(x -> prime(x, n), P.data))

# prime a PEPS with specified indices
function ITensors.prime(indices::Array{<:Index,1}, P::PEPS, n::Integer=1)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Would be nice to implement a more direct prime(siteinds, ::PEPS, ...) function to complement prime(linkinds, ::PEPS, ...). I'll add that as an issue.

function primeinds(tensor)
prime_inds = [ind for ind in inds(tensor) if ind in indices]
return replaceinds(tensor, prime_inds => prime(prime_inds, n))
end
return PEPS(map(x -> primeinds(x), P.data))
end

# prime linkinds of a PEPS
function ITensors.prime(::typeof(linkinds), P::PEPS, n::Integer=1)
return PEPS(mapinds(x -> prime(x, n), linkinds, P.data))
end

function ITensors.addtags(::typeof(linkinds), P::PEPS, args...)
return PEPS(addtags(linkinds, P.data, args...))
end

function ITensors.removetags(::typeof(linkinds), P::PEPS, args...)
return PEPS(removetags(linkinds, P.data, args...))
end

ITensors.data(P::PEPS) = P.data

split_network(P::PEPS) = PEPS(split_network(data(P)))

function ITensors.commoninds(p1::PEPS, p2::PEPS)
return mapreduce(a -> commoninds(a...), vcat, zip(p1.data, p2.data))
end

# Get the tensor network of <peps|peps'>
function inner_network(peps::PEPS, peps_prime::PEPS)
return vcat(vcat(peps.data...), vcat(peps_prime.data...))
Expand Down Expand Up @@ -117,3 +143,13 @@ function flatten(v::Array{<:PEPS})
tensor_list = [vcat(peps.data...) for peps in v]
return vcat(tensor_list...)
end

function insert_projectors(peps::PEPS, center, cutoff=1e-15, maxdim=100)
# Square the tensor network
psi_bra = addtags(linkinds, dag.(peps.data), "bra")
psi_ket = addtags(linkinds, peps.data, "ket")
tn = psi_bra .* psi_ket
bmps = boundary_mps(tn; cutoff=cutoff, maxdim=maxdim)
tn_split, pl, pr = insert_projectors(tn, bmps; center=center)
return tn_split, vcat(reduce(vcat, pl), reduce(vcat, pr))
end
3 changes: 2 additions & 1 deletion src/Optimizations/Optimizations.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,10 @@ module Optimizations

using ITensors

export gradient_descent, generate_inner_network
export gradient_descent, generate_inner_network, rayleigh_quotient

include("peps.jl")
include("itensor_network.jl")
include("run.jl")
include("optimizers.jl")

Expand Down
15 changes: 15 additions & 0 deletions src/Optimizations/itensor_network.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
using ChainRulesCore
using ..ITensorNetworks
using ..ITensorNetworks: split_network, default_projector_center

function ChainRulesCore.rrule(
::typeof(split_network),
Comment thread
LinjianMa marked this conversation as resolved.
tn::Matrix{ITensor};
projector_center=default_projector_center(tn),
)
function pullback(dtn_split::Matrix{ITensor})
dtn = map(t -> replaceprime(t, 1 => 0), dtn_split)
return (NoTangent(), dtn, NoTangent())
end
return split_network(tn; projector_center=projector_center), pullback
end
85 changes: 82 additions & 3 deletions src/Optimizations/peps.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,15 @@ using AutoHOOT, ChainRulesCore, Zygote
using ..ITensorAutoHOOT
using ..ITensorNetworks
using ITensors: setinds
using ..ITensorNetworks: PEPS, inner_network, flatten
using ..ITensorNetworks: PEPS, inner_network, flatten, insert_projectors, split_network
using ..ITensorAutoHOOT: batch_tensor_contraction

broadcast_notangent(a) = broadcast(_ -> NoTangent(), a)

function ChainRulesCore.rrule(::typeof(ITensors.data), P::PEPS)
return P.data, d_data -> (NoTangent(), PEPS(d_data))
end

function ChainRulesCore.rrule(::typeof(PEPS), data::Matrix{ITensor})
return PEPS(data), dpeps -> (NoTangent(), dpeps.data)
end
Expand All @@ -13,13 +19,41 @@ function ChainRulesCore.rrule(::typeof(ITensors.prime), P::PEPS, n::Integer=1)
return prime(P, n), dprime -> (NoTangent(), prime(dprime, -n), NoTangent())
end

function ChainRulesCore.rrule(
::typeof(ITensors.addtags), ::typeof(linkinds), P::PEPS, args...
)
function pullback(dtag_peps)
dP = ITensors.removetags(linkinds, dtag_peps, args...)
return (NoTangent(), NoTangent(), dP, broadcast_notangent(args)...)
end
return ITensors.addtags(linkinds, P, args...), pullback
end

function ChainRulesCore.rrule(
::typeof(ITensors.removetags), ::typeof(linkinds), P::PEPS, args...
)
function pullback(dtag_peps)
dP = ITensors.addtags(linkinds, dtag_peps, args...)
return (NoTangent(), NoTangent(), dP, broadcast_notangent(args)...)
end
return ITensors.removetags(linkinds, P, args...), pullback
end

function ChainRulesCore.rrule(
::typeof(ITensors.prime), ::typeof(linkinds), P::PEPS, n::Integer=1
)
return prime(linkinds, P, n),
dprime -> (NoTangent(), NoTangent(), prime(linkinds, dprime, -n), NoTangent())
end

function ChainRulesCore.rrule(
::typeof(ITensors.prime), indices::Array{<:Index,1}, P::PEPS, n::Integer=1
)
primeinds = [prime(ind, n) for ind in indices]
return prime(indices, P, n),
dprime -> (NoTangent(), NoTangent(), prime(primeinds, dprime, -n), NoTangent())
end

function ChainRulesCore.rrule(::typeof(flatten), v::Array{<:PEPS})
size_list = [size(peps.data) for peps in v]
function adjoint_pullback(dt)
Expand Down Expand Up @@ -68,9 +102,28 @@ end
peps::PEPS, peps_prime::PEPS, peps_prime_ham::PEPS, Hlocal::Array
)

function generate_inner_network(
peps::PEPS,
peps_prime::PEPS,
peps_prime_ham::PEPS,
projectors::Array{<:ITensor,1},
Hlocal::Array,
)
network_list = generate_inner_network(peps, peps_prime, peps_prime_ham, Hlocal)
return map(network -> vcat(network, projectors), network_list)
end

@non_differentiable generate_inner_network(
peps::PEPS,
peps_prime::PEPS,
peps_prime_ham::PEPS,
projectors::Array{<:ITensor,1},
Hlocal::Array,
)

function rayleigh_quotient(inners::Array)
self_inner = inners[length(inners)][]
expectations = sum(inners[1:(length(inners) - 1)])[]
self_inner = inners[end][]
expectations = sum(inners[1:(end - 1)])[]
return expectations / self_inner
end

Expand All @@ -86,3 +139,29 @@ function loss_grad_wrap(peps::PEPS, Hlocal::Array)
loss_w_grad(peps::PEPS) = loss(peps), gradient(loss, peps)[1]
return loss_w_grad
end

@non_differentiable insert_projectors(peps::PEPS, center, cutoff, maxdim)

@non_differentiable ITensors.commoninds(p1::PEPS, p2::PEPS)

function loss_grad_wrap(peps::PEPS, Hlocal::Array, ::typeof(insert_projectors))
center = (div(size(peps.data)[1] - 1, 2) + 1, :)
function loss(peps::PEPS)
tn_split, projectors = insert_projectors(peps, center)
peps_bra = addtags(linkinds, peps, "bra")
peps_ket = addtags(linkinds, peps, "ket")
sites = commoninds(peps_bra, peps_ket)
peps_bra_split = split_network(peps_bra)
peps_ket_split = split_network(peps_ket)
peps_ket_split_ham = prime(sites, peps_ket_split)
# generate network
network_list = generate_inner_network(
peps_bra_split, peps_ket_split, peps_ket_split_ham, projectors, Hlocal
)
variables = flatten([peps_bra_split, peps_ket_split, peps_ket_split_ham])
inners = batch_tensor_contraction(network_list, variables...)
return rayleigh_quotient(inners)
end
loss_w_grad(peps::PEPS) = loss(peps), gradient(loss, peps)[1]
return loss_w_grad
end
44 changes: 33 additions & 11 deletions src/Optimizations/run.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,18 @@ using OptimKit
using ..ITensorNetworks
using ..ITensorNetworks: broadcast_add, broadcast_minus, broadcast_mul, broadcast_inner

function gradient_descent(peps::PEPS, loss_w_grad; stepsize::Float64, num_sweeps::Int)
# gradient descent iterations
losses = []
for iter in 1:num_sweeps
l, g = loss_w_grad(peps)
print("The rayleigh quotient at iteraton $iter is $l\n")
peps = broadcast_minus(peps, broadcast_mul(stepsize, g))
push!(losses, l)
end
return losses
end

"""Update PEPS based on gradient descent
Parameters
----------
Expand All @@ -15,21 +27,19 @@ An array containing Rayleigh quotient losses after each iteration.
"""
function gradient_descent(peps::PEPS, Hlocal::Array; stepsize::Float64, num_sweeps::Int)
loss_w_grad = loss_grad_wrap(peps, Hlocal)
# gradient descent iterations
losses = []
for iter in 1:num_sweeps
l, g = loss_w_grad(peps)
print("The rayleigh quotient at iteraton $iter is $l\n")
peps = broadcast_minus(peps, broadcast_mul(stepsize, g))
push!(losses, l)
end
return losses
return gradient_descent(peps, loss_w_grad; stepsize=stepsize, num_sweeps=num_sweeps)
end

function OptimKit.optimize(peps::PEPS, Hlocal::Array; num_sweeps::Int, method="GD")
function gradient_descent(
peps::PEPS, Hlocal::Array, ::typeof(insert_projectors); stepsize::Float64, num_sweeps::Int
)
loss_w_grad = loss_grad_wrap(peps, Hlocal, insert_projectors)
return gradient_descent(peps, loss_w_grad; stepsize=stepsize, num_sweeps=num_sweeps)
end

function OptimKit.optimize(peps::PEPS, loss_w_grad; num_sweeps::Int, method="GD")
@assert(method in ["GD", "LBFGS", "CG"])
inner(x, peps1, peps2) = broadcast_inner(peps1, peps2)
loss_w_grad = loss_grad_wrap(peps, Hlocal)
scale(peps, alpha) = broadcast_mul(alpha, peps)
add(peps1, peps2, alpha) = broadcast_add(peps1, broadcast_mul(alpha, peps2))
retract(peps1, peps2, alpha) = (add(peps1, peps2, alpha), peps2)
Expand All @@ -48,3 +58,15 @@ function OptimKit.optimize(peps::PEPS, Hlocal::Array; num_sweeps::Int, method="G
)
return history[:, 1]
end

function OptimKit.optimize(peps::PEPS, Hlocal::Array; num_sweeps::Int, method="GD")
loss_w_grad = loss_grad_wrap(peps, Hlocal)
return optimize(peps, loss_w_grad; num_sweeps=num_sweeps, method=method)
end

function OptimKit.optimize(
peps::PEPS, Hlocal::Array, ::typeof(insert_projectors); num_sweeps::Int, method="GD"
)
loss_w_grad = loss_grad_wrap(peps, Hlocal, insert_projectors)
return optimize(peps, loss_w_grad; num_sweeps=num_sweeps, method=method)
end
Loading