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.
- feectools/__init__.py +0 -0
- feectools/accelerate/__init__.py +0 -0
- feectools/accelerate/accelerate.py +220 -0
- feectools/accelerate/compile_psydac.mk +52 -0
- feectools/api/__init__.py +0 -0
- feectools/api/essential_bc.py +122 -0
- feectools/api/fem_bilinear_form.py +2226 -0
- feectools/api/fem_common.py +286 -0
- feectools/api/fem_sum_form.py +123 -0
- feectools/api/settings.py +82 -0
- feectools/core/__init__.py +11 -0
- feectools/core/bsplines.py +1107 -0
- feectools/core/bsplines_kernels.py +1349 -0
- feectools/core/field_evaluation_kernels.py +5015 -0
- feectools/core/tests/__init__.py +0 -0
- feectools/core/tests/test_bsplines.py +263 -0
- feectools/core/tests/test_bsplines_kernel.py +40 -0
- feectools/core/tests/test_bsplines_pyccel.py +752 -0
- feectools/ddm/__init__.py +3 -0
- feectools/ddm/basic.py +78 -0
- feectools/ddm/blocking_data_exchanger.py +348 -0
- feectools/ddm/cart.py +1835 -0
- feectools/ddm/interface_data_exchanger.py +122 -0
- feectools/ddm/mpi.py +109 -0
- feectools/ddm/nonblocking_data_exchanger.py +331 -0
- feectools/ddm/partition.py +207 -0
- feectools/ddm/petsc.py +112 -0
- feectools/ddm/tests/__init__.py +0 -0
- feectools/ddm/tests/test_cart_1d.py +138 -0
- feectools/ddm/tests/test_cart_2d.py +164 -0
- feectools/ddm/tests/test_cart_3d.py +158 -0
- feectools/ddm/tests/test_multicart_2d.py +173 -0
- feectools/ddm/tests/test_partition.py +124 -0
- feectools/ddm/utilities.py +24 -0
- feectools/feec/__init__.py +0 -0
- feectools/feec/derivatives.py +780 -0
- feectools/feec/dof_kernels.py +210 -0
- feectools/feec/global_geometric_projectors.py +1073 -0
- feectools/feec/hodge.py +148 -0
- feectools/fem/__init__.py +0 -0
- feectools/fem/basic.py +465 -0
- feectools/fem/grid.py +181 -0
- feectools/fem/partitioning.py +344 -0
- feectools/fem/projectors.py +160 -0
- feectools/fem/splines.py +559 -0
- feectools/fem/tensor.py +1393 -0
- feectools/fem/tests/__init__.py +0 -0
- feectools/fem/tests/analytical_profiles_1d.py +100 -0
- feectools/fem/tests/analytical_profiles_base.py +34 -0
- feectools/fem/tests/splines_error_bounds.py +155 -0
- feectools/fem/tests/test_spline_histopolation.py +120 -0
- feectools/fem/tests/test_spline_interpolation.py +182 -0
- feectools/fem/tests/test_splines.py +184 -0
- feectools/fem/tests/test_splines_par.py +46 -0
- feectools/fem/tests/test_vector_spaces.py +150 -0
- feectools/fem/tests/utilities.py +47 -0
- feectools/fem/vector.py +729 -0
- feectools/linalg/__init__.py +0 -0
- feectools/linalg/basic.py +1386 -0
- feectools/linalg/block.py +1451 -0
- feectools/linalg/direct_solvers.py +201 -0
- feectools/linalg/fft.py +258 -0
- feectools/linalg/kernels/__init__.py +0 -0
- feectools/linalg/kernels/axpy_kernels.py +57 -0
- feectools/linalg/kernels/inner_kernels.py +100 -0
- feectools/linalg/kernels/matvec_kernels.py +206 -0
- feectools/linalg/kernels/stencil2IJV_kernels.py +227 -0
- feectools/linalg/kernels/stencil2coo_kernels.py +179 -0
- feectools/linalg/kernels/transpose_kernels.py +263 -0
- feectools/linalg/kron.py +911 -0
- feectools/linalg/solvers.py +1914 -0
- feectools/linalg/sparse.py +114 -0
- feectools/linalg/stencil.py +2923 -0
- feectools/linalg/stencil_dot_kernels.py +317 -0
- feectools/linalg/stencil_transpose_kernels.py +372 -0
- feectools/linalg/tests/__init__.py +0 -0
- feectools/linalg/tests/test_block.py +1588 -0
- feectools/linalg/tests/test_fft.py +106 -0
- feectools/linalg/tests/test_kron_stencil_matrix.py +114 -0
- feectools/linalg/tests/test_linalg.py +1065 -0
- feectools/linalg/tests/test_matrix_free.py +128 -0
- feectools/linalg/tests/test_solvers.py +213 -0
- feectools/linalg/tests/test_stencil_interface_matrix.py +379 -0
- feectools/linalg/tests/test_stencil_vector.py +1036 -0
- feectools/linalg/tests/test_stencil_vector_space.py +440 -0
- feectools/linalg/topetsc.py +522 -0
- feectools/linalg/utilities.py +200 -0
- feectools/utilities/__init__.py +0 -0
- feectools/utilities/quadratures.py +113 -0
- feectools/utilities/utils.py +166 -0
- feectools/version.py +1 -0
- feectools-0.1.0.dist-info/METADATA +66 -0
- feectools-0.1.0.dist-info/RECORD +98 -0
- feectools-0.1.0.dist-info/WHEEL +5 -0
- feectools-0.1.0.dist-info/entry_points.txt +3 -0
- feectools-0.1.0.dist-info/licenses/AUTHORS +22 -0
- feectools-0.1.0.dist-info/licenses/LICENSE +21 -0
- 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
|