Skip to content
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.10.2"
version = "0.10.3"
authors = ["ITensor developers <support@itensor.org> and contributors"]

[workspace]
Expand Down
72 changes: 21 additions & 51 deletions src/apply/apply_operators.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,11 @@ 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,
outputnames, replacedimnames
using ITensorBase: AbstractITensor, apply, dimnames, inputnames, operator, replacedimnames
using LinearAlgebra: norm
using MatrixAlgebraKit: eigh_full, project_hermitian, qr_compact, svd_trunc
using MatrixAlgebraKit: project_hermitian, qr_compact, svd_trunc
using NamedGraphs: boundary_edges
using TensorAlgebra.MatrixAlgebra: invsqrth_safe, sqrth_safe
using TensorAlgebra.MatrixAlgebra: sqrth_invsqrth_safe, sqrth_safe

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

Expand Down Expand Up @@ -205,36 +204,6 @@ 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 All @@ -260,13 +229,13 @@ function apply_gate_bp_nsite!(
normalize, kwargs...
)
v = only(vs)
ψv = ITB.apply(op, state[v])
ψv = apply(op, state[v])
if normalize
gauges = [
message_root(env[e])
for e in boundary_edges(state, vs; dir = :in)
sqrt_messages = [
sqrth_safe(project_hermitian(env[e])) for
e in boundary_edges(state, vs; dir = :in)
]
ψv /= norm(prod([[ψv]; gauges]))
ψv /= norm(foldl((ψ, m) -> apply(m, ψ), sqrt_messages; init = ψv))
end
dest[v] = ψv
return dest
Expand All @@ -279,17 +248,19 @@ function apply_gate_bp_nsite!(
)
v1, v2 = vs
edges_in = boundary_edges(state, vs; dir = :in)
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))
roots_v1 =
[sqrth_invsqrth_safe(project_hermitian(env[e])) for e in edges_in if dst(e) == v1]
roots_v2 =
[sqrth_invsqrth_safe(project_hermitian(env[e])) for e in edges_in if dst(e) == v2]
sqrt_messages_v1, invsqrt_messages_v1 = first.(roots_v1), last.(roots_v1)
sqrt_messages_v2, invsqrt_messages_v2 = first.(roots_v2), last.(roots_v2)

ψ_v1 = prod([[state[v1]]; gauges_v1])
ψ_v2 = prod([[state[v2]]; gauges_v2])
ψ_v1 = foldl((ψ, m) -> apply(m, ψ), sqrt_messages_v1; init = state[v1])
ψ_v2 = foldl((ψ, m) -> apply(m, ψ), sqrt_messages_v2; init = state[v2])

Q_v1, R_v1 = qr_compact(ψ_v1, setdiff(dimnames(ψ_v1), dimnames(ψ_v2), dimnames(op)))
Q_v2, R_v2 = qr_compact(ψ_v2, setdiff(dimnames(ψ_v2), dimnames(ψ_v1), dimnames(op)))
op_R_v1v2 = ITB.apply(op, R_v1 * R_v2)
op_R_v1v2 = apply(op, R_v1 * R_v2)
U_v1, S, U_v2 = svd_trunc(op_R_v1v2, setdiff(dimnames(R_v1), dimnames(R_v2)); trunc)
if normalize
S = S / norm(S)
Expand All @@ -299,15 +270,14 @@ function apply_gate_bp_nsite!(
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])
dest[v1] = foldl((ψ, m) -> apply(m, ψ), invsqrt_messages_v1; init = Q_v1 * R_v1)
dest[v2] = foldl((ψ, m) -> apply(m, ψ), invsqrt_messages_v2; init = Q_v2 * R_v2)

# 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,)
replacedimnames(conj(R_v1), name_v1 => name_v2) * R_v1, (name_v2,), (name_v1,)
)
env[v2 => v1] = operator(
replacedimnames(conj(R_v2), name_v1 => name_v2) * R_v2, (name_v1,), (name_v2,)
replacedimnames(conj(R_v2), name_v1 => name_v2) * R_v2, (name_v2,), (name_v1,)
)
return dest
end
9 changes: 4 additions & 5 deletions src/beliefpropagation/beliefpropagation.jl
Original file line number Diff line number Diff line change
Expand Up @@ -270,14 +270,13 @@ function message_update!(algorithm::SimpleMessageUpdate, cache, factors, edge)
end

# `NormNetwork`: the message is a doubled (ket/bra) bond operator. Contracting a plain vertex factor
# with the incoming messages leaves the surviving bond legs dangling, so assign the ket/bra pairing
# the norm network gives this edge (the same convention as `similar_message_environment`). Normalize
# by the trace, which is sign-correct on fermionic bonds where the entrywise `sum` can flip the
# odd-parity block's sign.
# with the incoming messages leaves the surviving bond legs dangling, so assign the bra/ket pairing
# the norm network gives this edge (the same convention as `similar_message_environment`), in which
# the message is positive semidefinite and its trace is a positive normalization.
function message_update!(algorithm::SimpleMessageUpdate, cache, factors::NormNetwork, edge)
new_tensor = updated_message(algorithm, cache, factors, edge)
new_message = operator(
new_tensor, linknames(KetView(factors), edge), linknames(BraView(factors), edge)
new_tensor, linknames(BraView(factors), edge), linknames(KetView(factors), edge)
)
if algorithm.normalize
message_norm = tr(new_message)
Expand Down
10 changes: 4 additions & 6 deletions src/beliefpropagation/messagecache.jl
Original file line number Diff line number Diff line change
Expand Up @@ -203,14 +203,12 @@ function similar_message_environment(nn::NormNetwork)
ketview = KetView(nn)

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

branames = linknames(braview, edge)
braaxis = unnamed.(linkaxes(braview, edge))

# 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)
# Bra leg = operator output, ket leg = input, the bipartition in which the message
# is positive semidefinite.
message = similar_operator(ketview[vertex], braaxis, branames, ketnames)

return edge => message
end
Expand Down
2 changes: 2 additions & 0 deletions test/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ OMEinsumContractionOrders = "6f22d1fd-8eed-4bb7-9776-e7d684900715"
QuadGK = "1fd47b50-473d-5c70-9696-f719f8f3bcdc"
Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c"
StableRNGs = "860ef19b-820b-49d6-a774-d7a799459cd3"
TensorAlgebra = "68bd88dc-f39d-4e12-b2ca-f046b68fcc6a"
TensorKitSectors = "13a9c161-d5da-41f0-bcbd-e1a08ae0647f"
Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40"

Expand All @@ -37,5 +38,6 @@ OMEinsumContractionOrders = "1"
QuadGK = "2.11.2"
Random = "1.10"
StableRNGs = "1"
TensorAlgebra = "0.20"
TensorKitSectors = "0.3"
Test = "1.10"
23 changes: 19 additions & 4 deletions test/test_apply_operator.jl
Original file line number Diff line number Diff line change
@@ -1,12 +1,13 @@
using GradedArrays: U1, gradedrange
using Graphs: dst, edges, src, vertices
using ITensorBase: ITensorBase as ITB, Index, name, operator, setname, uniquename
using ITensorBase: Index, apply, name, operator, setname, uniquename
using ITensorNetworksNext: NormNetwork, apply_operator, apply_operators, insertlink!,
message_environment, tensornetwork
using MatrixAlgebraKit: svd_trunc, truncrank
using NamedGraphs: named_cycle_graph, named_path_graph
using Random: AbstractRNG
using StableRNGs: StableRNG
using TensorAlgebra.MatrixAlgebra: sqrth_invsqrth_safe
using TensorKitSectors: FermionParity
using Test: @test, @testset

Expand Down Expand Up @@ -60,7 +61,7 @@ end
randn_operator(rng, T, (site_axes[2], site_axes[3])),
)
gated, _ = apply_operator(gate, network, env)
@test prod(gated) ≈ ITB.apply(gate, prod(network)) rtol = eps(real(T))^(1 / 3)
@test prod(gated) ≈ apply(gate, prod(network)) rtol = eps(real(T))^(1 / 3)
end
end

Expand All @@ -71,7 +72,7 @@ end
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))
gated_full = apply(gate, prod(network))
left = [name(site_axes[v]) for v in 1:2]
U, S, Vt = svd_trunc(gated_full, left; trunc = truncrank(k))
gated, _ = apply_operator(gate, network, env; trunc = truncrank(k))
Expand All @@ -88,7 +89,21 @@ end
g1 = randn_operator(rng, T, (site_axes[2], site_axes[3]))
g2 = randn_operator(rng, T, (site_axes[3], site_axes[4]))
gated, _ = apply_operators([g1, g2], network, env)
@test prod(gated) ≈ ITB.apply(g2, ITB.apply(g1, prod(network))) rtol =
@test prod(gated) ≈ apply(g2, apply(g1, prod(network))) rtol =
eps(real(T))^(1 / 3)
end

@testset "message roots gauge a vertex and undo it" begin
rng = StableRNG(123)
g = named_path_graph(N)
site_axes = Dict(v => Index(site_range) for v in vertices(g))
network, env = random_state(rng, T, g, site_axes; nlayers = 2, trunc = truncrank(4))
rtol = eps(real(T))^(1 / 3)

for (edge, v) in ((2 => 3, 3), (3 => 2, 2))
sqrt_message, invsqrt_message = sqrth_invsqrth_safe(env[edge])
@test apply(invsqrt_message, apply(sqrt_message, network[v])) ≈ network[v] rtol =
rtol
end
end
end
35 changes: 29 additions & 6 deletions test/test_beliefpropagation.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,19 +2,22 @@ import AlgorithmsInterface as AI
using Base.Broadcast: materialize
using DataGraphs: DataGraphs, DataGraph, edge_data, edge_data_type
using Dictionaries: Dictionary, dictionary, set!
using GradedArrays: U1, gradedrange
using GradedArrays: U1, gradedrange, isdual
using Graphs: AbstractGraph, add_vertex!, dst, edges, has_edge, has_vertex, nv, rem_edge!,
src, vertices
using ITensorBase: Greedy, ITensor, Index, inds, name, noprime, outputnames, prime
using ITensorBase:
Greedy, ITensor, Index, apply, inds, name, noprime, outputnames, prime, state
using ITensorNetworksNext: ITensorNetworksNext, Exact, ITensorNetwork, MessageCache,
NormNetwork, SimpleMessageUpdate, StopWhenConverged, beliefpropagation,
bethe_free_energy, contract_network, contraction_order, edge_scalar, factor_tensors,
incoming_messages, insertlink!, linkinds, message_environment, messagecache,
region_scalar, subgraph, tensornetwork, updated_message, vertex_scalar, vertex_scalars
using LinearAlgebra: LinearAlgebra
bethe_free_energy, bratensor, contract_network, contraction_order, edge_scalar,
factor_tensors, incoming_messages, insertlink!, kettensor, linkaxes, linkinds,
message_environment, messagecache, region_scalar, subgraph, tensornetwork,
updated_message, vertex_scalar, vertex_scalars
using LinearAlgebra: LinearAlgebra, norm, tr
using NamedGraphs: NamedEdge, all_edges, incident_edges, named_comb_tree, named_grid,
named_path_graph, vertextype
using StableRNGs: StableRNG
using TensorAlgebra.MatrixAlgebra: sqrth_invsqrth_safe
using TensorKitSectors: FermionParity
using Test: @test, @testset

Expand Down Expand Up @@ -286,6 +289,26 @@ end
z_exact = (ket * conj(ket))[]
z_bp = exp(bethe_free_energy(nn, cache))
@test z_bp ≈ z_exact rtol = eps(real(T))^(1 / 3)

for edge in edges(cache)
msg = cache[edge]
@test real(tr(msg)) > 0
sqrt_msg, invsqrt_msg = sqrth_invsqrth_safe(msg)
v = dst(edge)
@test apply(invsqrt_msg, apply(sqrt_msg, network[v])) ≈ network[v] rtol =
eps(real(T))^(1 / 3)
end

@test isdual(only(linkaxes(network, 1 => 2))) !=
isdual(only(linkaxes(network, 4 => 3)))
ones = message_environment(one, nn)
for (edge, rest) in ((1 => 2, 2:4), (4 => 3, 1:3))
layers =
[[kettensor(nn, v) for v in rest]; [bratensor(nn, v) for v in rest]]
z_rest = contract_network([state(ones[edge]); layers])[]
@test z_rest ≈ norm(prod([network[v] for v in rest]))^2 rtol =
eps(real(T))^(1 / 3)
end
end
end

Expand Down
Loading
Loading