feectools 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (98) hide show
  1. feectools/__init__.py +0 -0
  2. feectools/accelerate/__init__.py +0 -0
  3. feectools/accelerate/accelerate.py +220 -0
  4. feectools/accelerate/compile_psydac.mk +52 -0
  5. feectools/api/__init__.py +0 -0
  6. feectools/api/essential_bc.py +122 -0
  7. feectools/api/fem_bilinear_form.py +2226 -0
  8. feectools/api/fem_common.py +286 -0
  9. feectools/api/fem_sum_form.py +123 -0
  10. feectools/api/settings.py +82 -0
  11. feectools/core/__init__.py +11 -0
  12. feectools/core/bsplines.py +1107 -0
  13. feectools/core/bsplines_kernels.py +1349 -0
  14. feectools/core/field_evaluation_kernels.py +5015 -0
  15. feectools/core/tests/__init__.py +0 -0
  16. feectools/core/tests/test_bsplines.py +263 -0
  17. feectools/core/tests/test_bsplines_kernel.py +40 -0
  18. feectools/core/tests/test_bsplines_pyccel.py +752 -0
  19. feectools/ddm/__init__.py +3 -0
  20. feectools/ddm/basic.py +78 -0
  21. feectools/ddm/blocking_data_exchanger.py +348 -0
  22. feectools/ddm/cart.py +1835 -0
  23. feectools/ddm/interface_data_exchanger.py +122 -0
  24. feectools/ddm/mpi.py +109 -0
  25. feectools/ddm/nonblocking_data_exchanger.py +331 -0
  26. feectools/ddm/partition.py +207 -0
  27. feectools/ddm/petsc.py +112 -0
  28. feectools/ddm/tests/__init__.py +0 -0
  29. feectools/ddm/tests/test_cart_1d.py +138 -0
  30. feectools/ddm/tests/test_cart_2d.py +164 -0
  31. feectools/ddm/tests/test_cart_3d.py +158 -0
  32. feectools/ddm/tests/test_multicart_2d.py +173 -0
  33. feectools/ddm/tests/test_partition.py +124 -0
  34. feectools/ddm/utilities.py +24 -0
  35. feectools/feec/__init__.py +0 -0
  36. feectools/feec/derivatives.py +780 -0
  37. feectools/feec/dof_kernels.py +210 -0
  38. feectools/feec/global_geometric_projectors.py +1073 -0
  39. feectools/feec/hodge.py +148 -0
  40. feectools/fem/__init__.py +0 -0
  41. feectools/fem/basic.py +465 -0
  42. feectools/fem/grid.py +181 -0
  43. feectools/fem/partitioning.py +344 -0
  44. feectools/fem/projectors.py +160 -0
  45. feectools/fem/splines.py +559 -0
  46. feectools/fem/tensor.py +1393 -0
  47. feectools/fem/tests/__init__.py +0 -0
  48. feectools/fem/tests/analytical_profiles_1d.py +100 -0
  49. feectools/fem/tests/analytical_profiles_base.py +34 -0
  50. feectools/fem/tests/splines_error_bounds.py +155 -0
  51. feectools/fem/tests/test_spline_histopolation.py +120 -0
  52. feectools/fem/tests/test_spline_interpolation.py +182 -0
  53. feectools/fem/tests/test_splines.py +184 -0
  54. feectools/fem/tests/test_splines_par.py +46 -0
  55. feectools/fem/tests/test_vector_spaces.py +150 -0
  56. feectools/fem/tests/utilities.py +47 -0
  57. feectools/fem/vector.py +729 -0
  58. feectools/linalg/__init__.py +0 -0
  59. feectools/linalg/basic.py +1386 -0
  60. feectools/linalg/block.py +1451 -0
  61. feectools/linalg/direct_solvers.py +201 -0
  62. feectools/linalg/fft.py +258 -0
  63. feectools/linalg/kernels/__init__.py +0 -0
  64. feectools/linalg/kernels/axpy_kernels.py +57 -0
  65. feectools/linalg/kernels/inner_kernels.py +100 -0
  66. feectools/linalg/kernels/matvec_kernels.py +206 -0
  67. feectools/linalg/kernels/stencil2IJV_kernels.py +227 -0
  68. feectools/linalg/kernels/stencil2coo_kernels.py +179 -0
  69. feectools/linalg/kernels/transpose_kernels.py +263 -0
  70. feectools/linalg/kron.py +911 -0
  71. feectools/linalg/solvers.py +1914 -0
  72. feectools/linalg/sparse.py +114 -0
  73. feectools/linalg/stencil.py +2923 -0
  74. feectools/linalg/stencil_dot_kernels.py +317 -0
  75. feectools/linalg/stencil_transpose_kernels.py +372 -0
  76. feectools/linalg/tests/__init__.py +0 -0
  77. feectools/linalg/tests/test_block.py +1588 -0
  78. feectools/linalg/tests/test_fft.py +106 -0
  79. feectools/linalg/tests/test_kron_stencil_matrix.py +114 -0
  80. feectools/linalg/tests/test_linalg.py +1065 -0
  81. feectools/linalg/tests/test_matrix_free.py +128 -0
  82. feectools/linalg/tests/test_solvers.py +213 -0
  83. feectools/linalg/tests/test_stencil_interface_matrix.py +379 -0
  84. feectools/linalg/tests/test_stencil_vector.py +1036 -0
  85. feectools/linalg/tests/test_stencil_vector_space.py +440 -0
  86. feectools/linalg/topetsc.py +522 -0
  87. feectools/linalg/utilities.py +200 -0
  88. feectools/utilities/__init__.py +0 -0
  89. feectools/utilities/quadratures.py +113 -0
  90. feectools/utilities/utils.py +166 -0
  91. feectools/version.py +1 -0
  92. feectools-0.1.0.dist-info/METADATA +66 -0
  93. feectools-0.1.0.dist-info/RECORD +98 -0
  94. feectools-0.1.0.dist-info/WHEEL +5 -0
  95. feectools-0.1.0.dist-info/entry_points.txt +3 -0
  96. feectools-0.1.0.dist-info/licenses/AUTHORS +22 -0
  97. feectools-0.1.0.dist-info/licenses/LICENSE +21 -0
  98. feectools-0.1.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,201 @@
1
+ # coding: utf-8
2
+ # Copyright 2018 Jalal Lakhlili, Yaman Güçlü
3
+
4
+ from abc import abstractmethod
5
+ import numpy as np
6
+ from scipy.linalg.lapack import dgbtrf, dgbtrs, sgbtrf, sgbtrs, cgbtrf, cgbtrs, zgbtrf, zgbtrs
7
+ from scipy.sparse import spmatrix
8
+ from scipy.sparse.linalg import splu
9
+
10
+ from feectools.linalg.basic import LinearSolver
11
+
12
+ __all__ = ('BandedSolver', 'SparseSolver')
13
+
14
+ #===============================================================================
15
+ class BandedSolver(LinearSolver):
16
+ """
17
+ Solve the equation Ax = b for x, assuming A is banded matrix.
18
+
19
+ Parameters
20
+ ----------
21
+ u : integer
22
+ Number of non-zero upper diagonal.
23
+
24
+ l : integer
25
+ Number of non-zero lower diagonal.
26
+
27
+ bmat : nd-array
28
+ Banded matrix.
29
+
30
+ """
31
+ def __init__(self, u, l, bmat, transposed=False):
32
+
33
+ self._u = u
34
+ self._l = l
35
+ self._transposed = transposed
36
+
37
+ # ... LU factorization
38
+ if bmat.dtype == np.float32:
39
+ self._factor_function = sgbtrf
40
+ self._solver_function = sgbtrs
41
+ elif bmat.dtype == np.float64:
42
+ self._factor_function = dgbtrf
43
+ self._solver_function = dgbtrs
44
+ elif bmat.dtype == np.complex64:
45
+ self._factor_function = cgbtrf
46
+ self._solver_function = cgbtrs
47
+ elif bmat.dtype == np.complex128:
48
+ self._factor_function = zgbtrf
49
+ self._solver_function = zgbtrs
50
+ else:
51
+ msg = f'Cannot create a BandedSolver for bmat.dtype = {bmat.dtype}'
52
+ raise NotImplementedError(msg)
53
+
54
+ self._bmat, self._ipiv, self._finfo = self._factor_function(bmat, l, u)
55
+
56
+ self._sinfo = None
57
+
58
+ self._space = np.ndarray
59
+ self._dtype = bmat.dtype
60
+
61
+ @property
62
+ def finfo(self):
63
+ return self._finfo
64
+
65
+ @property
66
+ def sinfo(self):
67
+ return self._sinfo
68
+
69
+ #--------------------------------------
70
+ # Abstract interface
71
+ #--------------------------------------
72
+ @property
73
+ def space(self):
74
+ return self._space
75
+
76
+ def transpose(self):
77
+ cls = type(self)
78
+ obj = super().__new__(cls)
79
+
80
+ obj._u = self._l
81
+ obj._l = self._u
82
+ obj._bmat = self._bmat
83
+ obj._ipiv = self._ipiv
84
+ obj._finfo = self._finfo
85
+ obj._factor_function = self._factor_function
86
+ obj._solver_function = self._solver_function
87
+ obj._sinfo = None
88
+ obj._space = self._space
89
+ obj._dtype = self._dtype
90
+ obj._transposed = not self._transposed
91
+
92
+ return obj
93
+
94
+ #...
95
+ def solve(self, rhs, out=None):
96
+ """
97
+ Solves for the given right-hand side.
98
+
99
+ Parameters
100
+ ----------
101
+ rhs : ndarray
102
+ The right-hand sides to solve for. The vectors are assumed to be given in C-contiguous order,
103
+ i.e. if multiple right-hand sides are given, then rhs is a two-dimensional array with the 0-th
104
+ index denoting the number of the right-hand side, and the 1-st index denoting the element inside
105
+ a right-hand side.
106
+
107
+ out : ndarray | NoneType
108
+ Output vector. If given, it has to have the same shape and datatype as rhs.
109
+ """
110
+ assert rhs.T.shape[0] == self._bmat.shape[1]
111
+
112
+ transposed = self._transposed
113
+
114
+ if out is None:
115
+ preout, self._sinfo = self._solver_function(self._bmat, self._l, self._u, rhs.T, self._ipiv,
116
+ trans=transposed)
117
+ out = preout.T
118
+
119
+ else:
120
+ assert out.shape == rhs.shape
121
+ assert out.dtype == rhs.dtype
122
+
123
+ # support in-place operations
124
+ if rhs is not out:
125
+ out[:] = rhs
126
+
127
+ # TODO: handle non-contiguous views?
128
+
129
+ # we want FORTRAN-contiguous data (default is assumed to be C contiguous)
130
+ _, self._sinfo = self._solver_function(self._bmat, self._l, self._u, out.T, self._ipiv, overwrite_b=True,
131
+ trans=transposed)
132
+
133
+ return out
134
+
135
+ #===============================================================================
136
+ class SparseSolver (LinearSolver):
137
+ """
138
+ Solve the equation Ax = b for x, assuming A is scipy sparse matrix.
139
+
140
+ Parameters
141
+ ----------
142
+ spmat : scipy.sparse.spmatrix
143
+ Generic sparse matrix.
144
+
145
+ """
146
+ def __init__(self, spmat, transposed=False):
147
+
148
+ assert isinstance(spmat, spmatrix)
149
+
150
+ self._space = np.ndarray
151
+ self._splu = splu(spmat.tocsc())
152
+ self._transposed = transposed
153
+
154
+ #--------------------------------------
155
+ # Abstract interface
156
+ #--------------------------------------
157
+ @property
158
+ def space(self):
159
+ return self._space
160
+
161
+ def transpose(self):
162
+ cls = type(self)
163
+ obj = super().__new__(cls)
164
+
165
+ obj._space = self._space
166
+ obj._splu = self._splu
167
+ obj._transposed = not self._transposed
168
+
169
+ return obj
170
+
171
+ #...
172
+ def solve(self, rhs, out=None):
173
+ """
174
+ Solves for the given right-hand side.
175
+
176
+ Parameters
177
+ ----------
178
+ rhs : ndarray
179
+ The right-hand sides to solve for. The vectors are assumed to be given in C-contiguous order,
180
+ i.e. if multiple right-hand sides are given, then rhs is a two-dimensional array with the 0-th
181
+ index denoting the number of the right-hand side, and the 1-st index denoting the element inside
182
+ a right-hand side.
183
+
184
+ out : ndarray | NoneType
185
+ Output vector. If given, it has to have the same shape and datatype as rhs.
186
+ """
187
+
188
+ assert rhs.T.shape[0] == self._splu.shape[1]
189
+ transposed = self._transposed
190
+
191
+ if out is None:
192
+ out = self._splu.solve(rhs.T, trans='T' if transposed else 'N').T
193
+
194
+ else:
195
+ assert out.shape == rhs.shape
196
+ assert out.dtype == rhs.dtype
197
+
198
+ # currently no in-place solve exposed
199
+ out[:] = self._splu.solve(rhs.T, trans='T' if transposed else 'N').T
200
+
201
+ return out
@@ -0,0 +1,258 @@
1
+ from feectools.linalg.basic import LinearOperator, LinearSolver
2
+ from feectools.linalg.stencil import StencilVectorSpace
3
+ from feectools.linalg.kron import KroneckerLinearSolver
4
+
5
+ import numpy as np
6
+ import scipy.fft as scifft
7
+ import os
8
+
9
+ class DistributedFFTBase(LinearOperator):
10
+ """
11
+ A base class for the distributed FFT, DCT and DST.
12
+ Internally calls a KroneckerLinearSolver on a solver which just applies the FFT or some other function.
13
+
14
+ Parameters
15
+ ----------
16
+ space : StencilVectorSpace
17
+ The vector space needed for the KroneckerLinearSolver internally.
18
+
19
+ function : callable | list/tuple of callables
20
+ A list/tuple of callables function, each with one parameter x which applies some function in-place on x.
21
+ The function at position i is applied to the i-th tensor direction.
22
+ If only a single callable is given, it is used for all directions.
23
+ """
24
+ def toarray(self):
25
+ raise NotImplementedError('toarray() is not defined for DistributedFFTBase.')
26
+
27
+ def tosparse(self):
28
+ raise NotImplementedError('tosparse() is not defined for DistributedFFTBase.')
29
+
30
+ # Possible additions for the future:
31
+ # * split off the LinearSolver class when used with the space ndarray (as used in the KroneckerLinearSolver),
32
+ # and make it state if it works in-place (or if it needs temporary memory), and what its optimal
33
+ # size is (FFT might work faster with padding)
34
+ # * include FFTW support (e.g. pyfftw)
35
+
36
+ class OneDimSolver(LinearSolver):
37
+ """
38
+ A one-dimensional solver which just applies a given function.
39
+
40
+ Parameters
41
+ ----------
42
+ function : Callable
43
+ The given function.
44
+ """
45
+ def __init__(self, function):
46
+ self._function = function
47
+
48
+ @property
49
+ def space(self):
50
+ return np.ndarray
51
+
52
+ def transpose(self):
53
+ raise NotImplementedError('transpose() is not implemented for OneDimSolvers')
54
+
55
+ def solve(self, rhs, out=None):
56
+ if out is None:
57
+ out = np.empty_like(rhs)
58
+
59
+ if out is not rhs:
60
+ out[:] = rhs
61
+
62
+ self._function(out)
63
+
64
+ return out
65
+
66
+ def __init__(self, space, functions):
67
+ assert isinstance(space, StencilVectorSpace)
68
+ if isinstance(functions, list) or isinstance(functions, tuple):
69
+ solvers = [DistributedFFTBase.OneDimSolver(function) for function in functions]
70
+ else:
71
+ onedimsolver = DistributedFFTBase.OneDimSolver(functions)
72
+ solvers = [onedimsolver] * space.ndim
73
+ self._isolver = KroneckerLinearSolver(space, space, solvers)
74
+
75
+ # ...
76
+ @property
77
+ def domain(self):
78
+ return self._isolver.space
79
+
80
+ # ...
81
+ @property
82
+ def codomain(self):
83
+ return self._isolver.space
84
+
85
+ # ...
86
+ @property
87
+ def dtype( self ):
88
+ return self._isolver.dtype
89
+
90
+ # ...
91
+ def dot(self, v, out=None):
92
+ # just call the KroneckerLinearSolver
93
+ return self._isolver.solve(v, out=out)
94
+
95
+ def transpose(self, conjugate=False):
96
+ raise NotImplementedError()
97
+
98
+ # IMPORTANT NOTE: All of these scifft.fft functions currently trust that overwrite_x=True will yield an in-place fft...
99
+ # (this is not completely given to hold forever, so in case these tests fail in some future version, change this)
100
+
101
+ class DistributedFFT(DistributedFFTBase):
102
+ """
103
+ Equals an n-dimensional FFT operation, except that it works on a distributed/parallel StencilVector.
104
+
105
+ Internally calls scipy.fft.fft for each direction.
106
+
107
+ Parameters
108
+ ----------
109
+ space : StencilVectorSpace
110
+ The space the n-dimensional FFT should be run on. Must have a complex data type (i.e. space.dtype.kind == 'c').
111
+
112
+ norm : str
113
+ Specifies the normalization factor. See the documentation of the corresponding scipy.fft.fft parameter.
114
+
115
+ workers : Union[int, NoneType]
116
+ Specifies the number of worker threads. By default set to the number of OpenMP threads, if given.
117
+ See also the documentation of the corresponding scipy.fft.fft parameter.
118
+ """
119
+ def __init__(self, space, norm=None, workers=os.environ.get('OMP_NUM_THREADS', None)):
120
+ # only allow complex data types
121
+ assert isinstance(space, StencilVectorSpace)
122
+ assert np.dtype(space.dtype).kind == 'c'
123
+ workers = int(workers) if workers is not None else None
124
+
125
+ super().__init__(space, lambda out: scifft.fft(
126
+ out, axis=1, overwrite_x=True, workers=workers, norm=norm))
127
+
128
+
129
+ class DistributedIFFT(DistributedFFTBase):
130
+ """
131
+ Equals an n-dimensional IFFT operation, except that it works on a distributed/parallel StencilVector.
132
+
133
+ Internally calls scipy.fft.ifft for each direction.
134
+
135
+ Parameters
136
+ ----------
137
+ space : StencilVectorSpace
138
+ The space the n-dimensional IFFT should be run on. Must have a complex data type (i.e. space.dtype.kind == 'c').
139
+
140
+ norm : str
141
+ Specifies the normalization factor. See the documentation of the corresponding scipy.fft.ifft parameter.
142
+
143
+ workers : Union[int, NoneType]
144
+ Specifies the number of worker threads. By default set to the number of OpenMP threads, if given.
145
+ See also the documentation of the corresponding scipy.fft.ifft parameter.
146
+ """
147
+ def __init__(self, space, norm=None, workers=os.environ.get('OMP_NUM_THREADS', None)):
148
+ # only allow complex data types
149
+ assert isinstance(space, StencilVectorSpace)
150
+ assert np.dtype(space.dtype).kind == 'c'
151
+ workers = int(workers) if workers is not None else None
152
+
153
+ super().__init__(space, lambda out: scifft.ifft(
154
+ out, axis=1, overwrite_x=True, workers=workers, norm=norm))
155
+
156
+ class DistributedDCT(DistributedFFTBase):
157
+ """
158
+ Equals an n-dimensional DCT operation, except that it works on a distributed/parallel StencilVector.
159
+
160
+ Internally calls scipy.fft.dct for each direction.
161
+
162
+ Parameters
163
+ ----------
164
+ space : StencilVectorSpace
165
+ The space the n-dimensional DCT should be run on.
166
+
167
+ norm : str
168
+ Specifies the normalization factor. See the documentation of the corresponding scipy.fft.dct parameter.
169
+
170
+ workers : Union[int, NoneType]
171
+ Specifies the number of worker threads. By default set to the number of OpenMP threads, if given.
172
+ See also the documentation of the corresponding scipy.fft.dct parameter.
173
+
174
+ ttype : int
175
+ The DCT type to use. (the name of this parameter in the underlying method is actually `type`).
176
+ """
177
+ def __init__(self, space, norm=None, workers=os.environ.get('OMP_NUM_THREADS', None), ttype=2):
178
+ workers = int(workers) if workers is not None else None
179
+ super().__init__(space, lambda out: scifft.dct(
180
+ out, axis=1, overwrite_x=True, workers=workers, norm=norm, type=ttype))
181
+
182
+ class DistributedIDCT(DistributedFFTBase):
183
+ """
184
+ Equals an n-dimensional IDCT operation, except that it works on a distributed/parallel StencilVector.
185
+
186
+ Internally calls scipy.fft.idct for each direction.
187
+
188
+ Parameters
189
+ ----------
190
+ space : StencilVectorSpace
191
+ The space the n-dimensional IDCT should be run on.
192
+
193
+ norm : str
194
+ Specifies the normalization factor. See the documentation of the corresponding scipy.fft.idct parameter.
195
+
196
+ workers : Union[int, NoneType]
197
+ Specifies the number of worker threads. By default set to the number of OpenMP threads, if given.
198
+ See also the documentation of the corresponding scipy.fft.idct parameter.
199
+
200
+ ttype : int
201
+ The DCT type to use. (the name of this parameter in the underlying method is actually `type`).
202
+ """
203
+ def __init__(self, space, norm=None, workers=os.environ.get('OMP_NUM_THREADS', None), ttype=2):
204
+ workers = int(workers) if workers is not None else None
205
+ super().__init__(space, lambda out: scifft.idct(
206
+ out, axis=1, overwrite_x=True, workers=workers, norm=norm, type=ttype))
207
+
208
+ class DistributedDST(DistributedFFTBase):
209
+ """
210
+ Equals an n-dimensional DST operation, except that it works on a distributed/parallel StencilVector.
211
+
212
+ Internally calls scipy.fft.dst for each direction.
213
+
214
+ Parameters
215
+ ----------
216
+ space : StencilVectorSpace
217
+ The space the n-dimensional DST should be run on.
218
+
219
+ norm : str
220
+ Specifies the normalization factor. See the documentation of the corresponding scipy.fft.dst parameter.
221
+
222
+ workers : Union[int, NoneType]
223
+ Specifies the number of worker threads. By default set to the number of OpenMP threads, if given.
224
+ See also the documentation of the corresponding scipy.fft.dst parameter.
225
+
226
+ ttype : int
227
+ The DCT type to use. (the name of this parameter in the underlying method is actually `type`).
228
+ """
229
+ def __init__(self, space, norm=None, workers=os.environ.get('OMP_NUM_THREADS', None), ttype=2):
230
+ workers = int(workers) if workers is not None else None
231
+ super().__init__(space, lambda out: scifft.dst(
232
+ out, axis=1, overwrite_x=True, workers=workers, norm=norm, type=ttype))
233
+
234
+ class DistributedIDST(DistributedFFTBase):
235
+ """
236
+ Equals an n-dimensional IDST operation, except that it works on a distributed/parallel StencilVector.
237
+
238
+ Internally calls scipy.fft.idst for each direction.
239
+
240
+ Parameters
241
+ ----------
242
+ space : StencilVectorSpace
243
+ The space the n-dimensional IDST should be run on.
244
+
245
+ norm : str
246
+ Specifies the normalization factor. See the documentation of the corresponding scipy.fft.idst parameter.
247
+
248
+ workers : Union[int, NoneType]
249
+ Specifies the number of worker threads. By default set to the number of OpenMP threads, if given.
250
+ See also the documentation of the corresponding scipy.fft.idst parameter.
251
+
252
+ ttype : int
253
+ The DCT type to use. (the name of this parameter in the underlying method is actually `type`).
254
+ """
255
+ def __init__(self, space, norm=None, workers=os.environ.get('OMP_NUM_THREADS', None), ttype=2):
256
+ workers = int(workers) if workers is not None else None
257
+ super().__init__(space, lambda out: scifft.idst(
258
+ out, axis=1, overwrite_x=True, workers=workers, norm=norm, type=ttype))
File without changes
@@ -0,0 +1,57 @@
1
+ from typing import TypeVar
2
+
3
+ T = TypeVar('T', float, complex)
4
+
5
+ #========================================================================================================
6
+ def axpy_1d(alpha: 'T', x: 'T[:]', y: 'T[:]'):
7
+ """
8
+ Kernel for computing y = alpha * x + y.
9
+
10
+ Parameters
11
+ ----------
12
+ alpha : float | complex
13
+ Scaling coefficient.
14
+
15
+ x, y : 1D Numpy arrays of (float | complex) data
16
+ Data of the vectors.
17
+ """
18
+ n1, = x.shape
19
+ for i1 in range(n1):
20
+ y[i1] += alpha * x[i1]
21
+
22
+ #========================================================================================================
23
+ def axpy_2d(alpha: 'T', x: 'T[:,:]', y: 'T[:,:]'):
24
+ """
25
+ Kernel for computing y = alpha * x + y.
26
+
27
+ Parameters
28
+ ----------
29
+ alpha : float | complex
30
+ Scaling coefficient.
31
+
32
+ x, y : 2D Numpy arrays of (float | complex) data
33
+ Data of the vectors.
34
+ """
35
+ n1, n2 = x.shape
36
+ for i1 in range(n1):
37
+ for i2 in range(n2):
38
+ y[i1, i2] += alpha * x[i1, i2]
39
+
40
+ #========================================================================================================
41
+ def axpy_3d(alpha: 'T', x: 'T[:,:,:]', y: 'T[:,:,:]'):
42
+ """
43
+ Kernel for computing y = alpha * x + y.
44
+
45
+ Parameters
46
+ ----------
47
+ alpha : float | complex
48
+ Scaling coefficient.
49
+
50
+ x, y : 3D Numpy arrays of (float | complex) data
51
+ Data of the vectors.
52
+ """
53
+ n1, n2, n3 = x.shape
54
+ for i1 in range(n1):
55
+ for i2 in range(n2):
56
+ for i3 in range(n3):
57
+ y[i1, i2, i3] += alpha * x[i1, i2, i3]
@@ -0,0 +1,100 @@
1
+
2
+ #!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!#
3
+ #!!!!!!!!!!!!!!!!!!! WARNING !!!!!!!!!!!!!!!!!!!#
4
+ #!!!!!!! Conjugate on the first argument !!!!!!!#
5
+ #!!!!!!!!!! This will need an update !!!!!!!!!!!#
6
+ #!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!#
7
+
8
+ from typing import TypeVar
9
+
10
+ T = TypeVar('T', float, complex)
11
+
12
+ #==============================================================================
13
+ def inner_1d(v1: 'T[:]', v2: 'T[:]', nghost0: 'int64'):
14
+ """
15
+ Kernel for computing the inner product (case of two 1D vectors).
16
+
17
+ Parameters
18
+ ----------
19
+ v1, v2 : 1D NumPy array
20
+ Data of the vectors from which we are computing the inner product.
21
+
22
+ nghost0 : int
23
+ Number of ghost cells of the arrays along the index 0.
24
+
25
+ Returns
26
+ -------
27
+ res : scalar
28
+ Scalar (real or complex) containing the result of the inner product.
29
+ """
30
+ shape0, = v1.shape
31
+
32
+ res = v1[0] - v1[0]
33
+ for i0 in range(nghost0, shape0 - nghost0):
34
+ res += v1[i0].conjugate() * v2[i0]
35
+
36
+ return res
37
+
38
+ #==============================================================================
39
+ def inner_2d(v1: 'T[:,:]', v2: 'T[:,:]', nghost0: 'int64', nghost1: 'int64'):
40
+ """
41
+ Kernel for computing the inner product (case of two 2D vectors).
42
+
43
+ Parameters
44
+ ----------
45
+ v1, v2 : 2D NumPy array
46
+ Data of the vectors from which we are computing the inner product.
47
+
48
+ nghost0 : int
49
+ Number of ghost cells of the arrays along the index 0.
50
+
51
+ nghost1 : int
52
+ Number of ghost cells of the arrays along the index 1.
53
+
54
+ Returns
55
+ -------
56
+ res : scalar
57
+ Scalar (real or complex) containing the result of the inner product.
58
+ """
59
+ shape0, shape1 = v1.shape
60
+
61
+ res = v1[0, 0] - v1[0, 0]
62
+ for i0 in range(nghost0, shape0 - nghost0):
63
+ for i1 in range(nghost1, shape1 - nghost1):
64
+ res += v1[i0, i1].conjugate() * v2[i0, i1]
65
+
66
+ return res
67
+
68
+ #==============================================================================
69
+ def inner_3d(v1: 'T[:,:,:]', v2: 'T[:,:,:]', nghost0: 'int64', nghost1: 'int64', nghost2: 'int64'):
70
+ """
71
+ Kernel for computing the inner product (case of two 3D vectors).
72
+
73
+ Parameters
74
+ ----------
75
+ v1, v2 : 3D NumPy array
76
+ Data of the vectors from which we are computing the inner product.
77
+
78
+ nghost0 : int
79
+ Number of ghost cells of the arrays along the index 0.
80
+
81
+ nghost1 : int
82
+ Number of ghost cells of the arrays along the index 1.
83
+
84
+ nghost2 : int
85
+ Number of ghost cells of the arrays along the index 2.
86
+
87
+ Returns
88
+ -------
89
+ res : scalar
90
+ Scalar (real or complex) containing the result of the inner product.
91
+ """
92
+ shape0, shape1, shape2 = v1.shape
93
+
94
+ res = v1[0, 0, 0] - v1[0, 0, 0]
95
+ for i0 in range(nghost0, shape0 - nghost0):
96
+ for i1 in range(nghost1, shape1 - nghost1):
97
+ for i2 in range(nghost2, shape2 - nghost2):
98
+ res += v1[i0, i1, i2].conjugate() * v2[i0, i1, i2]
99
+
100
+ return res