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,911 @@
1
+ #coding = utf-8
2
+ from functools import reduce
3
+
4
+ import numpy as np
5
+ from scipy.sparse import kron
6
+ from scipy.sparse import coo_matrix
7
+
8
+ from feectools.linalg.basic import LinearOperator, LinearSolver
9
+ from feectools.linalg.stencil import StencilVectorSpace, StencilVector, StencilMatrix
10
+
11
+ __all__ = ('KroneckerStencilMatrix',
12
+ 'KroneckerLinearSolver',
13
+ 'KroneckerDenseMatrix',
14
+ 'kronecker_solve')
15
+
16
+ #==============================================================================
17
+ class KroneckerStencilMatrix(LinearOperator):
18
+ """
19
+ Kronecker product of 1D stencil matrices.
20
+
21
+ Parameters
22
+ ----------
23
+ V : StencilVectorSpace
24
+ The domain.
25
+
26
+ W : StencilVectorSpace
27
+ The codomain.
28
+
29
+ args : list of StencilMatrix
30
+ Factors of the Kronecker product (one for each dimension).
31
+
32
+ """
33
+
34
+ def __init__(self, V, W, *args):
35
+
36
+ assert isinstance(V, StencilVectorSpace)
37
+ assert isinstance(W, StencilVectorSpace)
38
+
39
+ for i,A in enumerate(args):
40
+ assert isinstance(A, LinearOperator)
41
+ assert A.domain.ndim == 1
42
+ assert A.domain.npts[0] == V.npts[i]
43
+
44
+ self._domain = V
45
+ self._codomain = W
46
+ self._mats = args
47
+ self._ndim = len(args)
48
+
49
+ #--------------------------------------
50
+ # Abstract interface
51
+ #--------------------------------------
52
+ @property
53
+ def domain( self ):
54
+ return self._domain
55
+
56
+ # ...
57
+ @property
58
+ def codomain( self ):
59
+ return self._codomain
60
+
61
+ # ...
62
+ @property
63
+ def dtype( self ):
64
+ return self.domain.dtype
65
+
66
+ # ...
67
+ @property
68
+ def ndim( self ):
69
+ return self._ndim
70
+
71
+ # ...
72
+ @property
73
+ def mats( self ):
74
+ return self._mats
75
+
76
+ # ...
77
+ def dot(self, x, out=None):
78
+
79
+ dot = np.dot
80
+
81
+ assert isinstance(x, StencilVector)
82
+ assert x.space is self.domain
83
+
84
+ # Necessary if vector space is periodic or distributed across processes
85
+ if not x.ghost_regions_in_sync:
86
+ x.update_ghost_regions()
87
+
88
+ if out is not None:
89
+ assert isinstance(out, StencilVector)
90
+ assert out.space is self.codomain
91
+ else:
92
+ out = StencilVector(self.codomain)
93
+
94
+ starts = self._codomain.starts
95
+ ends = self._codomain.ends
96
+ pads = self._codomain.pads
97
+ shifts = self._codomain.shifts
98
+
99
+ mats = self.mats
100
+ nrows = tuple(e-s+1 for s,e in zip(starts, ends))
101
+ pnrows = tuple(2*p+1 for p in pads)
102
+
103
+ for ii in np.ndindex(*nrows):
104
+ v = 0.
105
+ xx = tuple(i+p*s for i,p,s in zip(ii, pads, shifts))
106
+
107
+ for jj in np.ndindex(*pnrows):
108
+ i_mats = [mat._data[s, j] for s,j,mat in zip(xx, jj, mats)]
109
+ ii_jj = tuple(i+j+(s-1)*p for i,j,p,s in zip(ii, jj, pads, shifts))
110
+ v += x._data[ii_jj] * np.prod(i_mats)
111
+
112
+ out._data[xx] = v
113
+
114
+ # IMPORTANT: flag that ghost regions are not up-to-date
115
+ out.ghost_regions_in_sync = False
116
+ return out
117
+
118
+ # ...
119
+ def copy(self):
120
+ mats = [m.copy() for m in self.mats]
121
+ return KroneckerStencilMatrix(self.domain, self.codomain, *mats)
122
+
123
+ # ...
124
+ def __neg__(self):
125
+ mats = [-self.mats[0], *(m.copy() for m in self.mats[1:])]
126
+ return KroneckerStencilMatrix(self.domain, self.codomain, *mats)
127
+
128
+ # ...
129
+ def __mul__(self, a):
130
+ mats = [*(m.copy() for m in self.mats[:-1]), self.mats[-1] * a]
131
+ return KroneckerStencilMatrix(self.domain, self.codomain, *mats)
132
+
133
+ # ...
134
+ def __imul__(self, a):
135
+ self.mats[-1] *= a
136
+ return self
137
+
138
+ #--------------------------------------
139
+ # Other properties/methods
140
+ #--------------------------------------
141
+
142
+ def __getitem__(self, key):
143
+ pads = self._codomain.pads
144
+ rows = key[:self.ndim]
145
+ cols = key[self.ndim:]
146
+ mats = self.mats
147
+ elements = [A[i,j] for A,i,j in zip(mats, rows, cols)]
148
+ return np.prod(elements)
149
+
150
+ def tostencil(self):
151
+
152
+ mats = self.mats
153
+ ssc = self.codomain.starts
154
+ eec = self.codomain.ends
155
+ ssd = self.domain.starts
156
+ eed = self.domain.ends
157
+ pads = [A.pads[0] for A in self.mats]
158
+ xpads = self.domain.pads
159
+
160
+ # Number of rows in matrix (along each dimension)
161
+ nrows = [ed-s+1 for s,ed in zip(ssd, eed)]
162
+ nrows_extra = [0 if ec<=ed else ec-ed for ec,ed in zip(eec,eed)]
163
+
164
+ # create the stencil matrix
165
+ M = StencilMatrix(self.domain, self.codomain, pads=tuple(pads))
166
+
167
+ mats = [mat._data for mat in mats]
168
+
169
+ self._tostencil(M._data, mats, nrows, nrows_extra, pads, xpads)
170
+ return M
171
+
172
+ @staticmethod
173
+ def _tostencil(M, mats, nrows, nrows_extra, pads, xpads):
174
+
175
+ ndiags = [2*p + 1 for p in pads]
176
+ diff = [xp-p for xp,p in zip(xpads, pads)]
177
+ ndim = len(nrows)
178
+
179
+ for xx in np.ndindex( *nrows ):
180
+
181
+ ii = tuple(xp + x for xp, x in zip(xpads, xx) )
182
+
183
+ for kk in np.ndindex( *ndiags ):
184
+
185
+ values = [mat[i,k] for mat,i,k in zip(mats, ii, kk)]
186
+ M[(*ii, *kk)] = np.prod(values)
187
+
188
+ # handle partly-multiplied rows
189
+ new_nrows = nrows.copy()
190
+ for d,er in enumerate(nrows_extra):
191
+
192
+ rows = new_nrows.copy()
193
+ del rows[d]
194
+
195
+ for n in range(er):
196
+ for xx in np.ndindex(*rows):
197
+ xx = list(xx)
198
+ xx.insert(d, nrows[d]+n)
199
+
200
+ ii = tuple(x+xp for x,xp in zip(xx, xpads))
201
+ ee = [max(x-l+1,0) for x,l in zip(xx, nrows)]
202
+ jj = tuple( slice(x+d, x+d+2*p+1-e) for x,p,d,e in zip(xx, pads, diff, ee) )
203
+ ndiags = [2*p + 1-e for p,e in zip(pads,ee)]
204
+ kk = [slice(None,diag) for diag in ndiags]
205
+ ii_kk = tuple( list(ii) + kk )
206
+
207
+ for kk in np.ndindex( *ndiags ):
208
+ values = [mat[i,k] for mat,i,k in zip(mats, ii, kk)]
209
+ M[(*ii, *kk)] = np.prod(values)
210
+ new_nrows[d] += er
211
+
212
+ def tosparse(self):
213
+ return reduce(kron, (m.tosparse() for m in self.mats))
214
+
215
+ def toarray(self):
216
+ return self.tosparse().toarray()
217
+
218
+ def transpose(self, conjugate=False):
219
+ mats_tr = [Mi.transpose(conjugate=conjugate) for Mi in self.mats]
220
+ return KroneckerStencilMatrix(self.codomain, self.domain, *mats_tr)
221
+
222
+ #==============================================================================
223
+ class KroneckerDenseMatrix(LinearOperator):
224
+ """
225
+ Kronecker product of 1D dense matrices.
226
+
227
+ Parameters
228
+ ----------
229
+ V : StencilVectorSpace
230
+ The domain.
231
+
232
+ W : StencilVectorSpace
233
+ The codomain.
234
+
235
+ args : list of ndarray
236
+ Factors of the Kronecker product (one for each dimension).
237
+
238
+ """
239
+
240
+ def __init__(self, V, W, *args , with_pads=False):
241
+
242
+ assert isinstance(V, StencilVectorSpace)
243
+ assert isinstance(W, StencilVectorSpace)
244
+ assert V.pads == W.pads
245
+
246
+ for i,A in enumerate(args):
247
+ assert isinstance(A, np.ndarray)
248
+ if with_pads:
249
+ assert A.shape[1] == V.npts[i] + 2*V.pads[i]
250
+ else:
251
+ assert A.shape[1] == V.npts[i]
252
+
253
+ if not with_pads:
254
+ args = [np.pad(a,p) for a,p in zip(args, W.pads)]
255
+
256
+ self._domain = V
257
+ self._codomain = W
258
+ self._mats = list(args)
259
+ self._ndim = len(args)
260
+
261
+ #--------------------------------------
262
+ # Abstract interface
263
+ #--------------------------------------
264
+ @property
265
+ def domain(self):
266
+ return self._domain
267
+
268
+ # ...
269
+ @property
270
+ def codomain(self):
271
+ return self._codomain
272
+
273
+ # ...
274
+ @property
275
+ def dtype(self):
276
+ return self.domain.dtype
277
+
278
+ # ...
279
+ @property
280
+ def ndim(self):
281
+ return self._ndim
282
+
283
+ # ...
284
+ @property
285
+ def mats(self):
286
+ return self._mats
287
+
288
+ # ...
289
+ def dot(self, x, out=None):
290
+
291
+ dot = np.dot
292
+
293
+ assert isinstance(x, StencilVector)
294
+ assert x.space is self.domain
295
+
296
+ # Necessary if vector space is periodic or distributed across processes
297
+ if not x.ghost_regions_in_sync:
298
+ x.update_ghost_regions()
299
+
300
+ if out is not None:
301
+ assert isinstance(out, StencilVector)
302
+ assert out.space is self.codomain
303
+ else:
304
+ out = StencilVector(self.codomain)
305
+
306
+ d_starts = self._domain.starts
307
+ d_ends = self._domain.ends
308
+ c_starts = self._codomain.starts
309
+ c_ends = self._codomain.ends
310
+ pads = self._codomain.pads
311
+ mats = self.mats
312
+
313
+ nrows = tuple(e-s+1 for s,e in zip(c_starts, c_ends))
314
+ ncols = tuple(e-s+1+2*p for s,e,p in zip(d_starts, d_ends, pads))
315
+ kk = tuple(slice(s, s+nc) for nc,s in zip(ncols, d_starts))
316
+ x_data = x._data.ravel()
317
+ out_data = out._data
318
+
319
+ for xx in np.ndindex(*nrows):
320
+ ii = tuple(x+p for x,p in zip(xx,pads))
321
+ i_mats = [mat[i+s, k] for i,s,k,mat in zip(ii, c_starts, kk, mats)]
322
+ out_data[ii] = np.dot(x_data, np.outer(*i_mats).ravel())
323
+
324
+ # IMPORTANT: flag that ghost regions are not up-to-date
325
+ out.ghost_regions_in_sync = False
326
+ return out
327
+
328
+ # ...
329
+ def copy(self):
330
+ mats = [m.copy() for m in self.mats]
331
+ return KroneckerDenseMatrix(self.domain, self.codomain, *mats, with_pads=True)
332
+
333
+ # ...
334
+ def __neg__(self):
335
+ mats = [-self.mats[0], *(m.copy() for m in self.mats[1:])]
336
+ return KroneckerDenseMatrix(self.domain, self.codomain, *mats, with_pads=True)
337
+
338
+ # ...
339
+ def __mul__(self, a):
340
+ mats = [*(m.copy() for m in self.mats[:-1]), self.mats[-1] * a]
341
+ return KroneckerDenseMatrix(self.domain, self.codomain, *mats, with_pads=True)
342
+
343
+ # ...
344
+ def __rmul__(self, a):
345
+ mats = [a * self.mats[0], *(m.copy() for m in self.mats[1:])]
346
+ return KroneckerDenseMatrix(self.domain, self.codomain, *mats, with_pads=True)
347
+
348
+ # ...
349
+ def __imul__(self, a):
350
+ self.mats[-1] *= a
351
+ return self
352
+
353
+ #--------------------------------------
354
+ # Other properties/methods
355
+ #--------------------------------------
356
+
357
+ def tosparse(self, **kwargs):
358
+ return coo_matrix(reduce(kron, (m[p:-p,p:-p] for m,p in zip(self.mats, self.domain.pads))))
359
+
360
+ def toarray(self):
361
+ return reduce(kron, (m[p:-p,p:-p] for m,p in zip(self.mats, self.domain.pads)))
362
+
363
+ def transpose(self, conjugate=False):
364
+ mats = [Mi.conj() for Mi in self.mats] if conjugate else self.mats
365
+ mats_tr = [Mi.T for Mi in mats]
366
+ return KroneckerDenseMatrix(self.codomain, self.domain, *mats_tr, with_pads=True)
367
+
368
+ def exchange_assembly_data( self ):
369
+ pass
370
+
371
+ def set_backend(self, backend, precompiled=False):
372
+ pass
373
+ #==============================================================================
374
+ class KroneckerLinearSolver(LinearOperator):
375
+ """
376
+ A solver for Ax=b, where A is a Kronecker matrix from arbirary dimension d,
377
+ defined by d solvers. We also need information about the space of b.
378
+
379
+ Parameters
380
+ ----------
381
+ V : StencilVectorSpace
382
+ The space b will live in; i.e. which gives us information about
383
+ the distribution of the right-hand side b.
384
+
385
+ W : StencilVectorSpace
386
+ The space x will live in; i.e. which gives us information about
387
+ the distribution of the unknown vector x.
388
+
389
+ solvers : list of LinearSolver
390
+ The components of A in each dimension.
391
+
392
+ Attributes
393
+ ----------
394
+ domain : StencilVectorSpace
395
+ The space of the rhs vector b.
396
+
397
+ codomain : StencilVectorSpace
398
+ The space of the unknown vector x.
399
+ """
400
+ def __init__(self, V, W, solvers):
401
+ assert isinstance(V, StencilVectorSpace)
402
+ assert isinstance(W, StencilVectorSpace)
403
+ assert hasattr( solvers, '__iter__' )
404
+ for solver in solvers:
405
+ assert isinstance(solver, LinearSolver)
406
+
407
+ assert V.ndim == len(solvers)
408
+ assert W.ndim == len(solvers)
409
+ assert V.npts == W.npts
410
+
411
+ # general arguments
412
+ self._domain = V
413
+ self._codomain = W
414
+ self._solvers = solvers
415
+ self._parallel = self._domain.parallel
416
+ self._dtype = self._codomain._dtype
417
+ if self._parallel:
418
+ self._mpi_type = self._domain._mpi_type
419
+ else:
420
+ self._mpi_type = None
421
+ self._ndim = self._codomain.ndim
422
+
423
+ # compute and setup solver arguments
424
+ self._setup_solvers()
425
+
426
+ # compute reordering permutations between the steps
427
+ self._setup_permutations()
428
+
429
+ # for now: allocate temporary arrays here (can be removed later)
430
+ self._temp1, self._temp2 = self._allocate_temps()
431
+
432
+ def _setup_solvers(self):
433
+ """
434
+ Computes the distribution of elements and sets up the solvers
435
+ (which potentially utilize MPI).
436
+ """
437
+ # slice sizes
438
+ starts = np.array(self._domain.starts)
439
+ ends = np.array(self._domain.ends) + 1
440
+ self._slice = tuple([slice(s, e) for s,e in zip(starts, ends)])
441
+
442
+ # local and global sizes
443
+ nglobals = self._domain.npts
444
+ nlocals = ends - starts
445
+ self._localsize = np.prod(nlocals)
446
+ mglobals = self._localsize // nlocals
447
+ self._nlocals = nlocals
448
+
449
+ # solver passes (and mlocal size)
450
+ solver_passes = [None] * self._ndim
451
+
452
+ tempsize = self._localsize
453
+ self._allserial = True
454
+ for i in range(self._ndim):
455
+ # decide for each direction individually, if we should
456
+ # use a serial or a parallel/distributed solver
457
+ # useful e.g. if we have little data in some directions
458
+ # (and thus no data distributed there)
459
+
460
+ if not self._parallel or self._domain.cart.subcomm[i].size <= 1:
461
+ # serial solve
462
+ solver_passes[i] = KroneckerLinearSolver.KroneckerSolverSerialPass(
463
+ self._solvers[i], nglobals[i], mglobals[i])
464
+ else:
465
+ # for the parallel case, use Alltoallv
466
+ solver_passes[i] = KroneckerLinearSolver.KroneckerSolverParallelPass(
467
+ self._solvers[i], self._domain._mpi_type, i,
468
+ self._domain.cart, mglobals[i], nglobals[i], nlocals[i], self._localsize)
469
+
470
+ # we have a parallel solve pass now, so we are not completely local any more
471
+ self._allserial = False
472
+
473
+ # update memory requirements
474
+ tempsize = max(tempsize, solver_passes[i].required_memory())
475
+
476
+ # we want to start with the last dimension
477
+ self._solver_passes = list(reversed(solver_passes))
478
+ self._tempsize = tempsize
479
+
480
+ def _setup_permutations(self):
481
+ """
482
+ Creates the permutations and matrix shapes which occur during reordering
483
+ the data for the Kronecker solve operations.
484
+ """
485
+
486
+ # we use a single permutation for all steps
487
+ # it is: (n, 1, 2, ..., n-1)
488
+ self._perm = np.arange(self._ndim)
489
+ self._perm[1:] = self._perm[:-1]
490
+ self._perm[0] = self._ndim - 1
491
+
492
+ # side note: we tried out other permutations:
493
+ # swapping one dimension with the last one each time showed a bad performance...
494
+
495
+ # re-order the shapes based on the permutations
496
+ self._shapes = [None] * self._ndim
497
+ self._shapes[0] = self._nlocals
498
+ for i in range(1, self._ndim):
499
+ self._shapes[i] = self._shapes[i-1][self._perm]
500
+
501
+ def _allocate_temps(self):
502
+ """
503
+ Allocates all temporary data needed for the solve operation.
504
+ """
505
+ temp1 = np.empty((self._tempsize,), dtype=self._dtype)
506
+ if self._ndim <= 1 and self._allserial:
507
+ # if ndim==1 and we have no parallelism,
508
+ # we can avoid allocating a second temp array
509
+ temp2 = None
510
+ else:
511
+ temp2 = np.empty((self._tempsize,), dtype=self._dtype)
512
+ return temp1, temp2
513
+
514
+ @property
515
+ def domain(self):
516
+ return self._domain
517
+
518
+ @property
519
+ def codomain(self):
520
+ return self._codomain
521
+
522
+ @property
523
+ def dtype(self):
524
+ return None
525
+
526
+ def toarray(self):
527
+ raise NotImplementedError('toarray() is not defined for KroneckerLinearSolvers.')
528
+
529
+ def tosparse(self):
530
+ raise NotImplementedError('tosparse() is not defined for KroneckerLinearSolvers.')
531
+
532
+ def transpose(self, conjugate=False):
533
+ new_domain = self._codomain
534
+ new_codomain = self._domain
535
+ new_solvers = [solver.transpose() for solver in self._solvers]
536
+ return KroneckerLinearSolver(new_domain, new_codomain, new_solvers)
537
+
538
+ def dot(self, v, out=None):
539
+ return self.solve(v, out=out)
540
+
541
+ @property
542
+ def solvers(self):
543
+ """
544
+ Returns an immutable view onto references to the one-dimensional solvers.
545
+ """
546
+ return tuple(self._solvers)
547
+
548
+ def solve(self, rhs, out=None):
549
+ """
550
+ Solves Ax=b where A is a Kronecker product matrix (and represented as such),
551
+ and b is a suitable vector.
552
+ """
553
+
554
+ # type checks
555
+ assert rhs.space is self._domain
556
+
557
+ if out is not None:
558
+ assert isinstance( out, StencilVector )
559
+ assert out.space is self._codomain
560
+ else:
561
+ out = StencilVector( rhs.space )
562
+
563
+ inslice = rhs[self._slice]
564
+ outslice = out[self._slice]
565
+
566
+ # call the actual kernel
567
+ self._solve_nd(inslice, outslice)
568
+
569
+ out.update_ghost_regions()
570
+ return out
571
+
572
+ def _solve_nd(self, inslice, outslice):
573
+ """
574
+ The internal solve loop. Can handle arbitrary dimensions.
575
+ """
576
+ temp1 = self._temp1
577
+ temp2 = self._temp2
578
+
579
+ # copy input
580
+ self._inslice_to_temp(inslice, temp1)
581
+
582
+ # internal passes
583
+ for i in range(self._ndim - 1):
584
+ # solve direction
585
+ self._solver_passes[i].solve_pass(temp1, temp2)
586
+
587
+ # reorder and swap
588
+ self._reorder_temp_to_temp(temp1, temp2, i)
589
+ temp1, temp2 = temp2, temp1
590
+
591
+ # last pass
592
+ self._solver_passes[-1].solve_pass(temp1, temp2)
593
+
594
+ # copy to output
595
+ self._reorder_temp_to_outslice(temp1, outslice)
596
+
597
+ def _inslice_to_temp(self, inslice, target):
598
+ """
599
+ Copies data to an internal, 1-dimensional temporary array.
600
+ Does not allocate any new array.
601
+ """
602
+ targetview = target[:self._localsize]
603
+ targetview.shape = inslice.shape
604
+
605
+ targetview[:] = inslice
606
+
607
+ def _reorder_temp_to_temp(self, source, target, i):
608
+ """
609
+ Reorders the dimensions of the temporary arrays, and copies data from one to another.
610
+ Does not allocate any new array.
611
+ """
612
+ sourceview = source[:self._localsize]
613
+ sourceview.shape = self._shapes[i]
614
+
615
+ targetview = target[:self._localsize]
616
+ targetview.shape = self._shapes[i+1]
617
+
618
+ targetview[:] = sourceview.transpose(self._perm)
619
+
620
+ def _reorder_temp_to_outslice(self, source, outslice):
621
+ """
622
+ Reorders the dimensions of the temporary array for a final time, and copies it to the output.
623
+ Does not allocate any new array.
624
+ """
625
+ sourceview = source[:self._localsize]
626
+ sourceview.shape = self._shapes[-1]
627
+
628
+ outslice[:] = sourceview.transpose(self._perm)
629
+
630
+ class KroneckerSolverSerialPass:
631
+ """
632
+ Solves a linear equation for several right-hand sides at the same time,
633
+ given that the data is already in memory.
634
+
635
+ Parameters
636
+ ----------
637
+ solver : BandedSolver or SparseSolver
638
+ The internally used solver class.
639
+
640
+ nglobal : int
641
+ The length of the dimension which we want to solve for.
642
+
643
+ mglobal : int
644
+ The number of right-hand sides we want to solve. Equals the product of the
645
+ number of dimensions which we do NOT want to solve for
646
+ (when squashing all these dimensions into a single one).
647
+ I.e. mglobal*nglobal is the total data size.
648
+ """
649
+ def __init__(self, solver, nglobal, mglobal):
650
+ self._numrhs = mglobal
651
+ self._dimrhs = nglobal
652
+ self._datasize = nglobal*mglobal
653
+ self._solver = solver
654
+ self._view = None
655
+
656
+ def required_memory(self):
657
+ """
658
+ Returns the required memory for this operation. Minimum size for the workmem and tempmem parameters.
659
+ """
660
+ return self._datasize
661
+
662
+ def solve_pass(self, workmem, tempmem):
663
+ """
664
+ Solves the data available in workmem, assuming that all data is available locally.
665
+
666
+ Parameters
667
+ ----------
668
+ workmem : ndarray
669
+ The data which is to be solved. It is a one-dimensional ndarray
670
+ which contains all columns contiguously ordered in memory one after another.
671
+ Its minimum size is also given by `self.required_mem()`.
672
+
673
+ tempmem : ndarray
674
+ Ignored, it exists for compatibility with the parallel solver.
675
+ """
676
+ # reshape necessary memory in column-major
677
+ view = workmem[:self._datasize]
678
+ view.shape = (self._numrhs,self._dimrhs)
679
+
680
+ # call solver in in-place mode
681
+ self._solver.solve(view, out=view)
682
+
683
+ class KroneckerSolverParallelPass:
684
+ """
685
+ Solves a linear equation for several right-hand sides at the same time,
686
+ using an Alltoallv operation to distribute the data.
687
+
688
+ The parameters use the form of n and m; here n denotes the
689
+ length of the dimension we want to solve for, and m is the
690
+ length of all other dimensions, multiplied with each other.
691
+ These n and m are then suffixed with local and global,
692
+ denoting how much of them we have (or want to have) locally.
693
+ So, nglobal is the dimension of the columns we want to solve,
694
+ nlocal is the part we have on our local processor. mglobal is
695
+ the number of right-hand sides to solve in the whole communicator,
696
+ and mlocal is the number of right-hand sides we will solve on our
697
+ local processor.
698
+
699
+ Parameters
700
+ ----------
701
+ solver : BandedSolver or SparseSolver
702
+ The internally used solver class.
703
+
704
+ mpi_type : MPI type
705
+ The MPI type of the space. Used for the Alltoallv.
706
+
707
+ i : int
708
+ The index of the dimension.
709
+
710
+ cart : CartDecomposition
711
+ The cartesian decomposition we use.
712
+
713
+ mglobal : int
714
+ The number of right-hand sides we want to solve. Equals the product of the
715
+ number of dimensions which we do NOT want to solve for
716
+ (when squashing all these dimensions into a single one).
717
+ I.e. mglobal*nglobal is the total data size in our communicator (not on the whole grid though).
718
+
719
+ nglobal : int
720
+ The length of the dimension which we want to solve.
721
+ (the total length, not the one we have on this process)
722
+
723
+ nlocal : int
724
+ The length of the part of the dimension to solve which is located on this process already.
725
+
726
+ localsize : int
727
+ The size of data on our local process.
728
+ Equals mlocal * nlocal (given that we know the former).
729
+ """
730
+
731
+ # To understand the following, here is a short explaination. Consider two processes like this:
732
+ #
733
+ # Pr1 | Pr2
734
+ # 0 1 | 2 3
735
+ # 4 5 | 6 7
736
+ # 8 9 | A B
737
+ # C D | E F
738
+ #
739
+ # i.e. Pr1 has 0 1 4 5 8 9 C D; Pr2 has 2 3 6 7 A B E F
740
+ #
741
+ # We now would like to get each line on at least one process. So, we do an AlltoAll like this:
742
+ #
743
+ # Pr1 | Pr2
744
+ # 0 1 | 2 3 | to Pr1
745
+ # 4 5 | 6 7 | to Pr1
746
+ # ------------------
747
+ # 8 9 | A B | to Pr2
748
+ # C D | E F | to Pr2
749
+ #
750
+ # But the data is transported per process, i.e. we get in this order:
751
+ # 0 1 4 5 2 3 6 7 on Pr1
752
+ # 8 9 C D A B E F on Pr2
753
+ #
754
+ # so we still need to re-order (i.e. partially transpose) locally to finally get what we want.
755
+ # 0 1 2 3 4 5 6 7 on Pr1
756
+ # 8 9 A B C D E F on Pr2
757
+ #
758
+
759
+ # NOTE: ideas for future improvements, if this is too slow:
760
+ #
761
+ # * Use MPI composite Datatypes (i.e. MPI contiguous and vector).
762
+ # This may improve performance, depending on the implementation
763
+ # (therefore, a library-level switch or similar would be an option here).
764
+ # Mainly, we could push what happens in _blocked_to_contiguous and
765
+ # _contiguous_to_blocked methods into the MPI implementation.
766
+ #
767
+ # * Use Alltoall instead of Alltoallv, when applicable, since it might be faster as well.
768
+ # This if for example the case, if the cartesian communicator (cart argument)
769
+ # assigns the same number of data points to all processes.
770
+ # (i.e. global_ends[i] - global_starts[i] is constant)
771
+ # Then, we only need mlocal to be constant (except the last element) as well.
772
+ #
773
+
774
+ def __init__(self, solver, mpi_type, i, cart, mglobal, nglobal, nlocal, localsize):
775
+ self._nglobal = nglobal
776
+
777
+ # cartesian distribution
778
+ comm = cart.subcomm[i]
779
+ cartend = cart.global_ends[i] + 1
780
+ cartstart = cart.global_starts[i]
781
+ cartsize = cartend - cartstart
782
+
783
+ # source MPI sizes and disps
784
+ # distribute the data like
785
+ # (N+1, N+1, ..., N+1, N, N, ...)
786
+ # where N = floor(mglobaldata / comm.size)
787
+ mlocal_pre = mglobal // comm.size
788
+ mlocal_add = mglobal % comm.size
789
+ sourcesizes = np.full((comm.size,), mlocal_pre, dtype=int)
790
+ sourcesizes[:mlocal_add] += 1
791
+ mlocal = sourcesizes[comm.rank]
792
+ sourcesizes *= nlocal
793
+
794
+ # disps, created from the sizes
795
+ sourcedisps = np.zeros((comm.size+1,), dtype=int)
796
+ np.cumsum(sourcesizes, out=sourcedisps[1:])
797
+ sourcedisps = sourcedisps[:-1]
798
+
799
+ # target MPI sizes and disps
800
+ # (mlocal is the same over all processes in the communicator)
801
+ targetsizes = cartsize * mlocal
802
+ targetdisps = cartstart * mlocal
803
+
804
+ # setting all arguments to keep
805
+ self._mlocal = mlocal
806
+ self._localsize = localsize
807
+ self._datasize = mlocal * nglobal
808
+ self._source_transfer = (sourcesizes, sourcedisps)
809
+ self._target_transfer = (targetsizes, targetdisps)
810
+ self._mpi_type = mpi_type
811
+ self._cartstart = cartstart
812
+ self._cartend = cartend
813
+ self._comm = comm
814
+ self._serialsolver = KroneckerLinearSolver.KroneckerSolverSerialPass(solver, nglobal, mlocal)
815
+
816
+ def required_memory(self):
817
+ """
818
+ Returns the required memory for this operation. Minimum size for the workmem and tempmem parameters.
819
+ """
820
+ return max(self._datasize, self._localsize)
821
+
822
+ def _blocked_to_contiguous(self, blocked, contiguous):
823
+ """
824
+ Copies from a blocked view to a contiguous view.
825
+ Equals roughly a partial transpose, if the block sizes in the cartesian grid are the same.
826
+ """
827
+ blocked_view = blocked[:self._datasize]
828
+ blocked_view.shape = (self._mlocal,self._nglobal)
829
+ for start, end in zip(self._cartstart, self._cartend):
830
+ contiguouspart = contiguous[start*self._mlocal:end*self._mlocal]
831
+ contiguouspart.shape = (self._mlocal,end-start)
832
+ blocked_view[:,start:end] = contiguouspart
833
+
834
+ def _contiguous_to_blocked(self, blocked, contiguous):
835
+ """
836
+ Copies from a contiguous view to a blocked view.
837
+ Equals roughly a partial transpose, if the block sizes in the cartesian grid are the same.
838
+ """
839
+ blocked_view = blocked[:self._datasize]
840
+ blocked_view.shape = (self._mlocal,self._nglobal)
841
+ for start, end in zip(self._cartstart, self._cartend):
842
+ contiguouspart = contiguous[start*self._mlocal:end*self._mlocal]
843
+ contiguouspart.shape = (self._mlocal,end-start)
844
+ contiguouspart[:] = blocked_view[:,start:end]
845
+
846
+ def solve_pass(self, workmem, tempmem):
847
+ """
848
+ Solves the data available in workmem in a distributed manner, using MPI_Alltoallv.
849
+
850
+ Parameters
851
+ ----------
852
+ workmem : ndarray
853
+ The data which is used for solving.
854
+ All columns to be solved are ordered contiguously.
855
+ Its minimum size is given by `self.required_mem()`
856
+
857
+ tempmem : ndarray
858
+ Temporary array of the same minimum size as workmem.
859
+ """
860
+ # preparation
861
+ sourceargs = [workmem[:self._localsize], self._source_transfer, self._mpi_type]
862
+ targetargs = [tempmem[:self._datasize], self._target_transfer, self._mpi_type]
863
+
864
+ # parts of stripes -> blocked stripes
865
+ self._comm.Alltoallv(sourceargs, targetargs)
866
+
867
+ # blocked stripes -> ordered stripes
868
+ self._blocked_to_contiguous(workmem, tempmem)
869
+
870
+ # actual solve (source contains the data)
871
+ self._serialsolver.solve_pass(workmem, tempmem)
872
+
873
+ # ordered stripes -> blocked stripes
874
+ self._contiguous_to_blocked(workmem, tempmem)
875
+
876
+ # blocked stripes -> parts of stripes
877
+ self._comm.Alltoallv(targetargs, sourceargs)
878
+
879
+ #==============================================================================
880
+ def kronecker_solve(solvers, rhs, out=None):
881
+ """
882
+ Solve linear system Ax=b with A=kron( A_n, A_{n-1}, ..., A_2, A_1 ), given
883
+ $n$ separate linear solvers $L_n$ for the 1D problems $A_n x_n = b_n$:
884
+
885
+ x_n = L_n.solve( b_n )
886
+
887
+ Parameters
888
+ ----------
889
+ solvers : list( LinearSolver )
890
+ List of linear solvers along each direction: [L_1, L_2, ..., L_n].
891
+
892
+ rhs : StencilVector
893
+ Right hand side vector of linear system Ax=b.
894
+
895
+ """
896
+ # all these feasability checks are again performed in the KroneckerLinearSolver class
897
+ assert hasattr(solvers, '__iter__')
898
+ for solver in solvers:
899
+ assert isinstance(solver, LinearSolver)
900
+
901
+ assert isinstance(rhs, StencilVector)
902
+ assert rhs.space.ndim == len(solvers)
903
+
904
+ if out is not None:
905
+ assert isinstance(out, StencilVector)
906
+ assert out.space is rhs.space
907
+ else:
908
+ out = StencilVector(rhs.space)
909
+
910
+ kronsolver = KroneckerLinearSolver(rhs.space, rhs.space, solvers)
911
+ return kronsolver.solve(rhs, out=out)