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
@@ -0,0 +1,2923 @@
1
+ # coding: utf-8
2
+ #
3
+ # Copyright 2018 Yaman Güçlü
4
+
5
+ import os
6
+ import warnings
7
+
8
+ import numpy as np
9
+
10
+ from types import MappingProxyType
11
+ from scipy.sparse import coo_matrix, diags as sp_diags
12
+
13
+ from feectools.ddm.mpi import mpi as MPI
14
+ from feectools.linalg.basic import VectorSpace, Vector, LinearOperator
15
+ from feectools.ddm.cart import find_mpi_type, CartDecomposition, InterfaceCartDecomposition
16
+ from feectools.ddm.utilities import get_data_exchanger
17
+ from feectools.api.settings import PSYDAC_BACKENDS
18
+
19
+ from feectools.linalg.kernels.axpy_kernels import axpy_1d, axpy_2d, axpy_3d
20
+ from feectools.linalg.kernels.inner_kernels import inner_1d, inner_2d, inner_3d
21
+ from feectools.linalg.kernels.matvec_kernels import matvec_1d, matvec_2d, matvec_3d
22
+ from feectools.linalg.kernels.transpose_kernels import transpose_1d, transpose_2d, transpose_3d
23
+ from feectools.linalg.kernels.transpose_kernels import interface_transpose_1d, interface_transpose_2d, interface_transpose_3d
24
+ from feectools.linalg.kernels.stencil2coo_kernels import stencil2coo_1d_F, stencil2coo_2d_F, stencil2coo_3d_F
25
+ from feectools.linalg.kernels.stencil2coo_kernels import stencil2coo_1d_C, stencil2coo_2d_C, stencil2coo_3d_C
26
+
27
+
28
+ __all__ = (
29
+ 'StencilVectorSpace',
30
+ 'StencilVector',
31
+ 'StencilMatrix',
32
+ 'StencilInterfaceMatrix'
33
+ )
34
+
35
+ #===============================================================================
36
+ # Dictionary used to select correct kernel functions based on dimensionality
37
+ kernels = {
38
+ 'axpy' : (None, axpy_1d, axpy_2d, axpy_3d),
39
+ 'inner' : (None, inner_1d, inner_2d, inner_3d),
40
+ 'matvec': (None, matvec_1d, matvec_2d, matvec_3d),
41
+ 'transpose': (None, transpose_1d, transpose_2d, transpose_3d),
42
+ 'interface_transpose': (None, interface_transpose_1d, interface_transpose_2d, interface_transpose_3d),
43
+ 'stencil2coo': {'F': (None, stencil2coo_1d_F, stencil2coo_2d_F, stencil2coo_3d_F),
44
+ 'C': (None, stencil2coo_1d_C, stencil2coo_2d_C, stencil2coo_3d_C)}
45
+ }
46
+
47
+ #===============================================================================
48
+ def compute_diag_len(pads, shifts_domain, shifts_codomain, return_padding=False):
49
+ """
50
+ Compute the diagonal length and the padding of the stencil matrix for each direction,
51
+ using the shifts of the domain and the codomain.
52
+
53
+ Parameters
54
+ ----------
55
+ pads : tuple-like (int)
56
+ Padding along each direction.
57
+
58
+ shifts_domain : tuple_like (int)
59
+ Shifts of the domain along each direction.
60
+
61
+ shifts_codomain : tuple_like (int)
62
+ Shifts of the codomain along each direction.
63
+
64
+ return_padding : bool
65
+ Return the new padding if True.
66
+
67
+ Returns
68
+ -------
69
+ n : (int)
70
+ Diagonal length of the stencil matrix.
71
+
72
+ ep : (int)
73
+ Padding that constitutes the starting index of the non zero elements.
74
+ """
75
+ n = ((np.ceil((pads+1)/shifts_codomain)-1)*shifts_domain).astype('int')
76
+ ep = -np.minimum(0, n-pads)
77
+ n = n + ep + pads + 1
78
+ if return_padding:
79
+ return n.astype('int'), ep.astype('int')
80
+ else:
81
+ return n.astype('int')
82
+
83
+ #===============================================================================
84
+ class StencilVectorSpace(VectorSpace):
85
+ """
86
+ Vector space for n-dimensional stencil format. Two different initializations
87
+ are possible:
88
+
89
+ - serial : StencilVectorSpace(npts, pads, periods, shifts=None, starts=None, ends=None, dtype=float)
90
+ - parallel: StencilVectorSpace(cart, dtype=float)
91
+
92
+ Parameters
93
+ ----------
94
+ npts : tuple-like (int)
95
+ Number of entries along each direction
96
+ (= global dimensions of vector space).
97
+
98
+ pads : tuple-like (int)
99
+ Padding p along each direction needed for the ghost regions.
100
+
101
+ periods : tuple-like (bool)
102
+ Periodicity along each direction.
103
+
104
+ shifts : tuple-like (int)
105
+ shift m of the coefficients in each direction.
106
+
107
+ starts : tuple-like (int)
108
+ Index of the first coefficient local to the space in each direction.
109
+
110
+ ends : tuple-like (int)
111
+ Index of the last coefficient local to the space in each direction.
112
+
113
+ dtype : type
114
+ Type of scalar entries.
115
+
116
+ cart : feectools.ddm.cart.CartDecomposition
117
+ Tensor-product grid decomposition according to MPI Cartesian topology.
118
+
119
+ """
120
+
121
+ def __init__(self, cart, dtype=float):
122
+
123
+ assert isinstance(cart, (CartDecomposition, InterfaceCartDecomposition))
124
+
125
+ # Sequential attributes
126
+ self._parallel = cart.is_parallel
127
+ self._cart = cart
128
+ self._ndim = cart._ndims
129
+ self._npts = cart.npts
130
+ self._pads = cart.pads
131
+ self._periods = cart.periods
132
+ self._shifts = cart.shifts
133
+ self._dtype = dtype
134
+ self._starts = cart.starts
135
+ self._ends = cart.ends
136
+
137
+ # The shape of the allocated numpy array
138
+ self._shape = cart.shape
139
+ self._parent_starts = cart.parent_starts
140
+ self._parent_ends = cart.parent_ends
141
+ self._mpi_type = find_mpi_type(dtype)
142
+
143
+ # The dictionary follows the structure {(axis, ext): StencilVectorSpace()}
144
+ # where axis and ext represent the boundary shared by two patches
145
+ self._interfaces = {}
146
+ self._interfaces_readonly = MappingProxyType(self._interfaces)
147
+
148
+ # Parallel attributes
149
+ if cart.is_parallel and not cart.is_comm_null:
150
+ self._mpi_type = find_mpi_type(dtype)
151
+ if isinstance(cart, InterfaceCartDecomposition):
152
+ # TODO : Check if this line really change the ._shape
153
+ self._shape = cart.get_interface_communication_infos(cart.axis)['gbuf_recv_shape'][0]
154
+ else:
155
+ self._synchronizer = get_data_exchanger(cart, dtype , assembly=True, blocking=False)
156
+
157
+ # Select kernel for AXPY operation
158
+ if self._ndim in [1, 2, 3]:
159
+ self._axpy_func = kernels['axpy'][self._ndim]
160
+ else:
161
+ self._axpy_func = self._axpy_python
162
+ self._axpy_work = self.zeros() # work array
163
+
164
+ # Select kernel for inner product
165
+ if self._ndim in [1, 2, 3]:
166
+ self._inner_func = kernels['inner'][self._ndim]
167
+ else:
168
+ self._inner_func = self._inner_python
169
+
170
+ # Constant arguments for inner product: total number of ghost cells
171
+ self._inner_consts = tuple(np.int64(p * s) for p, s in zip(self._pads, self._shifts))
172
+
173
+ # TODO [YG, 06.09.2023]: print warning if pure Python functions are used
174
+
175
+ #--------------------------------------
176
+ # Pure Python methods for backup
177
+ #--------------------------------------
178
+ def _axpy_python(self, a, x, y):
179
+ w = self._axpy_work
180
+ x.copy(out=w) # w <- x
181
+ w *= a # w <- a * x
182
+ y += w # y <- a * x + y
183
+
184
+ @staticmethod
185
+ def _inner_python(v1, v2, nghost):
186
+ index = tuple(slice(ng, -ng) for ng in nghost)
187
+ return np.vdot(v1[index].flat, v2[index].flat)
188
+
189
+ #--------------------------------------
190
+ # Abstract interface
191
+ #--------------------------------------
192
+ @property
193
+ def dimension(self):
194
+ """ The dimension of a vector space V is the cardinality
195
+ (i.e. the number of vectors) of a basis of V over its base field.
196
+ """
197
+ return np.prod(self._npts)
198
+
199
+ # ...
200
+ @property
201
+ def dtype(self):
202
+ return self._dtype
203
+
204
+ # ...
205
+ def zeros(self):
206
+ """
207
+ Get a copy of the null element of the StencilVectorSpace V.
208
+
209
+ Returns
210
+ -------
211
+ null : StencilVector
212
+ A new vector object with all components equal to zero.
213
+
214
+ """
215
+ return StencilVector(self)
216
+
217
+ #...
218
+ def inner(self, x, y):
219
+ """
220
+ Evaluate the inner vector product between two vectors of this space V.
221
+
222
+ If the field of V is real, compute the classical scalar product.
223
+ If the field of V is complex, compute the classical sesquilinear
224
+ product with linearity on the second vector.
225
+
226
+ TODO [YG 01.05.2025]: Currently, the first vector is conjugated. We
227
+ want to reverse this behavior in order to align with the convention
228
+ of FEniCS.
229
+
230
+ Parameters
231
+ ----------
232
+ x : Vector
233
+ The first vector in the scalar product. In the case of a complex
234
+ field, the inner product is antilinear w.r.t. this vector (hence
235
+ this vector is conjugated).
236
+
237
+ y : Vector
238
+ The second vector in the scalar product. The inner product is
239
+ linear w.r.t. this vector.
240
+
241
+ Returns
242
+ -------
243
+ float | complex
244
+ The scalar product of the two vectors. Note that inner(x, x) is
245
+ a non-negative real number which is zero if and only if x = 0.
246
+
247
+ """
248
+
249
+ assert isinstance(x, StencilVector)
250
+ assert isinstance(y, StencilVector)
251
+ assert x.space is self
252
+ assert y.space is self
253
+
254
+ inner_func = self._inner_func
255
+ inner_args = (x._data, y._data, *self._inner_consts)
256
+
257
+ if self.parallel:
258
+ # Sometimes in the parallel case, we can get an empty vector that breaks our kernel
259
+ x._dot_send_data[0] = 0 if x._data.shape[0] == 0 else inner_func(*inner_args)
260
+ self.cart.global_comm.Allreduce((x._dot_send_data, self.mpi_type),
261
+ (x._dot_recv_data, self.mpi_type),
262
+ op=MPI.SUM )
263
+ return x._dot_recv_data[0]
264
+ else:
265
+ return inner_func(*inner_args)
266
+
267
+ # ...
268
+ def axpy(self, a, x, y):
269
+ """
270
+ Increment the vector y with the a-scaled vector x, i.e. y = a * x + y,
271
+ provided that x and y belong to the same vector space V (self).
272
+ The scalar value a may be real or complex, depending on the field of V.
273
+
274
+ Parameters
275
+ ----------
276
+ a : scalar
277
+ The scaling coefficient needed for the operation.
278
+
279
+ x : StencilVector
280
+ The vector which is not modified by this function.
281
+
282
+ y : StencilVector
283
+ The vector modified by this function (incremented by a * x).
284
+ """
285
+ assert isinstance(x, StencilVector)
286
+ assert isinstance(y, StencilVector)
287
+ assert x._space is self
288
+ assert y._space is self
289
+
290
+ if self.dtype == complex:
291
+ a = complex(a)
292
+ else:
293
+ if isinstance(a, complex):
294
+ raise TypeError('A complex scalar was given in a real case')
295
+ else:
296
+ a = float(a)
297
+
298
+ self._axpy_func(a, x._data, y._data)
299
+
300
+ for axis, ext in self.interfaces:
301
+ self._axpy_func(a, x._interface_data[axis, ext], y._interface_data[axis, ext])
302
+
303
+ x._sync = x._sync and y._sync
304
+
305
+ #--------------------------------------
306
+ # Other properties/methods
307
+ #--------------------------------------
308
+ @property
309
+ def mpi_type(self):
310
+ return self._mpi_type
311
+
312
+ @property
313
+ def shape(self):
314
+ return self._shape
315
+
316
+ @property
317
+ def parallel(self):
318
+ return self._parallel
319
+
320
+ # ...
321
+ @property
322
+ def cart(self):
323
+ return self._cart
324
+
325
+ # ...
326
+ @property
327
+ def npts(self):
328
+ return self._npts
329
+
330
+ # ...
331
+ @property
332
+ def starts(self):
333
+ return self._starts
334
+
335
+ # ...
336
+ @property
337
+ def ends(self):
338
+ return self._ends
339
+
340
+ # ...
341
+ @property
342
+ def parent_starts(self):
343
+ return self._parent_starts
344
+
345
+ # ...
346
+ @property
347
+ def parent_ends(self):
348
+ return self._parent_ends
349
+
350
+ # ...
351
+ @property
352
+ def pads(self):
353
+ return self._pads
354
+
355
+ # ...
356
+ @property
357
+ def periods(self):
358
+ return self._periods
359
+
360
+ # ...
361
+ @property
362
+ def shifts(self):
363
+ return self._shifts
364
+
365
+ # ...
366
+ @property
367
+ def ndim(self):
368
+ return self._ndim
369
+
370
+ @property
371
+ def interfaces(self):
372
+ return self._interfaces_readonly
373
+
374
+ def set_interface(self, axis, ext, cart):
375
+ """
376
+ Set the interface space along a given axis and extremity.
377
+
378
+ Parameters
379
+ ----------
380
+ axis : int
381
+ The axis of the new Interface space.
382
+
383
+ ext: {-1, 1}
384
+ The extremity of the new Interface space.
385
+
386
+ cart: CartDecomposition
387
+ The cart of the new space.
388
+ """
389
+
390
+ assert int(ext) in [-1, 1]
391
+ assert isinstance(cart, (CartDecomposition, InterfaceCartDecomposition))
392
+
393
+ if cart.is_comm_null:
394
+ return
395
+
396
+ # Create the interface space in the parallel case using the new cart
397
+ if isinstance(cart, InterfaceCartDecomposition):
398
+ # Case where the patches that share the interface are owned by different intra-communicators
399
+ space = StencilVectorSpace(cart, dtype=self.dtype)
400
+ self._interfaces[axis, ext] = space
401
+ else:
402
+ # Case where the patches that share the interface are owned by the same intra-communicator
403
+ if self.parent_ends[axis] is not None:
404
+ diff = min(1,self.parent_ends[axis]-self.ends[axis])
405
+ else:
406
+ diff = 0
407
+
408
+ starts = list(cart._starts)
409
+ ends = list(cart._ends)
410
+ parent_starts = list(cart._parent_starts)
411
+ parent_ends = list(cart._parent_ends)
412
+ if ext == 1:
413
+ starts[axis] = self.ends[axis]-self.pads[axis]+diff
414
+ if parent_starts[axis] is not None:
415
+ parent_starts[axis] = parent_ends[axis]-self.pads[axis]
416
+ else:
417
+ ends[axis] = self.pads[axis]-diff
418
+ if parent_ends[axis] is not None:
419
+ parent_ends[axis] = self.pads[axis]
420
+
421
+ cart = cart.change_starts_ends(tuple(starts), tuple(ends), tuple(parent_starts), tuple(parent_ends))
422
+
423
+ #TODO Check if we create object from it, otherwise its only purpose is to store some parameters which is innefficient
424
+ space = StencilVectorSpace(cart, self.dtype)
425
+
426
+ self._interfaces[axis, ext] = space
427
+
428
+ #===============================================================================
429
+ class StencilVector(Vector):
430
+ """
431
+ Vector in n-dimensional stencil format.
432
+
433
+ Parameters
434
+ ----------
435
+ V : feectools.linalg.stencil.StencilVectorSpace
436
+ Space to which the new vector belongs.
437
+
438
+ """
439
+ def __init__(self, V):
440
+
441
+ assert isinstance(V, StencilVectorSpace)
442
+
443
+ self._space = V
444
+ self._sizes = V.shape
445
+ self._ndim = len(V.npts)
446
+ self._data = np.zeros(V.shape, dtype=V.dtype)
447
+ self._dot_send_data = np.zeros((1,), dtype=V.dtype)
448
+ self._dot_recv_data = np.zeros((1,), dtype=V.dtype)
449
+ self._interface_data = {}
450
+ self._requests = None
451
+
452
+ # allocate data for the boundary that shares an interface
453
+ for axis, ext in V.interfaces:
454
+ self._interface_data[axis, ext] = np.zeros(V.interfaces[axis, ext].shape, dtype=V.dtype)
455
+
456
+ #prepare communications
457
+ if V.cart.is_parallel and not V.cart.is_comm_null and isinstance(V.cart, CartDecomposition):
458
+ self._requests = V._synchronizer.prepare_communications(self._data)
459
+
460
+ # TODO: distinguish between different directions
461
+ self._sync = False
462
+
463
+ #...
464
+ def __del__(self):
465
+ # Release memory of persistent MPI communication channels
466
+ if self._requests:
467
+ for request in self._requests:
468
+ request.Free()
469
+
470
+ #--------------------------------------
471
+ # Abstract interface
472
+ #--------------------------------------
473
+ @property
474
+ def space(self):
475
+ return self._space
476
+
477
+ # ...
478
+ def toarray(self, *, order='C', with_pads=False):
479
+ """
480
+ Return a numpy 1D array corresponding to the given StencilVector,
481
+ with or without pads.
482
+
483
+ Parameters
484
+ ----------
485
+ with_pads : bool
486
+ If True, include pads in output array (ignored in serial case).
487
+
488
+ order: {'C','F'}
489
+ Memory representation of the data ‘C’ for row-major ordering (C-style), ‘F’ column-major ordering (Fortran-style).
490
+
491
+ Returns
492
+ -------
493
+ array : numpy.ndarray
494
+ A copy of the data array collapsed into one dimension.
495
+
496
+ """
497
+
498
+ # In parallel case, call different functions based on 'with_pads' flag
499
+ if self.space.parallel:
500
+ if with_pads:
501
+ return self._toarray_parallel_with_pads(order=order)
502
+ else:
503
+ return self._toarray_parallel_no_pads(order=order)
504
+
505
+ # In serial case, ignore 'with_pads' flag
506
+ return self.toarray_local(order=order)
507
+
508
+ #...
509
+ def copy(self, out=None):
510
+ if self is out:
511
+ return self
512
+ w = out or StencilVector( self._space )
513
+ np.copyto(w._data, self._data, casting='no')
514
+ for axis, ext in self._space.interfaces:
515
+ np.copyto(w._interface_data[axis, ext], self._interface_data[axis, ext], casting='no')
516
+ w._sync = self._sync
517
+ return w
518
+
519
+ #...
520
+ def conjugate(self, out=None):
521
+ if out is not None:
522
+ assert isinstance(out, StencilVector)
523
+ assert out.space is self.space
524
+ else:
525
+ out = StencilVector(self.space)
526
+ np.conjugate(self._data, out=out._data, casting='no')
527
+ for axis, ext in self._space.interfaces:
528
+ np.conjugate(self._interface_data[axis, ext], out=out._interface_data[axis, ext], casting='no')
529
+ out._sync = self._sync
530
+ return out
531
+
532
+ #...
533
+ def __neg__(self):
534
+ w = StencilVector( self._space )
535
+ np.negative(self._data, out=w._data)
536
+ for axis, ext in self._space.interfaces:
537
+ np.negative(self._interface_data[axis, ext], out=w._interface_data[axis, ext])
538
+ w._sync = self._sync
539
+ return w
540
+
541
+ #...
542
+ def __mul__(self, a):
543
+ w = StencilVector( self._space )
544
+ np.multiply(self._data, a, out=w._data)
545
+ for axis, ext in self._space.interfaces:
546
+ np.multiply(self._interface_data[axis, ext], a, out=w._interface_data[axis, ext])
547
+ w._sync = self._sync
548
+ return w
549
+
550
+ #...
551
+ def __add__(self, v):
552
+ assert isinstance( v, StencilVector )
553
+ assert v._space is self._space
554
+ w = StencilVector( self._space )
555
+ np.add(self._data, v._data, out=w._data)
556
+ for axis, ext in self._space.interfaces:
557
+ np.add(self._interface_data[axis, ext], v._interface_data[axis, ext], out=w._interface_data[axis, ext])
558
+ w._sync = self._sync and v._sync
559
+ return w
560
+
561
+ #...
562
+ def __sub__(self, v):
563
+ assert isinstance( v, StencilVector )
564
+ assert v._space is self._space
565
+ w = StencilVector( self._space )
566
+ np.subtract(self._data, v._data, out=w._data)
567
+ for axis, ext in self._space.interfaces:
568
+ np.subtract(self._interface_data[axis, ext], v._interface_data[axis, ext], out=w._interface_data[axis, ext])
569
+ w._sync = self._sync and v._sync
570
+ return w
571
+
572
+ #...
573
+ def __imul__(self, a):
574
+ self._data *= a
575
+ for axis, ext in self._space.interfaces:
576
+ self._interface_data[axis, ext] *= a
577
+ return self
578
+
579
+ #...
580
+ def __iadd__(self, v):
581
+ assert isinstance( v, StencilVector )
582
+ assert v._space is self._space
583
+ self._data += v._data
584
+ for axis, ext in self._space.interfaces:
585
+ self._interface_data[axis, ext] += v._interface_data[axis, ext]
586
+ self._sync = v._sync and self._sync
587
+ return self
588
+
589
+ #...
590
+ def __isub__(self, v):
591
+ assert isinstance( v, StencilVector )
592
+ assert v._space is self._space
593
+ self._data -= v._data
594
+ for axis, ext in self._space.interfaces:
595
+ self._interface_data[axis, ext] -= v._interface_data[axis, ext]
596
+ self._sync = v._sync and self._sync
597
+ return self
598
+
599
+ #--------------------------------------
600
+ # Other properties/methods
601
+ #--------------------------------------
602
+ @property
603
+ def starts(self):
604
+ return self._space.starts
605
+
606
+ # ...
607
+ @property
608
+ def ends(self):
609
+ return self._space.ends
610
+
611
+ # ...
612
+ @property
613
+ def pads(self):
614
+ return self._space.pads
615
+
616
+ # ...
617
+ def __str__(self):
618
+ txt = '\n'
619
+ txt += '> starts :: {starts}\n'.format( starts= self.starts )
620
+ txt += '> ends :: {ends}\n' .format( ends = self.ends )
621
+ txt += '> pads :: {pads}\n' .format( pads = self.pads )
622
+ txt += '> data :: {data}\n' .format( data = self._data )
623
+ txt += '> sync :: {sync}\n' .format( sync = self._sync )
624
+ return txt
625
+
626
+ # ...
627
+ def toarray_local(self , *, order='C'):
628
+ """ return the local array without the padding"""
629
+ idx = tuple( slice(m*p,-m*p) if p != 0 else slice(0, None) for p,m in zip(self.pads, self.space.shifts) )
630
+ return self._data[idx].flatten( order=order)
631
+
632
+ # ...
633
+ def _toarray_parallel_no_pads(self, order='C'):
634
+ a = np.zeros( self.space.npts, self.dtype )
635
+ idx_from = tuple( slice(m*p,-m*p) if p != 0 else slice(0, None) for p,m in zip(self.pads, self.space.shifts) )
636
+ idx_to = tuple( slice(s,e+1) for s,e in zip(self.starts,self.ends) )
637
+ a[idx_to] = self._data[idx_from]
638
+ return a.flatten( order=order)
639
+
640
+ # ...
641
+ def _toarray_parallel_with_pads(self, order='C'):
642
+
643
+ pads = [m*p for m,p in zip(self.space.shifts, self.pads)]
644
+ # Step 0: create extended n-dimensional array with zero values
645
+ shape = tuple( n+2*p for n,p in zip( self.space.npts, pads ) )
646
+ a = np.zeros( shape, self.dtype )
647
+
648
+ # Step 1: write extended data chunk (local to process) onto array
649
+ idx = tuple( slice(s,e+2*p+1) for s,e,p in
650
+ zip( self.starts, self.ends, pads) )
651
+ a[idx] = self._data
652
+
653
+ # Step 2: if necessary, apply periodic boundary conditions to array
654
+ ndim = self.space.ndim
655
+
656
+ for direction in range( ndim ):
657
+
658
+ periodic = self.space.cart.periods[direction]
659
+ coord = self.space.cart.coords [direction]
660
+ nproc = self.space.cart.nprocs [direction]
661
+
662
+ if periodic:
663
+
664
+ p = pads[direction]
665
+
666
+ if p == 0:
667
+ continue
668
+
669
+ # Left-most process: copy data from left to right
670
+ if coord == 0:
671
+ idx_from = tuple(
672
+ (slice(None,p) if d == direction else slice(None))
673
+ for d in range( ndim )
674
+ )
675
+ idx_to = tuple(
676
+ (slice(-2*p,-p) if d == direction else slice(None))
677
+ for d in range( ndim )
678
+ )
679
+ a[idx_to] = a[idx_from]
680
+
681
+ # Right-most process: copy data from right to left
682
+ if coord == nproc-1:
683
+ idx_from = tuple(
684
+ (slice(-p,None) if d == direction else slice(None))
685
+ for d in range( ndim )
686
+ )
687
+ idx_to = tuple(
688
+ (slice(p,2*p) if d == direction else slice(None))
689
+ for d in range( ndim )
690
+ )
691
+ a[idx_to] = a[idx_from]
692
+
693
+ # Step 3: remove ghost regions from global array
694
+ idx = tuple( slice(p,-p) if p != 0 else slice(0, None) for p in pads )
695
+ out = a[idx]
696
+
697
+ # Step 4: return flattened array
698
+ return out.flatten( order=order)
699
+
700
+ #...
701
+ def topetsc(self):
702
+ """ Convert to petsc data structure.
703
+ """
704
+ from feectools.linalg.topetsc import vec_topetsc
705
+ vec = vec_topetsc( self )
706
+ return vec
707
+
708
+ # ...
709
+ def __getitem__(self, key):
710
+ index = self._getindex(key)
711
+ return self._data[index]
712
+
713
+ # ...
714
+ def __setitem__(self, key, value):
715
+ index = self._getindex(key)
716
+ self._data[index] = value
717
+
718
+ # ...
719
+ @property
720
+ def ghost_regions_in_sync(self):
721
+ return self._sync
722
+
723
+ # ...
724
+ # NOTE: this property must be set collectively
725
+ @ghost_regions_in_sync.setter
726
+ def ghost_regions_in_sync(self, value):
727
+ assert isinstance(value, bool)
728
+ self._sync = value
729
+
730
+ # ...
731
+ # TODO: maybe change name to 'exchange'
732
+ def update_ghost_regions(self):
733
+ """
734
+ Update ghost regions before performing non-local access to vector
735
+ elements (e.g. in matrix-vector product).
736
+
737
+ Parameters
738
+ ----------
739
+ direction : int
740
+ Single direction along which to operate (if not specified, all of them).
741
+
742
+ """
743
+
744
+ # Update interior ghost regions
745
+ if self.space.parallel:
746
+ if not self.space.cart.is_comm_null:
747
+ # PARALLEL CASE: fill in ghost regions with data from neighbors
748
+ self.space._synchronizer.start_update_ghost_regions(self._data, self._requests)
749
+ self.space._synchronizer. end_update_ghost_regions(self._data, self._requests)
750
+ else:
751
+ # SERIAL CASE: fill in ghost regions along periodic directions, otherwise set to zero
752
+ self._update_ghost_regions_serial()
753
+
754
+ # Update interface ghost regions
755
+ if self.space.parallel:
756
+
757
+ for axis, ext in self.space.interfaces:
758
+ V = self.space.interfaces[axis, ext]
759
+ if isinstance(V.cart, InterfaceCartDecomposition):
760
+ continue
761
+ slices = [slice(s, e+2*m*p+1) for s,e,m,p in zip(V.starts, V.ends, V.shifts, V.pads)]
762
+ self._interface_data[axis, ext][...] = self._data[tuple(slices)]
763
+ else:
764
+
765
+ for axis, ext in self.space.interfaces:
766
+ V = self.space.interfaces[axis, ext]
767
+ slices = [slice(s, e+2*m*p+1) for s,e,m,p in zip(V.starts, V.ends, V.shifts, V.pads)]
768
+ self._interface_data[axis, ext][...] = self._data[tuple(slices)]
769
+
770
+ # Flag ghost regions as up-to-date
771
+ self._sync = True
772
+
773
+ # ...
774
+ def _update_ghost_regions_serial(self):
775
+
776
+ ndim = self._space.ndim
777
+ for direction in range(ndim):
778
+ periodic = self._space.periods[direction]
779
+ p = self._space.pads [direction] * self._space.shifts[direction]
780
+
781
+ if p == 0:
782
+ continue
783
+
784
+ idx_front = [slice(None)] * direction
785
+ idx_back = [slice(None)] * (ndim-direction-1)
786
+
787
+ if periodic:
788
+ # Copy data from left to right
789
+ idx_from = tuple(idx_front + [slice( p, 2*p)] + idx_back)
790
+ idx_to = tuple(idx_front + [slice(-p,None)] + idx_back)
791
+ self._data[idx_to] = self._data[idx_from]
792
+
793
+ # Copy data from right to left
794
+ idx_from = tuple(idx_front + [slice(-2*p,-p)] + idx_back)
795
+ idx_to = tuple(idx_front + [slice(None, p)] + idx_back)
796
+ self._data[idx_to] = self._data[idx_from]
797
+
798
+ else:
799
+ # Set left ghost region to zero
800
+ idx_ghost = tuple(idx_front + [slice(None, p)] + idx_back)
801
+ self._data[idx_ghost] = 0
802
+
803
+ # Set right ghost region to zero
804
+ idx_ghost = tuple(idx_front + [slice(-p,None)] + idx_back)
805
+ self._data[idx_ghost] = 0
806
+
807
+ # ...
808
+ def exchange_assembly_data(self):
809
+ """
810
+ Exchange assembly data.
811
+ """
812
+
813
+ if self.space.parallel and not self.space.cart.is_comm_null:
814
+ # PARALLEL CASE: fill in ghost regions with data from neighbors
815
+ self.space._synchronizer.start_exchange_assembly_data(self._data)
816
+ self.space._synchronizer. end_exchange_assembly_data(self._data)
817
+ else:
818
+ # SERIAL CASE: fill in ghost regions along periodic directions, otherwise set to zero
819
+ self._exchange_assembly_data_serial()
820
+
821
+ ndim = self._space.ndim
822
+ for direction in range(ndim):
823
+ idx_front = [slice(None)] * direction
824
+ idx_back = [slice(None)] * (ndim-direction-1)
825
+
826
+ p = self._space.pads [direction]
827
+ m = self._space.shifts[direction]
828
+
829
+ if p == 0:
830
+ continue
831
+
832
+ idx_from = tuple(idx_front + [slice(-m*p,None) if (-m*p+p)!=0 else slice(-m*p,None)] + idx_back)
833
+ self._data[idx_from] = 0.
834
+ idx_from = tuple(idx_front + [slice(0,m*p)] + idx_back)
835
+ self._data[idx_from] = 0.
836
+
837
+ # ...
838
+ def _exchange_assembly_data_serial(self):
839
+
840
+ ndim = self._space.ndim
841
+ for direction in range(ndim):
842
+
843
+ periodic = self._space.periods[direction]
844
+ p = self._space.pads [direction]
845
+ m = self._space.shifts [direction]
846
+
847
+ if p == 0:
848
+ continue
849
+
850
+ if periodic:
851
+ idx_front = [slice(None)] * direction
852
+ idx_back = [slice(None)] * (ndim-direction-1)
853
+
854
+ # Copy data from left to right
855
+ idx_to = tuple(idx_front + [slice( m*p, m*p+p)] + idx_back)
856
+ idx_from = tuple(idx_front + [slice(-m*p,-m*p+p) if (-m*p+p)!=0 else slice(-m*p,None)] + idx_back)
857
+ self._data[idx_to] += self._data[idx_from]
858
+
859
+ #--------------------------------------
860
+ # Private methods
861
+ #--------------------------------------
862
+ def _getindex(self, key):
863
+
864
+ # TODO: check if we should ignore padding elements
865
+ if not isinstance(key, tuple):
866
+ key = (key,)
867
+ index = []
868
+ for (i,s,p,m) in zip(key, self.starts, self.pads,self.space.shifts):
869
+ if isinstance(i, slice):
870
+ start = None if i.start is None else i.start - s + m*p
871
+ stop = None if i.stop is None else i.stop - s + m*p
872
+ l = slice(start, stop, i.step)
873
+ else:
874
+ l = i - s + m*p
875
+ index.append(l)
876
+ return tuple(index)
877
+
878
+ #===============================================================================
879
+ class StencilMatrix(LinearOperator):
880
+ """
881
+ Matrix in n-dimensional stencil format.
882
+
883
+ This is a linear operator that maps elements of stencil vector space V to
884
+ elements of stencil vector space W.
885
+
886
+ For now we only accept V==W.
887
+
888
+ Parameters
889
+ ----------
890
+ V : feectools.linalg.stencil.StencilVectorSpace
891
+ Domain of the new linear operator.
892
+
893
+ W : feectools.linalg.stencil.StencilVectorSpace
894
+ Codomain of the new linear operator.
895
+
896
+ pads:
897
+
898
+ backend:
899
+
900
+ precompiled : bool
901
+ Whether to use precompiled kernels for .dot() and .transpose()
902
+ """
903
+ def __init__( self, V, W, pads=None , backend=None, precompiled=True):
904
+
905
+ assert isinstance(V, StencilVectorSpace)
906
+ assert isinstance(W, StencilVectorSpace)
907
+ assert W.pads == V.pads
908
+ if not W.dtype==V.dtype:
909
+ raise NotImplementedError("The domain and the codomain should have the same data type.")
910
+
911
+ if pads is not None:
912
+ for p,vp in zip(pads, V.pads):
913
+ assert p<=vp
914
+
915
+ self._pads = pads or tuple(V.pads)
916
+ dims = list(W.shape)
917
+ diags = [compute_diag_len(p, md, mc) for p,md,mc in zip(self._pads, V.shifts, W.shifts)]
918
+ self._data = np.zeros(dims+diags, dtype=W.dtype)
919
+ self._domain = V
920
+ self._codomain = W
921
+ self._ndim = len(dims)
922
+ self._backend = backend
923
+ self._precompiled = precompiled
924
+ self._is_T = False
925
+ self._diag_indices = None
926
+ self._requests = None
927
+
928
+ # Parallel attributes
929
+ if W.parallel:
930
+ if W.cart.is_comm_null:return
931
+ # Create data exchanger for ghost regions
932
+ self._synchronizer = get_data_exchanger(
933
+ cart = W.cart,
934
+ dtype = W.dtype,
935
+ coeff_shape = diags,
936
+ assembly = True
937
+ )
938
+
939
+ # Flag ghost regions as not up-to-date (conservative choice)
940
+ self._sync = False
941
+
942
+ # Prepare the arguments for the dot product method
943
+ nd = [(ej-sj+2*gp*mj-mj*p-gp)//mj*mi+1 for sj,ej,mj,mi,p,gp in zip(V.starts, V.ends, V.shifts, W.shifts, self._pads, V.pads)]
944
+ nc = [ei-si+1 for si,ei,mj,p in zip(W.starts, W.ends, V.shifts, self._pads)]
945
+
946
+ # Number of rows in matrix (along each dimension)
947
+ nrows = [min(ni, nj) for ni,nj in zip(nc, nd)]
948
+ nrows_extra = [max(0, ni-nj) for ni,nj in zip(nc, nd)]
949
+
950
+ args = {}
951
+ args['starts'] = tuple(V.starts)
952
+ args['nrows'] = tuple(nrows)
953
+ args['nrows_extra'] = tuple(nrows_extra)
954
+ args['gpads'] = tuple(V.pads)
955
+ args['pads'] = tuple(self._pads)
956
+ args['dm'] = tuple(V.shifts)
957
+ args['cm'] = tuple(W.shifts)
958
+ ndiags, _ = list(zip(*[compute_diag_len(p,mj,mi, return_padding=True) for p,mi,mj in zip(self._pads, W.shifts, V.shifts)]))
959
+ args['pad_imp'] = [gp*m+gp+1-n-s%m+p-gp for gp,m,n,s,p in zip(V.pads, V.shifts, ndiags, V.starts, self._pads)]
960
+ args['ndiags'] = ndiags
961
+
962
+ self._dotargs_null = args
963
+ self._dot = kernels['matvec'][self._ndim]
964
+
965
+ self._transpose_args = self._prepare_transpose_args()
966
+ self._transpose_func = kernels['transpose'][self._ndim]
967
+
968
+ if backend is None:
969
+ backend = PSYDAC_BACKENDS.get(os.environ.get('PSYDAC_BACKEND')) or PSYDAC_BACKENDS['python']
970
+ self.set_backend(backend, precompiled)
971
+
972
+ #--------------------------------------
973
+ # Abstract interface
974
+ #--------------------------------------
975
+ @property
976
+ def domain(self):
977
+ return self._domain
978
+
979
+ # ...
980
+ @property
981
+ def codomain(self):
982
+ return self._codomain
983
+
984
+ # ...
985
+ @property
986
+ def dtype(self):
987
+ return self._domain.dtype
988
+
989
+ # ...
990
+ def dot(self, v, out=None):
991
+ """
992
+ Return the matrix/vector product between self and v.
993
+ This function optimized this product.
994
+
995
+ Parameters
996
+ ----------
997
+ v : StencilVector
998
+ Vector of the domain of self needed for the Matrix/Vector product.
999
+
1000
+ out : StencilVector
1001
+ Vector of the codomain of self.
1002
+
1003
+ Returns
1004
+ -------
1005
+ out : StencilVector
1006
+ Vector of the codomain of self, contain the result of the product.
1007
+ """
1008
+
1009
+ assert isinstance(v, StencilVector)
1010
+ assert v.space is self.domain
1011
+
1012
+ if out is not None:
1013
+ assert isinstance( out, StencilVector )
1014
+ assert out.space is self.codomain
1015
+ else:
1016
+ out = StencilVector( self.codomain )
1017
+
1018
+ # Necessary if vector space is distributed across processes
1019
+ if not v.ghost_regions_in_sync:
1020
+ v.update_ghost_regions()
1021
+
1022
+ self._func(self._data, v._data, out._data, **self._args)
1023
+
1024
+ # IMPORTANT: flag that ghost regions are not up-to-date
1025
+ out.ghost_regions_in_sync = False
1026
+ return out
1027
+
1028
+ # ...
1029
+ def vdot( self, v, out=None):
1030
+ """
1031
+ Return the matrix/vector product between the conjugate of self and v.
1032
+ This function optimized this product.
1033
+
1034
+ Parameters
1035
+ ----------
1036
+ v : StencilVector
1037
+ Vector of the domain of self needed for the Matrix/Vector product
1038
+
1039
+ out : StencilVector
1040
+ Vector of the codomain of self
1041
+
1042
+ Returns
1043
+ -------
1044
+ out : StencilVector
1045
+ Vector of the codomain of self, contain the result of the product
1046
+ """
1047
+
1048
+ assert isinstance(v, StencilVector)
1049
+ assert v.space is self.domain
1050
+
1051
+ if out is not None:
1052
+ assert isinstance(out, StencilVector)
1053
+ assert out.space is self.codomain
1054
+ else:
1055
+ out = StencilVector( self.codomain )
1056
+
1057
+ # Necessary if vector space is distributed across processes
1058
+ if not v.ghost_regions_in_sync:
1059
+ v.update_ghost_regions()
1060
+
1061
+ # Instead of computing A_*x, this function computes (A*x_)_
1062
+ self._func(self._data, np.conjugate(v._data), out._data, **self._args)
1063
+ np.conjugate(out._data, out=out._data)
1064
+
1065
+ # IMPORTANT: flag that ghost regions are not up-to-date
1066
+ out.ghost_regions_in_sync = False
1067
+ return out
1068
+
1069
+ # ...
1070
+ def transpose(self, conjugate=False, out=None):
1071
+ """"
1072
+ Return the transposed StencilMatrix, or the Hermitian Transpose if conjugate==True
1073
+
1074
+ Parameters
1075
+ ----------
1076
+ conjugate : Bool(optional)
1077
+ True to get the Hermitian adjoint.
1078
+
1079
+ out : StencilMatrix(optional)
1080
+ Optional out for the transpose to avoid temporaries
1081
+ """
1082
+ # For clarity rename self
1083
+ M = self
1084
+
1085
+ # If necessary, update ghost regions of original matrix M
1086
+ if not M.ghost_regions_in_sync:
1087
+ M.update_ghost_regions()
1088
+
1089
+ # Create new matrix where domain and codomain are swapped
1090
+ if out is not None :
1091
+ assert isinstance(out, StencilMatrix)
1092
+ assert out.codomain == M.domain
1093
+ assert out.domain == M.codomain
1094
+
1095
+ else :
1096
+ out = StencilMatrix(M.codomain, M.domain, pads=self._pads, backend=self._backend, precompiled=self._precompiled)
1097
+
1098
+ # Call low-level '_transpose' function (works on Numpy arrays directly)
1099
+ if conjugate:
1100
+ self._transpose_func(np.conjugate(M._data), out._data, **self._transpose_args)
1101
+ else:
1102
+ self._transpose_func(M._data, out._data, **self._transpose_args)
1103
+ return out
1104
+
1105
+ # ...
1106
+ def toarray(self, **kwargs):
1107
+ """ Convert to Numpy 2D array. """
1108
+
1109
+ order = kwargs.pop('order', 'C')
1110
+ with_pads = kwargs.pop('with_pads', False)
1111
+
1112
+ if self.codomain.parallel and with_pads:
1113
+ coo = self._tocoo_parallel_with_pads(order=order)
1114
+ else:
1115
+ coo = self._tocoo_no_pads(order=order)
1116
+
1117
+ return coo.toarray()
1118
+
1119
+ # ...
1120
+ def tosparse(self, **kwargs):
1121
+ """ Convert to any Scipy sparse matrix format. """
1122
+
1123
+ order = kwargs.pop('order', 'C')
1124
+ with_pads = kwargs.pop('with_pads', False)
1125
+
1126
+ if self.codomain.parallel and with_pads:
1127
+ coo = self._tocoo_parallel_with_pads(order=order)
1128
+ else:
1129
+ coo = self._tocoo_no_pads(order=order)
1130
+
1131
+ return coo
1132
+
1133
+ #--------------------------------------
1134
+ # Overridden properties/methods
1135
+ #--------------------------------------
1136
+ def __neg__(self):
1137
+ return self.__mul__(-1)
1138
+
1139
+ # ...
1140
+ def __mul__(self, a):
1141
+ w = StencilMatrix(self._domain, self._codomain, self._pads, self._backend, precompiled=self._precompiled)
1142
+ w._data = self._data * a
1143
+ w._func = self._func
1144
+ w._args = self._args
1145
+ w._sync = self._sync
1146
+ return w
1147
+
1148
+ #...
1149
+ def __add__(self, m):
1150
+ if isinstance(m, StencilMatrix):
1151
+ #assert isinstance(m, StencilMatrix)
1152
+ assert m._domain is self._domain
1153
+ assert m._codomain is self._codomain
1154
+ assert m._pads == self._pads
1155
+
1156
+ if m._backend is not self._backend:
1157
+ msg = 'Adding two matrices with different backends is ambiguous - defaulting to backend of first addend'
1158
+ warnings.warn(msg, category=RuntimeWarning)
1159
+
1160
+ w = StencilMatrix(self._domain, self._codomain, self._pads, self._backend, precompiled=self._precompiled)
1161
+ w._data = self._data + m._data
1162
+ w._func = self._func
1163
+ w._args = self._args
1164
+ w._sync = self._sync and m._sync
1165
+ return w
1166
+ else:
1167
+ return LinearOperator.__add__(self, m)
1168
+
1169
+ #...
1170
+ def __sub__(self, m):
1171
+ if isinstance(m, StencilMatrix):
1172
+ #assert isinstance(m, StencilMatrix)
1173
+ assert m._domain is self._domain
1174
+ assert m._codomain is self._codomain
1175
+ assert m._pads == self._pads
1176
+
1177
+ if m._backend is not self._backend:
1178
+ msg = 'Subtracting two matrices with different backends is ambiguous - defaulting to backend of the matrix we subtract from'
1179
+ warnings.warn(msg, category=RuntimeWarning)
1180
+
1181
+ w = StencilMatrix(self._domain, self._codomain, self._pads, backend=self._backend, precompiled=self._precompiled)
1182
+ w._data = self._data - m._data
1183
+ w._func = self._func
1184
+ w._args = self._args
1185
+ w._sync = self._sync and m._sync
1186
+ return w
1187
+ else:
1188
+ return LinearOperator.__sub__(self, m)
1189
+
1190
+ #--------------------------------------
1191
+ # New properties/methods
1192
+ #--------------------------------------
1193
+
1194
+ # TODO: check if this method is really needed!!
1195
+ def conjugate(self, out=None):
1196
+ if out is not None:
1197
+ assert isinstance(out, StencilMatrix)
1198
+ assert out.domain is self.domain
1199
+ assert out.codomain is self.codomain
1200
+ else:
1201
+ out = StencilMatrix(self.domain, self.codomain, pads=self.pads, backend=self._backend, precompiled=self._precompiled)
1202
+ out._func = self._func
1203
+ out._args = self._args
1204
+ np.conjugate(self._data, out=out._data, casting='no')
1205
+ return out
1206
+
1207
+ # ...
1208
+ # TODO: check if this method is really needed!!
1209
+ def conj(self, out=None):
1210
+ return self.conjugate(out=out)
1211
+
1212
+ # ...
1213
+ @property
1214
+ def pads(self):
1215
+ return self._pads
1216
+
1217
+ # ...
1218
+ @property
1219
+ def backend(self):
1220
+ return self._backend
1221
+
1222
+ # ...
1223
+ def __getitem__(self, key):
1224
+ index = self._getindex( key )
1225
+ return self._data[index]
1226
+
1227
+ # ...
1228
+ def __setitem__(self, key, value):
1229
+ index = self._getindex( key )
1230
+ self._data[index] = value
1231
+
1232
+ #...
1233
+ def max(self):
1234
+ return self._data.max()
1235
+
1236
+ #...
1237
+ def copy(self, out = None):
1238
+ """
1239
+ Create a copy of self, that can potentially be stored in a given StencilMatrix.
1240
+
1241
+ Parameters
1242
+ ----------
1243
+ out : StencilMatrix(optional)
1244
+ The existing StencilMatrix in which we want to copy self.
1245
+ """
1246
+ if out is not None :
1247
+ assert isinstance(out, StencilMatrix)
1248
+ assert out.domain == self.domain
1249
+ assert out.codomain == self.codomain
1250
+ else :
1251
+ out = StencilMatrix( self.domain, self.codomain, self._pads, backend=self._backend, precompiled=self._precompiled )
1252
+ out._data[:] = self._data[:]
1253
+ out._func = self._func
1254
+ out._args = self._args
1255
+ return out
1256
+
1257
+ #...
1258
+ def __imul__(self, a):
1259
+ self._data *= a
1260
+ return self
1261
+
1262
+ #...
1263
+ def __iadd__(self, m):
1264
+ if isinstance(m, StencilMatrix):
1265
+ #assert isinstance(m, StencilMatrix)
1266
+ assert m._domain is self._domain
1267
+ assert m._codomain is self._codomain
1268
+ assert m._pads == self._pads
1269
+ self._data += m._data
1270
+ self._sync = m._sync and self._sync
1271
+ return self
1272
+ else:
1273
+ return LinearOperator.__add__(self, m)
1274
+
1275
+ #...
1276
+ def __isub__(self, m):
1277
+ if isinstance(m, StencilMatrix):
1278
+ #assert isinstance(m, StencilMatrix)
1279
+ assert m._domain is self._domain
1280
+ assert m._codomain is self._codomain
1281
+ assert m._pads == self._pads
1282
+ self._data -= m._data
1283
+ self._sync = m._sync and self._sync
1284
+ return self
1285
+ else:
1286
+ return LinearOperator.__sub__(self, m)
1287
+
1288
+ #...
1289
+ def __abs__(self):
1290
+ w = StencilMatrix( self._domain, self._codomain, self._pads, backend=self._backend, precompiled=self._precompiled )
1291
+ w._data = abs(self._data)
1292
+ w._func = self._func
1293
+ w._args = self._args
1294
+ w._sync = self._sync
1295
+ return w
1296
+
1297
+ #...
1298
+ def remove_spurious_entries(self):
1299
+ """
1300
+ If any dimension is NOT periodic, make sure that the corresponding
1301
+ periodic corners are set to zero.
1302
+
1303
+ """
1304
+ # TODO: access 'self._data' directly for increased efficiency
1305
+
1306
+ ndim = self._domain.ndim
1307
+
1308
+ for direction in range(ndim):
1309
+
1310
+ periodic = self._domain.periods[direction]
1311
+
1312
+ if not periodic:
1313
+
1314
+ nc = self._codomain.npts[direction]
1315
+ nd = self._domain.npts[direction]
1316
+
1317
+ s = self._codomain.starts[direction]
1318
+ e = self._codomain.ends [direction]
1319
+ p = self.pads [direction]
1320
+
1321
+ idx_front = [slice(None)]*direction
1322
+ idx_back = [slice(None)]*(ndim-direction-1)
1323
+
1324
+ # Top-right corner
1325
+ for i in range( max(0,s), min(p,e+1) ):
1326
+ index = tuple( idx_front + [i] + idx_back +
1327
+ idx_front + [slice(-p,-i)] + idx_back )
1328
+ self[index] = 0
1329
+
1330
+ # Bottom-left corner
1331
+ for i in range( max(nd-p,s), min(nc,e+1) ):
1332
+ index = tuple( idx_front + [i] + idx_back +
1333
+ idx_front + [slice(nd-i,p+1)] + idx_back )
1334
+ self[index] = 0
1335
+
1336
+ # ...
1337
+ def update_ghost_regions(self):
1338
+ """
1339
+ Update ghost regions before performing non-local access to matrix
1340
+ elements (e.g. in matrix transposition).
1341
+ """
1342
+ ndim = self._codomain.ndim
1343
+ parallel = self._codomain.parallel
1344
+
1345
+ if parallel:
1346
+ if not self._codomain.cart.is_comm_null:
1347
+ # PARALLEL CASE: fill in ghost regions with data from neighbors
1348
+ self._synchronizer.start_update_ghost_regions( self._data, self._requests )
1349
+ self._synchronizer.end_update_ghost_regions( self._data , self._requests)
1350
+ else:
1351
+ # SERIAL CASE: fill in ghost regions along periodic directions, otherwise set to zero
1352
+ self._update_ghost_regions_serial()
1353
+
1354
+ # Flag ghost regions as up-to-date
1355
+ self._sync = True
1356
+
1357
+ # ...
1358
+ def exchange_assembly_data(self):
1359
+ """
1360
+ Exchange assembly data.
1361
+ """
1362
+ ndim = self._codomain.ndim
1363
+ parallel = self._codomain.parallel
1364
+
1365
+ if self._codomain.parallel:
1366
+ # PARALLEL CASE: fill in ghost regions with data from neighbors
1367
+ self._synchronizer.start_exchange_assembly_data( self._data )
1368
+ self._synchronizer.end_exchange_assembly_data( self._data )
1369
+ else:
1370
+ # SERIAL CASE: fill in ghost regions along periodic directions, otherwise set to zero
1371
+ self._exchange_assembly_data_serial()
1372
+
1373
+ ndim = self._codomain.ndim
1374
+ for direction in range(ndim):
1375
+ idx_front = [slice(None)]*direction
1376
+ idx_back = [slice(None)]*(ndim-direction-1)
1377
+
1378
+ p = self._codomain.pads [direction]
1379
+ m = self._codomain.shifts[direction]
1380
+
1381
+ if p == 0:
1382
+ continue
1383
+
1384
+ idx_from = tuple( idx_front + [ slice(-m*p,None) if (-m*p+p)!=0 else slice(-m*p,None)] + idx_back )
1385
+ self._data[idx_from] = 0.
1386
+ idx_from = tuple( idx_front + [ slice(0,m*p)] + idx_back )
1387
+ self._data[idx_from] = 0.
1388
+
1389
+ # ...
1390
+ def _exchange_assembly_data_serial(self):
1391
+
1392
+ ndim = self._codomain.ndim
1393
+ for direction in range(ndim):
1394
+
1395
+ periodic = self._codomain.periods[direction]
1396
+ p = self._codomain.pads [direction]
1397
+ m = self._codomain.shifts[direction]
1398
+
1399
+ if p == 0:
1400
+ continue
1401
+
1402
+ if periodic:
1403
+ idx_front = [slice(None)]*direction
1404
+ idx_back = [slice(None)]*(ndim-direction-1)
1405
+
1406
+ # Copy data from left to right
1407
+ idx_to = tuple( idx_front + [slice( m*p, m*p+p)] + idx_back )
1408
+ idx_from = tuple( idx_front + [slice(-m*p,-m*p+p) if (-m*p+p)!=0 else slice(-m*p,None)] + idx_back )
1409
+ self._data[idx_to] += self._data[idx_from]
1410
+
1411
+ # ...
1412
+ def diagonal(self, *, inverse = False, sqrt = False, out = None):
1413
+ """
1414
+ Get the coefficients on the main diagonal as a StencilDiagonalMatrix object.
1415
+
1416
+ Parameters
1417
+ ----------
1418
+ inverse : bool
1419
+ If True, get the inverse of the diagonal. (Default: False).
1420
+ Can be combined with sqrt to get the inverse square root.
1421
+
1422
+ sqrt : bool
1423
+ If True, get the square root of the diagonal. (Default: False).
1424
+ Can be combined with inverse to get the inverse square root.
1425
+
1426
+ out : StencilDiagonalMatrix
1427
+ If provided, write the diagonal entries into this matrix. (Default: None).
1428
+
1429
+ Returns
1430
+ -------
1431
+ StencilDiagonalMatrix
1432
+ The matrix which contains the main diagonal of self (or its inverse).
1433
+
1434
+ """
1435
+ # Check `inverse` argument
1436
+ assert isinstance(inverse, bool)
1437
+
1438
+ # Determine domain and codomain of the StencilDiagonalMatrix
1439
+ V, W = self.domain, self.codomain
1440
+ if inverse:
1441
+ V, W = W, V
1442
+
1443
+ # Check `out` argument
1444
+ if out is not None:
1445
+ assert isinstance(out, StencilDiagonalMatrix)
1446
+ assert out.domain is V
1447
+ assert out.codomain is W
1448
+
1449
+
1450
+ # Extract diagonal data from self and identify output array
1451
+ diagonal_indices = self._get_diagonal_indices()
1452
+ diag = self._data[diagonal_indices]
1453
+ data = out._data if out else None
1454
+
1455
+ # Calculate entries of StencilDiagonalMatrix
1456
+ if inverse:
1457
+ data = np.divide(1, diag, out=data)
1458
+ elif out:
1459
+ np.copyto(data, diag)
1460
+ else:
1461
+ data = diag.copy()
1462
+
1463
+ if sqrt:
1464
+ np.sqrt(data, out=data)
1465
+
1466
+ # If needed create a new StencilDiagonalMatrix object
1467
+ if out is None:
1468
+ out = StencilDiagonalMatrix(V, W, data)
1469
+
1470
+ return out
1471
+
1472
+ # ...
1473
+ def topetsc(self):
1474
+ """ Convert to PETSc data structure.
1475
+ """
1476
+ from feectools.linalg.topetsc import mat_topetsc
1477
+ mat = mat_topetsc(self)
1478
+ return mat
1479
+
1480
+ #--------------------------------------
1481
+ # Private methods
1482
+ #--------------------------------------
1483
+
1484
+ def _getindex(self, key):
1485
+
1486
+ nd = self._ndim
1487
+ ii = key[:nd]
1488
+ kk = key[nd:]
1489
+
1490
+ index = []
1491
+
1492
+ for i,s,p,m in zip( ii, self._codomain.starts, self._codomain.pads, self._codomain.shifts ):
1493
+ x = self._shift_index( i, m*p-s )
1494
+ index.append( x )
1495
+
1496
+ for k,p in zip( kk, self._pads ):
1497
+ l = self._shift_index( k, p )
1498
+ index.append( l )
1499
+ return tuple(index)
1500
+
1501
+ # ...
1502
+ @staticmethod
1503
+ def _shift_index(index, shift):
1504
+ if isinstance( index, slice ):
1505
+ start = None if index.start is None else index.start + shift
1506
+ stop = None if index.stop is None else index.stop + shift
1507
+ return slice(start, stop, index.step)
1508
+ else:
1509
+ return index + shift
1510
+
1511
+ def tocoo_local(self, order='C'):
1512
+
1513
+ # Shortcuts
1514
+ sc = self._codomain.starts
1515
+ ec = self._codomain.ends
1516
+ pc = self._codomain.pads
1517
+
1518
+ sd = self._domain.starts
1519
+ ed = self._domain.ends
1520
+ pd = self._domain.pads
1521
+
1522
+ nd = self._ndim
1523
+
1524
+ nr = [e-s+1 +2*p for s,e,p in zip(sc, ec, pc)]
1525
+ nc = [e-s+1 +2*p for s,e,p in zip(sd, ed, pd)]
1526
+
1527
+ ravel_multi_index = np.ravel_multi_index
1528
+
1529
+ # COO storage
1530
+ rows = []
1531
+ cols = []
1532
+ data = []
1533
+
1534
+ local = tuple( [slice(p,-p) for p in pc] + [slice(None)] * nd )
1535
+
1536
+ dd = [pdi-ppi for pdi,ppi in zip(pd, self._pads)]
1537
+
1538
+ for (index, value) in np.ndenumerate( self._data[local] ):
1539
+
1540
+ # index = [i1-s1, i2-s2, ..., p1+j1-i1, p2+j2-i2, ...]
1541
+
1542
+ xx = index[:nd] # ii is local
1543
+ ll = index[nd:] # l=p+k
1544
+
1545
+ ii = [x+p for x,p in zip(xx, pc)]
1546
+ jj = [(l+i+d)%n for (i,l,d,n) in zip(xx,ll,dd,nc)]
1547
+
1548
+ I = ravel_multi_index( ii, dims=nr, order=order )
1549
+ J = ravel_multi_index( jj, dims=nc, order=order )
1550
+
1551
+ rows.append( I )
1552
+ cols.append( J )
1553
+ data.append( value )
1554
+
1555
+ M = coo_matrix(
1556
+ (data,(rows,cols)),
1557
+ shape = [np.prod(nr),np.prod(nc)],
1558
+ dtype = self._domain.dtype
1559
+ )
1560
+
1561
+ M.eliminate_zeros()
1562
+
1563
+ return M
1564
+
1565
+ #...
1566
+ def _tocoo_no_pads(self , order='C'):
1567
+
1568
+ # Shortcuts
1569
+ nr = self._codomain.npts
1570
+ nd = self._ndim
1571
+ nc = self._domain.npts
1572
+ ss = self._codomain.starts
1573
+ cpads = self._codomain.pads
1574
+ dm = self._domain.shifts
1575
+ cm = self._codomain.shifts
1576
+
1577
+ pp = [np.int64(compute_diag_len(p,mj,mi)-(p+1)) for p,mi,mj in zip(self._pads, cm, dm)]
1578
+
1579
+ # Range of data owned by local process (no ghost regions)
1580
+ local = tuple( [slice(mi*p,-mi*p) if p != 0 else slice(p, None) for p,mi in zip(cpads, cm)] + [slice(None)] * nd )
1581
+ size = self._data[local].size
1582
+
1583
+ # COO storage
1584
+ rows = np.zeros(size, dtype='int64')
1585
+ cols = np.zeros(size, dtype='int64')
1586
+ data = np.zeros(size, dtype=self.dtype)
1587
+ nrl = [np.int64(e-s+1) for s,e in zip(self.codomain.starts, self.codomain.ends)]
1588
+ ncl = [np.int64(i) for i in self._data.shape[nd:]]
1589
+ ss = [np.int64(i) for i in ss]
1590
+ nr = [np.int64(i) for i in nr]
1591
+ nc = [np.int64(i) for i in nc]
1592
+ dm = [np.int64(i) for i in dm]
1593
+ cm = [np.int64(i) for i in cm]
1594
+ cpads = [np.int64(i) for i in cpads]
1595
+ pp = [np.int64(i) for i in pp]
1596
+
1597
+ stencil2coo = kernels['stencil2coo'][order][nd]
1598
+
1599
+ ind = stencil2coo(self._data, data, rows, cols, *nrl, *ncl, *ss, *nr, *nc, *dm, *cm, *cpads, *pp)
1600
+ M = coo_matrix(
1601
+ (data[:ind],(rows[:ind],cols[:ind])),
1602
+ shape = [np.prod(nr),np.prod(nc)],
1603
+ dtype = self.dtype)
1604
+ return M
1605
+
1606
+ #...
1607
+ def _tocoo_parallel_with_pads(self , order='C'):
1608
+
1609
+ # If necessary, update ghost regions
1610
+ if not self.ghost_regions_in_sync:
1611
+ self.update_ghost_regions()
1612
+
1613
+ # Shortcuts
1614
+ nr = self._codomain.npts
1615
+ nc = self._domain.npts
1616
+ nd = self._ndim
1617
+
1618
+ ss = self._codomain.starts
1619
+ ee = self._codomain.ends
1620
+ pp = self._pads
1621
+ pc = self._codomain.pads
1622
+ pd = self._domain.pads
1623
+ cc = self._codomain.periods
1624
+
1625
+ ravel_multi_index = np.ravel_multi_index
1626
+
1627
+ # COO storage
1628
+ rows = []
1629
+ cols = []
1630
+ data = []
1631
+
1632
+ # List of rows (to avoid duplicate updates)
1633
+ I_list = []
1634
+
1635
+ # Shape of row and diagonal spaces
1636
+ xx_dims = self._data.shape[:nd]
1637
+ ll_dims = self._data.shape[nd:]
1638
+
1639
+ # Cycle over rows (x = p + i - s)
1640
+ for xx in np.ndindex( *xx_dims ):
1641
+
1642
+ # Compute row multi-index with simple shift
1643
+ ii = [s + x - p for (s, x, p) in zip(ss, xx, pc)]
1644
+
1645
+ # Apply periodicity where appropriate
1646
+ ii = [i - n if (c and i >= n and i - n < s) else
1647
+ i + n if (c and i < 0 and i + n > e) else i
1648
+ for (i, s, e, n, c) in zip(ii, ss, ee, nr, cc)]
1649
+
1650
+ # Compute row flat index
1651
+ # Exclude values outside global limits of matrix
1652
+ try:
1653
+ I = ravel_multi_index( ii, dims=nr, order=order )
1654
+ except ValueError:
1655
+ continue
1656
+
1657
+ # If I is a new row, append it to list of rows
1658
+ # DO NOT update same row twice!
1659
+ if I not in I_list:
1660
+ I_list.append( I )
1661
+ else:
1662
+ continue
1663
+
1664
+ # Cycle over diagonals (l = p + k)
1665
+ for ll in np.ndindex( *ll_dims ):
1666
+
1667
+ # Compute column multi-index (k = j - i)
1668
+ jj = [(i+l-p) % n for (i,l,n,p) in zip(ii,ll,nc,pp)]
1669
+
1670
+ # Compute column flat index
1671
+ J = ravel_multi_index( jj, dims=nc, order=order )
1672
+
1673
+ # Extract matrix value
1674
+ value = self._data[(*xx, *ll)]
1675
+
1676
+ # Append information to COO arrays
1677
+ rows.append( I )
1678
+ cols.append( J )
1679
+ data.append( value )
1680
+
1681
+ # Create Scipy COO matrix
1682
+ M = coo_matrix(
1683
+ (data,(rows,cols)),
1684
+ shape = [np.prod(nr), np.prod(nc)],
1685
+ dtype = self._domain.dtype
1686
+ )
1687
+
1688
+ M.eliminate_zeros()
1689
+
1690
+ return M
1691
+
1692
+ # ...
1693
+ @property
1694
+ def ghost_regions_in_sync(self):
1695
+ return self._sync
1696
+
1697
+ # ...
1698
+ # NOTE: this property must be set collectively
1699
+ @ghost_regions_in_sync.setter
1700
+ def ghost_regions_in_sync(self, value):
1701
+ assert isinstance(value, bool)
1702
+ self._sync = value
1703
+
1704
+ # ...
1705
+ def _update_ghost_regions_serial(self):
1706
+
1707
+ ndim = self._codomain.ndim
1708
+ for direction in range(self._codomain.ndim):
1709
+
1710
+ periodic = self._codomain.periods[direction]
1711
+ p = self._codomain.pads [direction]
1712
+
1713
+ if p == 0:
1714
+ continue
1715
+
1716
+ idx_front = [slice(None)]*direction
1717
+ idx_back = [slice(None)]*(ndim-direction-1 + ndim)
1718
+
1719
+ if periodic:
1720
+
1721
+ # Copy data from left to right
1722
+ idx_from = tuple(idx_front + [slice( p, 2*p)] + idx_back)
1723
+ idx_to = tuple(idx_front + [slice(-p,None)] + idx_back)
1724
+ self._data[idx_to] = self._data[idx_from]
1725
+
1726
+ # Copy data from right to left
1727
+ idx_from = tuple(idx_front + [slice(-2*p,-p)] + idx_back)
1728
+ idx_to = tuple(idx_front + [slice(None, p)] + idx_back)
1729
+ self._data[idx_to] = self._data[idx_from]
1730
+
1731
+ else:
1732
+
1733
+ # Set left ghost region to zero
1734
+ idx_ghost = tuple(idx_front + [slice(None, p)] + idx_back)
1735
+ self._data[idx_ghost] = 0
1736
+
1737
+ # Set right ghost region to zero
1738
+ idx_ghost = tuple(idx_front + [slice(-p,None)] + idx_back)
1739
+ self._data[idx_ghost] = 0
1740
+
1741
+ # ...
1742
+ def _prepare_transpose_args(self):
1743
+
1744
+ #prepare the arguments for the transpose method
1745
+ V = self.domain
1746
+ W = self.codomain
1747
+ ssc = W.starts
1748
+ eec = W.ends
1749
+ ssd = V.starts
1750
+ eed = V.ends
1751
+ pads = self._pads
1752
+ gpads = V.pads
1753
+
1754
+ dm = V.shifts
1755
+ cm = W.shifts
1756
+
1757
+ # Number of rows in the transposed matrix (along each dimension)
1758
+ nrows = [e-s+1 for s, e in zip(ssd, eed)]
1759
+ ncols = [e-s+2*m*p+1 for s, e, m, p in zip(ssc, eec, cm, gpads)]
1760
+
1761
+ pp = pads
1762
+ ndiags, starts = list(zip(*[compute_diag_len(p, mi, mj, return_padding=True) for p, mi, mj in zip(pp, cm, dm)]))
1763
+ ndiagsT, _ = list(zip(*[compute_diag_len(p, mj, mi, return_padding=True) for p, mi, mj in zip(pp, cm, dm)]))
1764
+
1765
+ diff = [gp-p for gp, p in zip(gpads, pp)]
1766
+
1767
+ sl = [(s if mi > mj else 0) + (s % mi + mi//mj if mi < mj else 0)+(s if mi == mj else 0)\
1768
+ for s, p, mi, mj in zip(starts, pp, cm, dm)]
1769
+
1770
+ si = [(mi * p - mi * (int(np.ceil((p + 1)/mj)) - 1) if mi > mj else 0) + \
1771
+ (mi * p - mi * (p//mi) + d * (mi - 1) if mi < mj else 0) + \
1772
+ (mj * p - mj * (p//mi) + d * (mi - 1) if mi == mj else 0)\
1773
+ for mi, mj, p, d in zip(cm, dm, pp, diff)]
1774
+
1775
+ sk = [n-1\
1776
+ + (-(p % mj) if mi > mj else 0)\
1777
+ + (-p + mj * (p//mi) if mi < mj else 0)\
1778
+ + (-p + mj * (p//mi) if mi == mj else 0)\
1779
+ for mi, mj, n, p in zip(cm, dm, ndiagsT, pp)]
1780
+
1781
+ args={}
1782
+ args['n'] = np.int64(nrows)
1783
+ args['nc'] = np.int64(ncols)
1784
+ args['gp'] = np.int64(gpads)
1785
+ args['p'] = np.int64(pp)
1786
+ args['dm'] = np.int64(dm)
1787
+ args['cm'] = np.int64(cm)
1788
+ args['nd'] = np.int64(ndiags)
1789
+ args['ndT'] = np.int64(ndiagsT)
1790
+ args['si'] = np.int64(si)
1791
+ args['sk'] = np.int64(sk)
1792
+ args['sl'] = np.int64(sl)
1793
+
1794
+ return args
1795
+
1796
+ # ...
1797
+ def set_backend(self, backend, precompiled):
1798
+ '''
1799
+ Define which kernels are called when using .dot() and .transpose()
1800
+
1801
+ Parameters
1802
+ ----------
1803
+ backend : str
1804
+ Psydac backend option.
1805
+
1806
+ precompiled : bool
1807
+ Whether to use precompiled kernels.
1808
+ '''
1809
+ self._backend = backend
1810
+ self._args = self._dotargs_null.copy()
1811
+
1812
+ if self._backend is None:
1813
+ for key, arg in self._args.items():
1814
+ self._args[key] = np.int64(arg)
1815
+ self._func = self._dot
1816
+ self._args.pop('pads')
1817
+ elif precompiled:
1818
+
1819
+ # print('Using precompiled matvec and transpose kernels ...')
1820
+
1821
+ from feectools.linalg import stencil_dot_kernels
1822
+ from feectools.linalg import stencil_transpose_kernels
1823
+
1824
+ # matvec kernel
1825
+ dot_func_name = 'matvec_' + str(self._ndim) + 'd_kernel'
1826
+ self._func = getattr(stencil_dot_kernels, dot_func_name)
1827
+
1828
+ # parameter for rectangular matrices
1829
+ add = [int(end_in >= end_out) for end_in, end_out in zip(self.domain.ends, self.codomain.ends)]
1830
+
1831
+ self._args = {}
1832
+ if self._ndim == 1:
1833
+ self._args['s_in'] = int(self.domain.starts[0])
1834
+ self._args['p_in'] = int(self.domain.pads[0])
1835
+ self._args['add'] = int(add[0])
1836
+ self._args['s_out'] = int(self.codomain.starts[0])
1837
+ self._args['e_out'] = int(self.codomain.ends[0])
1838
+ self._args['p_out'] = int(self.codomain.pads[0])
1839
+ else:
1840
+ self._args['s_in'] = np.array(self.domain.starts)
1841
+ self._args['p_in'] = np.array(self.domain.pads)
1842
+ self._args['add'] = np.array(add)
1843
+ self._args['s_out'] = np.array(self.codomain.starts)
1844
+ self._args['e_out'] = np.array(self.codomain.ends)
1845
+ self._args['p_out'] = np.array(self.codomain.pads)
1846
+
1847
+ # transpose kernel
1848
+ transp_func_name = 'transpose_' + str(self._ndim) + 'd_kernel'
1849
+
1850
+ self._transpose_func = getattr(stencil_transpose_kernels, transp_func_name)
1851
+
1852
+ # parameter for rectangular matrices
1853
+ add = [int(end_out >= end_in) for end_in, end_out in zip(self.domain.ends, self.codomain.ends)]
1854
+
1855
+ self._transpose_args = {}
1856
+ if self._ndim == 1:
1857
+ self._transpose_args['s_in'] = int(self.codomain.starts[0])
1858
+ self._transpose_args['p_in'] = int(self.codomain.pads[0])
1859
+ self._transpose_args['add'] = int(add[0])
1860
+ self._transpose_args['s_out'] = int(self.domain.starts[0])
1861
+ self._transpose_args['e_out'] = int(self.domain.ends[0])
1862
+ self._transpose_args['p_out'] = int(self.domain.pads[0])
1863
+ else:
1864
+ self._transpose_args['s_in'] = np.array(self.codomain.starts)
1865
+ self._transpose_args['p_in'] = np.array(self.codomain.pads)
1866
+ self._transpose_args['add'] = np.array(add)
1867
+ self._transpose_args['s_out'] = np.array(self.domain.starts)
1868
+ self._transpose_args['e_out'] = np.array(self.domain.ends)
1869
+ self._transpose_args['p_out'] = np.array(self.domain.pads)
1870
+ else:
1871
+ raise AttributeError(f'This is the tiny-psydac version - must use precompiled kernels (but {precompiled = })!')
1872
+ from feectools.api.ast.linalg import LinearOperatorDot
1873
+ if self.domain.parallel:
1874
+ comm = self.codomain.cart.comm
1875
+ if self.domain == self.codomain:
1876
+ # In this case nrows_extra[i] == 0 for all i
1877
+ dot = LinearOperatorDot(self._ndim,
1878
+ block_shape = (1,1),
1879
+ keys = ((0,0),),
1880
+ comm = comm,
1881
+ backend=frozenset(backend.items()),
1882
+ nrows_extra = (self._args['nrows_extra'],),
1883
+ gpads=(self._args['gpads'],),
1884
+ pads=(self._args['pads'],),
1885
+ dm = (self._args['dm'],),
1886
+ cm = (self._args['cm'],),
1887
+ dtype=self.dtype)
1888
+
1889
+ starts = self._args.pop('starts')
1890
+ nrows = self._args.pop('nrows')
1891
+
1892
+ self._args.pop('nrows_extra')
1893
+ self._args.pop('gpads')
1894
+ self._args.pop('pads')
1895
+ self._args.pop('dm')
1896
+ self._args.pop('cm')
1897
+
1898
+ for i in range(len(nrows)):
1899
+ self._args['s00_{i}'.format(i=i+1)] = np.int64(starts[i])
1900
+
1901
+ for i in range(len(nrows)):
1902
+ self._args['n00_{i}'.format(i=i+1)] = np.int64(nrows[i])
1903
+
1904
+ else:
1905
+ dot = LinearOperatorDot(self._ndim,
1906
+ block_shape = (1,1),
1907
+ keys = ((0,0),),
1908
+ comm = comm,
1909
+ backend=frozenset(backend.items()),
1910
+ gpads=(self._args['gpads'],),
1911
+ pads=(self._args['pads'],),
1912
+ dm = (self._args['dm'],),
1913
+ cm = (self._args['cm'],),
1914
+ dtype=self.dtype)
1915
+
1916
+ starts = self._args.pop('starts')
1917
+ nrows = self._args.pop('nrows')
1918
+ nrows_extra = self._args.pop('nrows_extra')
1919
+
1920
+ self._args.pop('gpads')
1921
+ self._args.pop('pads')
1922
+ self._args.pop('dm')
1923
+ self._args.pop('cm')
1924
+
1925
+ for i in range(len(nrows)):
1926
+ self._args['s00_{i}'.format(i=i+1)] = np.int64(starts[i])
1927
+
1928
+ for i in range(len(nrows)):
1929
+ self._args['n00_{i}'.format(i=i+1)] = np.int64(nrows[i])
1930
+
1931
+ for i in range(len(nrows)):
1932
+ self._args['ne00_{i}'.format(i=i+1)] = np.int64(nrows_extra[i])
1933
+
1934
+ else:
1935
+ dot = LinearOperatorDot(self._ndim,
1936
+ block_shape = (1,1),
1937
+ keys = ((0,0),),
1938
+ comm = None,
1939
+ backend=frozenset(backend.items()),
1940
+ starts = (tuple(self._args['starts']),),
1941
+ nrows=(tuple(self._args['nrows']),),
1942
+ nrows_extra=(self._args['nrows_extra'],),
1943
+ gpads=(self._args['gpads'],),
1944
+ pads=(self._args['pads'],),
1945
+ dm = (self._args['dm'],),
1946
+ cm = (self._args['cm'],),
1947
+ dtype=self.dtype)
1948
+ self._args.pop('nrows')
1949
+ self._args.pop('nrows_extra')
1950
+ self._args.pop('gpads')
1951
+ self._args.pop('pads')
1952
+ self._args.pop('starts')
1953
+ self._args.pop('dm')
1954
+ self._args.pop('cm')
1955
+
1956
+ self._args.pop('pad_imp')
1957
+ self._args.pop('ndiags')
1958
+ self._func = dot.func
1959
+
1960
+ # ...
1961
+ def _get_diagonal_indices(self):
1962
+ """
1963
+ Compute the indices which should be applied to self._data in order to
1964
+ get the matrix entries on the main diagonal. The result is also stored
1965
+ in self._diag_indices, and retrieved from there on successive calls.
1966
+
1967
+ Returns
1968
+ -------
1969
+ tuple[numpy.ndarray, ndim]
1970
+ The diagonal indices as a tuple of NumPy arrays of identical shape
1971
+ (n1, n2, n3, ...).
1972
+
1973
+ """
1974
+
1975
+ if self._diag_indices is None:
1976
+
1977
+ dp = self.domain.pads
1978
+ dm = self.domain.shifts
1979
+ cm = self.codomain.shifts
1980
+ ss = self.codomain.starts
1981
+ pp = [compute_diag_len(p, mj, mi) - p - 1 for p, mi, mj in zip(self._pads, cm, dm)]
1982
+ nrows = [e - s + 1 for s, e in zip(self.codomain.starts, self.codomain.ends)]
1983
+ ndim = self.domain.ndim
1984
+
1985
+ indices = [np.zeros(np.prod(nrows), dtype=int) for _ in range(2 * ndim)]
1986
+
1987
+ for l, xx in enumerate(np.ndindex(*nrows)):
1988
+ ii = [m * p + x for m, p, x in zip(dm, dp, xx)]
1989
+ jj = [p + x + s - ((x+s) // mi) * mj for x, mi, mj, p, s in zip(xx, cm, dm, pp, ss)]
1990
+ for k in range(ndim):
1991
+ indices[k][l] = ii[k]
1992
+ indices[k + ndim][l] = jj[k]
1993
+
1994
+ self._diag_indices = tuple(idx.reshape(nrows) for idx in indices)
1995
+
1996
+ return self._diag_indices
1997
+
1998
+ #===============================================================================
1999
+ class StencilDiagonalMatrix(LinearOperator):
2000
+ """
2001
+ Linear operator which operates between stencil vector spaces, and which can
2002
+ be represented by a matrix with non-zero entries only on its main diagonal.
2003
+ As such this operator is completely local and requires no data communication.
2004
+
2005
+ We assume that the vectors in the domain and the codomain have the same
2006
+ shape and are distributed in the same way.
2007
+
2008
+ Parameters
2009
+ ----------
2010
+ V : feectools.linalg.stencil.StencilVectorSpace
2011
+ Domain of the new linear operator.
2012
+
2013
+ W : feectools.linalg.stencil.StencilVectorSpace
2014
+ Codomain of the new linear operator.
2015
+
2016
+ """
2017
+ def __init__(self, V, W, data):
2018
+
2019
+ # Check domain and codomain
2020
+ assert isinstance(V, StencilVectorSpace)
2021
+ assert isinstance(W, StencilVectorSpace)
2022
+ assert V.starts == W.starts
2023
+ assert V.ends == W.ends
2024
+
2025
+ data = np.asarray(data)
2026
+
2027
+ # Check shape of provided data
2028
+ shape = tuple(e - s + 1 for s, e in zip(V.starts, V.ends))
2029
+ assert data.shape == shape
2030
+
2031
+ # Store info in object
2032
+ self._domain = V
2033
+ self._codomain = W
2034
+ self._data = data
2035
+
2036
+ #--------------------------------------
2037
+ # Abstract interface
2038
+ #--------------------------------------
2039
+ @property
2040
+ def domain(self):
2041
+ return self._domain
2042
+
2043
+ @property
2044
+ def codomain(self):
2045
+ return self._codomain
2046
+
2047
+ @property
2048
+ def dtype(self):
2049
+ return self._data.dtype
2050
+
2051
+ def tosparse(self):
2052
+ return sp_diags(self._data.ravel())
2053
+
2054
+ def toarray(self):
2055
+ return self._data.copy()
2056
+
2057
+ def dot(self, v, out=None):
2058
+
2059
+ assert isinstance(v, StencilVector)
2060
+ assert v.space is self.domain
2061
+
2062
+ if out is not None:
2063
+ assert isinstance(out, StencilVector)
2064
+ assert out.space is self.codomain
2065
+ else:
2066
+ out = self.codomain.zeros()
2067
+
2068
+ V = self.domain
2069
+ i = tuple(slice(s, e + 1) for s, e in zip(V.starts, V.ends))
2070
+ np.multiply(self._data, v[i], out=out[i])
2071
+
2072
+ out.ghost_regions_in_sync = False
2073
+
2074
+ return out
2075
+
2076
+ # ...
2077
+ # TODO [YG 22.01.2024]: idot function will require a dedicated kernel
2078
+ # ...
2079
+
2080
+ def transpose(self, *, conjugate=False, out=None):
2081
+
2082
+ assert isinstance(conjugate, bool)
2083
+
2084
+ if out is not None:
2085
+ assert isinstance(out, StencilDiagonalMatrix)
2086
+ assert out.domain is self.codomain
2087
+ assert out.codomain is self.domain
2088
+
2089
+ if not (conjugate and self.dtype is complex):
2090
+
2091
+ if out is None:
2092
+ data = self._data.copy()
2093
+ else:
2094
+ np.copyto(out._data, self._data, casting='no')
2095
+
2096
+ else:
2097
+
2098
+ if out is None:
2099
+ data = np.conjugate(self._data, casting='no')
2100
+ else:
2101
+ np.conjugate(self._data, out=out._data, casting='no')
2102
+
2103
+ if out is None:
2104
+ out = StencilDiagonalMatrix(self.codomain, self.domain, data)
2105
+
2106
+ return out
2107
+
2108
+ #--------------------------------------
2109
+ # Other properties/methods
2110
+ #--------------------------------------
2111
+ def copy(self, *, out=None):
2112
+
2113
+ if out is self:
2114
+ return self
2115
+
2116
+ if out is None:
2117
+ data = self._data.copy()
2118
+ out = StencilDiagonalMatrix(self.domain, self.codomain, data)
2119
+ else:
2120
+ assert isinstance(out, StencilDiagonalMatrix)
2121
+ assert out.domain is self.domain
2122
+ assert out.codomain is self.codomain
2123
+ np.copyto(out._data, self._data, casting='no')
2124
+
2125
+ return out
2126
+
2127
+ def diagonal(self, *, inverse = False, out = None):
2128
+ """
2129
+ Get the coefficients on the main diagonal as a StencilDiagonalMatrix object.
2130
+
2131
+ In the default case (inverse=False, out=None) self is returned.
2132
+
2133
+ Parameters
2134
+ ----------
2135
+ inverse : bool
2136
+ If True, get the inverse of the diagonal. (Default: False).
2137
+
2138
+ out : StencilDiagonalMatrix
2139
+ If provided, write the diagonal entries into this matrix. (Default: None).
2140
+
2141
+ Returns
2142
+ -------
2143
+ StencilDiagonalMatrix
2144
+ Either self, or another StencilDiagonalMatrix with the diagonal inverse.
2145
+
2146
+ """
2147
+ # Check `inverse` argument
2148
+ assert isinstance(inverse, bool)
2149
+
2150
+ # Determine domain and codomain of the `out` matrix
2151
+ V, W = self.domain, self.codomain
2152
+ if inverse:
2153
+ V, W = W, V
2154
+
2155
+ # Check `out` argument and identify `data` array of output vector
2156
+ if out is None:
2157
+ data = None
2158
+ else:
2159
+ assert isinstance(out, StencilDiagonalMatrix)
2160
+ assert out.domain is V
2161
+ assert out.codomain is W
2162
+ data = out._data
2163
+
2164
+ # Calculate entries, or set `out=self` in default case
2165
+ if inverse:
2166
+ data = np.divide(1, diag, out=data)
2167
+ elif out:
2168
+ np.copyto(data, diag)
2169
+ else:
2170
+ out = self
2171
+
2172
+ # If needed create a new StencilDiagonalMatrix object
2173
+ if out is None:
2174
+ out = StencilDiagonalMatrix(V, W, data)
2175
+
2176
+ return out
2177
+
2178
+ #===============================================================================
2179
+ # TODO [YG, 28.01.2021]:
2180
+ # - Check if StencilMatrix should be subclassed
2181
+ # - Reimplement magic methods (some are simply copied from StencilMatrix)
2182
+ def flip_axis(index, n):
2183
+ s = n - index.start-1
2184
+ e = n - index.stop-1 if n > index.stop else None
2185
+ return slice(s,e,-1)
2186
+
2187
+ class StencilInterfaceMatrix(LinearOperator):
2188
+ """
2189
+ Matrix in n-dimensional stencil format for an interface.
2190
+
2191
+ This is a linear operator that maps elements of stencil vector space V to
2192
+ elements of stencil vector space W.
2193
+
2194
+ Parameters
2195
+ ----------
2196
+ V : feectools.linalg.stencil.StencilVectorSpace
2197
+ Domain of the new linear operator.
2198
+
2199
+ W : feectools.linalg.stencil.StencilVectorSpace
2200
+ Codomain of the new linear operator.
2201
+
2202
+ s_d : int
2203
+ The starting index of the domain.
2204
+
2205
+ s_c : int
2206
+ The starting index of the codomain.
2207
+
2208
+ d_axis : int
2209
+ The axis of the Interface of the domain.
2210
+
2211
+ c_axis : int
2212
+ The axis of the Interface of the codomain.
2213
+
2214
+ d_ext : int
2215
+ The extremity of the domain Interface space.
2216
+ the values must be 1 or -1.
2217
+
2218
+ c_ext : int
2219
+ The extremity of the codomain Interface space.
2220
+ the values must be 1 or -1.
2221
+
2222
+ dim : int
2223
+ The axis of the interface.
2224
+
2225
+ pads: <list|tuple>
2226
+ Padding of the linear operator.
2227
+
2228
+ """
2229
+ def __init__(self, V, W, s_d, s_c, d_axis, c_axis, d_ext, c_ext, *, flip=None, pads=None, backend=None):
2230
+
2231
+ assert isinstance(V, StencilVectorSpace)
2232
+ assert isinstance(W, StencilVectorSpace)
2233
+ assert W.pads == V.pads
2234
+
2235
+ Vin = V.interfaces[d_axis, d_ext]
2236
+
2237
+ if pads is not None:
2238
+ for p,vp in zip(pads, Vin.pads):
2239
+ assert p<=vp
2240
+
2241
+ self._pads = pads or tuple(Vin.pads)
2242
+ dims = list(W.shape)
2243
+
2244
+ if W.parent_ends[c_axis] is not None:
2245
+ diff = min(1, W.parent_ends[c_axis]-W.ends[c_axis])
2246
+ else:
2247
+ diff = 0
2248
+
2249
+ dims[c_axis] = W.pads[c_axis] + 1-diff + 2*W.shifts[c_axis]*W.pads[c_axis]
2250
+ diags = [compute_diag_len(p, md, mc) for p,md,mc in zip(self._pads, Vin.shifts, W.shifts)]
2251
+ self._data = np.zeros(dims + diags, dtype=W.dtype)
2252
+
2253
+ # Parallel attributes
2254
+ if W.parallel and not isinstance(W.cart, InterfaceCartDecomposition):
2255
+ if W.cart.is_comm_null:return
2256
+ # Create data exchanger for ghost regions
2257
+ self._synchronizer = get_data_exchanger(
2258
+ cart = W.cart,
2259
+ dtype = W.dtype,
2260
+ coeff_shape = diags,
2261
+ assembly = True,
2262
+ axis = c_axis,
2263
+ shape = self._data.shape
2264
+ )
2265
+
2266
+ self._flip = tuple([1]*len(dims) if flip is None else flip)
2267
+ self._permutation = list(range(len(dims)))
2268
+ self._permutation[d_axis], self._permutation[c_axis] = self._permutation[c_axis], self._permutation[d_axis]
2269
+ self._domain = V
2270
+ self._codomain = W
2271
+ self._domain_axis = d_axis
2272
+ self._codomain_axis = c_axis
2273
+ self._domain_ext = d_ext
2274
+ self._codomain_ext = c_ext
2275
+ self._domain_start = s_d
2276
+ self._codomain_start = s_c
2277
+ self._ndim = len(dims)
2278
+ self._backend = None
2279
+
2280
+
2281
+ # Prepare the arguments for the dot product method
2282
+ nd = [(ej-sj+2*gp*mj-mj*p-gp)//mj*mi+1 for sj,ej,mj,mi,p,gp in zip(Vin.starts, Vin.ends, Vin.shifts, W.shifts, self._pads, Vin.pads)]
2283
+ nc = [ei-si+1 for si,ei,mj,p in zip(W.starts, W.ends, Vin.shifts, self._pads)]
2284
+
2285
+ # Number of rows in matrix (along each dimension)
2286
+ nrows = [min(ni,nj) for ni,nj in zip(nc, nd)]
2287
+ nrows_extra = [max(0,ni-nj) for ni,nj in zip(nc, nd)]
2288
+ nrows_extra[c_axis] = max(W.npts[c_axis]-Vin.npts[c_axis], 0)
2289
+ nrows[c_axis] = W.pads[c_axis] + 1-diff-nrows_extra[c_axis]
2290
+
2291
+
2292
+ args = {}
2293
+ args['starts'] = tuple(Vin.starts)
2294
+ args['nrows'] = tuple(nrows)
2295
+ args['nrows_extra'] = tuple(nrows_extra)
2296
+ args['gpads'] = tuple(Vin.pads)
2297
+ args['pads'] = tuple(self._pads)
2298
+ args['dm'] = tuple(Vin.shifts)
2299
+ args['cm'] = tuple(W.shifts)
2300
+ args['c_axis'] = c_axis
2301
+ args['d_start'] = self._domain_start
2302
+ args['c_start'] = self._codomain_start
2303
+ args['flip'] = self._flip
2304
+ args['permutation'] = self._permutation
2305
+
2306
+ self._dotargs_null = args
2307
+ self._args = args.copy()
2308
+ self._func = self._dot
2309
+
2310
+ self._transpose_args = self._prepare_transpose_args()
2311
+ self._transpose_func = kernels['interface_transpose'][self._ndim]
2312
+
2313
+ if backend is None:
2314
+ backend = PSYDAC_BACKENDS.get(os.environ.get('PSYDAC_BACKEND'))
2315
+
2316
+ if backend:
2317
+ self.set_backend(backend)
2318
+
2319
+ # Flag ghost regions as not up-to-date (conservative choice)
2320
+ self._sync = False
2321
+
2322
+ #--------------------------------------
2323
+ # Abstract interface
2324
+ #--------------------------------------
2325
+ @property
2326
+ def domain(self):
2327
+ return self._domain
2328
+
2329
+ # ...
2330
+ @property
2331
+ def codomain(self):
2332
+ return self._codomain
2333
+
2334
+ # ...
2335
+ @property
2336
+ def dtype(self):
2337
+ return self.domain.dtype
2338
+
2339
+ # ...
2340
+ def dot(self, v, out=None):
2341
+
2342
+ assert isinstance(v, StencilVector)
2343
+ assert v.space is self.domain
2344
+
2345
+ # Necessary if vector space is distributed across processes
2346
+
2347
+ if out is not None:
2348
+ assert isinstance(out, StencilVector)
2349
+ assert out.space is self.codomain
2350
+ out[(slice(None,None),)*v.space.ndim] = 0.
2351
+ else:
2352
+ out = StencilVector( self.codomain )
2353
+
2354
+ # Necessary if vector space is distributed across processes
2355
+ if not v.ghost_regions_in_sync and not v.space.parallel:
2356
+ v.update_ghost_regions()
2357
+
2358
+ self._func(self._data, v._interface_data[self._domain_axis, self._domain_ext], out._data, **self._args)
2359
+ # IMPORTANT: flag that ghost regions are not up-to-date
2360
+ out.ghost_regions_in_sync = False
2361
+ return out
2362
+
2363
+ # ...
2364
+ @staticmethod
2365
+ def _dot(mat, v, out, starts, nrows, nrows_extra, gpads, pads, dm, cm, c_axis, d_start, c_start, flip, permutation):
2366
+
2367
+ # Index for k=i-j
2368
+ nrows = list(nrows)
2369
+ ndim = len(v.shape)
2370
+ kk = [slice(None)]*ndim
2371
+ diff = [xp-p for xp,p in zip(gpads, pads)]
2372
+
2373
+ ndiags, _ = list(zip(*[compute_diag_len(p,mj,mi, return_padding=True) for p,mi,mj in zip(pads,cm,dm)]))
2374
+ bb = [p*m+p+1-n-s%m for p,m,n,s in zip(gpads, dm, ndiags, starts)]
2375
+ nn = v.shape
2376
+
2377
+ for xx in np.ndindex( *nrows ):
2378
+ ii = [ mi*pi + x for mi,pi,x in zip(cm, gpads, xx) ]
2379
+ jj = tuple( slice(b-d+(x+s%mj)//mi*mj,b-d+(x+s%mj)//mi*mj+n) for x,mi,mj,b,s,n,d in zip(xx,cm,dm,bb,starts,ndiags,diff) )
2380
+ jj = [flip_axis(i,n) if f==-1 else i for i,f,n in zip(jj,flip,nn)]
2381
+ jj = tuple(jj[i] for i in permutation)
2382
+ ii_kk = tuple( ii + kk )
2383
+
2384
+ ii[c_axis] += c_start
2385
+ out[tuple(ii)] = np.dot( mat[ii_kk].flat, v[jj].flat )
2386
+
2387
+
2388
+ new_nrows = nrows.copy()
2389
+ for d,er in enumerate(nrows_extra):
2390
+
2391
+ rows = new_nrows.copy()
2392
+ del rows[d]
2393
+
2394
+ for n in range(er):
2395
+ for xx in np.ndindex(*rows):
2396
+ xx = list(xx)
2397
+ xx.insert(d, nrows[d]+n)
2398
+
2399
+ ii = [mi*pi + x for mi,pi,x in zip(cm, gpads, xx)]
2400
+ ee = [max(x-l+1,0) for x,l in zip(xx, nrows)]
2401
+ jj = tuple( slice(b-d+(x+s%mj)//mi*mj, b-d+(x+s%mj)//mi*mj+n-e) for x,mi,mj,d,e,b,s,n in zip(xx, cm, dm, diff, ee, bb, starts, ndiags) )
2402
+ jj = [flip_axis(i,n) if f==-1 else i for i,f,n in zip(jj, flip, nn)]
2403
+ jj = tuple(jj[i] for i in permutation)
2404
+ kk = [slice(None,n-e) for n,e in zip(ndiags, ee)]
2405
+ ii_kk = tuple( ii + kk )
2406
+ ii[c_axis] += c_start
2407
+ out[tuple(ii)] = np.dot( mat[ii_kk].flat, v[jj].flat )
2408
+
2409
+ new_nrows[d] += er
2410
+
2411
+ # ...
2412
+ def transpose( self, conjugate=False, out=None):
2413
+ """ Create new StencilInterfaceMatrix Mt, where domain and codomain are swapped
2414
+ with respect to original matrix M, and Mt_{ij} = M_{ji}.
2415
+ """
2416
+
2417
+ # For clarity rename self
2418
+ M = self
2419
+
2420
+ if out is None:
2421
+ # Create new matrix where domain and codomain are swapped
2422
+
2423
+ out = StencilInterfaceMatrix(M.codomain, M.domain, M.codomain_start, M.domain_start, M.codomain_axis, M.domain_axis, M.codomain_ext, M.domain_ext,
2424
+ flip=M.flip, pads=M.pads, backend=M.backend)
2425
+
2426
+ # Call low-level '_transpose' function (works on Numpy arrays directly)
2427
+ if conjugate:
2428
+ M._transpose_func(np.conjugate(M._data), out._data, **M._transpose_args)
2429
+ else:
2430
+ M._transpose_func(M._data, out._data, **M._transpose_args)
2431
+ return out
2432
+
2433
+ def _prepare_transpose_args(self):
2434
+
2435
+ #prepare the arguments for the transpose method
2436
+ V = self.domain
2437
+ W = self.codomain
2438
+ ssc = W.starts
2439
+ eec = W.ends
2440
+ ssd = V.interfaces[self._domain_axis, self._domain_ext].starts
2441
+ eed = V.interfaces[self._domain_axis, self._domain_ext].ends
2442
+ pads = self._pads
2443
+ gpads = V.pads
2444
+ dm = V.shifts
2445
+ cm = W.shifts
2446
+ dim = self._codomain_axis
2447
+
2448
+ # Number of rows in the transposed matrix (along each dimension)
2449
+ nrows = [e-s+1 for s,e in zip(ssd, eed)]
2450
+ ncols = [e-s+1+2*m*p for s, e, m, p in zip(ssc, eec, cm, gpads)]
2451
+
2452
+ pp = pads
2453
+ ndiags, starts = list(zip(*[compute_diag_len(p,mi,mj, return_padding=True) for p, mi, mj in zip(pp, cm, dm)]))
2454
+ ndiagsT, _ = list(zip(*[compute_diag_len(p,mj,mi, return_padding=True) for p, mi, mj in zip(pp, cm, dm)]))
2455
+
2456
+ diff = [gp-p for gp, p in zip(gpads, pp)]
2457
+
2458
+ sl = [(s if mi > mj else 0) + (s % mi + mi//mj if mi < mj else 0)+(s if mi == mj else 0)\
2459
+ for s, p, mi, mj in zip(starts, pp, cm, dm)]
2460
+
2461
+ si = [(mi * p - mi * (int(np.ceil((p + 1)/mj)) - 1) if mi > mj else 0) + \
2462
+ (mi * p - mi * (p//mi) + d * (mi - 1) if mi < mj else 0) + \
2463
+ (mj * p - mj * (p//mi) + d * (mi - 1) if mi == mj else 0)\
2464
+ for mi, mj, p, d in zip(cm, dm, pp, diff)]
2465
+
2466
+ sk = [n - 1\
2467
+ + (-(p % mj) if mi > mj else 0)\
2468
+ + (-p + mj * (p//mi) if mi < mj else 0)\
2469
+ + (-p + mj * (p//mi) if mi == mj else 0)\
2470
+ for mi, mj, n, p in zip(cm, dm, ndiagsT, pp)]
2471
+
2472
+
2473
+ if V.parent_ends[dim] is not None:
2474
+ diff_r = min(1, V.parent_ends[dim] - V.ends[dim])
2475
+ else:
2476
+ diff_r = 0
2477
+
2478
+ if W.parent_ends[dim] is not None:
2479
+ diff_c = min(1, W.parent_ends[dim] - W.ends[dim])
2480
+ else:
2481
+ diff_c = 0
2482
+
2483
+ nrows[dim] = pads[dim] + 1 - diff_r
2484
+ ncols[dim] = pads[dim] + 1 - diff_c + 2*cm[dim]*pads[dim]
2485
+
2486
+ args = {}
2487
+ args['n'] = np.int64(nrows)
2488
+ args['nc'] = np.int64(ncols)
2489
+ args['gp'] = np.int64(gpads)
2490
+ args['p'] = np.int64(pp)
2491
+ args['dm'] = np.int64(dm)
2492
+ args['cm'] = np.int64(cm)
2493
+ args['nd'] = np.int64(ndiags)
2494
+ args['ndT'] = np.int64(ndiagsT)
2495
+ args['si'] = np.int64(si)
2496
+ args['sk'] = np.int64(sk)
2497
+ args['sl'] = np.int64(sl)
2498
+
2499
+ return args
2500
+
2501
+ # ...
2502
+ def toarray(self, **kwargs):
2503
+
2504
+ order = kwargs.pop('order', 'C')
2505
+ with_pads = kwargs.pop('with_pads', False)
2506
+
2507
+ if self.codomain.parallel and with_pads:
2508
+ coo = self._tocoo_parallel_with_pads()
2509
+ else:
2510
+ coo = self._tocoo_no_pads()
2511
+
2512
+ return coo.toarray()
2513
+
2514
+ # ...
2515
+ def tosparse(self, **kwargs):
2516
+
2517
+ order = kwargs.pop('order', 'C')
2518
+ with_pads = kwargs.pop('with_pads', False)
2519
+
2520
+ if self.codomain.parallel and with_pads:
2521
+ coo = self._tocoo_parallel_with_pads()
2522
+ else:
2523
+ coo = self._tocoo_no_pads()
2524
+
2525
+ return coo
2526
+
2527
+ #...
2528
+ def copy(self):
2529
+ M = StencilInterfaceMatrix( self._domain, self._codomain,
2530
+ self._domain_start, self._codomain_start,
2531
+ self._domain_axis, self._codomain_axis,
2532
+ self._domain_ext, self._codomain_ext,
2533
+ flip=self._flip, pads=self._pads,
2534
+ backend=self._backend )
2535
+ M._data[:] = self._data[:]
2536
+ return M
2537
+
2538
+ # ...
2539
+ def __neg__(self):
2540
+ return self.__mul__(-1)
2541
+
2542
+ #...
2543
+ def __mul__(self, a):
2544
+ w = self.copy()
2545
+ w._data *= a
2546
+ w._sync = self._sync
2547
+ return w
2548
+
2549
+ #...
2550
+ def __add__(self, m):
2551
+ raise NotImplementedError('TODO: StencilInterfaceMatrix.__add__')
2552
+
2553
+ #...
2554
+ def __sub__(self, m):
2555
+ raise NotImplementedError('TODO: StencilInterfaceMatrix.__sub__')
2556
+
2557
+ #...
2558
+ def __imul__(self, a):
2559
+ self._data *= a
2560
+
2561
+ #...
2562
+ def __iadd__(self, m):
2563
+ raise NotImplementedError('TODO: StencilInterfaceMatrix.__iadd__')
2564
+
2565
+ #...
2566
+ def __isub__(self, m):
2567
+ raise NotImplementedError('TODO: StencilInterfaceMatrix.__isub__')
2568
+
2569
+ #--------------------------------------
2570
+ # Other properties/methods
2571
+ #--------------------------------------
2572
+
2573
+ # ...
2574
+ @property
2575
+ def domain_axis(self):
2576
+ return self._domain_axis
2577
+
2578
+ # ...
2579
+ @property
2580
+ def codomain_axis(self):
2581
+ return self._codomain_axis
2582
+
2583
+ # ...
2584
+ @property
2585
+ def domain_ext(self):
2586
+ return self._domain_ext
2587
+
2588
+ # ...
2589
+ @property
2590
+ def codomain_ext(self):
2591
+ return self._codomain_ext
2592
+
2593
+ # ...
2594
+ @property
2595
+ def domain_start(self):
2596
+ return self._domain_start
2597
+
2598
+ # ...
2599
+ @property
2600
+ def codomain_start(self):
2601
+ return self._codomain_start
2602
+
2603
+ # ...
2604
+ @property
2605
+ def dim(self):
2606
+ return self._ndim
2607
+
2608
+ # ...
2609
+ @property
2610
+ def flip(self):
2611
+ return self._flip
2612
+
2613
+ # ...
2614
+ @property
2615
+ def permutation(self):
2616
+ return self._permutation
2617
+
2618
+ # ...
2619
+ @property
2620
+ def pads(self):
2621
+ return self._pads
2622
+
2623
+ # ...
2624
+ def __getitem__(self, key):
2625
+ index = self._getindex( key )
2626
+ return self._data[index]
2627
+
2628
+ # ...
2629
+ def __setitem__(self, key, value):
2630
+ index = self._getindex( key )
2631
+ self._data[index] = value
2632
+
2633
+ #...
2634
+ def max(self):
2635
+ return self._data.max()
2636
+
2637
+ # ...
2638
+ @property
2639
+ def backend(self):
2640
+ return self._backend
2641
+
2642
+ #--------------------------------------
2643
+ # Private methods
2644
+ #--------------------------------------
2645
+ def _getindex(self, key):
2646
+
2647
+ nd = self._ndim
2648
+ ii = key[:nd]
2649
+ kk = key[nd:]
2650
+
2651
+ index = []
2652
+
2653
+ for i,s,p in zip(ii, self._codomain.starts, self._codomain.pads):
2654
+ x = self._shift_index(i, p-s)
2655
+ index.append(x)
2656
+
2657
+ for k,p in zip(kk, self._pads):
2658
+ l = self._shift_index(k, p)
2659
+ index.append(l)
2660
+
2661
+ return tuple(index)
2662
+
2663
+ # ...
2664
+ @staticmethod
2665
+ def _shift_index(index, shift):
2666
+ if isinstance(index, slice):
2667
+ start = None if index.start is None else index.start + shift
2668
+ stop = None if index.stop is None else index.stop + shift
2669
+ return slice(start, stop, index.step)
2670
+ else:
2671
+ return index + shift
2672
+
2673
+ #...
2674
+ def _tocoo_no_pads(self):
2675
+ # Shortcuts
2676
+ nr = self.codomain.npts
2677
+ nc = self.domain.npts
2678
+ ss = self.codomain.starts
2679
+ pp = self.codomain.pads
2680
+ nd = len(pp)
2681
+
2682
+ dim = self._codomain_axis
2683
+
2684
+ flip = self.flip
2685
+ permutation = self.permutation
2686
+ c_start = self.codomain_start
2687
+ d_start = self.domain_start
2688
+ dm = self.domain.shifts
2689
+ cm = self.codomain.shifts
2690
+
2691
+ ravel_multi_index = np.ravel_multi_index
2692
+
2693
+ # COO storage
2694
+ rows = []
2695
+ cols = []
2696
+ data = []
2697
+ # Range of data owned by local process (no ghost regions)
2698
+ local = tuple( [slice(m*p,-m*p) if p != 0 else slice(0, None) for m,p in zip(cm, pp)] + [slice(None)] * nd )
2699
+ pp = [compute_diag_len(p,mj,mi)-(p+1) for p,mi,mj in zip(self._pads, cm, dm)]
2700
+
2701
+ for (index,value) in np.ndenumerate( self._data[local] ):
2702
+ if value:
2703
+ # index = [i1, i2, ..., p1+j1-i1, p2+j2-i2, ...]
2704
+
2705
+ xx = index[:nd] # x=i-s
2706
+ ll = index[nd:] # l=p+k
2707
+
2708
+ ii = [s+x for s,x in zip(ss,xx)]
2709
+ di = [i//m for i,m in zip(ii,cm)]
2710
+
2711
+ jj = [(i*m+l-p)%n for (i,m,l,n,p) in zip(di,dm,ll,nc,pp)]
2712
+
2713
+ ii[dim] += c_start
2714
+ jj[dim] += d_start
2715
+
2716
+ jj = [n-j-1 if f==-1 else j for j,f,n in zip(jj,flip,nc)]
2717
+
2718
+ jj = [jj[i] for i in permutation]
2719
+
2720
+ I = ravel_multi_index(ii, dims=nr, order='C')
2721
+ J = ravel_multi_index(jj, dims=nc, order='C')
2722
+
2723
+ rows.append(I)
2724
+ cols.append(J)
2725
+ data.append(value)
2726
+
2727
+ M = coo_matrix(
2728
+ (data,(rows,cols)),
2729
+ shape = [np.prod(nr),np.prod(nc)],
2730
+ dtype = self.domain.dtype)
2731
+
2732
+ return M
2733
+
2734
+ # ...
2735
+ @property
2736
+ def ghost_regions_in_sync(self):
2737
+ return self._sync
2738
+
2739
+ # ...
2740
+ # NOTE: this property must be set collectively
2741
+ @ghost_regions_in_sync.setter
2742
+ def ghost_regions_in_sync(self, value):
2743
+ assert isinstance(value, bool)
2744
+ self._sync = value
2745
+
2746
+ # ...
2747
+ def _update_ghost_regions_serial(self, direction: int):
2748
+
2749
+ if direction is None:
2750
+ for d in range(self._codomain.ndim):
2751
+ self._update_ghost_regions_serial(d)
2752
+ return
2753
+
2754
+ ndim = self._codomain.ndim
2755
+ periodic = self._codomain.periods[direction]
2756
+ p = self._codomain.pads [direction]
2757
+
2758
+ if p == 0:
2759
+ return
2760
+
2761
+ idx_front = [slice(None)] * direction
2762
+ idx_back = [slice(None)] * (ndim-direction-1)
2763
+
2764
+ if periodic:
2765
+
2766
+ # Copy data from left to right
2767
+ idx_from = tuple(idx_front + [slice( p, 2*p)] + idx_back)
2768
+ idx_to = tuple(idx_front + [slice(-p,None)] + idx_back)
2769
+ self._data[idx_to] = self._data[idx_from]
2770
+
2771
+ # Copy data from right to left
2772
+ idx_from = tuple(idx_front + [slice(-2*p,-p)] + idx_back)
2773
+ idx_to = tuple(idx_front + [slice(None, p)] + idx_back)
2774
+ self._data[idx_to] = self._data[idx_from]
2775
+
2776
+ else:
2777
+
2778
+ # Set left ghost region to zero
2779
+ idx_ghost = tuple(idx_front + [slice(None, p)] + idx_back)
2780
+ self._data[idx_ghost] = 0
2781
+
2782
+ # Set right ghost region to zero
2783
+ idx_ghost = tuple(idx_front + [slice(-p,None)] + idx_back)
2784
+ self._data[idx_ghost] = 0
2785
+
2786
+ # ...
2787
+ def exchange_assembly_data(self):
2788
+ """
2789
+ Exchange assembly data.
2790
+ """
2791
+ ndim = self._codomain.ndim
2792
+ parallel = self._codomain.parallel
2793
+
2794
+ if self._codomain.parallel:
2795
+ # PARALLEL CASE: fill in ghost regions with data from neighbors
2796
+ self._synchronizer.start_exchange_assembly_data(self._data)
2797
+ self._synchronizer. end_exchange_assembly_data(self._data)
2798
+ else:
2799
+ # SERIAL CASE: fill in ghost regions along periodic directions, otherwise set to zero
2800
+ self._exchange_assembly_data_serial()
2801
+
2802
+ # ...
2803
+ def _exchange_assembly_data_serial(self):
2804
+
2805
+ ndim = self._codomain.ndim
2806
+ for direction in range(ndim):
2807
+ if direction == self._codomain_axis:
2808
+ continue
2809
+ periodic = self._codomain.periods[direction]
2810
+ p = self._codomain.pads [direction]
2811
+ m = self._codomain.shifts [direction]
2812
+
2813
+ if periodic:
2814
+ idx_front = [slice(None)] * direction
2815
+ idx_back = [slice(None)] * (ndim-direction-1)
2816
+
2817
+ # Copy data from left to right
2818
+ idx_to = tuple(idx_front + [slice( m*p, m*p+p)] + idx_back)
2819
+ idx_from = tuple(idx_front + [slice(-m*p,-m*p+p) if (-m*p+p)!=0 else slice(-m*p, None)] + idx_back)
2820
+ self._data[idx_to] += self._data[idx_from]
2821
+
2822
+ # ...
2823
+ def set_backend(self, backend, precompiled=False):
2824
+ raise AttributeError(f'This is the tiny-psydac version - must use precompiled kernels (but {precompiled = })!')
2825
+ from feectools.api.ast.linalg import LinearOperatorDot
2826
+
2827
+ self._backend = backend
2828
+ self._args = self._dotargs_null.copy()
2829
+
2830
+ if self._backend is None:
2831
+ self._func = self._dot
2832
+ else:
2833
+ if self.domain.parallel:
2834
+
2835
+ comm = self.domain.interfaces[self._domain_axis, self._domain_ext].cart.local_comm
2836
+
2837
+ if self.domain == self.codomain:
2838
+ # In this case nrows_extra[i] == 0 for all i
2839
+ dot = LinearOperatorDot(self._ndim,
2840
+ block_shape = (1,1),
2841
+ keys = ((0,0),),
2842
+ comm = comm,
2843
+ backend=frozenset(backend.items()),
2844
+ nrows_extra=(self._args['nrows_extra'],),
2845
+ gpads=(self._args['gpads'],),
2846
+ pads=(self._args['pads'],),
2847
+ dm = (self._args['dm'],),
2848
+ cm = (self._args['cm'],),
2849
+ interface=True,
2850
+ flip_axis=self._flip,
2851
+ interface_axis=self._codomain_axis,
2852
+ d_start=(self._domain_start,),
2853
+ c_start=(self._codomain_start,),
2854
+ dtype= self.dtype)
2855
+
2856
+ starts = self._args.pop('starts')
2857
+ nrows = self._args.pop('nrows')
2858
+
2859
+ self._args = {}
2860
+ for i in range(len(nrows)):
2861
+ self._args['s00_{i}'.format(i=i+1)] = np.int64(starts[i])
2862
+
2863
+ for i in range(len(nrows)):
2864
+ self._args['n00_{i}'.format(i=i+1)] = np.int64(nrows[i])
2865
+
2866
+ else:
2867
+ dot = LinearOperatorDot(self._ndim,
2868
+ block_shape = (1,1),
2869
+ keys = ((0,0),),
2870
+ comm = comm,
2871
+ backend=frozenset(backend.items()),
2872
+ gpads=(self._args['gpads'],),
2873
+ pads=(self._args['pads'],),
2874
+ dm = (self._args['dm'],),
2875
+ cm = (self._args['cm'],),
2876
+ interface=True,
2877
+ flip_axis=self._flip,
2878
+ interface_axis=self._codomain_axis,
2879
+ d_start=(self._domain_start,),
2880
+ c_start=(self._codomain_start,),
2881
+ dtype= self.dtype)
2882
+
2883
+ starts = self._args.pop('starts')
2884
+ nrows = self._args.pop('nrows')
2885
+ nrows_extra = self._args.pop('nrows_extra')
2886
+
2887
+ self._args = {}
2888
+
2889
+ for i in range(len(nrows)):
2890
+ self._args['s00_{i}'.format(i=i+1)] = np.int64(starts[i])
2891
+
2892
+ for i in range(len(nrows)):
2893
+ self._args['n00_{i}'.format(i=i+1)] = np.int64(nrows[i])
2894
+
2895
+ for i in range(len(nrows)):
2896
+ self._args['ne00_{i}'.format(i=i+1)] = np.int64(nrows_extra[i])
2897
+
2898
+ else:
2899
+ dot = LinearOperatorDot(self._ndim,
2900
+ block_shape = (1,1),
2901
+ keys = ((0,0),),
2902
+ comm = None,
2903
+ backend=frozenset(backend.items()),
2904
+ starts = (tuple(self._args['starts']),),
2905
+ nrows=(self._args['nrows'],),
2906
+ nrows_extra=(self._args['nrows_extra'],),
2907
+ gpads=(self._args['gpads'],),
2908
+ pads=(self._args['pads'],),
2909
+ dm = (self._args['dm'],),
2910
+ cm = (self._args['cm'],),
2911
+ interface=True,
2912
+ flip_axis=self._flip,
2913
+ interface_axis=self._codomain_axis,
2914
+ d_start=(self._domain_start,),
2915
+ c_start=(self._codomain_start,),
2916
+ dtype= self.dtype)
2917
+
2918
+ self._args = {}
2919
+
2920
+ self._func = dot.func
2921
+
2922
+ #===============================================================================
2923
+ del VectorSpace, Vector