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/abstracttensornetwork.jl b/src/abstracttensornetwork.jl index 4ea4c34a..3a314df3 100644 --- a/src/abstracttensornetwork.jl +++ b/src/abstracttensornetwork.jl @@ -4,7 +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: dimnames, inds, name, named, nametype, prime, uniquename, unnamedtype +using ITensorBase: ITensorOperator, dimnames, inds, inputnames, name, named, nametype, + prime, uniquename, unnamedtype using LinearAlgebra: LinearAlgebra using MacroTools: @capture using NamedGraphs: @@ -130,3 +131,23 @@ function insertlink!(tn::AbstractGraph, e) return tn end + +function operator_support(tn::AbstractGraph, op::ITensorOperator) + support = Set{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 f824b9d7..f82528d5 100644 --- a/src/apply/apply_operators.jl +++ b/src/apply/apply_operators.jl @@ -2,9 +2,9 @@ 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 LinearAlgebra: norm +using ITensorBase: ITensorBase as ITB, AbstractITensor, apply, dimnames, inputnames, + operator, outputnames, replacedimnames +using LinearAlgebra: norm, normalize! using MatrixAlgebraKit: eigh_full, project_hermitian, qr_compact, svd_trunc using NamedGraphs: boundary_edges using TensorAlgebra.MatrixAlgebra: invsqrth_safe, sqrth_safe @@ -239,12 +239,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!( @@ -256,29 +259,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 = ITB.apply(op, state[v]) + vertex = only(vertices) + ψv = apply(op, state[vertex]) if normalize gauges = [ message_root(env[e]) - for e in boundary_edges(state, vs; dir = :in) + for e in boundary_edges(state, vertices; dir = :in) ] ψv /= norm(prod([[ψv]; gauges])) 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) 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)) @@ -289,11 +292,11 @@ function apply_gate_bp_nsite!( 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) - end + + normalize && normalize!(S) + name_v1, name_v2 = dimnames(S) sqrt_S = sqrth_safe(S, (name_v1,), (name_v2,); atol = 0, rtol = 0) R_v1 = replacedimnames(U_v1 * sqrt_S, name_v2 => name_v1) diff --git a/src/beliefpropagation/messagecache.jl b/src/beliefpropagation/messagecache.jl index fc6bc611..91f6aa23 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,45 +155,49 @@ 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) return mapreduce(vertex -> vertex_scalar(factors, messages, vertex), *, region) end +# (log|∏terms|, sign(∏terms)) +function sumlogabs(terms) + T = eltype(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. 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 sumlog(numerator_terms) - sumlog(denominator_terms) end # ===================================== NormNetwork ====================================== # diff --git a/src/tensornetwork.jl b/src/tensornetwork.jl index ba96462e..1cb3cbc9 100644 --- a/src/tensornetwork.jl +++ b/src/tensornetwork.jl @@ -116,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 dimnames(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 = dimnames(tensor) +function update_tensornetwork_metadata!(tn, vertex, tensor) + oldnames = isassigned(tn, vertex) ? dimnames(tn[vertex]) : Set() + newnames = dimnames(tensor) - oldinds = get(mapview(dimnames, 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 @@ -160,8 +182,6 @@ function set!_tensornetwork(tn::ITensorNetwork, vertex, tensor) end end - set!(tn.tensors, vertex, tensor) - return tn end @@ -171,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/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_tensornetwork.jl b/test/test_tensornetwork.jl index ab333b15..b9e318b9 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 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]) + + # 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]) + + # 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,))