feectools 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (98) hide show
  1. feectools/__init__.py +0 -0
  2. feectools/accelerate/__init__.py +0 -0
  3. feectools/accelerate/accelerate.py +220 -0
  4. feectools/accelerate/compile_psydac.mk +52 -0
  5. feectools/api/__init__.py +0 -0
  6. feectools/api/essential_bc.py +122 -0
  7. feectools/api/fem_bilinear_form.py +2226 -0
  8. feectools/api/fem_common.py +286 -0
  9. feectools/api/fem_sum_form.py +123 -0
  10. feectools/api/settings.py +82 -0
  11. feectools/core/__init__.py +11 -0
  12. feectools/core/bsplines.py +1107 -0
  13. feectools/core/bsplines_kernels.py +1349 -0
  14. feectools/core/field_evaluation_kernels.py +5015 -0
  15. feectools/core/tests/__init__.py +0 -0
  16. feectools/core/tests/test_bsplines.py +263 -0
  17. feectools/core/tests/test_bsplines_kernel.py +40 -0
  18. feectools/core/tests/test_bsplines_pyccel.py +752 -0
  19. feectools/ddm/__init__.py +3 -0
  20. feectools/ddm/basic.py +78 -0
  21. feectools/ddm/blocking_data_exchanger.py +348 -0
  22. feectools/ddm/cart.py +1835 -0
  23. feectools/ddm/interface_data_exchanger.py +122 -0
  24. feectools/ddm/mpi.py +109 -0
  25. feectools/ddm/nonblocking_data_exchanger.py +331 -0
  26. feectools/ddm/partition.py +207 -0
  27. feectools/ddm/petsc.py +112 -0
  28. feectools/ddm/tests/__init__.py +0 -0
  29. feectools/ddm/tests/test_cart_1d.py +138 -0
  30. feectools/ddm/tests/test_cart_2d.py +164 -0
  31. feectools/ddm/tests/test_cart_3d.py +158 -0
  32. feectools/ddm/tests/test_multicart_2d.py +173 -0
  33. feectools/ddm/tests/test_partition.py +124 -0
  34. feectools/ddm/utilities.py +24 -0
  35. feectools/feec/__init__.py +0 -0
  36. feectools/feec/derivatives.py +780 -0
  37. feectools/feec/dof_kernels.py +210 -0
  38. feectools/feec/global_geometric_projectors.py +1073 -0
  39. feectools/feec/hodge.py +148 -0
  40. feectools/fem/__init__.py +0 -0
  41. feectools/fem/basic.py +465 -0
  42. feectools/fem/grid.py +181 -0
  43. feectools/fem/partitioning.py +344 -0
  44. feectools/fem/projectors.py +160 -0
  45. feectools/fem/splines.py +559 -0
  46. feectools/fem/tensor.py +1393 -0
  47. feectools/fem/tests/__init__.py +0 -0
  48. feectools/fem/tests/analytical_profiles_1d.py +100 -0
  49. feectools/fem/tests/analytical_profiles_base.py +34 -0
  50. feectools/fem/tests/splines_error_bounds.py +155 -0
  51. feectools/fem/tests/test_spline_histopolation.py +120 -0
  52. feectools/fem/tests/test_spline_interpolation.py +182 -0
  53. feectools/fem/tests/test_splines.py +184 -0
  54. feectools/fem/tests/test_splines_par.py +46 -0
  55. feectools/fem/tests/test_vector_spaces.py +150 -0
  56. feectools/fem/tests/utilities.py +47 -0
  57. feectools/fem/vector.py +729 -0
  58. feectools/linalg/__init__.py +0 -0
  59. feectools/linalg/basic.py +1386 -0
  60. feectools/linalg/block.py +1451 -0
  61. feectools/linalg/direct_solvers.py +201 -0
  62. feectools/linalg/fft.py +258 -0
  63. feectools/linalg/kernels/__init__.py +0 -0
  64. feectools/linalg/kernels/axpy_kernels.py +57 -0
  65. feectools/linalg/kernels/inner_kernels.py +100 -0
  66. feectools/linalg/kernels/matvec_kernels.py +206 -0
  67. feectools/linalg/kernels/stencil2IJV_kernels.py +227 -0
  68. feectools/linalg/kernels/stencil2coo_kernels.py +179 -0
  69. feectools/linalg/kernels/transpose_kernels.py +263 -0
  70. feectools/linalg/kron.py +911 -0
  71. feectools/linalg/solvers.py +1914 -0
  72. feectools/linalg/sparse.py +114 -0
  73. feectools/linalg/stencil.py +2923 -0
  74. feectools/linalg/stencil_dot_kernels.py +317 -0
  75. feectools/linalg/stencil_transpose_kernels.py +372 -0
  76. feectools/linalg/tests/__init__.py +0 -0
  77. feectools/linalg/tests/test_block.py +1588 -0
  78. feectools/linalg/tests/test_fft.py +106 -0
  79. feectools/linalg/tests/test_kron_stencil_matrix.py +114 -0
  80. feectools/linalg/tests/test_linalg.py +1065 -0
  81. feectools/linalg/tests/test_matrix_free.py +128 -0
  82. feectools/linalg/tests/test_solvers.py +213 -0
  83. feectools/linalg/tests/test_stencil_interface_matrix.py +379 -0
  84. feectools/linalg/tests/test_stencil_vector.py +1036 -0
  85. feectools/linalg/tests/test_stencil_vector_space.py +440 -0
  86. feectools/linalg/topetsc.py +522 -0
  87. feectools/linalg/utilities.py +200 -0
  88. feectools/utilities/__init__.py +0 -0
  89. feectools/utilities/quadratures.py +113 -0
  90. feectools/utilities/utils.py +166 -0
  91. feectools/version.py +1 -0
  92. feectools-0.1.0.dist-info/METADATA +66 -0
  93. feectools-0.1.0.dist-info/RECORD +98 -0
  94. feectools-0.1.0.dist-info/WHEEL +5 -0
  95. feectools-0.1.0.dist-info/entry_points.txt +3 -0
  96. feectools-0.1.0.dist-info/licenses/AUTHORS +22 -0
  97. feectools-0.1.0.dist-info/licenses/LICENSE +21 -0
  98. feectools-0.1.0.dist-info/top_level.txt +1 -0
feectools/ddm/cart.py ADDED
@@ -0,0 +1,1835 @@
1
+ # coding: utf-8
2
+
3
+ import os
4
+ import numpy as np
5
+ from itertools import product
6
+
7
+ from feectools.ddm.mpi import mpi as MPI
8
+ from feectools.ddm.mpi import MockMPI
9
+ from feectools.ddm.partition import compute_dims, partition_procs_per_patch
10
+
11
+
12
+ __all__ = ('find_mpi_type',
13
+ 'MultiPatchDomainDecomposition',
14
+ 'DomainDecomposition',
15
+ 'CartDecomposition',
16
+ 'InterfaceCartDecomposition',
17
+ 'create_interfaces_cart')
18
+
19
+ #===============================================================================
20
+ def find_mpi_type( dtype ):
21
+ """
22
+ Find correct MPI datatype that corresponds to user-provided datatype.
23
+
24
+ Parameters
25
+ ----------
26
+ dtype : [type | str | numpy.dtype | mpi4py.MPI.Datatype]
27
+ Datatype for which the corresponding MPI datatype is requested.
28
+
29
+ Returns
30
+ -------
31
+ mpi_type : mpi4py.MPI.Datatype
32
+ MPI datatype to be used for communication.
33
+
34
+ """
35
+ if not isinstance(MPI, MockMPI):
36
+ if isinstance( dtype, MPI.Datatype ):
37
+ mpi_type = dtype
38
+ else:
39
+ nt = np.dtype( dtype )
40
+ mpi_type = MPI._typedict[nt.char]
41
+ else:
42
+ mpi_type = np.dtype( dtype )
43
+
44
+ return mpi_type
45
+
46
+ class MultiPatchDomainDecomposition:
47
+ """
48
+ Cartesian decomposition of multiple N-Cube grids.
49
+ This is built on top of an MPI communicator decomposed into smaller disjoint intra-communicators
50
+ assigned to each N-Cube grid to construct a multi-dimensional
51
+ Cartesian topology.
52
+
53
+ Parameters
54
+ ----------
55
+ ncells : list of list of int
56
+ The number of cells in each direction for each grid.
57
+
58
+ periods: list of bool
59
+ The periodicity of the domain in each direction for each grid.
60
+
61
+ comm : MPI.Comm
62
+ MPI communicator that will be used to spawn the grids.
63
+
64
+ num_threads: int
65
+ Number of threads used by one MPI rank.
66
+ """
67
+ def __init__(self, ncells, periods, comm=None, num_threads=None):
68
+
69
+ assert len( ncells ) == len( periods )
70
+ if not isinstance(MPI, MockMPI) and comm is not None:
71
+ assert isinstance( comm, MPI.Comm )
72
+ num_threads = num_threads if num_threads else int(os.environ.get('OMP_NUM_THREADS', 1))
73
+
74
+ # Store input arguments
75
+ self._ncells = tuple( ncells )
76
+ self._periods = tuple( periods )
77
+ self._num_threads = num_threads
78
+ self._comm = comm
79
+
80
+ # ...
81
+ self._npatches = len( ncells )
82
+
83
+ # ...
84
+ if comm is None:
85
+ size = 1
86
+ rank = 0
87
+ else:
88
+ size = comm.Get_size()
89
+ rank = comm.Get_rank()
90
+
91
+ sizes, rank_ranges = partition_procs_per_patch(self._ncells, size)
92
+
93
+ self._rank = rank
94
+ self._size = size
95
+ self._sizes = tuple( sizes )
96
+ self._rank_ranges = tuple( rank_ranges )
97
+
98
+
99
+ global_group = comm.group if comm is not None else None
100
+ owned_groups = []
101
+
102
+ local_groups = [None]*self._npatches
103
+ local_communicators = [None]*self._npatches
104
+
105
+ for i,r in enumerate(rank_ranges):
106
+ if rank>=r[0] and rank<=r[1]:
107
+ if comm is not None:
108
+ local_groups[i] = global_group.Range_incl([[r[0], r[1], 1]])
109
+ local_communicators[i] = comm.Create_group(local_groups[i], i)
110
+ owned_groups.append(i)
111
+ else:
112
+ local_communicators[i] = MPI.COMM_NULL
113
+
114
+ domains = [DomainDecomposition(nc, P, comm=subcomm, global_comm=comm, num_threads=num_threads, size=size)\
115
+ for nc,P,subcomm, size in zip(ncells, periods, local_communicators, sizes)]
116
+
117
+
118
+ self._local_groups = tuple(local_groups)
119
+ self._local_communicators = tuple(local_communicators)
120
+ self._owned_groups = tuple(owned_groups)
121
+ self._domains = tuple(domains)
122
+
123
+ @property
124
+ def ncells( self ):
125
+ return self._ncells
126
+
127
+ @property
128
+ def periods( self ):
129
+ return self._periods
130
+
131
+ @property
132
+ def size( self ):
133
+ return self._size
134
+
135
+ @property
136
+ def rank( self ):
137
+ return self._rank
138
+
139
+ @property
140
+ def sizes( self ):
141
+ return self._sizes
142
+
143
+ @property
144
+ def rank_ranges( self ):
145
+ return self._rank_ranges
146
+
147
+ @property
148
+ def local_groups( self ):
149
+ return self._local_groups
150
+
151
+ @property
152
+ def local_communicators( self ):
153
+ return self._local_communicators
154
+
155
+ @property
156
+ def owned_groups( self ):
157
+ return self._owned_groups
158
+
159
+ @property
160
+ def domains( self ):
161
+ return self._domains
162
+
163
+ @property
164
+ def num_threads( self ):
165
+ return self._num_threads
166
+
167
+ @property
168
+ def comm( self ):
169
+ return self._comm
170
+
171
+ class DomainDecomposition:
172
+ """
173
+ Cartesian decomposition of an N-Cube grid.
174
+ This is built on top of an MPI communicator with multi-dimensional
175
+ Cartesian topology.
176
+
177
+ Parameters
178
+ ----------
179
+ ncells : list of int
180
+ The number of cells in each direction.
181
+
182
+ periods: list of bool
183
+ The periodcity of the domain in each direction.
184
+
185
+ comm : MPI.Comm|None
186
+ MPI communicator that will be used to spawn a new Cartesian communicator.
187
+ In the serial case comm == None.
188
+
189
+ global_comm : MPI.Comm|None
190
+ MPI global communicator that contains all the processes owned by comm.
191
+ In the serial case comm == None.
192
+
193
+ num_threads: int|None
194
+ Number of threads used by one MPI rank.
195
+
196
+ size: int|None
197
+ The number of processes assigned to the domain.
198
+ This information is needed when comm is None (sequential case) or comm == MPI.COMM_NULL (MPI rank does not own the domain),
199
+ to be able to calculate global_element_starts and global_element_ends.
200
+
201
+ mpi_dims_mask: list of bool
202
+ True if the dimension is to be used in the domain decomposition (=default for each dimension).
203
+ If mpi_dims_mask[i]=False, the i-th dimension will not be decomposed.
204
+
205
+ """
206
+
207
+ def __init__(self, ncells, periods, comm=None, global_comm=None, num_threads=None, size=None, mpi_dims_mask=None):
208
+
209
+ # Check input arguments
210
+ # TODO: check that arguments are identical across all processes
211
+ assert len( ncells ) == len( periods )
212
+ assert all( n >=1 for n in ncells )
213
+ assert all( isinstance( period, bool ) for period in periods )
214
+ if isinstance(MPI, MockMPI):
215
+ comm = None
216
+ else:
217
+ if comm is not None:
218
+ assert isinstance( comm, MPI.Comm )
219
+
220
+
221
+ self._ncells = tuple ( ncells )
222
+ self._periods = tuple ( periods )
223
+ self._comm = comm
224
+ self._global_comm = comm if global_comm is None else global_comm
225
+ self._comm_cart = comm
226
+ self._num_threads = num_threads if num_threads else int(os.environ.get('OMP_NUM_THREADS', 1))
227
+
228
+ # ...
229
+ if comm is None:
230
+ self._size = 1
231
+ self._rank = 0
232
+ elif self.is_comm_null:
233
+ assert size is not None
234
+ self._size = size
235
+ self._rank = -1
236
+ else:
237
+ self._size = comm.Get_size()
238
+ self._rank = comm.Get_rank()
239
+
240
+ self._ndims = len(ncells)
241
+ nprocs, block_shape = compute_dims( self._size, self._ncells, mpi_dims_mask=mpi_dims_mask )
242
+
243
+ self._nprocs = nprocs
244
+
245
+ # Store arrays with all the starts and ends along each direction for every process
246
+ self._global_element_starts = [None]*self._ndims
247
+ self._global_element_ends = [None]*self._ndims
248
+ global_shapes = [None]*self._ndims
249
+ for axis in range( self._ndims ):
250
+ n = ncells[axis]
251
+ d = nprocs[axis]
252
+ s = n//d
253
+ global_shapes[axis] = np.array([s]*d)
254
+ global_shapes[axis][:n%d] += 1
255
+
256
+ self._global_element_ends [axis] = np.cumsum(global_shapes[axis])-1
257
+ self._global_element_starts[axis] = np.array( [0] + [e+1 for e in self._global_element_ends[axis][:-1]] )
258
+
259
+ if self.is_comm_null:return
260
+
261
+ if comm is None:
262
+ # compute the coords for all processes
263
+ self._global_coords = np.array([np.unravel_index(rank, nprocs) for rank in range(self._size)])
264
+ self._coords = self._global_coords[self._rank]
265
+ self._rank_in_topo = 0
266
+ self._ranks_in_topo = np.array([0])
267
+ else:
268
+ # Create a MPI cart
269
+ self._comm_cart = comm.Create_cart(
270
+ dims = self._nprocs,
271
+ periods = self._periods,
272
+ reorder = False,
273
+ )
274
+
275
+ # Know my coordinates in the topology
276
+ self._rank_in_topo = self._comm_cart.Get_rank()
277
+ self._coords = self._comm_cart.Get_coords( rank=self._rank_in_topo )
278
+ self._ranks_in_topo = np.array(self._comm_cart.group.Translate_ranks(list(range(self._comm_cart.size)), comm.group))
279
+
280
+ # Start/end values of global indices (without ghost regions)
281
+ self._starts = tuple( self._global_element_starts[axis][c] for axis,c in zip(range(self._ndims), self._coords) )
282
+ self._ends = tuple( self._global_element_ends [axis][c] for axis,c in zip(range(self._ndims), self._coords) )
283
+
284
+ self._local_ncells = tuple(e-s+1 for s,e in zip(self._starts, self._ends))
285
+
286
+ if comm is None:return
287
+
288
+ # Create (N-1)-dimensional communicators within the Cartesian topology
289
+ self._subcomm = [None]*self._ndims
290
+ for i in range(self._ndims):
291
+ remain_dims = [i==j for j in range( self._ndims )]
292
+ self._subcomm[i] = self._comm_cart.Sub( remain_dims )
293
+
294
+ #---------------------------------------------------------------------------
295
+ # Global properties (same for each process)
296
+ #---------------------------------------------------------------------------
297
+ @property
298
+ def ndim( self ):
299
+ return self._ndims
300
+
301
+ @property
302
+ def ncells( self ):
303
+ return self._ncells
304
+
305
+ @property
306
+ def periods( self ):
307
+ return self._periods
308
+
309
+ @property
310
+ def size( self ):
311
+ return self._size
312
+
313
+ @property
314
+ def rank( self ):
315
+ return self._rank
316
+
317
+ @property
318
+ def num_threads( self ):
319
+ return self._num_threads
320
+
321
+ @property
322
+ def comm( self ):
323
+ return self._comm
324
+
325
+ @property
326
+ def comm_cart( self ):
327
+ return self._comm_cart
328
+
329
+ @property
330
+ def global_comm( self ):
331
+ return self._global_comm
332
+
333
+ @property
334
+ def nprocs( self ):
335
+ return self._nprocs
336
+
337
+ @property
338
+ def global_element_starts( self ):
339
+ return self._global_element_starts
340
+
341
+ @property
342
+ def global_element_ends( self ):
343
+ return self._global_element_ends
344
+
345
+ @property
346
+ def is_comm_null( self ):
347
+ return self.comm == MPI.COMM_NULL
348
+
349
+ @property
350
+ def is_parallel( self ):
351
+ return self._comm is not None
352
+
353
+ @property
354
+ def ranks_in_topo( self ):
355
+ return self._ranks_in_topo
356
+ #---------------------------------------------------------------------------
357
+ # Local properties
358
+ #---------------------------------------------------------------------------
359
+ @property
360
+ def starts( self ):
361
+ return self._starts
362
+
363
+ @property
364
+ def ends( self ):
365
+ return self._ends
366
+
367
+ @property
368
+ def coords( self ):
369
+ return self._coords
370
+
371
+ @property
372
+ def subcomm( self ):
373
+ return self._subcomm
374
+
375
+ @property
376
+ def local_ncells( self ):
377
+ return self._local_ncells
378
+
379
+ #---------------------------------------------------------------------------
380
+ def coords_exist( self, coords ):
381
+ return all( P or (0 <= c < d) for P,c,d in zip( self._periods, coords, self._nprocs ) )
382
+
383
+ def refine(self, ncells, global_element_starts, global_element_ends):
384
+ """ Create the new Cartesian decomposition of the refined domain.
385
+
386
+ Parameters
387
+ ----------
388
+ ncells : list or tuple of int
389
+ Number of cells of refined space.
390
+
391
+ global_starts: list of list of int
392
+ The starts of the coefficients for every process along each direction.
393
+
394
+ global_ends: list of list of int
395
+ The ends of the coefficients for every process along each direction.
396
+
397
+ Returns
398
+ -------
399
+ domain : CartDecomposition
400
+ Cartesian decomposition of the refined domain.
401
+ """
402
+
403
+ # Check input arguments
404
+ assert len( ncells ) == len( self.ncells )
405
+ assert all(nc>=snc for nc, snc in zip(ncells, self.ncells))
406
+
407
+ domain = DomainDecomposition(self.ncells, self.periods, comm=self.comm,
408
+ global_comm=self.global_comm, num_threads=self.num_threads,
409
+ size=self.size)
410
+ domain._ncells = tuple ( ncells )
411
+
412
+ # Store arrays with all the starts and ends along each direction for every process
413
+ domain._global_element_starts = tuple(global_element_starts)
414
+ domain._global_element_ends = tuple(global_element_ends)
415
+ if self.is_comm_null:return domain
416
+
417
+ # Start/end values of global indices (without ghost regions)
418
+ domain._starts = tuple( domain._global_element_starts[axis][c] for axis,c in zip(range(self._ndims), self._coords) )
419
+ domain._ends = tuple( domain._global_element_ends [axis][c] for axis,c in zip(range(self._ndims), self._coords) )
420
+
421
+ domain._local_ncells = tuple(e-s+1 for s,e in zip(self._starts, self._ends))
422
+ return domain
423
+
424
+ #==================================================================================
425
+ class CartDecomposition():
426
+ """
427
+ Cartesian decomposition of a tensor-product grid of spline coefficients.
428
+ This is built on top of an MPI communicator with multi-dimensional
429
+ Cartesian topology.
430
+
431
+ Parameters
432
+ ----------
433
+
434
+ domain_decomposition : DomainDecomposition
435
+ The Domain partition.
436
+
437
+ npts : list or tuple of int
438
+ Number of coefficients in the global grid along each dimension.
439
+
440
+ global_starts: list of list of int
441
+ The starts of the global points for every process along each direction.
442
+
443
+ global_ends: list of list of int
444
+ The ends of the global points for every process along each direction.
445
+
446
+ pads : list or tuple of int
447
+ Padding along each grid dimension.
448
+ In 1D, this is the number of extra coefficients added at each boundary
449
+ of the local domain to permit non-local operations with compact support;
450
+ this concept extends to multiple dimensions through a tensor product.
451
+
452
+ shifts: list or tuple of int
453
+ Shifts along each grid dimension.
454
+ It takes values bigger or equal to one, it represents the multiplicity of each knot.
455
+
456
+ """
457
+ def __init__( self, domain_decomposition, npts, global_starts, global_ends, pads, shifts ):
458
+
459
+ # Check input arguments
460
+ # TODO: check that arguments are identical across all processes
461
+ assert len( npts ) == len( global_starts ) == len( global_ends ) == len( pads ) == len(shifts)
462
+ assert all( n >=1 for n in npts )
463
+ assert min(min(gs) for gs in global_starts) >= 0
464
+ assert min(min(ge) for ge in global_ends ) >= 0
465
+ assert all( p >=0 for p in pads )
466
+
467
+ # Store input arguments
468
+ self._domain_decomposition = domain_decomposition
469
+ self._npts = tuple( npts )
470
+ self._global_starts = tuple( [ np.asarray(gs) for gs in global_starts] )
471
+ self._global_ends = tuple( [ np.asarray(ge) for ge in global_ends] )
472
+ self._pads = tuple( pads )
473
+ self._shifts = tuple( shifts )
474
+ self._periods = domain_decomposition.periods
475
+ self._ndims = len( npts )
476
+ self._comm = domain_decomposition.comm
477
+ self._comm_cart = domain_decomposition.comm_cart
478
+ self._local_comm = domain_decomposition.comm
479
+ self._global_comm = domain_decomposition.global_comm
480
+ self._num_threads = domain_decomposition.num_threads
481
+ self._starts = (0,)*self._ndims
482
+ self._ends = (-1,)*self._ndims
483
+ self._shape = (0,)*self._ndims
484
+ self._parent_starts = (None,)*self._ndims
485
+ self._parent_ends = (None,)*self._ndims
486
+
487
+ if self._comm == MPI.COMM_NULL:
488
+ return
489
+
490
+ self._size = domain_decomposition.size
491
+ self._rank = domain_decomposition.rank
492
+ self._nprocs = domain_decomposition.nprocs
493
+ # ...
494
+
495
+ # Know my coordinates in the topology
496
+ self._coords = domain_decomposition.coords
497
+
498
+ # Start/end values of global indices (without ghost regions)
499
+ self._starts = tuple( self._global_starts[axis][c] for axis,c in zip(range(self._ndims), self._coords) )
500
+ self._ends = tuple( self._global_ends [axis][c] for axis,c in zip(range(self._ndims), self._coords) )
501
+
502
+ # List of 1D global indices (without ghost regions)
503
+ self._grids = tuple( range(s,e+1) for s,e in zip( self._starts, self._ends ) )
504
+
505
+ # Compute shape of local arrays in topology (with ghost regions)
506
+ self._shape = tuple( e-s+1+2*m*p for s,e,p,m in zip( self._starts, self._ends, self._pads, shifts ) )
507
+
508
+ # # Extended grids with ghost regions
509
+ # self._extended_grids = tuple( range(s-m*p,e+m*p+1) for s,e,p,m in zip( self._starts, self._ends, self._pads, shifts ) )
510
+
511
+ self._petsccart = None
512
+
513
+ if self._comm is None:return
514
+
515
+ # Create (N-1)-dimensional communicators within the Cartesian topology
516
+ self._subcomm = domain_decomposition.subcomm
517
+
518
+ # dict to store information for communicating with neighbors
519
+ self._shift_info = {}
520
+
521
+ # # dict to store information for communicating with neighbors using non blocking communications
522
+ self._shift_info_non_blocking = {}
523
+
524
+ #---------------------------------------------------------------------------
525
+ # Global properties (same for each process)
526
+ #---------------------------------------------------------------------------
527
+ @property
528
+ def ndim( self ):
529
+ return self._ndims
530
+
531
+ @property
532
+ def npts( self ):
533
+ return self._npts
534
+
535
+ @property
536
+ def pads( self ):
537
+ return self._pads
538
+
539
+ @property
540
+ def periods( self ):
541
+ return self._periods
542
+
543
+ @property
544
+ def shifts( self ):
545
+ return self._shifts
546
+
547
+ @property
548
+ def reorder( self ):
549
+ return self._reorder
550
+
551
+ @property
552
+ def comm( self ):
553
+ return self._comm
554
+
555
+ @property
556
+ def local_comm( self ):
557
+ """ The intra subcommunicator used by this class."""
558
+ return self._local_comm
559
+
560
+ @property
561
+ def global_comm( self ):
562
+ """ The intra-communicator passed by the user, usualy it's MPI.COMM_WORLD."""
563
+ return self._global_comm
564
+
565
+ @property
566
+ def comm_cart( self ):
567
+ """ Intra-communicator with a Cartesian topology."""
568
+ return self._comm_cart
569
+
570
+ @property
571
+ def nprocs( self ):
572
+ """ Number of processes in each dimension."""
573
+ return self._nprocs
574
+
575
+ @property
576
+ def reverse_axis(self):
577
+ """ The axis of the reversed Cartesian topology."""
578
+ return self._reverse_axis
579
+
580
+ @property
581
+ def global_starts( self ):
582
+ """ The starts of all the processes in the cartesian decomposition."""
583
+ return self._global_starts
584
+
585
+ @property
586
+ def global_ends( self ):
587
+ """ The ends of all the processes in the cartesian decomposition."""
588
+ return self._global_ends
589
+
590
+ @property
591
+ def is_comm_null( self ):
592
+ return self.comm == MPI.COMM_NULL
593
+
594
+ @property
595
+ def is_parallel( self ):
596
+ return self._comm is not None
597
+
598
+ @property
599
+ def num_threads( self ):
600
+ return self._num_threads
601
+
602
+ @property
603
+ def domain_decomposition( self ):
604
+ return self._domain_decomposition
605
+
606
+ #---------------------------------------------------------------------------
607
+ # Local properties
608
+ #---------------------------------------------------------------------------
609
+ @property
610
+ def starts( self ):
611
+ return self._starts
612
+
613
+ @property
614
+ def ends( self ):
615
+ return self._ends
616
+
617
+ @property
618
+ def parent_starts( self ):
619
+ return self._parent_starts
620
+
621
+ @property
622
+ def parent_ends( self ):
623
+ return self._parent_ends
624
+
625
+ @property
626
+ def coords( self ):
627
+ return self._coords
628
+
629
+ @property
630
+ def shape( self ):
631
+ return self._shape
632
+
633
+ # TODO check if the property ranks_in_topo is still defined
634
+ @property
635
+ def ranks_in_topo( self ):
636
+ return self._ranks_in_topo
637
+
638
+ @property
639
+ def subcomm( self ):
640
+ return self._subcomm
641
+
642
+ #---------------------------------------------------------------------------
643
+ def topetsc( self ):
644
+ """ Convert the cart to a petsc cart.
645
+ """
646
+ if self._petsccart is None:
647
+ from feectools.ddm.petsc import PetscCart
648
+ self._petsccart = PetscCart(self)
649
+ return self._petsccart
650
+
651
+ #---------------------------------------------------------------------------
652
+ def coords_exist( self, coords ):
653
+
654
+ return all( P or (0 <= c < d) for P,c,d in zip( self._periods, coords, self._nprocs ) )
655
+
656
+ #---------------------------------------------------------------------------
657
+ def get_shift_info( self, direction, disp ):
658
+
659
+ if len(self._shift_info) == 0:
660
+ for axis in range( self._ndims ):
661
+ for d in [-1,1]:
662
+ self._shift_info[ axis, d ] = \
663
+ self._compute_shift_info( axis, d )
664
+
665
+ return self._shift_info[ direction, disp ]
666
+
667
+ #---------------------------------------------------------------------------
668
+ def get_shift_info_non_blocking( self, shift ):
669
+
670
+ if len(self._shift_info_non_blocking) == 0:
671
+ zero_shift = tuple( [0]*self._ndims )
672
+ for sh in product( [-1,0,1], repeat=self._ndims ):
673
+ if sh == zero_shift:
674
+ continue
675
+ self._shift_info_non_blocking[sh] = self._compute_shift_info_non_blocking( sh )
676
+
677
+ return self._shift_info_non_blocking[ shift ]
678
+
679
+ #---------------------------------------------------------------------------
680
+ def get_shared_memory_subdivision( self, shape ):
681
+
682
+ assert len(shape) == self._ndims
683
+
684
+ try:
685
+ nthreads , block_shape = compute_dims( self._num_threads, shape , min_blocksizes=[2*p for p in self._pads], try_uniform=True)
686
+ except ValueError:
687
+ print("Cannot compute dimensions with given input values!")
688
+ self.comm.Abort(1)
689
+
690
+ # compute the coords for all threads
691
+ coords_from_rank = np.array([np.unravel_index(rank, nthreads) for rank in range(self._num_threads)])
692
+ rank_from_coords = np.zeros([n+1 for n in nthreads], dtype=int)
693
+ for r in range(self._num_threads):
694
+ c = coords_from_rank[r]
695
+ rank_from_coords[tuple(c)] = r
696
+
697
+ # rank_from_coords is not used in the current version of the assembly code
698
+ # it's used in the commented second version, where we don't use a global barrier, but needs more checks to work
699
+
700
+ for i in range(self._ndims):
701
+ ind = [slice(None,None)]*self._ndims
702
+ ind[i] = nthreads[i]
703
+ rank_from_coords[tuple(ind)] = self._num_threads
704
+
705
+ # Store arrays with all the starts and ends along each direction for every thread
706
+ thread_global_starts = [None]*self._ndims
707
+ thread_global_ends = [None]*self._ndims
708
+ for axis in range( self._ndims ):
709
+ n = shape[axis]
710
+ d = nthreads[axis]
711
+ thread_global_starts[axis] = np.array( [( c *n)//d for c in range( d )] )
712
+ thread_global_ends [axis] = np.array( [((c+1)*n)//d-1 for c in range( d )] )
713
+
714
+ return coords_from_rank, rank_from_coords, thread_global_starts, thread_global_ends, self._num_threads
715
+
716
+ #---------------------------------------------------------------------------
717
+ def reduce_grid(self, global_starts, global_ends):
718
+ """
719
+ Returns a new CartDecomposition object with a coarser grid from the original one
720
+ we do that by giving a new global_starts and global_ends of the coefficients
721
+ in each dimension.
722
+
723
+ Parameters
724
+ ----------
725
+ global_starts : list/tuple
726
+ the list of the new global_starts in each dimesion.
727
+
728
+ global_ends : list/tuple
729
+ the list of the new global_ends in each dimesion.
730
+
731
+ """
732
+ # Make a copy
733
+ # cart = CartDecomposition(self.npts, self.pads, self.periods, self.reorder, comm=self.comm)
734
+ cart = CartDecomposition(self.domain_decomposition, tuple(end[-1] + 1 for end in global_ends), global_starts, global_ends, self.pads, self.shifts)
735
+ # cart._npts = tuple(end[-1] + 1 for end in global_ends)
736
+
737
+ cart._ndims = self._ndims
738
+
739
+ # Create a 2D MPI cart
740
+ cart._comm_cart = self._comm_cart
741
+
742
+ # Know my coordinates in the topology
743
+ cart._coords = self._coords
744
+
745
+ # Start/end values of global indices (without ghost regions)
746
+ cart._starts = tuple( starts[i] for i,starts in zip( self._coords, global_starts) )
747
+ cart._ends = tuple( ends[i] for i,ends in zip( self._coords, global_ends ) )
748
+
749
+ # List of 1D global indices (without ghost regions)
750
+ cart._grids = tuple( range(s,e+1) for s,e in zip( cart._starts, cart._ends ) )
751
+
752
+ # Compute shape of local arrays in topology (with ghost regions)
753
+ cart._shape = tuple( e-s+1+2*p for s,e,p in zip( cart._starts, cart._ends, cart._pads ) )
754
+
755
+ # Extended grids with ghost regions
756
+ cart._extended_grids = tuple( range(s-p,e+p+1) for s,e,p in zip( cart._starts, cart._ends, cart._pads ) )
757
+
758
+ # N-dimensional global indices with ghost regions
759
+ cart._extended_indices = product( *cart._extended_grids )
760
+
761
+ # Compute/store information for communicating with neighbors
762
+ cart._shift_info = {}
763
+ for dimension in range( cart._ndims ):
764
+ for disp in [-1,1]:
765
+ cart._shift_info[ dimension, disp ] = \
766
+ cart._compute_shift_info( dimension, disp )
767
+
768
+ # Store arrays with all the starts and ends along each direction
769
+ cart._global_starts = global_starts
770
+ cart._global_ends = global_ends
771
+
772
+ return cart
773
+
774
+ #---------------------------------------------------------------------------
775
+ def reduce_npts( self, npts, global_starts, global_ends, shifts):
776
+ """
777
+ Compute the cart of the reduced space.
778
+
779
+ Parameters
780
+ ----------
781
+ npts : list or tuple of int
782
+ Number of coefficients in the global grid along each dimension.
783
+
784
+ global_starts : list of list of int
785
+ The starts of the global points for every process along each direction.
786
+
787
+ global_ends : list of list of int
788
+ The ends of the global points for every process along each direction.
789
+
790
+ shifts : list or tuple of int
791
+ Shifts along each grid dimension.
792
+ It takes values bigger or equal to one, it represents the multiplicity of each knot.
793
+
794
+ Returns
795
+ -------
796
+ v: CartDecomposition
797
+ The reduced cart.
798
+
799
+ """
800
+
801
+ cart = CartDecomposition(self.domain_decomposition, npts, global_starts, global_ends, self.pads, shifts)
802
+ cart._parent_starts = self.starts
803
+ cart._parent_ends = self.ends
804
+ return cart
805
+
806
+ #---------------------------------------------------------------------------
807
+ def change_starts_ends( self, starts, ends, parent_starts, parent_ends):
808
+ """ Create a slice of the cart based on the new starts and ends.
809
+ WARNING! this function should be used carefully,
810
+ as it might generate errors if it was not used properly in the communication process.
811
+ """
812
+ cart = CartDecomposition(self._domain_decomposition, self._npts, self._global_starts, self._global_ends, self._pads, self._shifts)
813
+
814
+ assert self.comm is None or self.comm.size == 1
815
+
816
+ # Start/end values of global indices (without ghost regions)
817
+ cart._starts = tuple(starts)
818
+ cart._ends = tuple(ends)
819
+
820
+ # List of 1D global indices (without ghost regions)
821
+ cart._grids = tuple( range(s,e+1) for s,e in zip( cart._starts, cart._ends ) )
822
+
823
+ # Compute shape of local arrays in topology (with ghost regions)
824
+ cart._shape = tuple( e-s+1+2*m*p for s,e,p,m in zip( cart._starts, cart._ends, cart._pads, cart._shifts ) )
825
+
826
+ cart._parent_starts = parent_starts
827
+ cart._parent_ends = parent_ends
828
+
829
+ if self._comm is None:return cart
830
+ # Compute/store information for communicating with neighbors
831
+ cart._shift_info = {}
832
+ for dimension in range( cart._ndims ):
833
+ for disp in [-1,1]:
834
+ cart._shift_info[ dimension, disp ] = \
835
+ cart._compute_shift_info( dimension, disp )
836
+
837
+ return cart
838
+
839
+ #---------------------------------------------------------------------------
840
+ def _compute_shift_info( self, direction, disp ):
841
+
842
+ assert( 0 <= direction < self._ndims )
843
+ assert( isinstance( disp, int ) )
844
+
845
+ # reorder = self.reverse_axis == direction
846
+ # Process ranks for data shifting with MPI_SENDRECV
847
+ (rank_source, rank_dest) = self.comm_cart.Shift( direction, disp )
848
+
849
+ # if reorder:
850
+ # (rank_source, rank_dest) = (rank_dest, rank_source)
851
+
852
+ # Mesh info info along given direction
853
+ s = self._starts[direction]
854
+ e = self._ends [direction]
855
+ p = self._pads [direction]
856
+ m = self._shifts[direction]
857
+
858
+ # Shape of send/recv subarrays
859
+ buf_shape = np.array( self._shape )
860
+ buf_shape[direction] = m*p
861
+
862
+ # Start location of send/recv subarrays
863
+ send_starts = np.zeros( self._ndims, dtype=int )
864
+ recv_starts = np.zeros( self._ndims, dtype=int )
865
+ send_assembly_starts = np.zeros( self._ndims, dtype=int )
866
+ recv_assembly_starts = np.zeros( self._ndims, dtype=int )
867
+
868
+ if disp > 0:
869
+ recv_starts[direction] = 0
870
+ send_starts[direction] = e-s+1
871
+ recv_assembly_starts[direction] = 0
872
+ send_assembly_starts[direction] = e-s+1+m*p
873
+ elif disp < 0:
874
+ recv_starts[direction] = e-s+1+m*p
875
+ send_starts[direction] = m*p
876
+ recv_assembly_starts[direction] = e-s+1+m*p
877
+ send_assembly_starts[direction] = 0
878
+
879
+ # Store all information into dictionary
880
+ info = {'rank_dest' : rank_dest,
881
+ 'rank_source' : rank_source,
882
+ 'buf_shape' : tuple( buf_shape ),
883
+ 'send_starts' : tuple( send_starts ),
884
+ 'recv_starts' : tuple( recv_starts ),
885
+ 'send_assembly_starts': tuple( send_assembly_starts ),
886
+ 'recv_assembly_starts': tuple( recv_assembly_starts )}
887
+ return info
888
+
889
+ #---------------------------------------------------------------------------
890
+ def _compute_shift_info_non_blocking( self, shift ):
891
+
892
+ assert( len( shift ) == self._ndims )
893
+
894
+ # Compute coordinates of destination and source
895
+ coords_dest = [c+h for c,h in zip( self._coords, shift )]
896
+ coords_source = [c-h for c,h in zip( self._coords, shift )]
897
+
898
+ # Convert coordinates to rank, taking care of non-periodic dimensions
899
+ if self.coords_exist( coords_dest ):
900
+ rank_dest = self._comm_cart.Get_cart_rank( coords_dest )
901
+ else:
902
+ rank_dest = MPI.PROC_NULL
903
+
904
+ if len([i for i in shift if i==0]) == 2 and rank_dest != MPI.PROC_NULL:
905
+ direction = [i for i,s in enumerate(shift) if s != 0][0]
906
+ comm = self._subcomm[direction]
907
+ local_dest_rank = self._comm_cart.group.Translate_ranks(np.array([rank_dest]), comm.group)[0]
908
+ else:
909
+ local_dest_rank = rank_dest
910
+ comm = self._comm_cart
911
+
912
+ if self.coords_exist( coords_source ):
913
+ rank_source = self._comm_cart.Get_cart_rank( coords_source )
914
+ else:
915
+ rank_source = MPI.PROC_NULL
916
+
917
+ if len([i for i in shift if i==0]) == 2 and rank_source != MPI.PROC_NULL:
918
+ direction = [i for i,s in enumerate(shift) if s != 0][0]
919
+ comm = self._subcomm[direction]
920
+ local_source_rank = self._comm_cart.group.Translate_ranks(np.array([rank_source]), comm.group)[0]
921
+ else:
922
+ local_source_rank = rank_source
923
+ comm = self._comm_cart
924
+
925
+ # Compute information for exchanging ghost cell data
926
+ buf_shape = []
927
+ send_starts = []
928
+ recv_starts = []
929
+ for s,e,m,p,h in zip( self._starts, self._ends, self._shifts, self._pads, shift ):
930
+
931
+ if h == 0:
932
+ buf_length = e-s+1
933
+ recv_start = m*p
934
+ send_start = m*p
935
+
936
+ elif h == 1:
937
+ buf_length = m*p
938
+ recv_start = 0
939
+ send_start = e-s+1
940
+
941
+ elif h == -1:
942
+ buf_length = m*p
943
+ recv_start = e-s+1+m*p
944
+ send_start = m*p
945
+
946
+ buf_shape .append( buf_length )
947
+ send_starts.append( send_start )
948
+ recv_starts.append( recv_start )
949
+
950
+ # Compute unique identifier for messages traveling along 'shift'
951
+ tag = sum( (h%3)*(3**n) for h,n in zip(shift, range(self._ndims)) )
952
+
953
+ # Store all information into dictionary
954
+ info = {'rank_dest' : rank_dest,
955
+ 'rank_source' : rank_source,
956
+ 'local_dest_rank' : local_dest_rank,
957
+ 'local_source_rank' : local_source_rank,
958
+ 'comm' : comm,
959
+ 'tag' : tag,
960
+ 'buf_shape' : tuple( buf_shape ),
961
+ 'send_starts' : tuple( send_starts ),
962
+ 'recv_starts' : tuple( recv_starts )}
963
+
964
+ # return dictionary
965
+ return info
966
+
967
+ #===============================================================================
968
+ class InterfaceCartDecomposition:
969
+ """
970
+ The Cartesian decomposition of an interface constructed from the Cartesian decomposition of the patches that shares an interface.
971
+ This is built using a new inter-communicator between the cartesian grids.
972
+
973
+ Parameters
974
+ ----------
975
+
976
+ cart_minus: CartDecomposition
977
+ The cartesian decomposition of the minus patch.
978
+
979
+ cart_plus: CartDecomposition
980
+ The cartesian decomposition of the plus patch.
981
+
982
+ comm : mpi4py.MPI.Comm
983
+ MPI communicator that will be used to spawn the cart decomposition
984
+
985
+ axes: list of ints
986
+ The axes of the patches that share the interface.
987
+
988
+ exts: list of ints
989
+ The extremities of the patches that share the interface.
990
+
991
+ ranks_in_topo:
992
+ The ranks of the processes that share the interface.
993
+
994
+ local_groups: list of MPI.Group
995
+ The groups that constucts the patches that share the interface.
996
+
997
+ local_communicators: list of intra-communicators
998
+ The communicators of the patches that share the interface.
999
+
1000
+ root_ranks: list of ints
1001
+ The root ranks in the global communicator of the patches.
1002
+
1003
+ requests: list of MPI.Request
1004
+ the requests of the communications between the cartesian topologies that share the interface.
1005
+
1006
+ """
1007
+ def __init__(self, cart_minus, cart_plus, comm, axes, exts, ranks_in_topo, local_groups, local_communicators, root_ranks, requests, reduce_elements=False):
1008
+
1009
+ domain_decomposition_minus = cart_minus.domain_decomposition
1010
+ domain_decomposition_plus = cart_plus.domain_decomposition
1011
+ global_starts_minus = cart_minus.global_starts
1012
+ global_starts_plus = cart_plus.global_starts
1013
+ global_ends_minus = cart_minus.global_ends
1014
+ global_ends_plus = cart_plus.global_ends
1015
+
1016
+ npts_minus = cart_minus.npts
1017
+ npts_plus = cart_plus.npts
1018
+ pads_minus = cart_minus.pads
1019
+ pads_plus = cart_plus.pads
1020
+ shifts_minus = cart_minus.shifts
1021
+ shifts_plus = cart_plus.shifts
1022
+
1023
+ periods_minus = domain_decomposition_minus.periods
1024
+ periods_plus = domain_decomposition_plus.periods
1025
+ axis_minus, axis_plus = axes
1026
+ ext_minus, ext_plus = exts
1027
+ size_minus, size_plus = len(ranks_in_topo[0]), len(ranks_in_topo[1])
1028
+
1029
+ assert axis_minus == axis_plus
1030
+ num_threads = domain_decomposition_minus.num_threads
1031
+
1032
+ root_rank_minus, root_rank_plus = root_ranks
1033
+ local_comm_minus, local_comm_plus = local_communicators
1034
+ ranks_in_topo_minus, ranks_in_topo_plus = ranks_in_topo
1035
+
1036
+ self._cart_minus = cart_minus
1037
+ self._cart_plus = cart_plus
1038
+ self._ndims = len( npts_minus )
1039
+ self._domain_decomposition_minus = domain_decomposition_minus
1040
+ self._domain_decomposition_plus = domain_decomposition_plus
1041
+ self._npts_minus = npts_minus
1042
+ self._npts_plus = npts_plus
1043
+ self._global_starts_minus = global_starts_minus
1044
+ self._global_starts_plus = global_starts_plus
1045
+ self._global_ends_minus = global_ends_minus
1046
+ self._global_ends_plus = global_ends_plus
1047
+ self._pads_minus = pads_minus
1048
+ self._pads_plus = pads_plus
1049
+ self._periods_minus = periods_minus
1050
+ self._periods_plus = periods_plus
1051
+ self._shifts_minus = shifts_minus
1052
+ self._shifts_plus = shifts_plus
1053
+ self._axis = axis_minus
1054
+ self._ext_minus = ext_minus
1055
+ self._ext_plus = ext_plus
1056
+ self._shape = (0,)*self._ndims
1057
+ self._comm = comm
1058
+ self._local_comm_minus = local_comm_minus
1059
+ self._local_comm_plus = local_comm_plus
1060
+ self._root_rank_minus = root_rank_minus
1061
+ self._root_rank_plus = root_rank_plus
1062
+ self._ranks_in_topo_minus = ranks_in_topo_minus
1063
+ self._ranks_in_topo_plus = ranks_in_topo_plus
1064
+ self._local_group_minus = local_groups[0]
1065
+ self._local_group_plus = local_groups[1]
1066
+ self._local_rank_minus = None
1067
+ self._local_rank_plus = None
1068
+ self._intercomm = MPI.COMM_NULL
1069
+ self._num_threads = num_threads
1070
+
1071
+ if comm == MPI.COMM_NULL:
1072
+ return
1073
+
1074
+ if local_comm_minus != MPI.COMM_NULL :
1075
+ self._local_rank_minus = local_comm_minus.rank
1076
+
1077
+ if local_comm_plus != MPI.COMM_NULL:
1078
+ self._local_rank_plus = local_comm_plus.rank
1079
+
1080
+ nprocs_minus, block_shape = compute_dims( size_minus, domain_decomposition_minus.ncells )
1081
+ nprocs_plus, block_shape = compute_dims( size_plus, domain_decomposition_plus.ncells )
1082
+
1083
+ self._nprocs_minus = nprocs_minus
1084
+ self._nprocs_plus = nprocs_plus
1085
+
1086
+ if requests:MPI.Request.Waitall(requests)
1087
+ dtype = find_mpi_type('int64')
1088
+ if local_comm_minus != MPI.COMM_NULL and reduce_elements == False:
1089
+ local_comm_minus.Bcast((ranks_in_topo_plus,ranks_in_topo_plus.size, dtype), root=0)
1090
+
1091
+ if local_comm_plus != MPI.COMM_NULL and reduce_elements == False:
1092
+ local_comm_plus.Bcast((ranks_in_topo_minus,ranks_in_topo_minus.size, dtype), root=0)
1093
+
1094
+ self._coords_from_rank_minus = np.array([np.unravel_index(rank, nprocs_minus) for rank in range(size_minus)])
1095
+ self._coords_from_rank_plus = np.array([np.unravel_index(rank, nprocs_plus) for rank in range(size_plus)])
1096
+
1097
+ rank_from_coords_minus = np.zeros(nprocs_minus, dtype=int)
1098
+ rank_from_coords_plus = np.zeros(nprocs_plus, dtype=int)
1099
+
1100
+ for r in range(size_minus):
1101
+ rank_from_coords_minus[tuple(self._coords_from_rank_minus[r])] = r
1102
+
1103
+ for r in range(size_plus):
1104
+ rank_from_coords_plus[tuple(self._coords_from_rank_plus[r])] = r
1105
+
1106
+ index_minus = [slice(None, None)]*len(npts_minus)
1107
+ index_plus = [slice(None, None)]*len(npts_minus)
1108
+ index_minus[axis_minus] = 0 if ext_minus == -1 else -1
1109
+ index_plus[axis_plus] = 0 if ext_plus == -1 else -1
1110
+
1111
+ self._boundary_ranks_minus = rank_from_coords_minus[tuple(index_minus)].ravel()
1112
+ self._boundary_ranks_plus = rank_from_coords_plus[tuple(index_plus)].ravel()
1113
+
1114
+ boundary_group_minus = local_groups[0].Incl(self._boundary_ranks_minus)
1115
+ boundary_group_plus = local_groups[1].Incl(self._boundary_ranks_plus)
1116
+
1117
+ comm_minus = comm.Create_group(boundary_group_minus)
1118
+ comm_plus = comm.Create_group(boundary_group_plus)
1119
+
1120
+ root_minus = boundary_group_minus.Translate_ranks([0], comm.group)[0]
1121
+ root_plus = boundary_group_plus.Translate_ranks([0], comm.group)[0]
1122
+
1123
+ procs_index_minus = local_groups[0].Translate_ranks(self._boundary_ranks_minus, boundary_group_minus)
1124
+ procs_index_plus = local_groups[1].Translate_ranks(self._boundary_ranks_plus, boundary_group_plus)
1125
+
1126
+ # Reorder procs ranks from 0 to local_group.size-1
1127
+ self._boundary_ranks_minus = self._boundary_ranks_minus[procs_index_minus]
1128
+ self._boundary_ranks_plus = self._boundary_ranks_plus[procs_index_plus]
1129
+
1130
+ if root_minus != root_plus:
1131
+ if not comm_minus == MPI.COMM_NULL:
1132
+ self._intercomm = comm_minus.Create_intercomm(0, comm, root_plus)
1133
+ self._local_comm = comm_minus
1134
+
1135
+ elif not comm_plus == MPI.COMM_NULL:
1136
+ self._intercomm = comm_plus.Create_intercomm(0, comm, root_minus)
1137
+ self._local_comm = comm_plus
1138
+
1139
+ if self._intercomm == MPI.COMM_NULL:
1140
+ return
1141
+
1142
+ self._local_boundary_ranks_minus = local_groups[0].Translate_ranks(self._boundary_ranks_minus, boundary_group_minus)
1143
+ self._local_boundary_ranks_plus = local_groups[1].Translate_ranks(self._boundary_ranks_plus, boundary_group_plus)
1144
+
1145
+ # high = self._local_rank_plus is not None
1146
+ # self._intercomm = self._intercomm.Merge(high=high)
1147
+
1148
+ if self._local_rank_minus is not None:
1149
+ # Store input arguments
1150
+ self._npts = tuple( npts_minus )
1151
+ self._pads = tuple( pads_minus )
1152
+ self._periods = tuple( periods_minus )
1153
+ self._shifts = tuple( shifts_minus )
1154
+ self._dims = nprocs_minus
1155
+
1156
+ self._global_starts = self._global_starts_minus
1157
+ self._global_ends = self._global_ends_minus
1158
+
1159
+ # Start/end values of global indices (without ghost regions)
1160
+ self._coords = self._coords_from_rank_minus[self._local_rank_minus]
1161
+ self._starts = tuple( self._global_starts[d][c] for d,c in enumerate(self._coords) )
1162
+ self._ends = tuple( self._global_ends [d][c] for d,c in enumerate(self._coords) )
1163
+ self._domain_decomposition = domain_decomposition_minus
1164
+
1165
+ if self._local_rank_plus is not None:
1166
+ # Store input arguments
1167
+ self._npts = tuple( npts_plus )
1168
+ self._pads = tuple( pads_plus )
1169
+ self._periods = tuple( periods_plus )
1170
+ self._shifts = tuple( shifts_plus )
1171
+ self._dims = nprocs_plus
1172
+
1173
+ self._global_starts = self._global_starts_plus
1174
+ self._global_ends = self._global_ends_plus
1175
+
1176
+ # Start/end values of global indices (without ghost regions)
1177
+ self._coords = self._coords_from_rank_plus[self._local_rank_plus]
1178
+ self._starts = tuple( self._global_starts[d][c] for d,c in enumerate(self._coords) )
1179
+ self._ends = tuple( self._global_ends [d][c] for d,c in enumerate(self._coords) )
1180
+ self._domain_decomposition = domain_decomposition_plus
1181
+ # List of 1D global indices (without ghost regions)
1182
+ self._grids = tuple( range(s,e+1) for s,e in zip( self._starts, self._ends ) )
1183
+
1184
+ self._petsccart = None
1185
+ self._parent_starts = tuple([None]*self._ndims)
1186
+ self._parent_ends = tuple([None]*self._ndims)
1187
+ self._parent_npts_minus = tuple([None]*self._ndims)
1188
+ self._parent_npts_plus = tuple([None]*self._ndims)
1189
+ self._get_minus_starts_ends = None
1190
+ self._get_plus_starts_ends = None
1191
+
1192
+ self._interface_communication_infos = {}
1193
+
1194
+ #---------------------------------------------------------------------------
1195
+ # Global properties (same for each process)
1196
+ #---------------------------------------------------------------------------
1197
+ @property
1198
+ def ndims( self ):
1199
+ """Number of dimensions."""
1200
+ return self._ndims
1201
+
1202
+ @property
1203
+ def domain_decomposition_minus( self ):
1204
+ """ The DomainDecomposition of the minus patch."""
1205
+ return self._domain_decomposition_minus
1206
+
1207
+ @property
1208
+ def domain_decomposition_plus( self ):
1209
+ """ The DomainDecomposition of the plus patch."""
1210
+ return self._domain_decomposition_plus
1211
+
1212
+ @property
1213
+ def npts_minus( self ):
1214
+ """Number of points in the minus side of an interface."""
1215
+ return self._npts_minus
1216
+
1217
+ @property
1218
+ def npts_plus( self ):
1219
+ """Number of points in the plus side of an interface."""
1220
+ return self._npts_plus
1221
+
1222
+ @property
1223
+ def pads_minus( self ):
1224
+ """ padding in the minus side of an interface."""
1225
+ return self._pads_minus
1226
+
1227
+ @property
1228
+ def pads_plus( self ):
1229
+ """ padding in the plus side of an interface."""
1230
+ return self._pads_plus
1231
+
1232
+ @property
1233
+ def periods_minus( self ):
1234
+ """ Periodicity in the minus side of an interface."""
1235
+ return self._periods_minus
1236
+
1237
+ @property
1238
+ def periods_plus( self ):
1239
+ """ Periodicity in the plus side of an interface."""
1240
+ return self._periods_plus
1241
+
1242
+ @property
1243
+ def shifts_minus( self ):
1244
+ """ The shift values in the minus side of an interface."""
1245
+ return self._shifts_minus
1246
+
1247
+ @property
1248
+ def shifts_plus( self ):
1249
+ """ The shift values in the plus side of an interface."""
1250
+ return self._shifts_plus
1251
+
1252
+ @property
1253
+ def ext_minus( self ):
1254
+ """ the extremity of the boundary on the minus side of an interface."""
1255
+ return self._ext_minus
1256
+
1257
+ @property
1258
+ def ext_plus( self ):
1259
+ """ the extremity of the boundary on the plus side of an interface."""
1260
+ return self._ext_plus
1261
+
1262
+ @property
1263
+ def root_rank_minus( self ):
1264
+ """ The root rank of the intra-communicator defined in the minus patch."""
1265
+ return self._root_rank_minus
1266
+
1267
+ @property
1268
+ def root_rank_plus( self ):
1269
+ """ The root rank of the intra-communicator defined in the plus patch."""
1270
+ return self._root_rank_plus
1271
+
1272
+ @property
1273
+ def ranks_in_topo_minus( self ):
1274
+ """Array that maps the ranks in the intra-communicator on the minus patch to their rank in the corresponding Cartesian topology."""
1275
+ return self._ranks_in_topo_minus
1276
+
1277
+ @property
1278
+ def ranks_in_topo_plus( self ):
1279
+ """Array that maps the ranks in the intra-communicator on the plus patch to their rank in the corresponding Cartesian topology."""
1280
+ return self._ranks_in_topo_plus
1281
+
1282
+ @property
1283
+ def coords_from_rank_minus( self ):
1284
+ """ Array that maps the ranks of minus patch to their coordinates in the cartesian decomposition."""
1285
+ return self._coords_from_rank_minus
1286
+
1287
+ @property
1288
+ def coords_from_rank_plus( self ):
1289
+ """ Array that maps the ranks of plus patch to their coordinates in the cartesian decomposition."""
1290
+ return self._coords_from_rank_plus
1291
+
1292
+ @property
1293
+ def boundary_ranks_minus( self ):
1294
+ """ Array that contains the ranks defined on the boundary of the minus side of the interface."""
1295
+ return self._boundary_ranks_minus
1296
+
1297
+ @property
1298
+ def boundary_ranks_plus( self ):
1299
+ """ Array that contains the ranks defined on the boundary of the plus side of the interface."""
1300
+ return self._boundary_ranks_plus
1301
+
1302
+ @property
1303
+ def global_starts_minus( self ):
1304
+ """ The starts of all the processes in the cartesian decomposition defined on the minus patch."""
1305
+ return self._global_starts_minus
1306
+
1307
+ @property
1308
+ def global_starts_plus( self ):
1309
+ """ The starts of all the processes in the cartesian decomposition defined on the plus patch."""
1310
+ return self._global_starts_plus
1311
+
1312
+ @property
1313
+ def global_ends_minus( self ):
1314
+ """ The ends of all the processes in the cartesian decomposition defined on the minus patch."""
1315
+ return self._global_ends_minus
1316
+
1317
+ @property
1318
+ def global_ends_plus( self ):
1319
+ """ The ends of all the processes in the cartesian decomposition defined on the plus patch."""
1320
+ return self._global_ends_plus
1321
+
1322
+ @property
1323
+ def axis( self ):
1324
+ """ The axis of the interface."""
1325
+ return self._axis
1326
+
1327
+ @property
1328
+ def npts( self ):
1329
+ return self._npts
1330
+
1331
+ @property
1332
+ def pads( self ):
1333
+ return self._pads
1334
+
1335
+ @property
1336
+ def periods( self ):
1337
+ return self._periods
1338
+
1339
+ @property
1340
+ def shifts( self ):
1341
+ return self._shifts
1342
+
1343
+ @property
1344
+ def global_starts( self ):
1345
+ return self._global_starts
1346
+
1347
+ @property
1348
+ def global_ends( self ):
1349
+ return self._global_ends
1350
+
1351
+ @property
1352
+ def domain_decomposition( self ):
1353
+ return self._domain_decomposition
1354
+
1355
+ @property
1356
+ def comm( self ):
1357
+ return self._comm
1358
+
1359
+ @property
1360
+ def intercomm( self ):
1361
+ return self._intercomm
1362
+
1363
+ @property
1364
+ def is_comm_null( self ):
1365
+ return self._intercomm == MPI.COMM_NULL
1366
+
1367
+ @property
1368
+ def is_parallel( self ):
1369
+ return self._comm is not None
1370
+
1371
+ @property
1372
+ def num_threads( self ):
1373
+ return self._num_threads
1374
+ #---------------------------------------------------------------------------
1375
+ # Local properties
1376
+ #---------------------------------------------------------------------------
1377
+ @property
1378
+ def starts( self ):
1379
+ return self._starts
1380
+
1381
+ @property
1382
+ def ends( self ):
1383
+ return self._ends
1384
+
1385
+ @property
1386
+ def coords( self ):
1387
+ return self._coords
1388
+
1389
+ @property
1390
+ def shape( self ):
1391
+ return self._shape
1392
+
1393
+ @property
1394
+ def parent_starts( self ):
1395
+ return self._starts
1396
+
1397
+ @property
1398
+ def parent_ends( self ):
1399
+ return self._parent_ends
1400
+
1401
+ @property
1402
+ def local_group_minus( self ):
1403
+ """ The MPI Group of the ranks defined in the minus patch"""
1404
+ return self._local_group_minus
1405
+
1406
+ @property
1407
+ def local_group_plus( self ):
1408
+ """ The MPI Group of the ranks defined in the plus patch"""
1409
+ return self._local_group_plus
1410
+
1411
+ @property
1412
+ def local_comm_minus( self ):
1413
+ """ The MPI intra-subcommunicator defined in the minus patch"""
1414
+ return self._local_comm_minus
1415
+
1416
+ @property
1417
+ def local_comm_plus( self ):
1418
+ """ The MPI intra-subcommunicator defined in the plus patch"""
1419
+ return self._local_comm_plus
1420
+
1421
+ @property
1422
+ def local_comm( self ):
1423
+ """ The sub-communicator to which the process belongs, it can be local_comm_minus or local_comm_plus."""
1424
+ return self._local_comm
1425
+
1426
+ @property
1427
+ def local_rank_minus( self ):
1428
+ """ The rank of the process defined on minus side of the interface,
1429
+ the rank is undefined in the case where the process is defined in the plus side of the interface."""
1430
+ return self._local_rank_minus
1431
+
1432
+ @property
1433
+ def local_rank_plus( self ):
1434
+ """ The rank of the process defined on plus side of the interface,
1435
+ the rank is undefined in the case where the process is defined in the minus side of the interface."""
1436
+ return self._local_rank_plus
1437
+
1438
+ #---------------------------------------------------------------------------
1439
+ def reduce_npts( self, cart_minus, cart_plus):
1440
+ """ Compute the cart of the reduced space.
1441
+
1442
+ Parameters
1443
+ ----------
1444
+
1445
+ npts : list or tuple of int
1446
+ Number of coefficients in the global grid along each dimension.
1447
+
1448
+ global_starts: list of list of int
1449
+ The starts of the global points for every process along each direction.
1450
+
1451
+ global_ends: list of list of int
1452
+ The ends of the global points for every process along each direction.
1453
+
1454
+ shifts: list or tuple of int
1455
+ Shifts along each grid dimension.
1456
+ It takes values bigger or equal to one, it represents the multiplicity of each knot.
1457
+
1458
+ Returns
1459
+ -------
1460
+ v: CartDecomposition
1461
+ The reduced cart.
1462
+ """
1463
+
1464
+ comm = self.comm
1465
+ axes = [self.axis, self.axis]
1466
+ exts = [self.ext_minus, self.ext_plus]
1467
+ ranks_in_topo = [self.ranks_in_topo_minus, self.ranks_in_topo_plus]
1468
+ local_groups = [self.local_group_minus, self.local_group_plus]
1469
+ local_communicators = [self.local_comm_minus, self.local_comm_plus]
1470
+ root_ranks = [self.root_rank_minus, self.root_rank_plus]
1471
+ requests = []
1472
+ num_threads = self.num_threads
1473
+
1474
+ cart = InterfaceCartDecomposition(cart_minus, cart_plus, comm, axes, exts, ranks_in_topo, local_groups,
1475
+ local_communicators, root_ranks, requests, reduce_elements=True)
1476
+ cart._parent_starts = self.starts
1477
+ cart._parent_ends = self.ends
1478
+ cart._parent_npts_minus = self.npts_minus
1479
+ cart._parent_npts_plus = self.npts_plus
1480
+
1481
+ return cart
1482
+
1483
+ def set_interface_communication_infos( self, get_minus_starts_ends, get_plus_starts_ends ):
1484
+ self._interface_communication_infos[self._axis] = self._compute_interface_communication_infos_p2p(self._axis, get_minus_starts_ends, get_plus_starts_ends)
1485
+
1486
+ def get_interface_communication_infos( self, axis ):
1487
+ return self._interface_communication_infos[ axis ]
1488
+
1489
+ #---------------------------------------------------------------------------
1490
+ def _compute_interface_communication_infos( self, axis ):
1491
+
1492
+ if self._intercomm == MPI.COMM_NULL:
1493
+ return
1494
+
1495
+ # Mesh info
1496
+ npts_minus = self._npts_minus
1497
+ npts_plus = self._npts_plus
1498
+ p_npts_minus = self._parent_npts_minus
1499
+ p_npts_plus = self._parent_npts_plus
1500
+ pads_minus = self._pads_minus
1501
+ pads_plus = self._pads_plus
1502
+ shifts_minus = self._shifts_minus
1503
+ shifts_plus = self._shifts_plus
1504
+ ext_minus = self._ext_minus
1505
+ ext_plus = self._ext_plus
1506
+ indices = []
1507
+
1508
+ diff = 0
1509
+ if p_npts_minus[axis] is not None:
1510
+ diff = min(1,p_npts_minus[axis]-npts_minus[axis])
1511
+
1512
+ if self._local_rank_minus is not None:
1513
+ rank_minus = self._local_rank_minus
1514
+ coords = self._coords_from_rank_minus[rank_minus]
1515
+ starts = [self._global_starts_minus[d][c] for d,c in enumerate(coords)]
1516
+ ends = [self._global_ends_minus[d][c] for d,c in enumerate(coords)]
1517
+ send_shape = [e-s+1+2*m*p for s,e,m,p in zip(starts, ends, shifts_minus, pads_minus)]
1518
+ send_starts = [m*p for m,p in zip(shifts_minus, pads_minus)]
1519
+ m,p,s,e = shifts_minus[axis], pads_minus[axis], starts[axis], ends[axis]
1520
+ send_starts[axis] = m*p if ext_minus == -1 else m*p+e-s+1-p-1+diff
1521
+ starts[axis] = starts[axis] if ext_minus == -1 else ends[axis]-pads_minus[axis]+diff
1522
+ ends[axis] = starts[axis]+pads_minus[axis]-diff if ext_minus == -1 else ends[axis]
1523
+ send_buf_shape = [e-s+1 for s,e,p,m in zip(starts, ends, pads_minus, shifts_minus)]
1524
+
1525
+ # ...
1526
+ coords = self._coords_from_rank_plus[self._boundary_ranks_plus[0]]
1527
+ starts = [self._global_starts_plus[d][c] for d,c in enumerate(coords)]
1528
+ ends = [self._global_ends_plus[d][c] for d,c in enumerate(coords)]
1529
+
1530
+ recv_shape = [n+2*m*p for n,m,p in zip(npts_plus, shifts_plus, pads_plus)]
1531
+ recv_shape[axis] = pads_plus[axis]+1-diff + 2*shifts_plus[axis]*pads_plus[axis]
1532
+
1533
+ displacements = [0]*(len(self._boundary_ranks_plus)+1)
1534
+ recv_counts = [None]*len(self._boundary_ranks_plus)
1535
+ for k,b in enumerate(self._boundary_ranks_plus):
1536
+ coords = self._coords_from_rank_plus[b]
1537
+ starts = [self._global_starts_plus[d][c] for d,c in enumerate(coords)]
1538
+ ends = [self._global_ends_plus[d][c] for d,c in enumerate(coords)]
1539
+ starts[axis] = starts[axis] if ext_plus == -1 else ends[axis]-pads_plus[axis]+diff
1540
+ ends[axis] = starts[axis]+pads_plus[axis]-diff if ext_plus == -1 else ends[axis]
1541
+ shape_k = [e-s+1 for s,e in zip(starts, ends)]
1542
+ recv_counts[k] = np.prod(shape_k)
1543
+ ranges = [(s+p*m, p*m+e+1) for s,e,p,m in zip(starts, ends, pads_plus, shifts_plus)]
1544
+ ranges[axis] = (shifts_plus[axis]*pads_plus[axis], shifts_plus[axis]*pads_plus[axis]+shape_k[axis])
1545
+ indices += [np.ravel_multi_index( ii, dims=recv_shape, order='C' ) for ii in product(*[range(*a) for a in ranges])]
1546
+
1547
+ elif self._local_rank_plus is not None:
1548
+ rank_plus = self._local_rank_plus
1549
+ coords = self._coords_from_rank_plus[rank_plus]
1550
+ starts = [self._global_starts_plus[d][c] for d,c in enumerate(coords)]
1551
+ ends = [self._global_ends_plus[d][c] for d,c in enumerate(coords)]
1552
+ send_shape = [e-s+1+2*m*p for s,e,m,p in zip(starts, ends, shifts_plus, pads_plus)]
1553
+ send_starts = [m*p for m,p in zip(shifts_plus, pads_plus)]
1554
+ m,p,s,e = shifts_plus[axis], pads_plus[axis], starts[axis], ends[axis]
1555
+ send_starts[axis] = m*p if ext_plus == -1 else m*p+e-s+1-p-1+diff
1556
+ starts[axis] = starts[axis] if ext_plus == -1 else ends[axis]-pads_plus[axis]+diff
1557
+ ends[axis] = starts[axis]+pads_plus[axis]-diff if ext_plus == -1 else ends[axis]
1558
+ send_buf_shape = [e-s+1 for s,e,p,m in zip(starts, ends, pads_plus, shifts_plus)]
1559
+
1560
+ # ...
1561
+ coords = self._coords_from_rank_minus[self._boundary_ranks_minus[0]]
1562
+ starts = [self._global_starts_minus[d][c] for d,c in enumerate(coords)]
1563
+ ends = [self._global_ends_minus[d][c] for d,c in enumerate(coords)]
1564
+
1565
+ recv_shape = [n+2*m*p for n,m,p in zip(npts_minus, shifts_minus, pads_minus)]
1566
+ recv_shape[axis] = pads_minus[axis]+1-diff + 2*shifts_minus[axis]*pads_minus[axis]
1567
+
1568
+ displacements = [0]*(len(self._boundary_ranks_minus)+1)
1569
+ recv_counts = [None]*len(self._boundary_ranks_minus)
1570
+ for k,b in enumerate(self._boundary_ranks_minus):
1571
+ coords = self._coords_from_rank_minus[b]
1572
+ starts = [self._global_starts_minus[d][c] for d,c in enumerate(coords)]
1573
+ ends = [self._global_ends_minus[d][c] for d,c in enumerate(coords)]
1574
+ starts[axis] = starts[axis] if ext_minus == -1 else ends[axis]-pads_minus[axis]+diff
1575
+ ends[axis] = starts[axis]+pads_minus[axis]-diff if ext_minus == -1 else ends[axis]
1576
+ shape_k = [e-s+1 for s,e in zip(starts, ends)]
1577
+ recv_counts[k] = np.prod(shape_k)
1578
+ ranges = [(s+p*m, p*m+e+1) for s,e,p,m in zip(starts, ends, pads_minus, shifts_minus)]
1579
+ ranges[axis] = (shifts_minus[axis]*pads_minus[axis], shifts_minus[axis]*pads_minus[axis]+shape_k[axis])
1580
+ indices += [np.ravel_multi_index( ii, dims=recv_shape, order='C' ) for ii in product(*[range(*a) for a in ranges])]
1581
+
1582
+ displacements[1:] = np.cumsum(recv_counts)
1583
+ # Store all information into dictionary
1584
+ info = {'send_buf_shape' : tuple( send_buf_shape ),
1585
+ 'send_starts' : tuple( send_starts ),
1586
+ 'send_shape' : tuple( send_shape ),
1587
+ 'recv_shape' : tuple( recv_shape ),
1588
+ 'displacements' : tuple( displacements ),
1589
+ 'recv_counts' : tuple( recv_counts),
1590
+ 'indices' : indices}
1591
+
1592
+ return info
1593
+ #---------------------------------------------------------------------------
1594
+ def _compute_interface_communication_infos_p2p( self, axis , get_minus_starts_ends=None, get_plus_starts_ends=None):
1595
+
1596
+ if self._intercomm == MPI.COMM_NULL:
1597
+ return
1598
+
1599
+ # Mesh info
1600
+ npts_minus = self._npts_minus
1601
+ npts_plus = self._npts_plus
1602
+ p_npts_minus = self._parent_npts_minus
1603
+ p_npts_plus = self._parent_npts_plus
1604
+ pads_minus = self._pads_minus
1605
+ pads_plus = self._pads_plus
1606
+ shifts_minus = self._shifts_minus
1607
+ shifts_plus = self._shifts_plus
1608
+ ext_minus = self._ext_minus
1609
+ ext_plus = self._ext_plus
1610
+ indices = []
1611
+
1612
+ if get_minus_starts_ends is not None:
1613
+ self._get_minus_starts_ends = get_minus_starts_ends
1614
+
1615
+ if get_plus_starts_ends is not None:
1616
+ self._get_plus_starts_ends = get_plus_starts_ends
1617
+
1618
+ diff = 0
1619
+ if p_npts_minus[axis] is not None:
1620
+ diff = min(1,p_npts_minus[axis]-npts_minus[axis])
1621
+
1622
+ if self._local_rank_minus is not None:
1623
+ rank_minus = self._local_rank_minus
1624
+ coords = self._coords_from_rank_minus[rank_minus]
1625
+ starts_minus = [self._global_starts_minus[d][c] for d,c in enumerate(coords)]
1626
+ ends_minus = [self._global_ends_minus[d][c] for d,c in enumerate(coords)]
1627
+ starts_extended_minus = [s-m*p for s,m,p in zip(starts_minus, shifts_minus, pads_minus)]
1628
+ ends_extended_minus = [min(n-1,e+m*p) for e,m,p,n in zip(ends_minus, shifts_minus, pads_minus, npts_minus)]
1629
+ buf_shape = [e-s+1+2*m*p for s,e,m,p in zip(starts_minus, ends_minus, shifts_minus, pads_minus)]
1630
+ dest_ranks = []
1631
+ buf_send_shape = []
1632
+ gbuf_send_shape = []
1633
+ gbuf_send_starts = []
1634
+ for k,rank_plus in enumerate(self._boundary_ranks_plus):
1635
+ coords = self._coords_from_rank_plus[rank_plus]
1636
+ starts = [self._global_starts_plus[d][c] for d,c in enumerate(coords)]
1637
+ ends = [self._global_ends_plus[d][c] for d,c in enumerate(coords)]
1638
+ starts_m, ends_m = self._get_minus_starts_ends(starts, ends, npts_minus, npts_plus, axis, axis,
1639
+ ext_minus, ext_plus, pads_minus, pads_plus, shifts_minus, shifts_plus, diff)
1640
+ starts_inter = [max(s1,s2) for s1,s2 in zip(starts_minus, starts_m)]
1641
+ ends_inter = [min(e1,e2) for e1,e2 in zip(ends_minus, ends_m)]
1642
+ if any(s>e if i!=axis else False for i,(s,e) in enumerate(zip(starts_inter, ends_inter))):
1643
+ continue
1644
+
1645
+ starts_inter[axis] = starts_minus[axis] if ext_minus == -1 else ends_minus[axis]-pads_minus[axis]+diff
1646
+ ends_inter[axis] = starts_minus[axis]+pads_minus[axis]-diff if ext_minus == -1 else ends_minus[axis]
1647
+
1648
+ dest_ranks.append(self._local_boundary_ranks_plus[k])
1649
+ buf_send_shape.append([e-s+1 for s,e in zip(starts_inter, ends_inter)])
1650
+ gbuf_send_shape.append(buf_shape)
1651
+ gbuf_send_starts.append([si-s+m*p for si,s,m,p in zip(starts_inter, starts_minus, shifts_minus, pads_minus)])
1652
+
1653
+ buf_shape = [e-s+1+2*m*p for s,e,m,p in zip(starts_minus, ends_minus, shifts_plus, pads_plus)]
1654
+ buf_shape[axis] = 2*shifts_plus[axis]*pads_plus[axis] + pads_plus[axis]+1-diff
1655
+ source_ranks = []
1656
+ buf_recv_shape = []
1657
+ gbuf_recv_shape = []
1658
+ gbuf_recv_starts = []
1659
+ for k,rank_plus in enumerate(self._boundary_ranks_plus):
1660
+ coords = self._coords_from_rank_plus[rank_plus]
1661
+ starts = [self._global_starts_plus[d][c] for d,c in enumerate(coords)]
1662
+ ends = [self._global_ends_plus[d][c] for d,c in enumerate(coords)]
1663
+
1664
+ starts_inter = [max(s1,s2) for s1,s2 in zip(starts_extended_minus, starts)]
1665
+ ends_inter = [min(e1,e2) for e1,e2 in zip(ends_extended_minus, ends)]
1666
+ if any(s>e if i!=axis else False for i,(s,e) in enumerate(zip(starts_inter, ends_inter))):
1667
+ continue
1668
+
1669
+ starts_extended_minus[axis] = 0
1670
+ starts_inter[axis] = shifts_minus[axis]*pads_minus[axis]
1671
+ ends_inter[axis] = shifts_minus[axis]*pads_minus[axis] + pads_minus[axis]-diff
1672
+
1673
+ source_ranks.append(self._local_boundary_ranks_plus[k])
1674
+ buf_recv_shape.append([e-s+1 for s,e in zip(starts_inter, ends_inter)])
1675
+ gbuf_recv_shape.append(buf_shape)
1676
+ gbuf_recv_starts.append([si-s for si,s in zip(starts_inter, starts_extended_minus)])
1677
+
1678
+ elif self._local_rank_plus is not None:
1679
+ rank_plus = self._local_rank_plus
1680
+ coords = self._coords_from_rank_plus[rank_plus]
1681
+ starts_plus = [self._global_starts_plus[d][c] for d,c in enumerate(coords)]
1682
+ ends_plus = [self._global_ends_plus[d][c] for d,c in enumerate(coords)]
1683
+ starts_extended_plus = [s-m*p for s,m,p in zip(starts_plus, shifts_plus, pads_plus)]
1684
+ ends_extended_plus = [min(n-1, e+m*p) for e,m,p,n in zip(ends_plus, shifts_plus, pads_plus, npts_plus)]
1685
+ buf_shape = [e-s+1+2*m*p for s,e,m,p in zip(starts_plus, ends_plus, shifts_plus, pads_plus)]
1686
+ dest_ranks = []
1687
+ buf_send_shape = []
1688
+ gbuf_send_shape = []
1689
+ gbuf_send_starts = []
1690
+
1691
+ for k,rank_minus in enumerate(self._boundary_ranks_minus):
1692
+ coords = self._coords_from_rank_minus[rank_minus]
1693
+ starts = [self._global_starts_minus[d][c] for d,c in enumerate(coords)]
1694
+ ends = [self._global_ends_minus[d][c] for d,c in enumerate(coords)]
1695
+ starts_p, ends_p = self._get_plus_starts_ends(starts, ends, npts_minus, npts_plus, axis, axis,
1696
+ ext_minus, ext_plus, pads_minus, pads_plus, shifts_minus, shifts_plus, diff)
1697
+ starts_inter = [max(s1,s2) for s1,s2 in zip(starts_plus, starts_p)]
1698
+ ends_inter = [min(e1,e2) for e1,e2 in zip(ends_plus, ends_p)]
1699
+ if any(s>e if i!=axis else False for i,(s,e) in enumerate(zip(starts_inter, ends_inter))):
1700
+ continue
1701
+
1702
+ starts_inter[axis] = starts_plus[axis] if ext_plus == -1 else ends_plus[axis]-pads_plus[axis]+diff
1703
+ ends_inter[axis] = starts_plus[axis]+pads_plus[axis]-diff if ext_plus == -1 else ends_plus[axis]
1704
+
1705
+ dest_ranks.append(self._local_boundary_ranks_minus[k])
1706
+ buf_send_shape.append([e-s+1 for s,e in zip(starts_inter, ends_inter)])
1707
+ gbuf_send_shape.append(buf_shape)
1708
+ gbuf_send_starts.append([si-s+m*p for si,s,m,p in zip(starts_inter, starts_plus, shifts_plus, pads_plus)])
1709
+
1710
+ buf_shape = [e-s+1+2*m*p for s,e,m,p in zip(starts_plus, ends_plus, shifts_minus, pads_minus)]
1711
+ buf_shape[axis] = 2*shifts_minus[axis]*pads_minus[axis] + pads_minus[axis]+1-diff
1712
+ source_ranks = []
1713
+ buf_recv_shape = []
1714
+ gbuf_recv_shape = []
1715
+ gbuf_recv_starts = []
1716
+
1717
+ for k,rank_minus in enumerate(self._boundary_ranks_minus):
1718
+ coords = self._coords_from_rank_minus[rank_minus]
1719
+ starts = [self._global_starts_minus[d][c] for d,c in enumerate(coords)]
1720
+ ends = [self._global_ends_minus[d][c] for d,c in enumerate(coords)]
1721
+
1722
+ starts_inter = [max(s1,s2) for s1,s2 in zip(starts_extended_plus, starts)]
1723
+ ends_inter = [min(e1,e2) for e1,e2 in zip(ends_extended_plus, ends)]
1724
+ if any(s>e if i!=axis else False for i,(s,e) in enumerate(zip(starts_inter, ends_inter))):
1725
+ continue
1726
+
1727
+ starts_extended_plus[axis] = 0
1728
+ starts_inter[axis] = shifts_plus[axis]*pads_plus[axis]
1729
+ ends_inter[axis] = shifts_plus[axis]*pads_plus[axis] + pads_plus[axis]-diff
1730
+
1731
+ source_ranks.append(self._local_boundary_ranks_minus[k])
1732
+ buf_recv_shape.append([e-s+1 for s,e in zip(starts_inter, ends_inter)])
1733
+ gbuf_recv_shape.append(buf_shape)
1734
+ gbuf_recv_starts.append([si-s for si,s in zip(starts_inter, starts_extended_plus)])
1735
+
1736
+ # Store all information into dictionary
1737
+ info = {'dest_ranks' : tuple( dest_ranks ),
1738
+ 'buf_send_shape' : tuple( buf_send_shape ),
1739
+ 'gbuf_send_shape' : tuple( gbuf_send_shape ),
1740
+ 'gbuf_send_starts' : tuple( gbuf_send_starts ),
1741
+ 'source_ranks' : tuple( source_ranks ),
1742
+ 'buf_recv_shape' : tuple( buf_recv_shape),
1743
+ 'gbuf_recv_shape' : tuple( gbuf_recv_shape ),
1744
+ 'gbuf_recv_starts' : tuple( gbuf_recv_starts )
1745
+ }
1746
+
1747
+ return info
1748
+
1749
+ #===============================================================================
1750
+ def create_interfaces_cart(domain_decomposition, carts, interfaces, communication_info):
1751
+
1752
+ """
1753
+ This function Connects the Cartesian grids when they share an interface.
1754
+
1755
+ Parameters
1756
+ ----------
1757
+ domain_decomposition: MultiPatchDomainDecomposition
1758
+
1759
+ carts: list of CartDecomposition
1760
+ The cartesian decomposition of multiple grids.
1761
+
1762
+ interfaces: dict
1763
+ The connectivity of the grids.
1764
+ It contains the grids that share an interface along with their axes and extremities.
1765
+
1766
+ communication_info: tuple
1767
+ tuple of two functions that determines the communication info between two patches.
1768
+
1769
+ """
1770
+ assert isinstance(domain_decomposition, MultiPatchDomainDecomposition)
1771
+ assert isinstance(carts, (list, tuple))
1772
+
1773
+ periods = [cart.periods for cart in carts]
1774
+ num_threads = domain_decomposition.num_threads
1775
+ comm = domain_decomposition.comm
1776
+ global_group = comm.group
1777
+ local_groups = list(domain_decomposition.local_groups)
1778
+ rank_ranges = domain_decomposition.rank_ranges
1779
+ local_communicators = domain_decomposition.local_communicators
1780
+ owned_groups = domain_decomposition.owned_groups
1781
+
1782
+ interfaces_groups = {}
1783
+ interfaces_comm = {}
1784
+ interfaces_root_ranks = {}
1785
+ interfaces_carts = {}
1786
+
1787
+ for i,j in interfaces:
1788
+ interfaces_comm[i,j] = MPI.COMM_NULL
1789
+
1790
+ if i in owned_groups or j in owned_groups:
1791
+ if not local_groups[i]:
1792
+ local_groups[i] = global_group.Range_incl([[rank_ranges[i][0], rank_ranges[i][1], 1]])
1793
+ if not local_groups[j]:
1794
+ local_groups[j] = global_group.Range_incl([[rank_ranges[j][0], rank_ranges[j][1], 1]])
1795
+
1796
+ interfaces_groups[i,j] = local_groups[i].Union(local_groups[i], local_groups[j])
1797
+ interfaces_comm [i,j] = comm.Create_group(interfaces_groups[i,j])
1798
+ root_rank_i = local_groups[i].Translate_ranks([0], interfaces_groups[i,j])[0]
1799
+ root_rank_j = local_groups[j].Translate_ranks([0], interfaces_groups[i,j])[0]
1800
+ interfaces_root_ranks[i,j] = [root_rank_i, root_rank_j]
1801
+
1802
+ tag = lambda i,j,disp: (2+disp)*(i+j)
1803
+ dtype = find_mpi_type('int64')
1804
+
1805
+ for i,j in interfaces:
1806
+ req = []
1807
+ ranks_in_topo_i = None
1808
+ ranks_in_topo_j = None
1809
+ axis_i, ext_i = interfaces[i,j][0]
1810
+ axis_j, ext_j = interfaces[i,j][1]
1811
+ if interfaces_comm[i,j] != MPI.COMM_NULL:
1812
+ ranks_in_topo_i = domain_decomposition.domains[i].ranks_in_topo if i in owned_groups else np.full(local_groups[i].size, -1)
1813
+ ranks_in_topo_j = domain_decomposition.domains[j].ranks_in_topo if j in owned_groups else np.full(local_groups[j].size, -1)
1814
+
1815
+ if interfaces_comm[i,j].rank == interfaces_root_ranks[i,j][0]:
1816
+ req.append(interfaces_comm[i,j].Isend((ranks_in_topo_i, ranks_in_topo_i.size, dtype), interfaces_root_ranks[i,j][1], tag=tag(i,j,1)))
1817
+ req.append(interfaces_comm[i,j].Irecv((ranks_in_topo_j, ranks_in_topo_j.size, dtype), interfaces_root_ranks[i,j][1], tag=tag(i,j,-1)))
1818
+
1819
+ if interfaces_comm[i,j].rank == interfaces_root_ranks[i,j][1]:
1820
+ req.append(interfaces_comm[i,j].Isend((ranks_in_topo_j, ranks_in_topo_j.size, dtype), interfaces_root_ranks[i,j][0], tag=tag(i,j,-1)))
1821
+ req.append(interfaces_comm[i,j].Irecv((ranks_in_topo_i, ranks_in_topo_i.size, dtype), interfaces_root_ranks[i,j][0], tag=tag(i,j,1)))
1822
+
1823
+ interfaces_carts[i,j] = InterfaceCartDecomposition(carts[i], carts[j],
1824
+ comm=interfaces_comm[i,j],
1825
+ axes=[axis_i, axis_j], exts=[ext_i, ext_j],
1826
+ ranks_in_topo=[ranks_in_topo_i, ranks_in_topo_j],
1827
+ local_groups=[local_groups[i], local_groups[j]],
1828
+ local_communicators=[local_communicators[i], local_communicators[j]],
1829
+ root_ranks=interfaces_root_ranks[i,j],
1830
+ requests=req)
1831
+ if not interfaces_carts[i,j].is_comm_null:
1832
+ interfaces_carts[i,j].set_interface_communication_infos(*communication_info)
1833
+
1834
+
1835
+ return interfaces_carts