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,1451 @@
1
+ # coding: utf-8
2
+ #
3
+ # Copyright 2018 Jalal Lakhlili, Yaman Güçlü
4
+
5
+ import numpy as np
6
+
7
+ from types import MappingProxyType
8
+ from scipy.sparse import bmat, lil_matrix
9
+
10
+ from feectools.linalg.basic import VectorSpace, Vector, LinearOperator
11
+ from feectools.linalg.stencil import StencilMatrix
12
+ from feectools.ddm.cart import InterfaceCartDecomposition
13
+ from feectools.ddm.utilities import get_data_exchanger
14
+
15
+ __all__ = ('BlockVectorSpace', 'BlockVector', 'BlockLinearOperator')
16
+
17
+ #===============================================================================
18
+ class BlockVectorSpace(VectorSpace):
19
+ """
20
+ Product Vector Space V of two Vector Spaces (V1,V2) or more.
21
+
22
+ Parameters
23
+ ----------
24
+ *spaces : feectools.linalg.basic.VectorSpace
25
+ A list of Vector Spaces.
26
+
27
+ """
28
+ def __new__(cls, *spaces, connectivity=None):
29
+
30
+ # Check that all input arguments are vector spaces
31
+ if not all(isinstance(Vi, VectorSpace) for Vi in spaces):
32
+ raise TypeError('All input spaces must be VectorSpace objects')
33
+
34
+ # If no spaces are passed, raise an error
35
+ if len(spaces) == 0:
36
+ raise ValueError('Cannot create a BlockVectorSpace of zero spaces')
37
+
38
+ # If only one space is passed, return it without creating a new object
39
+ if len(spaces) == 1:
40
+ return spaces[0]
41
+
42
+ # Create a new BlockVectorSpace object
43
+ return VectorSpace.__new__(cls)
44
+
45
+ # ...
46
+ def __init__(self, *spaces, connectivity=None):
47
+
48
+ # Store spaces in a Tuple, because they will not be changed
49
+ self._spaces = tuple(spaces)
50
+
51
+ if all(np.dtype(s.dtype)==np.dtype(spaces[0].dtype) for s in spaces):
52
+ self._dtype = spaces[0].dtype
53
+ else:
54
+ raise NotImplementedError("The matrices domains don't have the same data type.")
55
+
56
+ self._connectivity = connectivity or {}
57
+ self._connectivity_readonly = MappingProxyType(self._connectivity)
58
+
59
+ #--------------------------------------
60
+ # Abstract interface
61
+ #--------------------------------------
62
+ @property
63
+ def dimension(self):
64
+ """
65
+ The dimension of a product space V = (V1, V2, ...] is the cardinality
66
+ (i.e. the number of vectors) of a basis of V over its base field.
67
+
68
+ """
69
+ return sum(Vi.dimension for Vi in self._spaces)
70
+
71
+ # ...
72
+ @property
73
+ def dtype(self):
74
+ return self._dtype
75
+
76
+ # ...
77
+ def zeros(self):
78
+ """
79
+ Get a copy of the null element of the product space V = [V1, V2, ...]
80
+
81
+ Returns
82
+ -------
83
+ null : BlockVector
84
+ A new vector object with all components equal to zero.
85
+
86
+ """
87
+ return BlockVector(self, [Vi.zeros() for Vi in self._spaces])
88
+
89
+ # ...
90
+ def inner(self, x, y):
91
+ """
92
+ Evaluate the inner vector product between two vectors of this space V.
93
+
94
+ If the field of V is real, compute the classical scalar product.
95
+ If the field of V is complex, compute the classical sesquilinear
96
+ product with linearity on the second vector.
97
+
98
+ TODO [YG 01.05.2025]: Currently, the first vector is conjugated. We
99
+ want to reverse this behavior in order to align with the convention
100
+ of FEniCS.
101
+
102
+ Parameters
103
+ ----------
104
+ x : Vector
105
+ The first vector in the scalar product. In the case of a complex
106
+ field, the inner product is antilinear w.r.t. this vector (hence
107
+ this vector is conjugated).
108
+
109
+ y : Vector
110
+ The second vector in the scalar product. The inner product is
111
+ linear w.r.t. this vector.
112
+
113
+ Returns
114
+ -------
115
+ float | complex
116
+ The scalar product of the two vectors. Note that inner(x, x) is
117
+ a non-negative real number which is zero if and only if x = 0.
118
+
119
+ """
120
+
121
+ assert isinstance(x, BlockVector)
122
+ assert isinstance(y, BlockVector)
123
+ assert x.space is self
124
+ assert y.space is self
125
+ return sum(Vi.inner(xi, yi) for Vi, xi, yi in zip(self.spaces, x.blocks, y.blocks))
126
+
127
+ #...
128
+ def axpy(self, a, x, y):
129
+ """
130
+ Increment the vector y with the a-scaled vector x, i.e. y = a * x + y,
131
+ provided that x and y belong to the same vector space V (self).
132
+ The scalar value a may be real or complex, depending on the field of V.
133
+
134
+ Parameters
135
+ ----------
136
+ a : scalar
137
+ The scaling coefficient needed for the operation.
138
+
139
+ x : BlockVector
140
+ The vector which is not modified by this function.
141
+
142
+ y : BlockVector
143
+ The vector modified by this function (incremented by a * x).
144
+ """
145
+
146
+ assert isinstance(x, BlockVector)
147
+ assert isinstance(y, BlockVector)
148
+ assert x.space is self
149
+ assert y.space is self
150
+
151
+ for Vi, xi, yi in zip(self.spaces, x.blocks, y.blocks):
152
+ Vi.axpy(a, xi, yi)
153
+
154
+ x._sync = x._sync and y._sync
155
+
156
+ #--------------------------------------
157
+ # Other properties/methods
158
+ #--------------------------------------
159
+ @property
160
+ def spaces(self):
161
+ return self._spaces
162
+
163
+ @property
164
+ def parallel(self):
165
+ """ Returns True if the memory is distributed."""
166
+ return self._spaces[0].parallel
167
+
168
+ @property
169
+ def starts(self):
170
+ return [s.starts for s in self._spaces]
171
+
172
+ @property
173
+ def ends(self):
174
+ return [s.ends for s in self._spaces]
175
+
176
+ @property
177
+ def pads(self):
178
+ return self._spaces[0].pads
179
+
180
+ @property
181
+ def n_blocks(self):
182
+ return len(self._spaces)
183
+
184
+ @property
185
+ def connectivity(self):
186
+ return self._connectivity_readonly
187
+
188
+ def __getitem__(self, key):
189
+ return self._spaces[key]
190
+
191
+ #===============================================================================
192
+ class BlockVector(Vector):
193
+ """
194
+ Block of Vectors, which is an element of a BlockVectorSpace.
195
+
196
+ Parameters
197
+ ----------
198
+ V : feectools.linalg.block.BlockVectorSpace
199
+ Space to which the new vector belongs.
200
+
201
+ blocks : list or tuple (feectools.linalg.basic.Vector)
202
+ List of Vector objects, belonging to the correct spaces (optional).
203
+
204
+ """
205
+ def __init__(self, V, blocks=None):
206
+
207
+ assert isinstance(V, BlockVectorSpace)
208
+ self._space = V
209
+
210
+ # We store the blocks in a List so that we can change them later.
211
+ if blocks:
212
+ # Verify that vectors belong to correct spaces and store them
213
+ assert isinstance(blocks, (list, tuple))
214
+ assert all((isinstance(b, Vector)) for b in blocks)
215
+ assert all((Vi is bi.space) for Vi,bi in zip(V.spaces, blocks))
216
+
217
+ self._blocks = list(blocks)
218
+ else:
219
+ # TODO: Each block is a 'zeros' vector of the correct space for now,
220
+ # but in the future we would like 'empty' vectors of the same space.
221
+ self._blocks = [Vi.zeros() for Vi in V.spaces]
222
+
223
+ # TODO: distinguish between different directions
224
+ self._sync = False
225
+
226
+ self._data_exchangers = {}
227
+ self._interface_buf = {}
228
+
229
+ if not V.parallel: return
230
+
231
+ # Prepare the data exchangers for the interface data
232
+ for i, j in V.connectivity:
233
+ ((axis_i, ext_i),(axis_j, ext_j)) = V.connectivity[i, j]
234
+
235
+ Vi = V.spaces[i]
236
+ Vj = V.spaces[j]
237
+ self._data_exchangers[i, j] = []
238
+
239
+ if isinstance(Vi, BlockVectorSpace) and isinstance(Vj, BlockVectorSpace):
240
+ # case of a system of equations
241
+ for k, (Vik, Vjk) in enumerate(zip(Vi.spaces, Vj.spaces)):
242
+ cart_i = Vik.cart
243
+ cart_j = Vjk.cart
244
+
245
+ if cart_i.is_comm_null and cart_j.is_comm_null: continue
246
+ if not cart_i.is_comm_null and not cart_j.is_comm_null: continue
247
+ if not (axis_i, ext_i) in Vik.interfaces: continue
248
+ cart_ij = Vik.interfaces[axis_i, ext_i].cart
249
+ assert isinstance(cart_ij, InterfaceCartDecomposition)
250
+ self._data_exchangers[i, j].append(get_data_exchanger(cart_ij, self.dtype))
251
+
252
+ elif not isinstance(Vi, BlockVectorSpace) and not isinstance(Vj, BlockVectorSpace):
253
+ # case of scalar equations
254
+ cart_i = Vi.cart
255
+ cart_j = Vj.cart
256
+ if cart_i.is_comm_null and cart_j.is_comm_null: continue
257
+ if not cart_i.is_comm_null and not cart_j.is_comm_null: continue
258
+ if not (axis_i, ext_i) in Vi.interfaces: continue
259
+
260
+ cart_ij = Vi.interfaces[axis_i, ext_i].cart
261
+ assert isinstance(cart_ij, InterfaceCartDecomposition)
262
+ self._data_exchangers[i, j].append(get_data_exchanger(cart_ij, self.dtype))
263
+ else:
264
+ raise NotImplementedError("This case is not treated")
265
+
266
+ for i, j in V.connectivity:
267
+ if len(self._data_exchangers.get((i, j), [])) == 0:
268
+ self._data_exchangers.pop((i, j), None)
269
+
270
+ #--------------------------------------
271
+ # Abstract interface
272
+ #--------------------------------------
273
+ @property
274
+ def space(self):
275
+ """ Vector space to which this vector belongs. """
276
+ return self._space
277
+
278
+ # ...
279
+ def toarray(self, order='C'):
280
+ """ Convert to Numpy 1D array. """
281
+ return np.concatenate([bi.toarray(order=order) for bi in self._blocks])
282
+
283
+ #...
284
+ def copy(self, out=None):
285
+ if self is out:
286
+ return self
287
+ w = out or BlockVector(self._space)#, [b.copy() for b in self._blocks])
288
+ for n, b in enumerate(self._blocks):
289
+ b.copy(out=w[n])
290
+ w._sync = self._sync
291
+ return w
292
+
293
+ #...
294
+ def conjugate(self, out=None):
295
+ if out is not None:
296
+ assert isinstance(out, BlockVector)
297
+ assert out.space is self.space
298
+ else:
299
+ out = BlockVector(self.space)
300
+
301
+ for (Lij, Lij_out) in zip(self.blocks, out.blocks):
302
+ Lij.conjugate(out=Lij_out)
303
+ out._sync = self._sync
304
+ return out
305
+
306
+ #...
307
+ def __neg__(self):
308
+ w = BlockVector(self._space, [-b for b in self._blocks])
309
+ w._sync = self._sync
310
+ return w
311
+
312
+ #...
313
+ def __mul__(self, a):
314
+ w = BlockVector(self._space, [b * a for b in self._blocks])
315
+ w._sync = self._sync
316
+ return w
317
+
318
+ #...
319
+ def __add__(self, v):
320
+ assert isinstance(v, BlockVector)
321
+ assert v._space is self._space
322
+ w = BlockVector(self._space, [b1 + b2 for b1, b2 in zip(self._blocks, v._blocks)])
323
+ w._sync = self._sync and v._sync
324
+ return w
325
+
326
+ #...
327
+ def __sub__(self, v):
328
+ assert isinstance(v, BlockVector)
329
+ assert v._space is self._space
330
+ w = BlockVector(self._space, [b1 - b2 for b1, b2 in zip(self._blocks, v._blocks)])
331
+ w._sync = self._sync and v._sync
332
+ return w
333
+
334
+ #...
335
+ def __imul__(self, a):
336
+ for b in self._blocks:
337
+ b *= a
338
+ return self
339
+
340
+ #...
341
+ def __iadd__(self, v):
342
+ assert isinstance(v, BlockVector)
343
+ assert v._space is self._space
344
+ for b1, b2 in zip(self._blocks, v._blocks):
345
+ b1 += b2
346
+ self._sync = self._sync and v._sync
347
+ return self
348
+
349
+ #...
350
+ def __isub__(self, v):
351
+ assert isinstance(v, BlockVector)
352
+ assert v._space is self._space
353
+ for b1, b2 in zip(self._blocks, v._blocks):
354
+ b1 -= b2
355
+ self._sync = self._sync and v._sync
356
+ return self
357
+
358
+ #--------------------------------------
359
+ # Other properties/methods
360
+ #--------------------------------------
361
+ @property
362
+ def blocks(self):
363
+ return tuple(self._blocks)
364
+
365
+ #...
366
+ @property
367
+ def n_blocks(self):
368
+ return len(self._blocks)
369
+
370
+ # ...
371
+ def __getitem__(self, key):
372
+ return self._blocks[key]
373
+
374
+ # ...
375
+ def __setitem__(self, key, value):
376
+ assert value.space == self.space[key]
377
+ assert isinstance(value, Vector)
378
+ self._blocks[key] = value
379
+
380
+ # ...
381
+ @property
382
+ def ghost_regions_in_sync(self):
383
+ return self._sync
384
+
385
+ # ...
386
+ # NOTE: this property must be set collectively
387
+ @ghost_regions_in_sync.setter
388
+ def ghost_regions_in_sync(self, value):
389
+ assert isinstance(value, bool)
390
+ self._sync = value
391
+ for vi in self.blocks:
392
+ vi.ghost_regions_in_sync = value
393
+
394
+ # ...
395
+ def update_ghost_regions(self):
396
+
397
+ req = self.start_update_interface_ghost_regions()
398
+
399
+ for vi in self.blocks:
400
+ vi.update_ghost_regions()
401
+
402
+ self.end_update_interface_ghost_regions(req)
403
+
404
+ # Flag ghost regions as up-to-date
405
+ self._sync = True
406
+
407
+ def start_update_interface_ghost_regions(self):
408
+ self._collect_interface_buf()
409
+ req = {}
410
+ for (i, j) in self._data_exchangers:
411
+ req[i, j] = [data_ex.start_update_ghost_regions(*bufs) for bufs, data_ex in zip(self._interface_buf[i, j], self._data_exchangers[i, j])]
412
+
413
+ return req
414
+
415
+ def end_update_interface_ghost_regions(self, req):
416
+
417
+ for (i, j) in self._data_exchangers:
418
+ for data_ex, bufs, req_ij in zip(self._data_exchangers[i, j], self._interface_buf[i, j], req[i, j]):
419
+ data_ex.end_update_ghost_regions(req_ij)
420
+
421
+ def _collect_interface_buf(self):
422
+ V = self.space
423
+ if not V.parallel:return
424
+ for i, j in V.connectivity:
425
+ if (i, j) not in self._data_exchangers:
426
+ continue
427
+ ((axis_i, ext_i), (axis_j, ext_j)) = V.connectivity[i, j]
428
+
429
+ Vi = V.spaces[i]
430
+ Vj = V.spaces[j]
431
+
432
+ # The process that owns the patch i will use block i to send data and receive in block j
433
+ self._interface_buf[i, j] = []
434
+ if isinstance(Vi, BlockVectorSpace) and isinstance(Vj, BlockVectorSpace):
435
+ # case of a system of equations
436
+ for k, (Vik, Vjk) in enumerate(zip(Vi.spaces, Vj.spaces)):
437
+
438
+ cart_i = Vik.cart
439
+ cart_j = Vjk.cart
440
+
441
+ buf = [None]*2
442
+ if cart_i.is_comm_null:
443
+ buf[0] = self._blocks[i]._blocks[k]._interface_data[axis_i, ext_i]
444
+ else:
445
+ buf[0] = self._blocks[i]._blocks[k]._data
446
+
447
+ if cart_j.is_comm_null:
448
+ buf[1] = self._blocks[j]._blocks[k]._interface_data[axis_j, ext_j]
449
+ else:
450
+ buf[1] = self._blocks[j]._blocks[k]._data
451
+
452
+ self._interface_buf[i,j].append(tuple(buf))
453
+ elif not isinstance(Vi, BlockVectorSpace) and not isinstance(Vj, BlockVectorSpace):
454
+ # case of scalar equations
455
+ cart_i = Vi.cart
456
+ cart_j = Vj.cart
457
+
458
+ if cart_i.is_comm_null:
459
+ read_buffer = self._blocks[i]._interface_data[axis_i, ext_i]
460
+ else:
461
+ read_buffer = self._blocks[i]._data
462
+
463
+ if cart_j.is_comm_null:
464
+ write_buffer = self._blocks[j]._interface_data[axis_j, ext_j]
465
+ else:
466
+ write_buffer = self._blocks[j]._data
467
+
468
+ self._interface_buf[i, j].append((read_buffer, write_buffer))
469
+
470
+ # ...
471
+ def exchange_assembly_data(self):
472
+ for vi in self.blocks:
473
+ vi.exchange_assembly_data()
474
+
475
+ # ...
476
+ def toarray_local(self, order='C'):
477
+ """ Convert to petsc Nest vector.
478
+ """
479
+
480
+ blocks = [v.toarray_local(order=order) for v in self._blocks]
481
+ return np.block([blocks])[0]
482
+
483
+ # ...
484
+ def topetsc(self):
485
+ """ Convert to petsc data structure.
486
+ """
487
+ from feectools.linalg.topetsc import vec_topetsc
488
+ vec = vec_topetsc( self )
489
+ return vec
490
+
491
+ #===============================================================================
492
+ class BlockLinearOperator(LinearOperator):
493
+ """
494
+ Linear operator that can be written as blocks of other Linear Operators.
495
+ Either the domain or the codomain of this operator, or both, should be of
496
+ class BlockVectorSpace.
497
+
498
+ Parameters
499
+ ----------
500
+ V1 : feectools.linalg.block.VectorSpace
501
+ Domain of the new linear operator.
502
+
503
+ V2 : feectools.linalg.block.VectorSpace
504
+ Codomain of the new linear operator.
505
+
506
+ blocks : dict | (list of lists) | (tuple of tuples)
507
+ LinearOperator objects (optional).
508
+
509
+ a) 'blocks' can be dictionary with
510
+ . key = tuple (i, j), where i and j are two integers >= 0
511
+ . value = corresponding LinearOperator Lij
512
+
513
+ b) 'blocks' can be list of lists (or tuple of tuples) where blocks[i][j]
514
+ is the LinearOperator Lij (if None, we assume null operator)
515
+
516
+ """
517
+ def __init__(self, V1, V2, blocks=None):
518
+
519
+ assert isinstance(V1, VectorSpace)
520
+ assert isinstance(V2, VectorSpace)
521
+
522
+ if not (isinstance(V1, BlockVectorSpace) or isinstance(V2, BlockVectorSpace)):
523
+ raise TypeError("Either domain or codomain must be of type BlockVectorSpace")
524
+
525
+ self._domain = V1
526
+ self._codomain = V2
527
+ self._blocks = {}
528
+
529
+ self._nrows = V2.n_blocks if isinstance(V2, BlockVectorSpace) else 1
530
+ self._ncols = V1.n_blocks if isinstance(V1, BlockVectorSpace) else 1
531
+
532
+ # Store blocks in dict (hence they can be manually changed later)
533
+ if blocks:
534
+
535
+ if isinstance(blocks, dict):
536
+ for (i, j), Lij in blocks.items():
537
+ self[i, j] = Lij
538
+
539
+ elif isinstance(blocks, (list, tuple)):
540
+ blocks = np.array(blocks, dtype=object)
541
+ for (i, j), Lij in np.ndenumerate(blocks):
542
+ self[i, j] = Lij
543
+
544
+ else:
545
+ raise ValueError( "Blocks can only be given as dict or 2D list/tuple." )
546
+
547
+ self._args = {}
548
+ self._blocks_as_args = self._blocks
549
+ self._increment = self._codomain.zeros()
550
+ self._args['inc'] = self._increment
551
+ self._args['n_rows'] = self._nrows
552
+ self._args['n_cols'] = self._ncols
553
+ self._func = self._dot
554
+ self._sync = False
555
+ self._backend = None
556
+
557
+ #--------------------------------------
558
+ # Abstract interface
559
+ #--------------------------------------
560
+ @property
561
+ def domain(self):
562
+ return self._domain
563
+
564
+ # ...
565
+ @property
566
+ def codomain(self):
567
+ return self._codomain
568
+
569
+ # ...
570
+ @property
571
+ def dtype(self):
572
+ return self.domain.dtype
573
+
574
+ def conjugate(self, out=None):
575
+ if out is not None:
576
+ assert isinstance(out, BlockLinearOperator)
577
+ assert out.domain is self.domain
578
+ assert out.codomain is self.codomain
579
+ else:
580
+ out = BlockLinearOperator(self.domain, self.codomain)
581
+
582
+ for (i, j), Lij in self._blocks.items():
583
+ assert isinstance(Lij, (StencilMatrix, BlockLinearOperator))
584
+ if out[i,j]==None:
585
+ out[i, j] = Lij.conjugate()
586
+ else:
587
+ Lij.conjugate(out=out[i,j])
588
+
589
+ return out
590
+
591
+ def conj(self, out=None):
592
+ return self.conjugate(out=out)
593
+
594
+ # NOTE [YG 27.03.2023]:
595
+ # NOTE as part of PR 279, this method was added to facilitate comparisons in tests,
596
+ # NOTE but then commented out as deemed unnecessary.
597
+ # def __eq__(self, B):
598
+ # """
599
+ # Return True if self and B are mathematically the same, else return False.
600
+ # Also returns False if at least one block is not the same object and the entries can't be accessed and compared using toarray().
601
+ #
602
+ # """
603
+ # assert isinstance(B, BlockLinearOperator)
604
+ #
605
+ # if self is B:
606
+ # return True
607
+ #
608
+ # nrows = self._nrows
609
+ # ncols = self._ncols
610
+ # if not ((B.n_block_cols == ncols) & (B.n_block_rows == nrows)):
611
+ # return False
612
+ #
613
+ # for i in range(nrows):
614
+ # for j in range(ncols):
615
+ # A_ij = self[i, j]
616
+ # B_ij = B[i, j]
617
+ # if not ( A_ij is B_ij ):
618
+ # if not (((A_ij is None) or (isinstance(A_ij, ZeroOperator))) & ((B_ij is None) or (isinstance(B_ij, ZeroOperator)))):
619
+ # if not ( np.array_equal(A_ij.toarray(), B_ij.toarray()) ):
620
+ # return False
621
+ # return True
622
+
623
+ # ...
624
+ def tosparse(self, **kwargs):
625
+ """ Convert to any Scipy sparse matrix format. """
626
+
627
+ # Shortcuts
628
+ nrows = self.n_block_rows
629
+ ncols = self.n_block_cols
630
+
631
+ # Utility functions: get domain of blocks on column j, get codomain of blocks on row i
632
+ block_domain = (lambda j: self.domain [j]) if ncols > 1 else (lambda j: self.domain)
633
+ block_codomain = (lambda i: self.codomain[i]) if nrows > 1 else (lambda i: self.codomain)
634
+
635
+ # Convert all blocks to Scipy sparse format
636
+ blocks_sparse = [[None for j in range(ncols)] for i in range(nrows)]
637
+ for i in range(nrows):
638
+ for j in range(ncols):
639
+ if (i, j) in self._blocks:
640
+ blocks_sparse[i][j] = self._blocks[i, j].tosparse(**kwargs)
641
+ else:
642
+ m = block_codomain(i).dimension
643
+ n = block_domain (j).dimension
644
+ blocks_sparse[i][j] = lil_matrix((m, n))
645
+
646
+ # Create sparse matrix from sparse blocks
647
+ M = bmat( blocks_sparse )
648
+ M.eliminate_zeros()
649
+
650
+ # Sanity check
651
+ assert M.shape[0] == self.codomain.dimension
652
+ assert M.shape[1] == self. domain.dimension
653
+
654
+ return M
655
+
656
+ # ...
657
+ def toarray(self, **kwargs):
658
+ """ Convert to Numpy 2D array. """
659
+ return self.tosparse(**kwargs).toarray()
660
+
661
+ # ...
662
+ def dot(self, v, out=None):
663
+
664
+ if self.n_block_cols == 1:
665
+ assert isinstance(v, Vector)
666
+ else:
667
+ assert isinstance(v, BlockVector)
668
+
669
+ assert v.space is self.domain
670
+
671
+ if out is not None:
672
+ if self.n_block_rows == 1:
673
+ assert isinstance(out, Vector)
674
+ else:
675
+ assert isinstance(out, BlockVector)
676
+
677
+ assert out.space is self.codomain
678
+ out *= 0.0
679
+ else:
680
+ out = self.codomain.zeros()
681
+
682
+ if not v.ghost_regions_in_sync:
683
+ v.update_ghost_regions()
684
+
685
+ self._func(self._blocks_as_args, v, out, **self._args)
686
+
687
+ out.ghost_regions_in_sync = False
688
+ return out
689
+
690
+ #...
691
+ @staticmethod
692
+ def _dot(blocks, v, out, n_rows, n_cols, inc):
693
+
694
+ if n_rows == 1:
695
+ for (_, j), L0j in blocks.items():
696
+ out += L0j.dot(v[j], out=inc)
697
+ elif n_cols == 1:
698
+ for (i, _), Li0 in blocks.items():
699
+ out[i] += Li0.dot(v, out=inc[i])
700
+ else:
701
+ for (i, j), Lij in blocks.items():
702
+ out[i] += Lij.dot(v[j], out=inc[i])
703
+
704
+ # ...
705
+ def transpose(self, conjugate=False, out=None):
706
+ """"
707
+ Return the transposed BlockLinearOperator, or the Hermitian Transpose if conjugate==True
708
+
709
+ Parameters
710
+ ----------
711
+ conjugate : Bool(optional)
712
+ True to get the Hermitian adjoint.
713
+
714
+ out : BlockLinearOperator(optional)
715
+ Optional out for the transpose to avoid temporaries
716
+ """
717
+ if out is not None:
718
+ assert isinstance(out, BlockLinearOperator)
719
+ assert out.codomain is self.domain
720
+ assert out.domain is self.codomain
721
+ for (i, j), Lij in self._blocks.items():
722
+ if out[j,i]==None:
723
+ out[j, i] = Lij.transpose(conjugate=conjugate)
724
+ else:
725
+ Lij.transpose(conjugate=conjugate, out=out[j,i])
726
+ else:
727
+ blocks, blocks_T = self.compute_interface_matrices_transpose()
728
+ blocks = {(j, i): b.transpose(conjugate=conjugate) for (i, j), b in blocks.items()}
729
+ blocks.update(blocks_T)
730
+ out = BlockLinearOperator(self.codomain, self.domain, blocks=blocks)
731
+
732
+ out.set_backend(self._backend)
733
+ return out
734
+
735
+ #--------------------------------------
736
+ # Overridden properties/methods
737
+ #--------------------------------------
738
+ def __neg__(self):
739
+ blocks = {ij: -Bij for ij, Bij in self._blocks.items()}
740
+ mat = BlockLinearOperator(self.domain, self.codomain, blocks=blocks)
741
+ if self._backend is not None:
742
+ mat._func = self._func
743
+ mat._args = self._args
744
+ mat._blocks_as_args = [mat._blocks[key]._data for key in self._blocks]
745
+ mat._backend = self._backend
746
+ return mat
747
+
748
+ # ...
749
+ def __mul__(self, a):
750
+ blocks = {ij: Bij * a for ij, Bij in self._blocks.items()}
751
+ mat = BlockLinearOperator(self.domain, self.codomain, blocks=blocks)
752
+ if self._backend is not None:
753
+ mat._func = self._func
754
+ mat._args = self._args
755
+ mat._blocks_as_args = [mat._blocks[key]._data for key in self._blocks]
756
+ mat._backend = self._backend
757
+ return mat
758
+
759
+ # ...
760
+ def __add__(self, M):
761
+ if not isinstance(M, BlockLinearOperator):
762
+ return LinearOperator.__add__(self, M)
763
+
764
+ assert M. domain is self.domain
765
+ assert M.codomain is self.codomain
766
+ blocks = {}
767
+ for ij in set(self._blocks.keys()) | set(M._blocks.keys()):
768
+ Bij = self[ij]
769
+ Mij = M[ij]
770
+ if Bij is None: blocks[ij] = Mij.copy()
771
+ elif Mij is None: blocks[ij] = Bij.copy()
772
+ else : blocks[ij] = Bij + Mij
773
+ mat = BlockLinearOperator(self.domain, self.codomain, blocks=blocks)
774
+ if len(mat._blocks) != len(self._blocks):
775
+ mat.set_backend(self._backend)
776
+ elif self._backend is not None:
777
+ mat._func = self._func
778
+ mat._args = self._args
779
+ mat._blocks_as_args = [mat._blocks[key]._data for key in self._blocks]
780
+ mat._backend = self._backend
781
+ return mat
782
+
783
+ # ...
784
+ def __sub__(self, M):
785
+ if not isinstance(M, BlockLinearOperator):
786
+ return LinearOperator.__sub__(self, M)
787
+
788
+ assert M. domain is self. domain
789
+ assert M.codomain is self.codomain
790
+ blocks = {}
791
+ for ij in set(self._blocks.keys()) | set(M._blocks.keys()):
792
+ Bij = self[ij]
793
+ Mij = M[ij]
794
+ if Bij is None: blocks[ij] = -Mij
795
+ elif Mij is None: blocks[ij] = Bij.copy()
796
+ else : blocks[ij] = Bij - Mij
797
+ mat = BlockLinearOperator(self.domain, self.codomain, blocks=blocks)
798
+ if len(mat._blocks) != len(self._blocks):
799
+ mat.set_backend(self._backend)
800
+ elif self._backend is not None:
801
+ mat._func = self._func
802
+ mat._args = self._args
803
+ mat._blocks_as_args = [mat._blocks[key]._data for key in self._blocks]
804
+ mat._backend = self._backend
805
+ return mat
806
+
807
+ #--------------------------------------
808
+ # New properties/methods
809
+ #--------------------------------------
810
+ def diagonal(self, *, inverse = False, sqrt = False, out = None):
811
+ """Get the coefficients on the main diagonal as another BlockLinearOperator object.
812
+
813
+ Parameters
814
+ ----------
815
+ inverse : bool
816
+ If True, get the inverse of the diagonal. (Default: False).
817
+ Can be combined with sqrt to get the inverse square root.
818
+
819
+ sqrt : bool
820
+ If True, get the square root of the diagonal. (Default: False).
821
+ Can be combined with inverse to get the inverse square root.
822
+
823
+ out : BlockLinearOperator
824
+ If provided, write the diagonal entries into this matrix. (Default: None).
825
+
826
+ Returns
827
+ -------
828
+ BlockLinearOperator
829
+ The matrix which contains the main diagonal of self (or its inverse).
830
+
831
+ """
832
+ # Determine domain and codomain of result
833
+ V, W = self.domain, self.codomain
834
+ if inverse:
835
+ V, W = W, V
836
+
837
+ # Check the `out` argument, if `None` create a new BlockLinearOperator
838
+ if out is not None:
839
+ assert isinstance(out, BlockLinearOperator)
840
+ assert out.domain is V
841
+ assert out.codomain is W
842
+
843
+ # Set any off-diagonal blocks to zero
844
+ for i, j in out.nonzero_block_indices:
845
+ if i != j:
846
+ out[i, j] = None
847
+ else:
848
+ out = BlockLinearOperator(V, W)
849
+
850
+ # Store the diagonal (or its inverse) into `out`
851
+ for i, j in self.nonzero_block_indices:
852
+ if i == j:
853
+ out[i, i] = self[i, i].diagonal(inverse = inverse, sqrt = sqrt, out = out[i, i])
854
+
855
+ return out
856
+
857
+ # ...
858
+ @property
859
+ def blocks(self):
860
+ """ Immutable 2D view (tuple of tuples) of the linear operator,
861
+ including the empty blocks as 'None' objects.
862
+ """
863
+ return tuple(
864
+ tuple(self._blocks.get((i, j), None) for j in range(self.n_block_cols))
865
+ for i in range(self.n_block_rows))
866
+
867
+ # ...
868
+ @property
869
+ def n_block_rows(self):
870
+ return self._nrows
871
+
872
+ # ...
873
+ @property
874
+ def n_block_cols(self):
875
+ return self._ncols
876
+
877
+ @property
878
+ def nonzero_block_indices(self):
879
+ """
880
+ Tuple of (i, j) pairs which identify the non-zero blocks:
881
+ i is the row index, j is the column index.
882
+ """
883
+ return tuple(self._blocks)
884
+
885
+ # ...
886
+ def update_ghost_regions(self):
887
+ for Lij in self._blocks.values():
888
+ Lij.update_ghost_regions()
889
+
890
+ # ...
891
+ def exchange_assembly_data(self):
892
+ for Lij in self._blocks.values():
893
+ Lij.exchange_assembly_data()
894
+
895
+ # ...
896
+ def remove_spurious_entries(self ):
897
+ for Lij in self._blocks.values():
898
+ Lij.remove_spurious_entries()
899
+
900
+ @property
901
+ def ghost_regions_in_sync(self):
902
+ return self._sync
903
+
904
+ @ghost_regions_in_sync.setter
905
+ def ghost_regions_in_sync( self, value ):
906
+ assert isinstance( value, bool )
907
+ self._sync = value
908
+ for Lij in self._blocks.values():
909
+ Lij.ghost_regions_in_sync = value
910
+
911
+ # ...
912
+ def __getitem__(self, key):
913
+
914
+ assert isinstance( key, tuple )
915
+ assert len( key ) == 2
916
+ assert 0 <= key[0] < self.n_block_rows
917
+ assert 0 <= key[1] < self.n_block_cols
918
+
919
+ return self._blocks.get( key, None )
920
+
921
+ # ...
922
+ def __setitem__(self, key, value):
923
+
924
+ assert isinstance( key, tuple )
925
+ assert len( key ) == 2
926
+ assert 0 <= key[0] < self.n_block_rows
927
+ assert 0 <= key[1] < self.n_block_cols
928
+
929
+ if value is None:
930
+ self._blocks.pop( key, None )
931
+ return
932
+
933
+ i,j = key
934
+ assert isinstance( value, LinearOperator )
935
+
936
+ # Check domain of rhs
937
+ if self.n_block_cols == 1:
938
+ assert value.domain is self.domain
939
+ else:
940
+ assert value.domain is self.domain[j]
941
+
942
+ # Check codomain of rhs
943
+ if self.n_block_rows == 1:
944
+ assert value.codomain is self.codomain
945
+ else:
946
+ assert value.codomain is self.codomain[i]
947
+
948
+ self._blocks[i,j] = value
949
+
950
+ # ...
951
+ def transform(self, operation):
952
+ """
953
+ Applies an operation on each block in this BlockLinearOperator.
954
+
955
+ Parameters
956
+ ----------
957
+ operation : LinearOperator -> LinearOperator
958
+ The operation which transforms each block.
959
+ """
960
+ blocks = {ij: operation(Bij) for ij, Bij in self._blocks.items()}
961
+ return BlockLinearOperator(self.domain, self.codomain, blocks=blocks)
962
+
963
+ # ...
964
+ def backend(self):
965
+ return self._backend
966
+
967
+ # ...
968
+ def copy(self, out=None):
969
+ """
970
+ Create a copy of self, that can potentially be stored in a given BlockLinearOperator.
971
+
972
+ Parameters
973
+ ----------
974
+ out : BlockLinearOperator(optional)
975
+ The existing BlockLinearOperator in which we want to copy self.
976
+
977
+ Returns
978
+ -------
979
+ BlockLinearOperator
980
+ The copy of `self`, either stored in the given BlockLinearOperator `out`
981
+ (if provided) or in a new one. In the corner case where `out=self` the
982
+ `self` object is immediately returned.
983
+ """
984
+ if out is not None:
985
+ if out is self:
986
+ return self
987
+ assert isinstance(out, BlockLinearOperator)
988
+ assert out.domain is self.domain
989
+ assert out.codomain is self.codomain
990
+ else:
991
+ out = BlockLinearOperator(self.domain, self.codomain)
992
+
993
+ for (i, j), Lij in self._blocks.items():
994
+ if out[i, j] is None:
995
+ out[i, j] = Lij.copy()
996
+ else:
997
+ Lij.copy(out = out[i, j])
998
+
999
+ out.set_backend(self._backend)
1000
+
1001
+ return out
1002
+
1003
+ # ...
1004
+ def __imul__(self, a):
1005
+ for Bij in self._blocks.values():
1006
+ Bij *= a
1007
+ return self
1008
+
1009
+ # ...
1010
+ def __iadd__(self, M):
1011
+ if not isinstance(M, BlockLinearOperator):
1012
+ return LinearOperator.__add__(self, M)
1013
+
1014
+ assert M. domain is self. domain
1015
+ assert M.codomain is self.codomain
1016
+
1017
+ for ij in set(self._blocks.keys()) | set(M._blocks.keys()):
1018
+
1019
+ Mij = M[ij]
1020
+ if Mij is None:
1021
+ continue
1022
+
1023
+ Bij = self[ij]
1024
+ if Bij is None:
1025
+ self[ij] = Mij.copy()
1026
+ else:
1027
+ Bij += Mij
1028
+
1029
+ return self
1030
+
1031
+ # ...
1032
+ def __isub__(self, M):
1033
+ if not isinstance(M, BlockLinearOperator):
1034
+ return LinearOperator.__sub__(self, M)
1035
+
1036
+ assert M. domain is self. domain
1037
+ assert M.codomain is self.codomain
1038
+
1039
+ for ij in set(self._blocks.keys()) | set(M._blocks.keys()):
1040
+
1041
+ Mij = M[ij]
1042
+ if Mij is None:
1043
+ continue
1044
+
1045
+ Bij = self[ij]
1046
+ if Bij is None:
1047
+ self[ij] = -Mij
1048
+ else:
1049
+ Bij -= Mij
1050
+
1051
+ return self
1052
+
1053
+ # ...
1054
+ def topetsc(self):
1055
+ """ Convert to petsc data structure.
1056
+ """
1057
+ from feectools.linalg.topetsc import mat_topetsc
1058
+ mat = mat_topetsc( self )
1059
+ return mat
1060
+
1061
+ def compute_interface_matrices_transpose(self):
1062
+ blocks = self._blocks.copy()
1063
+ blocks_T = {}
1064
+ if not self.codomain.parallel:
1065
+ return blocks, blocks_T
1066
+
1067
+ from feectools.ddm.mpi import mpi as MPI
1068
+ from feectools.linalg.stencil import StencilInterfaceMatrix
1069
+
1070
+ if not isinstance(self.codomain, BlockVectorSpace):
1071
+ return blocks, blocks_T
1072
+
1073
+ V = self.codomain
1074
+
1075
+ for i,j in V.connectivity:
1076
+ ((axis_i,ext_i), (axis_j,ext_j)) = V.connectivity[i,j]
1077
+
1078
+ Vi = V.spaces[i]
1079
+ Vj = V.spaces[j]
1080
+
1081
+ if isinstance(Vi, BlockVectorSpace) and isinstance(Vj, BlockVectorSpace):
1082
+ # case of a system of equations
1083
+ block_ij_exists = False
1084
+ blocks_T[j,i] = BlockLinearOperator(Vi, Vj)
1085
+ block_ij = blocks.get((i,j))._blocks.copy() if self[i,j] else None
1086
+ for k1,Vik1 in enumerate(Vi.spaces):
1087
+ for k2,Vjk2 in enumerate(Vj.spaces):
1088
+ cart_i = Vik1.cart
1089
+ cart_j = Vjk2.cart
1090
+
1091
+ if cart_i.is_comm_null and cart_j.is_comm_null:break
1092
+ if not cart_i.is_comm_null and not cart_j.is_comm_null:break
1093
+ if not (axis_i, ext_i) in Vik1.interfaces: break
1094
+ cart_ij = Vik1.interfaces[axis_i, ext_i].cart
1095
+ assert isinstance(cart_ij, InterfaceCartDecomposition)
1096
+
1097
+ if not cart_i.is_comm_null:
1098
+ if cart_ij.intercomm.rank == 0:
1099
+ root = MPI.ROOT
1100
+ else:
1101
+ root = MPI.PROC_NULL
1102
+
1103
+ else:
1104
+ root = 0
1105
+
1106
+ if not block_ij_exists:
1107
+ block_ij_exists = self[i,j] is not None
1108
+ block_ij_exists = cart_ij.intercomm.bcast(block_ij_exists, root= root) or block_ij_exists
1109
+
1110
+ if not block_ij_exists:break
1111
+ blocks.pop((i,j), None)
1112
+ block_ij_k1k2 = block_ij is not None and (k1,k2) in block_ij is not None
1113
+ block_ij_k1k2 = cart_ij.intercomm.bcast(block_ij_k1k2, root= root) or block_ij_k1k2
1114
+
1115
+ if block_ij_k1k2:
1116
+ if not cart_i.is_comm_null:
1117
+ block_ij_k1k2 = block_ij.pop((k1,k2))
1118
+ info = (block_ij_k1k2.domain_start, block_ij_k1k2.codomain_start, block_ij_k1k2.flip, block_ij_k1k2.pads)
1119
+ cart_ij.intercomm.bcast(info, root= root)
1120
+ else:
1121
+ info = cart_ij.intercomm.bcast(None, root=root)
1122
+ block_ij_k1k2 = StencilInterfaceMatrix(Vjk2, Vik1.interfaces[axis_i, ext_i], info[0], info[1], axis_j, axis_i, ext_j, ext_i, flip=info[2], pads=info[3])
1123
+ block_ji_k2k1 = StencilInterfaceMatrix(Vik1, Vjk2, info[1], info[0], axis_i, axis_j, ext_i, ext_j, flip=info[2], pads=info[3])
1124
+
1125
+ data_exchanger = get_data_exchanger(cart_ij, self.dtype, coeff_shape = block_ij_k1k2._data.shape[block_ij_k1k2._ndim:])
1126
+ data_exchanger.update_ghost_regions(array_minus=block_ij_k1k2._data)
1127
+
1128
+ if cart_i.is_comm_null:
1129
+ blocks_T[j,i][k2,k1] = block_ij_k1k2.transpose(out=block_ji_k2k1)
1130
+ else:
1131
+ continue
1132
+
1133
+ break
1134
+
1135
+ if (j,i) in blocks_T and len(blocks_T[j,i]._blocks) == 0:
1136
+ blocks_T.pop((j,i))
1137
+ if (i,j) in blocks and len(blocks[i,j]._blocks) == 0:
1138
+ blocks.pop((i,j))
1139
+
1140
+ block_ji_exists = False
1141
+ blocks_T[i,j] = BlockLinearOperator(Vj, Vi)
1142
+ block_ji = blocks.get((j,i))._blocks.copy() if self[j,i] else None
1143
+ for k1,Vik1 in enumerate(Vi.spaces):
1144
+ for k2,Vjk2 in enumerate(Vj.spaces):
1145
+ cart_i = Vik1.cart
1146
+ cart_j = Vjk2.cart
1147
+
1148
+ if cart_i.is_comm_null and cart_j.is_comm_null:break
1149
+ if not cart_i.is_comm_null and not cart_j.is_comm_null:break
1150
+ if not (axis_i, ext_i) in Vik1.interfaces: break
1151
+ interface_cart_i = Vik1.interfaces[axis_i, ext_i].cart
1152
+ interface_cart_j = Vjk2.interfaces[axis_j, ext_j].cart
1153
+ assert isinstance(interface_cart_i, InterfaceCartDecomposition)
1154
+ assert isinstance(interface_cart_j, InterfaceCartDecomposition)
1155
+
1156
+ if not cart_j.is_comm_null:
1157
+ if interface_cart_i.intercomm.rank == 0:
1158
+ root = MPI.ROOT
1159
+ else:
1160
+ root = MPI.PROC_NULL
1161
+
1162
+ else:
1163
+ root = 0
1164
+
1165
+ if not block_ji_exists:
1166
+ block_ji_exists = self[j,i] is not None
1167
+ block_ji_exists = interface_cart_i.intercomm.bcast(block_ji_exists, root= root) or block_ji_exists
1168
+
1169
+ if not block_ji_exists:break
1170
+ blocks.pop((j,i), None)
1171
+
1172
+ block_ji_k2k1 = block_ji is not None and (k2,k1) in block_ji is not None
1173
+ block_ji_k2k1 = interface_cart_i.intercomm.bcast(block_ji_k2k1, root= root) or block_ji_k2k1
1174
+
1175
+ if block_ji_k2k1:
1176
+ if not cart_j.is_comm_null:
1177
+ block_ji_k2k1 = block_ji.pop((k2,k1))
1178
+ info = (block_ji_k2k1.domain_start, block_ji_k2k1.codomain_start, block_ji_k2k1.flip, block_ji_k2k1.pads)
1179
+ interface_cart_i.intercomm.bcast(info, root= root)
1180
+ else:
1181
+ info = interface_cart_i.intercomm.bcast(None, root=root)
1182
+ block_ji_k2k1 = StencilInterfaceMatrix(Vik1, Vjk2.interfaces[axis_j, ext_j], info[0], info[1], axis_i, axis_j, ext_i, ext_j, flip=info[2], pads=info[3])
1183
+ block_ij_k1k2 = StencilInterfaceMatrix(Vjk2, Vik1, info[1], info[0], axis_j, axis_i, ext_j, ext_i, flip=info[2], pads=info[3])
1184
+
1185
+ interface_cart_i.comm.Barrier()
1186
+ data_exchanger = get_data_exchanger(interface_cart_j, self.dtype, coeff_shape = block_ji_k2k1._data.shape[block_ji_k2k1._ndim:])
1187
+
1188
+ data_exchanger.update_ghost_regions(array_plus=block_ji_k2k1._data)
1189
+
1190
+ if cart_j.is_comm_null:
1191
+ blocks_T[i,j][k1,k2] = block_ji_k2k1.transpose(out=block_ij_k1k2)
1192
+
1193
+ else:
1194
+ continue
1195
+
1196
+ break
1197
+
1198
+
1199
+ if (i,j) in blocks_T and len(blocks_T[i,j]._blocks) == 0:
1200
+ blocks_T.pop((i,j))
1201
+ if (j,i) in blocks and len(blocks[j,i]._blocks) == 0:
1202
+ blocks.pop((j,i))
1203
+
1204
+ elif not isinstance(Vi, BlockVectorSpace) and not isinstance(Vj, BlockVectorSpace):
1205
+
1206
+ # case of scalar equations
1207
+ cart_i = Vi.cart
1208
+ cart_j = Vj.cart
1209
+ if cart_i.is_comm_null and cart_j.is_comm_null:continue
1210
+ if not cart_i.is_comm_null and not cart_j.is_comm_null:continue
1211
+ if not (axis_i, ext_i) in Vi.interfaces: continue
1212
+ cart_ij = Vi.interfaces[axis_i, ext_i].cart
1213
+ assert isinstance(cart_ij, InterfaceCartDecomposition)
1214
+
1215
+ if not cart_i.is_comm_null:
1216
+ if cart_ij.intercomm.rank == 0:
1217
+ root = MPI.ROOT
1218
+ else:
1219
+ root = MPI.PROC_NULL
1220
+
1221
+ else:
1222
+ root = 0
1223
+
1224
+ block_ij_exists = self[i,j] is not None
1225
+ block_ij_exists = cart_ij.intercomm.bcast(block_ij_exists, root= root) or block_ij_exists
1226
+
1227
+ if block_ij_exists:
1228
+ if not cart_i.is_comm_null:
1229
+ block_ij = blocks.pop((i,j))
1230
+ info = (block_ij.domain_start, block_ij.codomain_start, block_ij.flip, block_ij.pads)
1231
+ cart_ij.intercomm.bcast(info, root= root)
1232
+ else:
1233
+ info = cart_ij.intercomm.bcast(None, root=root)
1234
+ block_ij = StencilInterfaceMatrix(Vj, Vi.interfaces[axis_i, ext_i], info[0], info[1], axis_j, axis_i, ext_j, ext_i, flip=info[2], pads=info[3])
1235
+ block_ji = StencilInterfaceMatrix(Vi, Vj, info[1], info[0], axis_i, axis_j, ext_i, ext_j, flip=info[2], pads=info[3])
1236
+
1237
+ data_exchanger = get_data_exchanger(cart_ij, self.dtype, coeff_shape = block_ij._data.shape[block_ij._ndim:])
1238
+ data_exchanger.update_ghost_regions(array_minus=block_ij._data)
1239
+
1240
+ if cart_i.is_comm_null:
1241
+ blocks_T[j,i] = block_ij.transpose(out=block_ji)
1242
+
1243
+ if not cart_j.is_comm_null:
1244
+ if cart_ij.intercomm.rank == 0:
1245
+ root = MPI.ROOT
1246
+ else:
1247
+ root = MPI.PROC_NULL
1248
+
1249
+ else:
1250
+ root = 0
1251
+
1252
+ block_ji_exists = self[j,i] is not None
1253
+ block_ji_exists = cart_ij.intercomm.bcast(block_ji_exists, root= root) or block_ji_exists
1254
+ if block_ji_exists:
1255
+ if not cart_j.is_comm_null:
1256
+ block_ji = blocks.pop((j,i))
1257
+ info = (block_ji.domain_start, block_ji.codomain_start, block_ji.flip, block_ji.pads)
1258
+ cart_ij.intercomm.bcast((block_ji.domain_start, block_ji.codomain_start, block_ji.flip, block_ji.pads), root= root)
1259
+ else:
1260
+ info = cart_ij.intercomm.bcast(None, root=root)
1261
+ block_ji = StencilInterfaceMatrix(Vi, Vj.interfaces[axis_j, ext_j], info[0], info[1], axis_i, axis_j, ext_i, ext_j, flip=info[2], pads=info[3])
1262
+ block_ij = StencilInterfaceMatrix(Vj, Vi, info[1], info[0], axis_j, axis_i, ext_j, ext_i, flip=info[2], pads=info[3])
1263
+
1264
+ data_exchanger = get_data_exchanger(cart_ij, self.dtype, coeff_shape = block_ji._data.shape[block_ji._ndim:])
1265
+ data_exchanger.update_ghost_regions(array_plus=block_ji._data)
1266
+
1267
+ if cart_j.is_comm_null:
1268
+ blocks_T[i,j] = block_ji.transpose(out=block_ij)
1269
+
1270
+ return blocks, blocks_T
1271
+
1272
+ def set_backend(self, backend, precompiled=False):
1273
+ if isinstance(self.domain, BlockVectorSpace) and isinstance(self.domain.spaces[0], BlockVectorSpace):
1274
+ return
1275
+
1276
+ if isinstance(self.codomain, BlockVectorSpace) and isinstance(self.codomain.spaces[0], BlockVectorSpace):
1277
+ return
1278
+
1279
+ if backend is None:return
1280
+ if backend is self._backend:return
1281
+
1282
+ raise AttributeError(f'This is the tiny-psydac version - must use precompiled kernels (but {precompiled = })!')
1283
+ from feectools.api.ast.linalg import LinearOperatorDot
1284
+ from feectools.linalg.stencil import StencilInterfaceMatrix, StencilMatrix
1285
+
1286
+ if not all(isinstance(b, (StencilMatrix, StencilInterfaceMatrix)) for b in self._blocks.values()):
1287
+ for b in self._blocks.values():
1288
+ b.set_backend(backend)
1289
+ return
1290
+
1291
+ block_shape = (self.n_block_rows, self.n_block_cols)
1292
+
1293
+ keys = self.nonzero_block_indices
1294
+ ndim = self._blocks[keys[0]]._ndim
1295
+ c_starts = []
1296
+ d_starts = []
1297
+
1298
+ interface = isinstance(self._blocks[keys[0]], StencilInterfaceMatrix)
1299
+ if interface:
1300
+ interface_axis = self._blocks[keys[0]]._codomain_axis
1301
+ d_ext = self._blocks[keys[0]]._domain_ext
1302
+ d_axis = self._blocks[keys[0]]._domain_axis
1303
+ flip_axis = self._blocks[keys[0]]._flip
1304
+ permutation = self._blocks[keys[0]]._permutation
1305
+
1306
+ for key in keys:
1307
+ c_starts.append(self._blocks[key]._codomain_start)
1308
+ d_starts.append(self._blocks[key]._domain_start)
1309
+
1310
+ c_starts = tuple(c_starts)
1311
+ d_starts = tuple(d_starts)
1312
+ else:
1313
+ interface_axis = None
1314
+ flip_axis = (1,)*ndim
1315
+ permutation = None
1316
+ c_starts = None
1317
+ d_starts = None
1318
+
1319
+ starts = []
1320
+ nrows = []
1321
+ nrows_extra = []
1322
+ gpads = []
1323
+ pads = []
1324
+ dm = []
1325
+ cm = []
1326
+ for key in keys:
1327
+ nrows.append(self._blocks[key]._dotargs_null['nrows'])
1328
+ nrows_extra.append(self._blocks[key]._dotargs_null['nrows_extra'])
1329
+ gpads.append(self._blocks[key]._dotargs_null['gpads'])
1330
+ pads.append(self._blocks[key]._dotargs_null['pads'])
1331
+ starts.append(self._blocks[key]._dotargs_null['starts'])
1332
+ cm.append(self._blocks[key]._dotargs_null['cm'])
1333
+ dm.append(self._blocks[key]._dotargs_null['dm'])
1334
+
1335
+ if self.domain.parallel:
1336
+ if interface:
1337
+ comm = self.domain.spaces[0].interfaces[d_axis, d_ext].cart.local_comm if isinstance(self.domain, BlockVectorSpace) else self.domain.interfaces[d_axis, d_ext].cart.local_comm
1338
+ else:
1339
+ comm = self.codomain.spaces[0].cart.comm if isinstance(self.codomain, BlockVectorSpace) else self.codomain.cart.comm
1340
+ if self.domain == self.codomain:
1341
+ # In this case nrows_extra[i] == 0 for all i
1342
+ dot = LinearOperatorDot(ndim,
1343
+ block_shape=block_shape,
1344
+ keys=keys,
1345
+ comm=comm,
1346
+ backend=frozenset(backend.items()),
1347
+ gpads=tuple(gpads),
1348
+ pads=tuple(pads),
1349
+ dm=tuple(dm),
1350
+ cm=tuple(cm),
1351
+ interface=interface,
1352
+ flip_axis=flip_axis,
1353
+ interface_axis=interface_axis,
1354
+ d_start=d_starts,
1355
+ c_start=c_starts,
1356
+ dtype=self._domain.dtype)
1357
+
1358
+ self._args = {}
1359
+ for k,key in enumerate(keys):
1360
+ key_str = ''.join(str(i) for i in key)
1361
+ starts_k = starts[k]
1362
+ for i in range(len(starts_k)):
1363
+ self._args['s{}_{}'.format(key_str, i+1)] = np.int64(starts_k[i])
1364
+
1365
+ for k,key in enumerate(keys):
1366
+ key_str = ''.join(str(i) for i in key)
1367
+ nrows_k = nrows[k]
1368
+ for i in range(len(nrows_k)):
1369
+ self._args['n{}_{}'.format(key_str, i+1)] = np.int64(nrows_k[i])
1370
+
1371
+
1372
+ for k,key in enumerate(keys):
1373
+ key_str = ''.join(str(i) for i in key)
1374
+ nrows_extra_k = nrows_extra[k]
1375
+ for i in range(len(nrows_extra_k)):
1376
+ self._args['ne{}_{}'.format(key_str, i+1)] = np.int64(nrows_extra_k[i])
1377
+
1378
+ else:
1379
+ dot = LinearOperatorDot(ndim,
1380
+ block_shape=block_shape,
1381
+ keys=keys,
1382
+ comm=comm,
1383
+ backend=frozenset(backend.items()),
1384
+ gpads=tuple(gpads),
1385
+ pads=tuple(pads),
1386
+ dm=tuple(dm),
1387
+ cm=tuple(cm),
1388
+ interface=interface,
1389
+ flip_axis=flip_axis,
1390
+ interface_axis=interface_axis,
1391
+ d_start=d_starts,
1392
+ c_start=c_starts,
1393
+ dtype=self._domain.dtype)
1394
+
1395
+ self._args = {}
1396
+
1397
+ for k,key in enumerate(keys):
1398
+ key_str = ''.join(str(i) for i in key)
1399
+ starts_k = starts[k]
1400
+ for i in range(len(starts_k)):
1401
+ self._args['s{}_{}'.format(key_str, i+1)] = np.int64(starts_k[i])
1402
+
1403
+ for k,key in enumerate(keys):
1404
+ key_str = ''.join(str(i) for i in key)
1405
+ nrows_k = nrows[k]
1406
+ for i in range(len(nrows_k)):
1407
+ self._args['n{}_{}'.format(key_str, i+1)] = np.int64(nrows_k[i])
1408
+
1409
+ for k,key in enumerate(keys):
1410
+ key_str = ''.join(str(i) for i in key)
1411
+ nrows_extra_k = nrows_extra[k]
1412
+ for i in range(len(nrows_extra_k)):
1413
+ self._args['ne{}_{}'.format(key_str, i+1)] = np.int64(nrows_extra_k[i])
1414
+
1415
+ else:
1416
+ dot = LinearOperatorDot(ndim,
1417
+ block_shape=block_shape,
1418
+ keys=keys,
1419
+ comm=None,
1420
+ backend=frozenset(backend.items()),
1421
+ starts=tuple(starts),
1422
+ nrows=tuple(nrows),
1423
+ nrows_extra=tuple(nrows_extra),
1424
+ gpads=tuple(gpads),
1425
+ pads=tuple(pads),
1426
+ dm=tuple(dm),
1427
+ cm=tuple(cm),
1428
+ interface=interface,
1429
+ flip_axis=flip_axis,
1430
+ interface_axis=interface_axis,
1431
+ d_start=d_starts,
1432
+ c_start=c_starts,
1433
+ dtype=self._domain.dtype)
1434
+ self._args = {}
1435
+
1436
+ self._blocks_as_args = [self._blocks[key]._data for key in keys]
1437
+ dot = dot.func
1438
+
1439
+ if interface:
1440
+ def func(blocks, v, out, **args):
1441
+ vs = [vi._interface_data[d_axis, d_ext] for vi in v.blocks] if isinstance(v, BlockVector) else [v._data]
1442
+ outs = [outi._data for outi in out.blocks] if isinstance(out, BlockVector) else [out._data]
1443
+ dot(*blocks, *vs, *outs, **args)
1444
+ else:
1445
+ def func(blocks, v, out, **args):
1446
+ vs = [vi._data for vi in v.blocks] if isinstance(v, BlockVector) else [v._data]
1447
+ outs = [outi._data for outi in out.blocks] if isinstance(out, BlockVector) else [out._data]
1448
+ dot(*blocks, *vs, *outs, **args)
1449
+
1450
+ self._func = func
1451
+ self._backend = backend