Skip to content
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
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
name = "ITensorNetworksNext"
uuid = "302f2e75-49f0-4526-aef7-d8ba550cb06c"
version = "0.9.9"
version = "0.9.10"
authors = ["ITensor developers <support@itensor.org> and contributors"]

[workspace]
Expand Down
60 changes: 45 additions & 15 deletions src/apply/apply_operators.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,12 @@ using .AlgorithmsInterfaceExtensions: AlgorithmsInterfaceExtensions as AIE
using AlgorithmsInterface: AlgorithmsInterface as AI
using Base: @kwdef
using Graphs: dst, src, vertices
using ITensorBase:
ITensorBase as ITB, AbstractITensor, dimnames, inputnames, operator, replacedimnames
using ITensorBase: ITensorBase as ITB, AbstractITensor, dimnames, inputnames, operator,
outputnames, replacedimnames
using LinearAlgebra: norm
using MatrixAlgebraKit: qr_compact, svd_trunc
using MatrixAlgebraKit: eigh_full, project_hermitian, qr_compact, svd_trunc
using NamedGraphs.GraphsExtensions: all_edges, boundary_edges
using TensorAlgebra.MatrixAlgebra: gram_eigh_full, gram_eigh_full_with_pinv
using TensorAlgebra.MatrixAlgebra: invsqrth_safe, sqrth_safe

# === Top-level user entry point ===

Expand Down Expand Up @@ -205,6 +205,36 @@ end

# === BP simple-update implementation ===

# The odd-parity sign leaves a fermionic message positive semidefinite in only one
# bipartition, so diagonalize it in the transposed (bra, ket) one, placing the
# eigenvectors on the bra and ket legs ready to sandwich a function of the eigenvalues.
# Shared by `message_root` and `message_gauge`.
function message_eigen(message)
ket, bra = outputnames(message), inputnames(message)
hermitian_message = project_hermitian(ITB.state(message), ket, bra)
d, v = eigh_full(hermitian_message, bra, ket)
name_d′, name_d = dimnames(d)
v_ket = conj(v)
v_bra = replacedimnames(v, only(ket) => only(bra), name_d => name_d′)
return d, v_bra, v_ket
end

# The balanced Hermitian root `v * √d * v'` of the message. This gauge works with
# fermions; the asymmetric gauge `√d * v'` has an issue that is under investigation.
function message_root(message)
d, v_bra, v_ket = message_eigen(message)
name_d′, name_d = dimnames(d)
return v_bra * sqrth_safe(d, (name_d′,), (name_d,)) * v_ket
end

# The message root paired with its inverse `v * √d⁻¹ * v'`, to gauge a bond and undo it.
function message_gauge(message)
d, v_bra, v_ket = message_eigen(message)
name_d′, name_d = dimnames(d)
return v_bra * sqrth_safe(d, (name_d′,), (name_d,)) * v_ket,
v_bra * invsqrth_safe(d, (name_d′,), (name_d,)) * v_ket
end

function apply_gate_bp!(
dest::AbstractITensorNetwork, op::AbstractITensor,
state::AbstractITensorNetwork, env; kwargs...
Expand Down Expand Up @@ -233,7 +263,7 @@ function apply_gate_bp_nsite!(
ψv = ITB.apply(op, state[v])
if normalize
gauges = [
conj(gram_eigh_full(env[e]))
message_root(env[e])
for e in boundary_edges(state, vs; dir = :in)
]
ψv /= norm(prod([[ψv]; gauges]))
Expand All @@ -249,12 +279,10 @@ function apply_gate_bp_nsite!(
)
v1, v2 = vs
edges_in = boundary_edges(state, vs; dir = :in)
grams_v1 =
[gram_eigh_full_with_pinv(env[e]) for e in edges_in if dst(e) == v1]
grams_v2 =
[gram_eigh_full_with_pinv(env[e]) for e in edges_in if dst(e) == v2]
gauges_v1, inv_gauges_v1 = conj.(first.(grams_v1)), conj.(last.(grams_v1))
gauges_v2, inv_gauges_v2 = conj.(first.(grams_v2)), conj.(last.(grams_v2))
siv_v1 = [message_gauge(env[e]) for e in edges_in if dst(e) == v1]
siv_v2 = [message_gauge(env[e]) for e in edges_in if dst(e) == v2]
gauges_v1, inv_gauges_v1 = first.(siv_v1), conj.(last.(siv_v1))
gauges_v2, inv_gauges_v2 = first.(siv_v2), conj.(last.(siv_v2))

ψ_v1 = prod([[state[v1]]; gauges_v1])
ψ_v2 = prod([[state[v2]]; gauges_v2])
Expand All @@ -267,17 +295,19 @@ function apply_gate_bp_nsite!(
S = S / norm(S)
end
name_v1, name_v2 = dimnames(S)
sqrt_S = sqrt(S, (name_v1,), (name_v2,))
sqrt_S = sqrth_safe(S, (name_v1,), (name_v2,); atol = 0, rtol = 0)
R_v1 = replacedimnames(U_v1 * sqrt_S, name_v2 => name_v1)
R_v2 = sqrt_S * U_v2

dest[v1] = prod([[Q_v1 * R_v1]; inv_gauges_v1])
dest[v2] = prod([[Q_v2 * R_v2]; inv_gauges_v2])

env[v1 => v2] = operator(conj(S), (name_v2,), (name_v1,))
# The graded contraction of a factor with its conjugate carries the odd-parity sign.
env[v1 => v2] = operator(
replacedimnames(conj(R_v1), name_v1 => name_v2) * R_v1, (name_v1,), (name_v2,)
)
env[v2 => v1] = operator(
conj(replacedimnames(S, name_v1 => name_v2, name_v2 => name_v1)),
(name_v2,), (name_v1,)
replacedimnames(conj(R_v2), name_v1 => name_v2) * R_v2, (name_v1,), (name_v2,)
)
return dest
end
27 changes: 23 additions & 4 deletions src/beliefpropagation/beliefpropagation.jl
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,9 @@ using .AlgorithmsInterfaceExtensions:
using AlgorithmsInterface: AlgorithmsInterface as AI
using DataGraphs: edge_data
using Graphs: AbstractEdge, edges, edgetype, has_edge, vertices
using ITensorBase: AbstractITensor
using LinearAlgebra: norm, normalize
using ITensorBase:
AbstractITensor, NamedTensorOperator, inputnames, operator, outputnames, state
using LinearAlgebra: norm, normalize, tr
using NamedGraphs.GraphsExtensions:
add_edges!, boundary_edges, forest_cover_edge_sequence, subgraph
using NamedGraphs.PartitionedGraphs: quotientvertices
Expand Down Expand Up @@ -243,10 +244,28 @@ function message_update!(algorithm::SimpleMessageUpdate, cache, factors, edge)
messages = collect(incoming_messages(cache, edge))
factor = factors[src(edge)]

new_message = contract_network([messages; [factor]]; alg = algorithm.contraction_alg)
# `contract_network` works on plain named arrays, so unwrap any operator messages to
# their underlying tensors before contracting (fermionic signs ride on the graded
# arrays, so nothing is lost).
message_tensors = state.(messages)
new_message = contract_network(
[message_tensors; [factor]]; alg = algorithm.contraction_alg
)

# `contract_network` drops the bra/ket operator structure, so restore it from the
# existing message. A doubled (ket/bra) message is then a bond operator: normalize by
# its trace, which is sign-correct for fermionic bonds (the entrywise `sum` can flip
# the odd-parity block's sign). A single-layer message stays a vector with no bra/ket
# pairing, so fall back to the entrywise sum there.
old_message = cache[edge]
if old_message isa NamedTensorOperator
new_message =
operator(new_message, outputnames(old_message), inputnames(old_message))
end

if algorithm.normalize
message_norm = sum(new_message)
message_norm =
new_message isa NamedTensorOperator ? tr(new_message) : sum(new_message)
if !iszero(message_norm)
new_message /= message_norm
end
Expand Down
11 changes: 6 additions & 5 deletions src/beliefpropagation/messagecache.jl
Original file line number Diff line number Diff line change
Expand Up @@ -209,13 +209,14 @@ function similar_message_environment(nn::NormNetwork)
ketview = KetView(nn)

ketnames = linknames(ketview, edge)
ketaxis = unnamed.(linkaxes(ketview, edge))

brainds = linkinds(braview, edge)
branames = name.(brainds)
braaxis = unnamed.(brainds)
branames = linknames(braview, edge)

# Message axis is conj to the tensor it points to.
message = similar_operator(ketview[vertex], braaxis, branames, ketnames)
# Bond leg (ket) = operator output, bra-layer leg = input. Built on the src-side ket
# axis, whose arrow is opposite the dst endpoint's bond, so the gauge contracts back
# into the destination state.
message = similar_operator(ketview[vertex], ketaxis, ketnames, branames)

return edge => message
end
Expand Down
4 changes: 3 additions & 1 deletion test/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ SparseArraysBase = "0d5efcca-f356-4864-8770-e1ed8d78f208"
StableRNGs = "860ef19b-820b-49d6-a774-d7a799459cd3"
Suppressor = "fd094767-a336-5f1f-9728-57cf17d0bbfb"
TensorAlgebra = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a"
TensorKitSectors = "13a9c161-d5da-41f0-bcbd-e1a08ae0647f"
Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40"

[sources.ITensorNetworksNext]
Expand All @@ -29,7 +30,7 @@ AlgorithmsInterface = "0.1"
Aqua = "0.8.14"
DataGraphs = "0.5"
Dictionaries = "0.4.5"
GradedArrays = "0.13.2"
GradedArrays = "0.14"
Graphs = "1.13.1"
ITensorBase = "0.13"
ITensorNetworksNext = "0.9"
Expand All @@ -44,4 +45,5 @@ SparseArraysBase = "0.10.4"
StableRNGs = "1"
Suppressor = "0.2.8"
TensorAlgebra = "0.16, 0.17"
TensorKitSectors = "0.3"
Test = "1.10"
70 changes: 15 additions & 55 deletions test/test_apply_operator.jl
Original file line number Diff line number Diff line change
@@ -1,19 +1,18 @@
using GradedArrays: U1, gradedrange
using Graphs: dst, edges, src, vertices
using ITensorBase: ITensorBase as ITB, Index, name, named, operator, replacedimnames,
setname, state, uniquename
using ITensorNetworksNext: BraView, ITensorNetwork, KetView, NormNetwork, apply_operator,
apply_operators, beliefpropagation, braname, insertlink!, linknames,
message_environment, messagecache, tensornetwork
using ITensorBase: ITensorBase as ITB, Index, name, operator, setname, uniquename
using ITensorNetworksNext: NormNetwork, apply_operator, apply_operators, insertlink!,
message_environment, tensornetwork
using MatrixAlgebraKit: svd_trunc, truncrank
using NamedGraphs.NamedGraphGenerators: named_cycle_graph, named_path_graph
using NamedGraphs: NamedGraph
using Random: AbstractRNG
using StableRNGs: StableRNG
using TensorKitSectors: FermionParity
using Test: @test, @testset

const spinone = Base.OneTo(3)
const spinone_u1 = gradedrange([U1(2) => 1, U1(0) => 1, U1(-2) => 1])
const fermion = gradedrange([FermionParity(0) => 2, FermionParity(1) => 2])

function randn_operator(rng::AbstractRNG, elt::Type, domain_namedaxes)
codomain_namedaxes = setname.(domain_namedaxes, uniquename.(name.(domain_namedaxes)))
Expand All @@ -22,6 +21,10 @@ function randn_operator(rng::AbstractRNG, elt::Type, domain_namedaxes)
return operator(data, name.(codomain_namedaxes), name.(domain_namedaxes))
end

# Build a random state by applying random gates layer by layer, carrying the belief
# propagation environment through the applications. The returned `env` is the environment
# the gate applications produced, ready to gauge the next application (belief-propagation
# convergence itself is covered separately in `test_beliefpropagation.jl`).
function random_state(rng::AbstractRNG, elt::Type, g, site_axes; nlayers, trunc)
network = tensornetwork(vertices(g)) do v
return randn(rng, elt, (site_axes[v],))
Expand All @@ -36,24 +39,11 @@ function random_state(rng::AbstractRNG, elt::Type, g, site_axes; nlayers, trunc)
gate = randn_operator(rng, elt, (site_axes[src(e)], site_axes[dst(e)]))
network, env = apply_operator(gate, network, env; trunc)
end
return network
return network, env
end

function operator_message_cache(nn::NormNetwork, messages)
return messagecache(keys(messages)) do edge
ketnames = linknames(KetView(nn), edge)
branames = linknames(BraView(nn), edge)

bramap = Dict(branames .=> Base.Fix1(braname, nn).(ketnames))

renamed_message = replacedimnames(name -> get(bramap, name, name), messages[edge])

return operator(renamed_message, branames, ketnames)
end
end

@testset "apply_operator (T=$T, $(nameof(typeof(site_range))))" for site_range in (
spinone, spinone_u1,
@testset "apply_operator (T=$T, $label)" for (label, site_range) in (
"spinone" => spinone, "spinone_u1" => spinone_u1, "fermion" => fermion,
),
T in (Float32, Float64, ComplexF64)

Expand All @@ -63,17 +53,7 @@ end
rng = StableRNG(123)
g = named_cycle_graph(N)
site_axes = Dict(v => Index(site_range) for v in vertices(g))
network = random_state(rng, T, g, site_axes; nlayers = 2, trunc = truncrank(4))

nn = NormNetwork(network)

env = beliefpropagation(
nn,
message_environment(msg -> state(fill!(msg, true)), nn);
stopping_criterion = (; maxiter = 100, tol = 1.0e-13)
)

env = operator_message_cache(nn, env)
network, env = random_state(rng, T, g, site_axes; nlayers = 2, trunc = truncrank(4))

for gate in (
randn_operator(rng, T, (site_axes[2],)),
Expand All @@ -88,17 +68,7 @@ end
rng = StableRNG(123)
g = named_path_graph(N)
site_axes = Dict(v => Index(site_range) for v in vertices(g))
network = random_state(rng, T, g, site_axes; nlayers = 2, trunc = truncrank(4))

nn = NormNetwork(network)

env = beliefpropagation(
nn,
message_environment(msg -> state(fill!(msg, true)), nn);
stopping_criterion = (; maxiter = 100, tol = 1.0e-13)
)

env = operator_message_cache(nn, env)
network, env = random_state(rng, T, g, site_axes; nlayers = 2, trunc = truncrank(4))

gate = randn_operator(rng, T, (site_axes[2], site_axes[3]))
gated_full = ITB.apply(gate, prod(network))
Expand All @@ -113,17 +83,7 @@ end
rng = StableRNG(123)
g = named_cycle_graph(N)
site_axes = Dict(v => Index(site_range) for v in vertices(g))
network = random_state(rng, T, g, site_axes; nlayers = 2, trunc = truncrank(4))

nn = NormNetwork(network)

env = beliefpropagation(
nn,
message_environment(msg -> state(fill!(msg, true)), nn);
stopping_criterion = (; maxiter = 100, tol = 1.0e-13)
)

env = operator_message_cache(nn, env)
network, env = random_state(rng, T, g, site_axes; nlayers = 2, trunc = truncrank(4))

g1 = randn_operator(rng, T, (site_axes[2], site_axes[3]))
g2 = randn_operator(rng, T, (site_axes[3], site_axes[4]))
Expand Down
Loading