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] diff --git a/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl b/src/AlgorithmsInterfaceExtensions/AlgorithmsInterfaceExtensions.jl index da7bf76a..7863604f 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 @@ -18,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 @@ -27,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!(problem, algorithm, state, substate) + finalize_substate!( + problem, algorithm, state, + subproblem, subalgorithm, substate + ) return state end 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 9515d084..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: - NamedUnitRange, inds, 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: @@ -131,3 +131,23 @@ function insertlink!(tn::AbstractGraph, e) return tn end + +function operator_support(tn::AbstractGraph, op::ITensorOperator) + support = Indices{vertextype(tn)}() + + for name in inputnames(op) + vertices = dimnamevertices(tn, name) + + if length(vertices) > 1 + throw( + ArgumentError( + "operator dim name $name associated with multiple vertices in tensor network." + ) + ) + end + + union!(support, vertices) + end + + return support +end diff --git a/src/apply/apply_operators.jl b/src/apply/apply_operators.jl index e029bcf9..0682df57 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: AbstractITensor, apply, 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 = operator_support(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,29 +228,29 @@ 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 - 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 - 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 - 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 = @@ -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) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index c7c66c4f..650a5858 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} @@ -142,57 +142,51 @@ 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 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 -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)) +function region_scalar(factors, messages, region) + return mapreduce(vertex -> vertex_scalar(factors, messages, vertex), *, region) +end - scalars = T[] +# (log|∏terms|, sign(∏terms)) +function sumlogabs(terms) + T = eltype(terms) - # Ignore repeated edges and their reverses. - for e in edges - if e in processed || reverse(e) in processed - continue - end - push!(processed, e) - push!(scalars, edge_scalar(cache, e)) - end - - return scalars + 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 region_scalar(factors, messages, region) - return mapreduce(vertex -> vertex_scalar(factors, messages, vertex), *, region) +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. -function bethe_free_energy(factors, messages) +function bethe_free_entropy(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 + denominator_terms = edge_scalars(factors, messages) if any(iszero, denominator_terms) return -Inf end - return sum(log.(numerator_terms)) - sum(log.(denominator_terms)) + return sumlog(numerator_terms) - sumlog(denominator_terms) end +bethe_free_energy(factors, messages) = -bethe_free_entropy(factors, messages) # ===================================== NormNetwork ====================================== # diff --git a/src/tensornetwork.jl b/src/tensornetwork.jl index a2216cb7..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) @@ -115,38 +116,60 @@ DataGraphs.is_edge_assigned(::ITensorNetwork, _edge) = false DataGraphs.get_vertex_data(tn::ITensorNetwork, v) = tn.tensors[v] +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 + 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_input(set_vertex_data!, 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_input(set_vertex_data!, 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{nametype(tn)}() + 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, oldnames, newnames) # 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(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, 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 +182,6 @@ function set!_tensornetwork(tn::ITensorNetwork, vertex, tensor) end end - set!(tn.tensors, vertex, tensor) - return tn end @@ -170,20 +191,14 @@ 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] +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) 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_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 diff --git a/test/test_beliefpropagation.jl b/test/test_beliefpropagation.jl index d759b51f..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, 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 + bethe_free_energy, bethe_free_entropy, bratensor, contract_network, contraction_order, + 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,46 @@ 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) + 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(tn, 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 +245,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 +267,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 +288,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 +327,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) diff --git a/test/test_tensornetwork.jl b/test/test_tensornetwork.jl index ab333b15..d9639475 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 ITensorNetworksNext: ITensorNetwork, has_ind, linkaxes, linkinds, linknames, siteaxes, - siteinds, sitenames, tensornetwork +using ITensorBase: Index, LazyITensor, inds, operator +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 @@ -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)) @@ -119,6 +125,34 @@ using Test: @test, @test_throws, @testset @test sitenames(tn, 3) == [s[3].name] end + @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))) + 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 issetequal(operator_support(tn, o1), Set([2])) + + o12 = operator(randn(2, 2, 2, 2), (Index(2), Index(2)), (s[1], s[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 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))],)) + @test_throws ArgumentError operator_support(tn, o_link) + end + @testset "`subgraph`" begin g = named_grid((3,))