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