11import ITensorNetworksNext. ITensorNetworksNextParallel as ITNNP
22using Dagger
3- using DataGraphs: DataGraphs, get_edge_data, get_vertex_data, is_edge_assigned,
3+ using DataGraphs: DataGraphs, edge_data, get_edge_data, get_vertex_data, is_edge_assigned,
44 is_vertex_assigned, set_edge_data!, set_vertex_data!, underlying_graph
5- using Dictionaries: Indices
5+ using Dictionaries: Dictionary, Indices, getindices
66using Graphs: AbstractEdge, AbstractGraph, dst, edges, src, vertices
77using ITensorNetworksNext: ITensorNetworksNext, BeliefPropagation, BeliefPropagationCache,
88 BeliefPropagationProblem, BeliefPropagationState, beliefpropagation,
9- forest_cover_edge_sequence, select_algorithm
9+ forest_cover_edge_sequence, select_algorithm, subcache
1010using NamedGraphs. GraphsExtensions: boundary_edges
1111using NamedGraphs. PartitionedGraphs: QuotientVertex, quotientedges, quotientvertices
1212using NamedGraphs: NamedGraphs
1313
14- function ITNNP. DaggerBeliefPropagationCache (network:: AbstractGraph )
14+ const DaggerBeliefPropagation = BeliefPropagation{<: ITNNP.DaggerNestedAlgorithm };
15+
16+ function ITNNP. DaggerBeliefPropagationCache (
17+ network:: AbstractGraph ;
18+ workers = nothing ,
19+ scopes = nothing
20+ )
1521 underlying_cache = BeliefPropagationCache (network)
1622
1723 keys = Indices (quotientvertices (underlying_cache))
1824
19- workers = Iterators . cycle (Dagger . Distributed . workers () )
20- worker_dict = similar (keys, Int)
25+ if isnothing (scopes )
26+ workers = isnothing (workers) ? Dagger . Distributed . workers () : workers
2127
22- for quotient_vertex in keys
23- worker, workers = Iterators. peel (workers)
24- worker_dict[quotient_vertex] = worker
28+ sorted_workers = Iterators. take (Iterators. cycle (workers), length (keys))
29+
30+ scopes = map (Dagger. ProcessScope, collect (sorted_workers))
31+ else
32+ if length (keys) != length (scopes)
33+ throw (
34+ ArgumentError (
35+ " Number of provided scopes must match the number of vertex partitions of underlying graph"
36+ )
37+ )
38+ end
2539 end
2640
41+ scope_dict = Dictionary (keys, scopes)
42+
2743 quotient_chunks = map (keys) do quotient_vertex
28- worker = worker_dict [quotient_vertex]
44+ scope = scope_dict [quotient_vertex]
2945 iterate = subcache (underlying_cache, quotient_vertex)
30- chunk = Dagger. @mutable worker = worker BeliefPropagationState (; iterate)
46+ chunk = Dagger. @mutable scope = scope BeliefPropagationState (; iterate)
3147 return chunk
3248 end
3349
@@ -42,7 +58,7 @@ function DataGraphs.is_vertex_assigned(bpc::ITNNP.DaggerBeliefPropagationCache,
4258 return is_vertex_assigned (bpc. underlying_cache, vertex)
4359end
4460function DataGraphs. is_edge_assigned (bpc:: ITNNP.DaggerBeliefPropagationCache , edge)
45- return is_edge_assigned (bpc. undelying_cache , edge)
61+ return is_edge_assigned (bpc. underlying_cache , edge)
4662end
4763
4864function DataGraphs. get_vertex_data (bpc:: ITNNP.DaggerBeliefPropagationCache , vertex)
@@ -52,7 +68,7 @@ function DataGraphs.get_edge_data(
5268 bpc:: ITNNP.DaggerBeliefPropagationCache ,
5369 edge:: AbstractEdge
5470 )
55- return get_edge_data (bpc. undelying_caches , edge)
71+ return get_edge_data (bpc. underlying_cache , edge)
5672end
5773
5874function DataGraphs. set_vertex_data! (bpc:: ITNNP.DaggerBeliefPropagationCache , val, vertex)
7389function ITensorNetworksNext. beliefpropagation_sweep (
7490 cache:: ITNNP.DaggerBeliefPropagationCache ;
7591 edges,
76- workers = workers (),
7792 kwargs...
7893 )
79- keys = collect (quotientvertices (cache))
80-
81- return ITNNP. dagger_algorithm (keys; keys, workers) do quotient_vertex
82- subcache = fetch (cache[quotient_vertex]). iterate
94+ return ITNNP. dagger_algorithm (quotientvertices (cache)) do quotient_vertex
95+ substate = fetch (cache[quotient_vertex])
96+ subcache = substate. iterate
8397
8498 subcache_edges = forest_cover_edge_sequence (subcache) ∩ edges
8599 incoming_edges = boundary_edges (cache, vertices (cache, quotient_vertex); dir = :in )
@@ -120,44 +134,61 @@ function ITNNP.get_subiterate(
120134end
121135
122136function AIE. set_substate! (
123- :: BeliefPropagationProblem ,
124- :: AIE.NestedAlgorithm ,
137+ problem :: BeliefPropagationProblem ,
138+ algorithm :: AIE.NestedAlgorithm ,
125139 state:: AIE.State ,
126140 substate:: ITNNP.DaggerState
127141 )
128142 dst_cache = state. iterate. iterate
129143
130144 state. iterate. maxdiff = 0.0
131145
132- for remote_result in substate. remote_results
133- get_maxdiff = dtask -> dtask. iterate. maxdiff
134- src_maxdiff = fetch (Dagger. @spawn get_maxdiff (remote_result))
135-
136- if src_maxdiff > state. iterate. maxdiff
137- state. iterate. maxdiff = src_maxdiff
138- end
146+ maxdiff_dtasks = map (substate. remote_results) do remote_result
147+ return Dagger. spawn (dtask -> dtask. iterate. maxdiff, remote_result)
139148 end
140149
141- function transfer_edges! (dst_chunk, src_chunk, edges)
142- src_subcache = src_chunk. iterate
143- dst_subcache = dst_chunk. iterate
144- for edge in edges
145- dst_subcache[edge] = src_subcache[edge]
146- end
147- return
150+ maxdiff = maximum (fetch, maxdiff_dtasks)
151+
152+ if maxdiff > state. iterate. maxdiff
153+ state. iterate. maxdiff = maxdiff
148154 end
149155
150156 transfer_dtasks = map (quotientedges (dst_cache)) do quotient_edge
151157 src_subcache = dst_cache[src (quotient_edge)]
152158 dst_subcache = dst_cache[dst (quotient_edge)]
153- return Dagger. @spawn transfer_edges! (
159+
160+ src_subcache = fetch (src_subcache)
161+
162+ return Dagger. spawn (
154163 dst_subcache,
155164 fetch (src_subcache),
156165 edges (dst_cache, quotient_edge)
157- )
166+ ) do dst, src, edges
167+ src_subcache = src. iterate
168+ dst_subcache = dst. iterate
169+ for edge in edges
170+ dst_subcache[edge] = src_subcache[edge]
171+ end
172+ end
158173 end
159174
160- wait .(transfer_dtasks)
175+ foreach (wait, transfer_dtasks)
176+
177+ return state
178+ end
179+
180+ function ITNNP. finalize_state! (
181+ :: BeliefPropagationProblem ,
182+ :: BeliefPropagation ,
183+ state:: ITNNP.DaggerState
184+ )
185+ dst_cache = state. iterate. iterate
186+
187+ for quotient_vertex in quotientvertices (dst_cache)
188+ substate = fetch (dst_cache[quotient_vertex])
189+ subcache = substate. iterate
190+ edge_data (dst_cache) .= edge_data (subcache)
191+ end
161192
162193 return state
163194end
0 commit comments