Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
22 commits
Select commit Hold shift + click to select a range
445f3c1
`add_edge!` and `rem_edge!` on the `ITensorNetwork` type now return f…
jack-dunham Jul 9, 2026
4dc749b
`dimnamevertices` for `ITensorNetwork` now returns empty set if index…
jack-dunham Jul 9, 2026
0b7e87d
New function `supportof` that gives the vertex support of an operator…
jack-dunham Jul 9, 2026
4bf1763
Refactor `ITensorNetwork` topology functions; add index precheck for …
jack-dunham Jul 9, 2026
a05bbb3
Refactor sum of log scalars into own function.
jack-dunham Jul 20, 2026
0e3a8c0
The `finalize_substate` function now dispatches on subsolve rather th…
jack-dunham Jul 20, 2026
26e9586
Upgrade to ITensorBase v0.13
jack-dunham Jul 21, 2026
973dcea
Fix rename of `vs` to `vertices` in function body.
jack-dunham Sep 8, 2026
abeaeca
Function `supportof` can now return an empty set; add tests.
jack-dunham Sep 8, 2026
0eddf91
Rename function `supportof` to `operator_support`.
jack-dunham Sep 14, 2026
b013f71
Fix variable names that refer to `inds` instead of the correct `names.`
jack-dunham Sep 14, 2026
f271445
Refactor and rename `sum_log_scalars` to `sumlog`.
jack-dunham Sep 14, 2026
caf406a
Replace `check_incoming_dimnames` with `check_input` and function dis…
jack-dunham Sep 14, 2026
59291bf
Formatting
jack-dunham Sep 14, 2026
40d030b
`finalize_substate!` now takes both solve and subsolve objects.
jack-dunham Sep 14, 2026
f6b23ec
Fix type stabilty in `vertex/edge_scalars`; `sumlogabs` now infers el…
jack-dunham Sep 14, 2026
9661ace
Misc updates and fixes.
jack-dunham Sep 17, 2026
70c15e9
Filter edges in `edge_scalars(cache)` method instead of method that t…
jack-dunham Sep 17, 2026
9f452ca
Rename `bethe_free_energy` to `bethe_free_entropy`; add `bethe_free_e…
jack-dunham Sep 28, 2026
a3a7f52
Functions `edge_scalar` and `edge_scalars` now take `factors` as firs…
jack-dunham Sep 28, 2026
8a9b645
Trim unsued import.
jack-dunham Sep 28, 2026
41eb83f
Minor version bump to 0.10.6.
jack-dunham Sep 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
name = "ITensorNetworksNext"
uuid = "302f2e75-49f0-4526-aef7-d8ba550cb06c"
version = "0.10.5"
version = "0.10.6"
authors = ["ITensor developers <support@itensor.org> and contributors"]

[workspace]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ module AlgorithmsInterfaceExtensions

using AlgorithmsInterface: AlgorithmsInterface as AI

abstract type NestedProblem <: AI.Problem end

# ============================ NestedAlgorithm =============================================

abstract type NestedAlgorithm <: AI.Algorithm end
Expand All @@ -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
Expand All @@ -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

Expand Down
1 change: 1 addition & 0 deletions src/ITensorNetworksNext.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
24 changes: 22 additions & 2 deletions src/abstracttensornetwork.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
37 changes: 20 additions & 17 deletions src/apply/apply_operators.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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!(
Expand All @@ -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 =
Expand All @@ -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)
Expand Down
62 changes: 28 additions & 34 deletions src/beliefpropagation/messagecache.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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)
Comment thread
jack-dunham marked this conversation as resolved.
end
bethe_free_energy(factors, messages) = -bethe_free_entropy(factors, messages)

# ===================================== NormNetwork ====================================== #

Expand Down
69 changes: 42 additions & 27 deletions src/tensornetwork.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -159,8 +182,6 @@ function set!_tensornetwork(tn::ITensorNetwork, vertex, tensor)
end
end

set!(tn.tensors, vertex, tensor)

return tn
end

Expand All @@ -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)}())

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this makes more sense, but was this inspired by a particular use case?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It is for consistency with the fallback method (which returns an empty set).

end

# PERF: fast lookup compared to `AbstractITensorNetwork` fallback.
has_dimname(tn::ITensorNetwork, name) = haskey(tn.dimname_vertices, name)
Expand Down
6 changes: 6 additions & 0 deletions src/utils.jl
Original file line number Diff line number Diff line change
@@ -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])
2 changes: 1 addition & 1 deletion test/test_algorithmsinterfaceextensions.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading
Loading