From 445f3c12220070df6223368ab319c0e21e4ff8fe Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Thu, 9 Jul 2026 13:20:54 -0400 Subject: [PATCH 01/22] `add_edge!` and `rem_edge!` on the `ITensorNetwork` type now return false instead of erroring This is inline with the `Graphs` behaviour. --- src/tensornetwork.jl | 14 +++----------- test/test_tensornetwork.jl | 10 ++++++++-- 2 files changed, 11 insertions(+), 13 deletions(-) diff --git a/src/tensornetwork.jl b/src/tensornetwork.jl index a2216cb7..dd51c0a2 100644 --- a/src/tensornetwork.jl +++ b/src/tensornetwork.jl @@ -170,17 +170,9 @@ function DataGraphs.underlying_graph_type(type::Type{<:ITensorNetwork{T, V}}) wh return fieldtype(type, :underlying_graph) end -function Graphs.rem_edge!(::ITensorNetwork, _edge) - return throw( - ErrorException("removing edges from the `ITensorNetwork` type is not supported.") - ) -end - -function Graphs.add_edge!(::ITensorNetwork, _edge) - return throw( - ErrorException("Adding edges to the `ITensorNetwork` type is not supported.") - ) -end +# Can't add/remove edges from `ITensorNetwork` as graph topology fixed by indices. +Graphs.rem_edge!(::ITensorNetwork, _edge) = false +Graphs.add_edge!(::ITensorNetwork, _edge) = false # PERF: fast lookup compared to `AbstractITensorNetwork` fallback. dimnamevertices(tn::ITensorNetwork, name) = tn.dimname_vertices[name] diff --git a/test/test_tensornetwork.jl b/test/test_tensornetwork.jl index ab333b15..bc060d64 100644 --- a/test/test_tensornetwork.jl +++ b/test/test_tensornetwork.jl @@ -48,8 +48,14 @@ using Test: @test, @test_throws, @testset @test_throws MethodError tn[e] = randn(2, 2) @test_throws MethodError tn[src(e) => dst(e)] = randn(2, 2) - # `rem_edge!` is intentionally unimplemented. - @test_throws ErrorException rem_edge!(tn, (1, 1) => (2, 1)) + # `rem_edge!` and `add_edge!` are intentionally unimplemented; they return + # `false` without modifying the network. + @test rem_edge!(tn, (1, 1) => (2, 1)) == false + @test has_edge(tn, (1, 1) => (2, 1)) + @test ne(tn) == 1 + @test add_edge!(tn, (2, 1) => (2, 2)) == false + @test !has_edge(tn, (2, 1) => (2, 2)) + @test ne(tn) == 1 tn[1, 1] = randn(Index(2)) tn[2, 1] = randn(Index(2)) From 4dc749ba764de22efe02b02ded19014314ff2d9e Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Thu, 9 Jul 2026 13:36:16 -0400 Subject: [PATCH 02/22] `dimnamevertices` for `ITensorNetwork` now returns empty set if index not in dictionary This is now consistant with the fallback defn of `dimnamevertices`. would error previously. --- src/tensornetwork.jl | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/tensornetwork.jl b/src/tensornetwork.jl index dd51c0a2..a298fa6e 100644 --- a/src/tensornetwork.jl +++ b/src/tensornetwork.jl @@ -175,7 +175,9 @@ Graphs.rem_edge!(::ITensorNetwork, _edge) = false Graphs.add_edge!(::ITensorNetwork, _edge) = false # PERF: fast lookup compared to `AbstractITensorNetwork` fallback. -dimnamevertices(tn::ITensorNetwork, name) = tn.dimname_vertices[name] +function dimnamevertices(tn::ITensorNetwork, name) + return get(tn.dimname_vertices, name, Set{vertextype(tn)}()) +end # PERF: fast lookup compared to `AbstractITensorNetwork` fallback. has_dimname(tn::ITensorNetwork, name) = haskey(tn.dimname_vertices, name) From 0b7e87d1be2aae55c8e744096d0944b4b89ab884 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Thu, 9 Jul 2026 13:36:40 -0400 Subject: [PATCH 03/22] New function `supportof` that gives the vertex support of an operator on a tensor network. --- src/abstracttensornetwork.jl | 24 ++++++++++++++++++++++-- src/apply/apply_operators.jl | 31 +++++++++++++++++-------------- 2 files changed, 39 insertions(+), 16 deletions(-) diff --git a/src/abstracttensornetwork.jl b/src/abstracttensornetwork.jl index 9515d084..d4218c78 100644 --- a/src/abstracttensornetwork.jl +++ b/src/abstracttensornetwork.jl @@ -4,8 +4,8 @@ using DataGraphs: DataGraphs, AbstractDataGraph, AbstractVertexDataGraph, edge_d using Dictionaries: Dictionary using Graphs: Graphs, AbstractEdge, AbstractGraph, add_edge!, add_vertex!, dst, edges, edgetype, ne, neighbors, nv, rem_edge!, src, vertices -using ITensorBase: - NamedUnitRange, inds, name, names, nametype, prime, uniquename, unnamedtype +using ITensorBase: ITensorOperator, NamedUnitRange, domainnames, inds, name, names, + nametype, prime, uniquename, unnamedtype using LinearAlgebra: LinearAlgebra using MacroTools: @capture using NamedGraphs: @@ -131,3 +131,23 @@ function insertlink!(tn::AbstractGraph, e) return tn end + +function supportof(tn::AbstractGraph, op::ITensorOperator) + support = Base.Generator(domainnames(op)) do name + vertices = dimnamevertices(tn, name) + + length(vertices) == 1 && return only(vertices) + + if length(vertices) == 0 + throw(ArgumentError("operator dim name $name not found in tensor network.")) + elseif length(vertices) > 1 + throw( + ArgumentError( + "operator dim name $name associated with multiple vertices in tensor network." + ) + ) + end + end + + return Set(support) +end diff --git a/src/apply/apply_operators.jl b/src/apply/apply_operators.jl index e029bcf9..80040e16 100644 --- a/src/apply/apply_operators.jl +++ b/src/apply/apply_operators.jl @@ -2,8 +2,8 @@ using .AlgorithmsInterfaceExtensions: AlgorithmsInterfaceExtensions as AIE using AlgorithmsInterface: AlgorithmsInterface as AI using Base: @kwdef using Graphs: dst, src, vertices -using ITensorBase: AbstractITensor, apply, inputnames, names, operator, rename -using LinearAlgebra: norm +using ITensorBase: ITensorBase, AbstractITensor, apply, domainnames, names, operator, rename +using LinearAlgebra: norm, normalize! using MatrixAlgebraKit: project_hermitian, qr_compact, svd_trunc using NamedGraphs: boundary_edges using TensorAlgebra.MatrixAlgebra: sqrth_invsqrth_safe, sqrth_safe @@ -208,12 +208,15 @@ function apply_gate_bp!( dest::AbstractITensorNetwork, op::AbstractITensor, state::AbstractITensorNetwork, env; kwargs... ) - op_in = inputnames(op) - vs = [v for v in vertices(state) if !isempty(intersect(op_in, sitenames(state, v)))] - isempty(vs) && throw( + vertices = supportof(state, op) + + isempty(vertices) && throw( ArgumentError("operator shares no indices with the tensor network") ) - return apply_gate_bp_nsite!(Val(length(vs)), dest, op, state, env, vs; kwargs...) + + N = Val(length(vertices)) + + return apply_gate_bp_nsite!(N, dest, op, state, env, vertices; kwargs...) end function apply_gate_bp_nsite!( @@ -225,11 +228,11 @@ end function apply_gate_bp_nsite!( ::Val{1}, dest::AbstractITensorNetwork, op::AbstractITensor, - state::AbstractITensorNetwork, env, vs; + state::AbstractITensorNetwork, env, vertices; normalize, kwargs... ) - v = only(vs) - ψv = apply(op, state[v]) + vertex = only(vertices) + ψv = apply(op, state[vertex]) if normalize sqrt_messages = [ sqrth_safe(project_hermitian(env[e])) for @@ -237,13 +240,13 @@ function apply_gate_bp_nsite!( ] ψv /= norm(foldl((ψ, m) -> apply(m, ψ), sqrt_messages; init = ψv)) end - dest[v] = ψv + dest[vertex] = ψv return dest end function apply_gate_bp_nsite!( ::Val{2}, dest::AbstractITensorNetwork, op::AbstractITensor, - state::AbstractITensorNetwork, env, vs; + state::AbstractITensorNetwork, env, vertices; trunc, normalize ) v1, v2 = vs @@ -262,9 +265,9 @@ function apply_gate_bp_nsite!( Q_v2, R_v2 = qr_compact(ψ_v2, setdiff(names(ψ_v2), names(ψ_v1), names(op))) op_R_v1v2 = apply(op, R_v1 * R_v2) U_v1, S, U_v2 = svd_trunc(op_R_v1v2, setdiff(names(R_v1), names(R_v2)); trunc) - if normalize - S = S / norm(S) - end + + normalize && normalize!(S) + name_v1, name_v2 = names(S) sqrt_S = sqrth_safe(S, (name_v1,), (name_v2,); atol = 0, rtol = 0) R_v1 = rename(U_v1 * sqrt_S, name_v2 => name_v1) From 4bf1763ba9ce4b6114df9aabd5ddfebb79912277 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Thu, 9 Jul 2026 15:14:47 -0400 Subject: [PATCH 04/22] Refactor `ITensorNetwork` topology functions; add index precheck for setting tensors --- src/tensornetwork.jl | 48 +++++++++++++++++++++++++++++++------------- 1 file changed, 34 insertions(+), 14 deletions(-) diff --git a/src/tensornetwork.jl b/src/tensornetwork.jl index a298fa6e..1f4992ae 100644 --- a/src/tensornetwork.jl +++ b/src/tensornetwork.jl @@ -115,38 +115,60 @@ DataGraphs.is_edge_assigned(::ITensorNetwork, _edge) = false DataGraphs.get_vertex_data(tn::ITensorNetwork, v) = tn.tensors[v] +function check_incoming_dimnames(tn, tensor, vertex) + for name in names(tensor) + vertices = get(tn.dimname_vertices, name, Set()) + if length(setdiff(vertices, Set([vertex]))) > 1 + throw( + ArgumentError( + "index $name can appear in at most one existing tensor" + ) + ) + end + end + return nothing +end + function DataGraphs.insert_vertex_data!(tn::ITensorNetwork, vertex, tensor) + check_incoming_dimnames(tn, tensor, vertex) add_vertex!(tn.underlying_graph, vertex) - set!_tensornetwork(tn, vertex, tensor) + update_tensornetwork_metadata!(tn, vertex, tensor) + insert!(tn.tensors, vertex, tensor) return tn end function DataGraphs.set_vertex_data!(tn::ITensorNetwork, tensor, vertex) - set!_tensornetwork(tn, vertex, tensor) + check_incoming_dimnames(tn, tensor, vertex) + update_tensornetwork_metadata!(tn, vertex, tensor) + set!(tn.tensors, vertex, tensor) return tn end -# "upsert" -function set!_tensornetwork(tn::ITensorNetwork, vertex, tensor) - newinds = names(tensor) +function update_tensornetwork_metadata!(tn, vertex, tensor) + oldnames = isassigned(tn, vertex) ? names(tn[vertex]) : Set() + newnames = names(tensor) - oldinds = get(mapview(names, tn.tensors), vertex, Set()) + update_tensornetwork_metadata!(tn, vertex, oldnames, newnames) + return tn +end + +function update_tensornetwork_metadata!(tn, vertex, oldinds, newinds) # Only have to deal with the indices that aren't shared. - for ind in symdiff(oldinds, newinds) - if ind in oldinds - delete_ind_edge!(tn, ind) - delete_ind_vertex!(tn, ind, vertex) + for name in symdiff(oldinds, newinds) + if name in oldinds + delete_ind_edge!(tn, name) + delete_ind_vertex!(tn, name, vertex) continue end # Now `ind` must be a new index that's not in `oldinds` - vertex_list = get!(tn.dimname_vertices, ind, Set()) + vertex_list = get!(tn.dimname_vertices, name, Set()) if length(vertex_list) > 1 throw( ArgumentError( - "index $ind can appear in at most one existing tensor, got $(length(vertex_list))." + "index $name can appear in at most one existing tensor, got $(length(vertex_list))." ) ) end @@ -159,8 +181,6 @@ function set!_tensornetwork(tn::ITensorNetwork, vertex, tensor) end end - set!(tn.tensors, vertex, tensor) - return tn end From a05bbb330d01b948a144a1e964d197ce1fbfba3c Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 20 Jul 2026 15:50:42 -0400 Subject: [PATCH 05/22] Refactor sum of log scalars into own function. Avoids some minor code duplication. --- src/beliefpropagation/messagecache.jl | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index c7c66c4f..cac0e709 100644 --- a/src/beliefpropagation/messagecache.jl +++ b/src/beliefpropagation/messagecache.jl @@ -175,23 +175,23 @@ function region_scalar(factors, messages, region) return mapreduce(vertex -> vertex_scalar(factors, messages, vertex), *, region) end +function sum_log_scalars(terms) + if any(t -> real(t) < 0, terms) + terms = complex.(terms) + end + return sum(log.(terms)) +end + # We need a graph structure here, so assume `factors` is a graph. function bethe_free_energy(factors, messages) numerator_terms = vertex_scalars(factors, messages) denominator_terms = edge_scalars(messages) - if any(t -> real(t) < 0, numerator_terms) - numerator_terms = complex.(numerator_terms) - end - if any(t -> real(t) < 0, denominator_terms) - denominator_terms = complex.(denominator_terms) - end - if any(iszero, denominator_terms) return -Inf end - return sum(log.(numerator_terms)) - sum(log.(denominator_terms)) + return sum_log_scalars(numerator_terms) - sum_log_scalars(denominator_terms) end # ===================================== NormNetwork ====================================== # From 0e3a8c09a922a8e8748affba932aba92e9e833de Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 20 Jul 2026 15:51:21 -0400 Subject: [PATCH 06/22] The `finalize_substate` function now dispatches on subsolve rather than solve. --- .../AlgorithmsInterfaceExtensions.jl | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl b/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl index da7bf76a..917fef9a 100644 --- a/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl +++ b/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl @@ -2,6 +2,8 @@ module AlgorithmsInterfaceExtensions using AlgorithmsInterface: AlgorithmsInterface as AI +abstract type NestedProblem <: AI.Problem end + # ============================ NestedAlgorithm ============================================= abstract type NestedAlgorithm <: AI.Algorithm end @@ -27,7 +29,7 @@ end function AI.step!(problem::AI.Problem, algorithm::NestedAlgorithm, state::AI.State) subproblem, subalgorithm, substate = initialize_subsolve(problem, algorithm, state) AI.solve!(subproblem, subalgorithm, substate) - finalize_substate!(problem, algorithm, state, substate) + finalize_substate!(subproblem, subalgorithm, substate, state) return state end From 26e95864005756f3e161387911a7fefb9d220e8e Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Tue, 21 Jul 2026 17:52:49 -0400 Subject: [PATCH 07/22] Upgrade to ITensorBase v0.13 Fix imports in `apply_operators.jl` --- src/abstracttensornetwork.jl | 4 ++-- src/apply/apply_operators.jl | 3 ++- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/src/abstracttensornetwork.jl b/src/abstracttensornetwork.jl index d4218c78..de30f3ac 100644 --- a/src/abstracttensornetwork.jl +++ b/src/abstracttensornetwork.jl @@ -4,7 +4,7 @@ using DataGraphs: DataGraphs, AbstractDataGraph, AbstractVertexDataGraph, edge_d using Dictionaries: Dictionary using Graphs: Graphs, AbstractEdge, AbstractGraph, add_edge!, add_vertex!, dst, edges, edgetype, ne, neighbors, nv, rem_edge!, src, vertices -using ITensorBase: ITensorOperator, NamedUnitRange, domainnames, inds, name, names, +using ITensorBase: ITensorOperator, NamedUnitRange, inputnames, inds, name, names, nametype, prime, uniquename, unnamedtype using LinearAlgebra: LinearAlgebra using MacroTools: @capture @@ -133,7 +133,7 @@ function insertlink!(tn::AbstractGraph, e) end function supportof(tn::AbstractGraph, op::ITensorOperator) - support = Base.Generator(domainnames(op)) do name + support = Base.Generator(inputnames(op)) do name vertices = dimnamevertices(tn, name) length(vertices) == 1 && return only(vertices) diff --git a/src/apply/apply_operators.jl b/src/apply/apply_operators.jl index 80040e16..858e6c9b 100644 --- a/src/apply/apply_operators.jl +++ b/src/apply/apply_operators.jl @@ -2,7 +2,8 @@ using .AlgorithmsInterfaceExtensions: AlgorithmsInterfaceExtensions as AIE using AlgorithmsInterface: AlgorithmsInterface as AI using Base: @kwdef using Graphs: dst, src, vertices -using ITensorBase: ITensorBase, AbstractITensor, apply, domainnames, names, operator, rename +using ITensorBase: ITensorBase as ITB, AbstractITensor, apply, inputnames, names, operator, + outputnames, rename using LinearAlgebra: norm, normalize! using MatrixAlgebraKit: project_hermitian, qr_compact, svd_trunc using NamedGraphs: boundary_edges From 973dceaec56816b7f69846fd1c189d5c15b77d9b Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Tue, 8 Sep 2026 09:04:34 -0400 Subject: [PATCH 08/22] Fix rename of `vs` to `vertices` in function body. --- src/apply/apply_operators.jl | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/apply/apply_operators.jl b/src/apply/apply_operators.jl index 858e6c9b..212cae45 100644 --- a/src/apply/apply_operators.jl +++ b/src/apply/apply_operators.jl @@ -237,7 +237,7 @@ function apply_gate_bp_nsite!( if normalize sqrt_messages = [ sqrth_safe(project_hermitian(env[e])) for - e in boundary_edges(state, vs; dir = :in) + e in boundary_edges(state, vertices; dir = :in) ] ψv /= norm(foldl((ψ, m) -> apply(m, ψ), sqrt_messages; init = ψv)) end @@ -250,8 +250,8 @@ function apply_gate_bp_nsite!( state::AbstractITensorNetwork, env, vertices; trunc, normalize ) - v1, v2 = vs - edges_in = boundary_edges(state, vs; dir = :in) + v1, v2 = vertices + edges_in = boundary_edges(state, vertices; dir = :in) roots_v1 = [sqrth_invsqrth_safe(project_hermitian(env[e])) for e in edges_in if dst(e) == v1] roots_v2 = From abeaeca5ae4f28c7045335a39e3018c24bd382f1 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Tue, 8 Sep 2026 10:35:18 -0400 Subject: [PATCH 09/22] Function `supportof` can now return an empty set; add tests. --- src/abstracttensornetwork.jl | 16 ++++++++-------- test/test_tensornetwork.jl | 32 ++++++++++++++++++++++++++++++-- 2 files changed, 38 insertions(+), 10 deletions(-) diff --git a/src/abstracttensornetwork.jl b/src/abstracttensornetwork.jl index de30f3ac..071ba714 100644 --- a/src/abstracttensornetwork.jl +++ b/src/abstracttensornetwork.jl @@ -4,7 +4,7 @@ using DataGraphs: DataGraphs, AbstractDataGraph, AbstractVertexDataGraph, edge_d using Dictionaries: Dictionary using Graphs: Graphs, AbstractEdge, AbstractGraph, add_edge!, add_vertex!, dst, edges, edgetype, ne, neighbors, nv, rem_edge!, src, vertices -using ITensorBase: ITensorOperator, NamedUnitRange, inputnames, inds, name, names, +using ITensorBase: ITensorOperator, NamedUnitRange, inds, inputnames, name, names, nametype, prime, uniquename, unnamedtype using LinearAlgebra: LinearAlgebra using MacroTools: @capture @@ -133,21 +133,21 @@ function insertlink!(tn::AbstractGraph, e) end function supportof(tn::AbstractGraph, op::ITensorOperator) - support = Base.Generator(inputnames(op)) do name - vertices = dimnamevertices(tn, name) + support = Set{vertextype(tn)}() - length(vertices) == 1 && return only(vertices) + for name in inputnames(op) + vertices = dimnamevertices(tn, name) - if length(vertices) == 0 - throw(ArgumentError("operator dim name $name not found in tensor network.")) - elseif length(vertices) > 1 + if length(vertices) > 1 throw( ArgumentError( "operator dim name $name associated with multiple vertices in tensor network." ) ) end + + union!(support, vertices) end - return Set(support) + return support end diff --git a/test/test_tensornetwork.jl b/test/test_tensornetwork.jl index bc060d64..29c600a8 100644 --- a/test/test_tensornetwork.jl +++ b/test/test_tensornetwork.jl @@ -2,9 +2,9 @@ using DataGraphs: DataGraph, assigned_edge_data, assigned_vertex_data, underlying_graph, vertex_data using Graphs: add_edge!, add_vertex!, dst, edges, edgetype, has_edge, has_vertex, is_directed, ne, nv, rem_edge!, rem_vertex!, src, vertices -using ITensorBase: Index, LazyITensor, inds +using ITensorBase: Index, LazyITensor, inds, operator using ITensorNetworksNext: ITensorNetwork, has_ind, linkaxes, linkinds, linknames, siteaxes, - siteinds, sitenames, tensornetwork + siteinds, sitenames, supportof, tensornetwork using NamedGraphs: convert_vertextype, incident_edges, named_grid, named_path_graph, similar_graph, subgraph, vertextype using Test: @test, @test_throws, @testset @@ -125,6 +125,34 @@ using Test: @test, @test_throws, @testset @test sitenames(tn, 3) == [s[3].name] end + @testset "`supportof`" begin + g = named_path_graph(3) + l = Dict(e => Index(2) for e in edges(g)) + l = merge(l, Dict(reverse(e) => l[e] for e in edges(g))) + s = Dict(v => Index(2) for v in vertices(g)) + tn = tensornetwork(vertices(g)) do v + is = map(e -> l[e], incident_edges(g, v)) + return randn((s[v], is...)) + end + + o1 = operator(randn(2, 2), (Index(2),), (s[2],)) + @test supportof(tn, o1) == Set([2]) + + o12 = operator(randn(2, 2, 2, 2), (Index(2), Index(2)), (s[1], s[2])) + @test supportof(tn, o12) == Set([1, 2]) + + # An input name that no tensor in the network carries contributes no vertex. + o_absent = operator(randn(2, 2), (Index(2),), (Index(2),)) + @test isempty(supportof(tn, o_absent)) + + o_partial = operator(randn(2, 2, 2, 2), (Index(2), Index(2)), (s[3], Index(2))) + @test supportof(tn, o_partial) == Set([3]) + + # A link index is carried by both endpoints of its edge. + o_link = operator(randn(2, 2), (Index(2),), (l[first(edges(g))],)) + @test_throws ArgumentError supportof(tn, o_link) + end + @testset "`subgraph`" begin g = named_grid((3,)) From 0eddf9185b22efad36ff0a1d92494f2ee3199bb8 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 14 Sep 2026 13:32:51 -0400 Subject: [PATCH 10/22] Rename function `supportof` to `operator_support`. --- src/abstracttensornetwork.jl | 2 +- src/apply/apply_operators.jl | 2 +- test/test_tensornetwork.jl | 14 +++++++------- 3 files changed, 9 insertions(+), 9 deletions(-) diff --git a/src/abstracttensornetwork.jl b/src/abstracttensornetwork.jl index 071ba714..92e899fe 100644 --- a/src/abstracttensornetwork.jl +++ b/src/abstracttensornetwork.jl @@ -132,7 +132,7 @@ function insertlink!(tn::AbstractGraph, e) return tn end -function supportof(tn::AbstractGraph, op::ITensorOperator) +function operator_support(tn::AbstractGraph, op::ITensorOperator) support = Set{vertextype(tn)}() for name in inputnames(op) diff --git a/src/apply/apply_operators.jl b/src/apply/apply_operators.jl index 212cae45..4615ba82 100644 --- a/src/apply/apply_operators.jl +++ b/src/apply/apply_operators.jl @@ -209,7 +209,7 @@ function apply_gate_bp!( dest::AbstractITensorNetwork, op::AbstractITensor, state::AbstractITensorNetwork, env; kwargs... ) - vertices = supportof(state, op) + vertices = operator_support(state, op) isempty(vertices) && throw( ArgumentError("operator shares no indices with the tensor network") diff --git a/test/test_tensornetwork.jl b/test/test_tensornetwork.jl index 29c600a8..3443f4af 100644 --- a/test/test_tensornetwork.jl +++ b/test/test_tensornetwork.jl @@ -4,7 +4,7 @@ using Graphs: add_edge!, add_vertex!, dst, edges, edgetype, has_edge, has_vertex is_directed, ne, nv, rem_edge!, rem_vertex!, src, vertices using ITensorBase: Index, LazyITensor, inds, operator using ITensorNetworksNext: ITensorNetwork, has_ind, linkaxes, linkinds, linknames, siteaxes, - siteinds, sitenames, supportof, tensornetwork + siteinds, sitenames, operator_support, tensornetwork using NamedGraphs: convert_vertextype, incident_edges, named_grid, named_path_graph, similar_graph, subgraph, vertextype using Test: @test, @test_throws, @testset @@ -125,7 +125,7 @@ using Test: @test, @test_throws, @testset @test sitenames(tn, 3) == [s[3].name] end - @testset "`supportof`" begin + @testset "`operator_support`" begin g = named_path_graph(3) l = Dict(e => Index(2) for e in edges(g)) l = merge(l, Dict(reverse(e) => l[e] for e in edges(g))) @@ -136,21 +136,21 @@ using Test: @test, @test_throws, @testset end o1 = operator(randn(2, 2), (Index(2),), (s[2],)) - @test supportof(tn, o1) == Set([2]) + @test operator_support(tn, o1) == Set([2]) o12 = operator(randn(2, 2, 2, 2), (Index(2), Index(2)), (s[1], s[2])) - @test supportof(tn, o12) == Set([1, 2]) + @test operator_support(tn, o12) == Set([1, 2]) # An input name that no tensor in the network carries contributes no vertex. o_absent = operator(randn(2, 2), (Index(2),), (Index(2),)) - @test isempty(supportof(tn, o_absent)) + @test isempty(operator_support(tn, o_absent)) o_partial = operator(randn(2, 2, 2, 2), (Index(2), Index(2)), (s[3], Index(2))) - @test supportof(tn, o_partial) == Set([3]) + @test operator_support(tn, o_partial) == Set([3]) # A link index is carried by both endpoints of its edge. o_link = operator(randn(2, 2), (Index(2),), (l[first(edges(g))],)) - @test_throws ArgumentError supportof(tn, o_link) + @test_throws ArgumentError operator_support(tn, o_link) end @testset "`subgraph`" begin From b013f7135c9853f14fb8a534bdee59fc31edefc2 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 14 Sep 2026 13:34:52 -0400 Subject: [PATCH 11/22] Fix variable names that refer to `inds` instead of the correct `names.` --- src/tensornetwork.jl | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/tensornetwork.jl b/src/tensornetwork.jl index 1f4992ae..c492f913 100644 --- a/src/tensornetwork.jl +++ b/src/tensornetwork.jl @@ -153,16 +153,16 @@ function update_tensornetwork_metadata!(tn, vertex, tensor) return tn end -function update_tensornetwork_metadata!(tn, vertex, oldinds, newinds) +function update_tensornetwork_metadata!(tn, vertex, oldnames, newnames) # Only have to deal with the indices that aren't shared. - for name in symdiff(oldinds, newinds) - if name in oldinds + for name in symdiff(oldnames, newnames) + if name in oldnames delete_ind_edge!(tn, name) delete_ind_vertex!(tn, name, vertex) continue end - # Now `ind` must be a new index that's not in `oldinds` + # Now `name` must be a new index that's not in `oldinds` vertex_list = get!(tn.dimname_vertices, name, Set()) if length(vertex_list) > 1 From f27144579123e750c6dd55e7a02b9f1e8bd08898 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 14 Sep 2026 13:38:57 -0400 Subject: [PATCH 12/22] Refactor and rename `sum_log_scalars` to `sumlog`. --- src/beliefpropagation/messagecache.jl | 22 ++++++++++++++++------ 1 file changed, 16 insertions(+), 6 deletions(-) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index cac0e709..ea2cdae0 100644 --- a/src/beliefpropagation/messagecache.jl +++ b/src/beliefpropagation/messagecache.jl @@ -175,11 +175,21 @@ function region_scalar(factors, messages, region) return mapreduce(vertex -> vertex_scalar(factors, messages, vertex), *, region) end -function sum_log_scalars(terms) - if any(t -> real(t) < 0, terms) - terms = complex.(terms) - end - return sum(log.(terms)) +# (log|∏terms|, sign(∏terms)) +function sumlogabs(terms) + + T = typeof(first(terms)) + + return mapreduce( + t -> (log(abs(t)), sign(t)), + ((d1, s1), (d2, s2)) -> (d1 + d2, s1 * s2), + terms; init = (zero(float(real(T))), one(T)) + ) +end + +function sumlog(terms) + d, s = sumlogabs(terms) + return s isa Real && s > 0 ? d : d + log(complex(s)) end # We need a graph structure here, so assume `factors` is a graph. @@ -191,7 +201,7 @@ function bethe_free_energy(factors, messages) return -Inf end - return sum_log_scalars(numerator_terms) - sum_log_scalars(denominator_terms) + return sumlog(numerator_terms) - sumlog(denominator_terms) end # ===================================== NormNetwork ====================================== # From caf406a5b306e6c892ab29c728b2ec4e8ffa8472 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 14 Sep 2026 13:42:24 -0400 Subject: [PATCH 13/22] Replace `check_incoming_dimnames` with `check_input` and function dispatch Convention from `MatrixAlgebraKit`. --- src/tensornetwork.jl | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/tensornetwork.jl b/src/tensornetwork.jl index c492f913..60175eb7 100644 --- a/src/tensornetwork.jl +++ b/src/tensornetwork.jl @@ -115,7 +115,7 @@ DataGraphs.is_edge_assigned(::ITensorNetwork, _edge) = false DataGraphs.get_vertex_data(tn::ITensorNetwork, v) = tn.tensors[v] -function check_incoming_dimnames(tn, tensor, vertex) +function check_input(::typeof(set_vertex_data!), tn, tensor, vertex) for name in names(tensor) vertices = get(tn.dimname_vertices, name, Set()) if length(setdiff(vertices, Set([vertex]))) > 1 @@ -130,7 +130,7 @@ function check_incoming_dimnames(tn, tensor, vertex) end function DataGraphs.insert_vertex_data!(tn::ITensorNetwork, vertex, tensor) - check_incoming_dimnames(tn, tensor, vertex) + check_input(set_vertex_data!, tn, tensor, vertex) add_vertex!(tn.underlying_graph, vertex) update_tensornetwork_metadata!(tn, vertex, tensor) insert!(tn.tensors, vertex, tensor) @@ -138,7 +138,7 @@ function DataGraphs.insert_vertex_data!(tn::ITensorNetwork, vertex, tensor) end function DataGraphs.set_vertex_data!(tn::ITensorNetwork, tensor, vertex) - check_incoming_dimnames(tn, tensor, vertex) + check_input(set_vertex_data!, tn, tensor, vertex) update_tensornetwork_metadata!(tn, vertex, tensor) set!(tn.tensors, vertex, tensor) return tn From 59291bf80cab9d9e48e8c75af07611fa43d72a90 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 14 Sep 2026 13:42:52 -0400 Subject: [PATCH 14/22] Formatting --- test/test_tensornetwork.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/test/test_tensornetwork.jl b/test/test_tensornetwork.jl index 3443f4af..b9e318b9 100644 --- a/test/test_tensornetwork.jl +++ b/test/test_tensornetwork.jl @@ -3,8 +3,8 @@ using DataGraphs: using Graphs: add_edge!, add_vertex!, dst, edges, edgetype, has_edge, has_vertex, is_directed, ne, nv, rem_edge!, rem_vertex!, src, vertices using ITensorBase: Index, LazyITensor, inds, operator -using ITensorNetworksNext: ITensorNetwork, has_ind, linkaxes, linkinds, linknames, siteaxes, - siteinds, sitenames, operator_support, tensornetwork +using ITensorNetworksNext: ITensorNetwork, has_ind, linkaxes, linkinds, linknames, + operator_support, siteaxes, siteinds, sitenames, tensornetwork using NamedGraphs: convert_vertextype, incident_edges, named_grid, named_path_graph, similar_graph, subgraph, vertextype using Test: @test, @test_throws, @testset From 40d030bf490af7f1684c845f839461f44640dd4a Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 14 Sep 2026 14:32:40 -0400 Subject: [PATCH 15/22] `finalize_substate!` now takes both solve and subsolve objects. --- .../AlgorithmsInterfaceExtensions.jl | 8 ++++++-- src/beliefpropagation/messagecache.jl | 1 - test/test_algorithmsinterfaceextensions.jl | 2 +- 3 files changed, 7 insertions(+), 4 deletions(-) diff --git a/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl b/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl index 917fef9a..7863604f 100644 --- a/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl +++ b/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl @@ -20,7 +20,8 @@ function initialize_subsolve( end function finalize_substate!( - problem::AI.Problem, algorithm::AI.Algorithm, state::AI.State, substate::AI.State + _problem::AI.Problem, _algorithm::AI.Algorithm, state::AI.State, + _subproblem::AI.Problem, _subalgorithm::AI.Algorithm, substate::AI.State ) state.iterate = substate.iterate return state @@ -29,7 +30,10 @@ end function AI.step!(problem::AI.Problem, algorithm::NestedAlgorithm, state::AI.State) subproblem, subalgorithm, substate = initialize_subsolve(problem, algorithm, state) AI.solve!(subproblem, subalgorithm, substate) - finalize_substate!(subproblem, subalgorithm, substate, state) + finalize_substate!( + problem, algorithm, state, + subproblem, subalgorithm, substate + ) return state end diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index ea2cdae0..268a22c0 100644 --- a/src/beliefpropagation/messagecache.jl +++ b/src/beliefpropagation/messagecache.jl @@ -177,7 +177,6 @@ end # (log|∏terms|, sign(∏terms)) function sumlogabs(terms) - T = typeof(first(terms)) return mapreduce( diff --git a/test/test_algorithmsinterfaceextensions.jl b/test/test_algorithmsinterfaceextensions.jl index 290ebb8a..e6b816a2 100644 --- a/test/test_algorithmsinterfaceextensions.jl +++ b/test/test_algorithmsinterfaceextensions.jl @@ -116,7 +116,7 @@ end # `finalize_substate!` copies the substate's iterate back into the # parent state. substate = AI.initialize_state(problem, algorithm; iterate = [42.0]) - AIE.finalize_substate!(problem, algorithm, state, substate) + AIE.finalize_substate!(problem, algorithm, state, problem, algorithm, substate) @test state.iterate == [42.0] end From f6b23ecbe7d80657ac79fe5faf99f33082e2a2a9 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 14 Sep 2026 17:52:17 -0400 Subject: [PATCH 16/22] Fix type stabilty in `vertex/edge_scalars`; `sumlogabs` now infers eltype from param --- src/beliefpropagation/messagecache.jl | 19 ++++++++----------- 1 file changed, 8 insertions(+), 11 deletions(-) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index 268a22c0..80e5a4ff 100644 --- a/src/beliefpropagation/messagecache.jl +++ b/src/beliefpropagation/messagecache.jl @@ -142,8 +142,10 @@ vertex_scalars(factors, messages) = vertex_scalars(factors, messages, keys(facto function vertex_scalars(factors::AbstractGraph, messages) return vertex_scalars(factors, messages, vertices(factors)) end +# `vertex_scalar` reads a number out of an `ITensor`, whose array field is untyped, so `map` would +# give element type `Any`; collecting the values instead picks up the type they actually have. function vertex_scalars(factors, messages, vertices) - return map(v -> vertex_scalar(factors, messages, v), vertices) + return [vertex_scalar(factors, messages, v) for v in vertices] end function edge_scalar(cache, edge) @@ -153,22 +155,17 @@ end edge_scalars(cache) = edge_scalars(cache, keys(cache)) function edge_scalars(cache, edges) - processed = Set{eltype(edges)}() - - T = Base.promote_op(edge_scalar, typeof(cache), eltype(edges)) - - scalars = T[] + unique_edges = Indices{eltype(edges)}() # Ignore repeated edges and their reverses. for e in edges - if e in processed || reverse(e) in processed + if e in unique_edges || reverse(e) in unique_edges continue end - push!(processed, e) - push!(scalars, edge_scalar(cache, e)) + insert!(unique_edges, e) end - return scalars + return [edge_scalar(cache, e) for e in unique_edges] end function region_scalar(factors, messages, region) @@ -177,7 +174,7 @@ end # (log|∏terms|, sign(∏terms)) function sumlogabs(terms) - T = typeof(first(terms)) + T = eltype(terms) return mapreduce( t -> (log(abs(t)), sign(t)), From 9661ace5240ebe3e13965f12c3d937140190582e Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Thu, 17 Sep 2026 13:01:04 -0400 Subject: [PATCH 17/22] Misc updates and fixes. --- src/ITensorNetworksNext.jl | 1 + src/abstracttensornetwork.jl | 2 +- src/beliefpropagation/messagecache.jl | 21 ++++++++++----------- src/tensornetwork.jl | 3 ++- src/utils.jl | 6 ++++++ test/test_tensornetwork.jl | 6 +++--- 6 files changed, 23 insertions(+), 16 deletions(-) create mode 100644 src/utils.jl diff --git a/src/ITensorNetworksNext.jl b/src/ITensorNetworksNext.jl index 33a93edb..fe1b2540 100644 --- a/src/ITensorNetworksNext.jl +++ b/src/ITensorNetworksNext.jl @@ -8,6 +8,7 @@ if VERSION >= v"1.11.0-DEV.469" ) end +include("utils.jl") include("select_algorithm.jl") include("AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl") include("abstracttensornetwork.jl") diff --git a/src/abstracttensornetwork.jl b/src/abstracttensornetwork.jl index 92e899fe..045b0793 100644 --- a/src/abstracttensornetwork.jl +++ b/src/abstracttensornetwork.jl @@ -133,7 +133,7 @@ function insertlink!(tn::AbstractGraph, e) end function operator_support(tn::AbstractGraph, op::ITensorOperator) - support = Set{vertextype(tn)}() + support = Indices{vertextype(tn)}() for name in inputnames(op) vertices = dimnamevertices(tn, name) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index 80e5a4ff..96d31ebc 100644 --- a/src/beliefpropagation/messagecache.jl +++ b/src/beliefpropagation/messagecache.jl @@ -4,8 +4,8 @@ using DataGraphs: DataGraphs, AbstractDataGraph, AbstractEdgeDataGraph, edge_dat using Dictionaries: Dictionary, getindices, set!, unset! using Graphs: AbstractGraph, connected_components, is_directed, is_tree using ITensorBase: state, unnamed -using NamedGraphs: AbstractNamedEdge, NamedDiGraph, NamedEdge, add_edges!, boundary_edges, - in_incident_edges, to_graph_index, vertextype +using NamedGraphs: AbstractNamedEdge, NamedDiGraph, NamedEdge, add_edges!, arrange_edge, + boundary_edges, in_incident_edges, to_graph_index, vertextype using SplitApplyCombine: mapmany struct MessageCache{T, V} <: AbstractEdgeDataGraph{T, V} @@ -145,7 +145,7 @@ end # `vertex_scalar` reads a number out of an `ITensor`, whose array field is untyped, so `map` would # give element type `Any`; collecting the values instead picks up the type they actually have. function vertex_scalars(factors, messages, vertices) - return [vertex_scalar(factors, messages, v) for v in vertices] + return narrow_map(v -> vertex_scalar(factors, messages, v), vertices) end function edge_scalar(cache, edge) @@ -155,17 +155,16 @@ end edge_scalars(cache) = edge_scalars(cache, keys(cache)) function edge_scalars(cache, edges) - unique_edges = Indices{eltype(edges)}() + seen = Indices{edgetype(cache)}() - # Ignore repeated edges and their reverses. - for e in edges - if e in unique_edges || reverse(e) in unique_edges - continue - end - insert!(unique_edges, e) + unique_edges = filter(edges) do edge + arranged = arrange_edge(edgetype(cache)(edge)) + arranged in seen && return false + insert!(seen, arranged) + return true end - return [edge_scalar(cache, e) for e in unique_edges] + return narrow_map(e -> edge_scalar(cache, e), unique_edges) end function region_scalar(factors, messages, region) diff --git a/src/tensornetwork.jl b/src/tensornetwork.jl index 60175eb7..24ee5a64 100644 --- a/src/tensornetwork.jl +++ b/src/tensornetwork.jl @@ -38,6 +38,7 @@ function ITensorNetwork{T, V}(tensors) where {T, V} return tn end +ITensorBase.nametype(tn::ITensorNetwork) = nametype(typeof(tn)) ITensorBase.nametype(::Type{<:ITensorNetwork{T, V, I}}) where {T, V, I} = I Graphs.vertices(tn::ITensorNetwork) = vertices(tn.underlying_graph) @@ -145,7 +146,7 @@ function DataGraphs.set_vertex_data!(tn::ITensorNetwork, tensor, vertex) end function update_tensornetwork_metadata!(tn, vertex, tensor) - oldnames = isassigned(tn, vertex) ? names(tn[vertex]) : Set() + oldnames = isassigned(tn, vertex) ? names(tn[vertex]) : Set{nametype(tn)}() newnames = names(tensor) update_tensornetwork_metadata!(tn, vertex, oldnames, newnames) diff --git a/src/utils.jl b/src/utils.jl new file mode 100644 index 00000000..abf666e4 --- /dev/null +++ b/src/utils.jl @@ -0,0 +1,6 @@ +using Dictionaries: AbstractIndices, Dictionary + +# `map` does something clever to figure out the element type of the output when not +# inferable, but the `map` overload on `Dictionary` does not, so we fix this here. +narrow_map(f, v) = map(f, v) +narrow_map(f, v::AbstractIndices) = Dictionary(v, [f(x) for x in v]) diff --git a/test/test_tensornetwork.jl b/test/test_tensornetwork.jl index b9e318b9..d9639475 100644 --- a/test/test_tensornetwork.jl +++ b/test/test_tensornetwork.jl @@ -136,17 +136,17 @@ using Test: @test, @test_throws, @testset end o1 = operator(randn(2, 2), (Index(2),), (s[2],)) - @test operator_support(tn, o1) == Set([2]) + @test issetequal(operator_support(tn, o1), Set([2])) o12 = operator(randn(2, 2, 2, 2), (Index(2), Index(2)), (s[1], s[2])) - @test operator_support(tn, o12) == Set([1, 2]) + @test issetequal(operator_support(tn, o12), Set([1, 2])) # An input name that no tensor in the network carries contributes no vertex. o_absent = operator(randn(2, 2), (Index(2),), (Index(2),)) @test isempty(operator_support(tn, o_absent)) o_partial = operator(randn(2, 2, 2, 2), (Index(2), Index(2)), (s[3], Index(2))) - @test operator_support(tn, o_partial) == Set([3]) + @test issetequal(operator_support(tn, o_partial), Set([3])) # A link index is carried by both endpoints of its edge. o_link = operator(randn(2, 2), (Index(2),), (l[first(edges(g))],)) From 70c15e9923ff9583e3bf72e93edfef9588036ca8 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Thu, 17 Sep 2026 13:39:14 -0400 Subject: [PATCH 18/22] Filter edges in `edge_scalars(cache)` method instead of method that takes edges directly. --- src/beliefpropagation/messagecache.jl | 19 ++++++++++--------- 1 file changed, 10 insertions(+), 9 deletions(-) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index 96d31ebc..760b04b8 100644 --- a/src/beliefpropagation/messagecache.jl +++ b/src/beliefpropagation/messagecache.jl @@ -152,21 +152,22 @@ function edge_scalar(cache, edge) return (cache[edge] * cache[reverse(edge)])[] end -edge_scalars(cache) = edge_scalars(cache, keys(cache)) +function edge_scalars(cache) + seen = Indices{keytype(cache)}() -function edge_scalars(cache, edges) - seen = Indices{edgetype(cache)}() - - unique_edges = filter(edges) do edge - arranged = arrange_edge(edgetype(cache)(edge)) - arranged in seen && return false - insert!(seen, arranged) + unique_edges = filter(keys(cache)) do edge + if edge in seen || reverse(edge) in seen + return false + end + insert!(seen, edge) return true end - return narrow_map(e -> edge_scalar(cache, e), unique_edges) + return edge_scalars(cache, unique_edges) end +edge_scalars(cache, edges) = narrow_map(e -> edge_scalar(cache, e), edges) + function region_scalar(factors, messages, region) return mapreduce(vertex -> vertex_scalar(factors, messages, vertex), *, region) end From 9f452cae1f9cd23fbfbd4fa19cf3e23b76f8490a Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 28 Sep 2026 10:29:35 -0400 Subject: [PATCH 19/22] Rename `bethe_free_energy` to `bethe_free_entropy`; add `bethe_free_energy` as its negative. The function returns the BP estimate of log Z, which is the free entropy; the free energy is -log Z. Adds tests for the relation between the two and for a vanishing edge overlap. Co-Authored-By: Claude Opus 5.5 --- src/beliefpropagation/messagecache.jl | 3 ++- test/test_beliefpropagation.jl | 29 ++++++++++++++++++++------- 2 files changed, 24 insertions(+), 8 deletions(-) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index 760b04b8..d91819d0 100644 --- a/src/beliefpropagation/messagecache.jl +++ b/src/beliefpropagation/messagecache.jl @@ -189,7 +189,7 @@ function sumlog(terms) end # We need a graph structure here, so assume `factors` is a graph. -function bethe_free_energy(factors, messages) +function bethe_free_entropy(factors, messages) numerator_terms = vertex_scalars(factors, messages) denominator_terms = edge_scalars(messages) @@ -199,6 +199,7 @@ function bethe_free_energy(factors, messages) return sumlog(numerator_terms) - sumlog(denominator_terms) end +bethe_free_energy(factors, messages) = -bethe_free_entropy(factors, messages) # ===================================== NormNetwork ====================================== # diff --git a/test/test_beliefpropagation.jl b/test/test_beliefpropagation.jl index d759b51f..61d37e3a 100644 --- a/test/test_beliefpropagation.jl +++ b/test/test_beliefpropagation.jl @@ -9,9 +9,9 @@ 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, bratensor, contract_network, contraction_order, edge_scalar, - factor_tensors, incoming_messages, insertlink!, kettensor, linkaxes, linkinds, - message_environment, messagecache, region_scalar, subgraph, tensornetwork, + bethe_free_energy, bethe_free_entropy, 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, @@ -145,6 +145,21 @@ end @test length(in_msgs) == 1 @test only(in_msgs) == bpc[3 => 2] end + @testset "Bethe free entropy and free energy" begin + g = named_path_graph(2) + l = Index(2) + tn = tensornetwork(v -> randn(l), vertices(g)) + + bpc = messagecache(edge -> ones(Tuple(linkinds(tn, edge))), all_edges(g)) + @test bethe_free_energy(tn, bpc) == -bethe_free_entropy(tn, bpc) + + bpc = messagecache(all_edges(g)) do edge + return edge == NamedEdge(1 => 2) ? [1.0, 0.0][l] : [0.0, 1.0][l] + end + @test iszero(edge_scalar(bpc, 1 => 2)) + @test bethe_free_entropy(tn, bpc) == -Inf + @test bethe_free_energy(tn, bpc) == Inf + end @testset "subgraph" begin g = named_grid((3,)) @@ -205,7 +220,7 @@ end cache = beliefpropagation( tn, messages; stopping_criterion = (; maxiter = 1) ) - z_bp = exp(bethe_free_energy(tn, cache)) + z_bp = exp(bethe_free_entropy(tn, cache)) z_exact = reduce(*, [tn[v] for v in vertices(g)])[] @test z_bp ≈ z_exact rtol = eps(real(T))^(1 / 3) @@ -227,7 +242,7 @@ end cache = beliefpropagation( tn, messages; stopping_criterion = (; maxiter = 1) ) - z_bp = exp(bethe_free_energy(tn, cache)) + z_bp = exp(bethe_free_entropy(tn, cache)) z_exact = reduce(*, [tn[v] for v in vertices(g)])[] @test z_bp ≈ z_exact rtol = eps(real(T))^(1 / 3) @@ -248,7 +263,7 @@ end stopping_criterion = (; maxiter = 10, tol = 1.0e-10) ) - z_bp = exp(bethe_free_energy(tn, cache)) + z_bp = exp(bethe_free_entropy(tn, cache)) @test z_bp ≈ 1.5^(n^2) end @@ -287,7 +302,7 @@ end # Belief propagation is exact on a tree, including on the fermionic norm network. ket = prod(network) z_exact = (ket * conj(ket))[] - z_bp = exp(bethe_free_energy(nn, cache)) + z_bp = exp(bethe_free_entropy(nn, cache)) @test z_bp ≈ z_exact rtol = eps(real(T))^(1 / 3) for edge in edges(cache) From a3a7f522e113d9c29405cc6252425f50c54ddc91 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 28 Sep 2026 10:36:54 -0400 Subject: [PATCH 20/22] Functions `edge_scalar` and `edge_scalars` now take `factors` as first arg. This is for consistency with `vertex_scalars`. --- src/beliefpropagation/messagecache.jl | 25 +++++------------ test/test_beliefpropagation.jl | 39 ++++++++++++++++++++++----- 2 files changed, 38 insertions(+), 26 deletions(-) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index d91819d0..650a5858 100644 --- a/src/beliefpropagation/messagecache.jl +++ b/src/beliefpropagation/messagecache.jl @@ -148,26 +148,13 @@ function vertex_scalars(factors, messages, vertices) return narrow_map(v -> vertex_scalar(factors, messages, v), vertices) end -function edge_scalar(cache, edge) - return (cache[edge] * cache[reverse(edge)])[] +# Takes factors as an unused argument for consistency with `vertex_scalar`. +edge_scalar(_factors, messages, edge) = (messages[edge] * messages[reverse(edge)])[] +edge_scalars(factors, messages) = edge_scalars(factors, messages, edges(factors)) +function edge_scalars(factors, messages, edges) + return narrow_map(e -> edge_scalar(factors, messages, e), edges) end -function edge_scalars(cache) - seen = Indices{keytype(cache)}() - - unique_edges = filter(keys(cache)) do edge - if edge in seen || reverse(edge) in seen - return false - end - insert!(seen, edge) - return true - end - - return edge_scalars(cache, unique_edges) -end - -edge_scalars(cache, edges) = narrow_map(e -> edge_scalar(cache, e), edges) - function region_scalar(factors, messages, region) return mapreduce(vertex -> vertex_scalar(factors, messages, vertex), *, region) end @@ -191,7 +178,7 @@ end # We need a graph structure here, so assume `factors` is a graph. function bethe_free_entropy(factors, messages) numerator_terms = vertex_scalars(factors, messages) - denominator_terms = edge_scalars(messages) + denominator_terms = edge_scalars(factors, messages) if any(iszero, denominator_terms) return -Inf diff --git a/test/test_beliefpropagation.jl b/test/test_beliefpropagation.jl index 61d37e3a..f838f5d3 100644 --- a/test/test_beliefpropagation.jl +++ b/test/test_beliefpropagation.jl @@ -3,16 +3,16 @@ using Base.Broadcast: materialize using DataGraphs: DataGraphs, DataGraph, edge_data, edge_data_type using Dictionaries: Dictionary, dictionary, set! using GradedArrays: U1, gradedrange, isdual -using Graphs: AbstractGraph, add_vertex!, dst, edges, has_edge, has_vertex, nv, rem_edge!, - src, vertices +using Graphs: AbstractGraph, add_vertex!, dst, edges, has_edge, has_vertex, ne, nv, + rem_edge!, src, vertices 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, bethe_free_entropy, 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 + edge_scalar, edge_scalars, 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 @@ -128,7 +128,7 @@ end # Vertex/edge/region scalars. @test vertex_scalar(tn, bpc, 2) isa ComplexF64 - @test edge_scalar(bpc, 1 => 2) isa Float64 + @test edge_scalar(tn, bpc, 1 => 2) isa Float64 @test region_scalar(tn, bpc, [1]) == vertex_scalar(tn, bpc, 1) @test region_scalar(tn, bpc, [2, 3]) == prod(vertex_scalars(tn, bpc, [2, 3])) @@ -145,6 +145,31 @@ end @test length(in_msgs) == 1 @test only(in_msgs) == bpc[3 => 2] end + @testset "Edge scalars" begin + rng = StableRNG(123) + g = named_path_graph(3) + l = Dict(e => Index(2) for e in edges(g)) + l = merge(l, Dict(reverse(e) => l[e] for e in edges(g))) + + tn = tensornetwork(vertices(g)) do v + is = map(e -> l[e], incident_edges(g, v)) + return randn(rng, Tuple(is)) + end + + bpc = messagecache(all_edges(g)) do edge + return randn(rng, Tuple(linkinds(tn, edge))) + end + + @test edge_scalar(tn, bpc, 1 => 2) == (bpc[1 => 2] * bpc[2 => 1])[] + @test edge_scalar(tn, bpc, 1 => 2) ≈ edge_scalar(tn, bpc, 2 => 1) + + scalars = edge_scalars(tn, bpc) + @test scalars isa Vector{Float64} + @test length(scalars) == ne(tn) + @test scalars == map(e -> edge_scalar(tn, bpc, e), edges(tn)) + + @test edge_scalars(tn, bpc, [2 => 3]) == [edge_scalar(tn, bpc, 2 => 3)] + end @testset "Bethe free entropy and free energy" begin g = named_path_graph(2) l = Index(2) @@ -156,7 +181,7 @@ end bpc = messagecache(all_edges(g)) do edge return edge == NamedEdge(1 => 2) ? [1.0, 0.0][l] : [0.0, 1.0][l] end - @test iszero(edge_scalar(bpc, 1 => 2)) + @test iszero(edge_scalar(tn, bpc, 1 => 2)) @test bethe_free_entropy(tn, bpc) == -Inf @test bethe_free_energy(tn, bpc) == Inf end From 8a9b645f44eba2119cb6a2b13eba6849e176acd0 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 28 Sep 2026 10:47:40 -0400 Subject: [PATCH 21/22] Trim unsued import. --- src/abstracttensornetwork.jl | 4 ++-- src/apply/apply_operators.jl | 3 +-- 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/src/abstracttensornetwork.jl b/src/abstracttensornetwork.jl index 045b0793..341d7a6d 100644 --- a/src/abstracttensornetwork.jl +++ b/src/abstracttensornetwork.jl @@ -4,8 +4,8 @@ using DataGraphs: DataGraphs, AbstractDataGraph, AbstractVertexDataGraph, edge_d using Dictionaries: Dictionary using Graphs: Graphs, AbstractEdge, AbstractGraph, add_edge!, add_vertex!, dst, edges, edgetype, ne, neighbors, nv, rem_edge!, src, vertices -using ITensorBase: ITensorOperator, NamedUnitRange, inds, inputnames, name, names, - nametype, prime, uniquename, unnamedtype +using ITensorBase: ITensorOperator, NamedUnitRange, inds, inputnames, name, names, nametype, + prime, uniquename, unnamedtype using LinearAlgebra: LinearAlgebra using MacroTools: @capture using NamedGraphs: diff --git a/src/apply/apply_operators.jl b/src/apply/apply_operators.jl index 4615ba82..0682df57 100644 --- a/src/apply/apply_operators.jl +++ b/src/apply/apply_operators.jl @@ -2,8 +2,7 @@ using .AlgorithmsInterfaceExtensions: AlgorithmsInterfaceExtensions as AIE using AlgorithmsInterface: AlgorithmsInterface as AI using Base: @kwdef using Graphs: dst, src, vertices -using ITensorBase: ITensorBase as ITB, AbstractITensor, apply, inputnames, names, operator, - outputnames, rename +using ITensorBase: AbstractITensor, apply, names, operator, rename using LinearAlgebra: norm, normalize! using MatrixAlgebraKit: project_hermitian, qr_compact, svd_trunc using NamedGraphs: boundary_edges From 41eb83fd6655b3d2436d3873375178c141412e98 Mon Sep 17 00:00:00 2001 From: Jack Dunham Date: Mon, 28 Sep 2026 14:48:08 -0400 Subject: [PATCH 22/22] Minor version bump to 0.10.6. --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Project.toml b/Project.toml index fa2c303a..4aafb8c9 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "ITensorNetworksNext" uuid = "302f2e75-49f0-4526-aef7-d8ba550cb06c" -version = "0.10.5" +version = "0.10.6" authors = ["ITensor developers and contributors"] [workspace]