-
Notifications
You must be signed in to change notification settings - Fork 3
ITensorNetwork type fixes and other minor refactors. #175
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
445f3c1
4dc749b
0b7e87d
4bf1763
a05bbb3
0e3a8c0
26e9586
973dcea
abeaeca
0eddf91
b013f71
f271445
caf406a
59291bf
40d030b
f6b23ec
9661ace
70c15e9
9f452ca
a3a7f52
8a9b645
41eb83f
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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)}()) | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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) | ||
|
|
||
| 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]) |
Uh oh!
There was an error while loading. Please reload this page.