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,1393 @@
1
+ # coding: utf-8
2
+
3
+ """
4
+ We assume here that a tensor space is the product of fem spaces whom basis are
5
+ of compact support
6
+
7
+ """
8
+ from feectools.ddm.mpi import mpi as MPI
9
+
10
+ import numpy as np
11
+ import itertools
12
+ import h5py
13
+ import os
14
+
15
+ from types import MappingProxyType
16
+
17
+ from feectools.linalg.stencil import StencilVectorSpace
18
+ from feectools.linalg.kron import kronecker_solve
19
+ from feectools.fem.basic import FemSpace, FemField
20
+ from feectools.fem.splines import SplineSpace
21
+ from feectools.fem.grid import FemAssemblyGrid
22
+ from feectools.fem.partitioning import create_cart, partition_coefficients
23
+ from feectools.ddm.cart import DomainDecomposition, CartDecomposition
24
+
25
+ from feectools.core.bsplines import (find_span,
26
+ basis_funs,
27
+ basis_funs_1st_der,
28
+ basis_ders_on_quad_grid,
29
+ elements_spans,
30
+ cell_index,
31
+ basis_ders_on_irregular_grid)
32
+
33
+ from feectools.core.field_evaluation_kernels import (eval_fields_1d_no_weights,
34
+ eval_fields_1d_irregular_no_weights,
35
+ eval_fields_1d_weighted,
36
+ eval_fields_1d_irregular_weighted,
37
+ eval_fields_2d_no_weights,
38
+ eval_fields_2d_irregular_no_weights,
39
+ eval_fields_2d_weighted,
40
+ eval_fields_2d_irregular_weighted,
41
+ eval_fields_3d_no_weights,
42
+ eval_fields_3d_irregular_no_weights,
43
+ eval_fields_3d_weighted,
44
+ eval_fields_3d_irregular_weighted)
45
+
46
+ __all__ = ('TensorFemSpace',)
47
+
48
+ #===============================================================================
49
+ class TensorFemSpace(FemSpace):
50
+ """
51
+ Tensor-product Finite Element space V.
52
+
53
+ Parameters
54
+ ----------
55
+ domain_decomposition : feectools.ddm.cart.DomainDecomposition
56
+
57
+ *spaces : feectools.fem.splines.SplineSpace
58
+ 1D finite element spaces.
59
+
60
+ coeff_space : feectools.linalg.stencil.StencilVectorSpace or None
61
+ The vector space to which the coefficients belong (optional).
62
+
63
+ cart : feectools.ddm.CartDecomposition or None
64
+ Object that contains all information about the Cartesian decomposition
65
+ of a tensor-product grid of coefficients.
66
+
67
+ dtype : {float, complex}
68
+ Data type of the coefficients.
69
+
70
+ Notes
71
+ -----
72
+ For now we assume that this tensor-product space can ONLY be constructed
73
+ from 1D spline spaces.
74
+
75
+ """
76
+
77
+ def __init__(self, domain_decomposition, *spaces, coeff_space=None, cart=None, dtype=float):
78
+
79
+ assert isinstance(domain_decomposition, DomainDecomposition)
80
+ assert all(isinstance(s, SplineSpace) for s in spaces)
81
+ assert dtype in (float, complex)
82
+ # TODO [YG 10.04.2024]: check if dtype test is too restrictive
83
+
84
+ # Handle optional arguments
85
+ if cart and coeff_space:
86
+ raise ValueError("Cannot provide both 'coeff_space' and 'cart' to constructor")
87
+ elif cart is not None:
88
+ assert isinstance(cart, CartDecomposition)
89
+ coeff_space = StencilVectorSpace(cart, dtype=dtype)
90
+ elif coeff_space is not None:
91
+ assert isinstance(coeff_space, StencilVectorSpace)
92
+ cart = coeff_space.cart
93
+ else:
94
+ cart = create_cart([domain_decomposition], [spaces])[0]
95
+ coeff_space = StencilVectorSpace(cart, dtype=dtype)
96
+
97
+ # Store some info
98
+ self._domain_decomposition = domain_decomposition
99
+ self._spaces = spaces
100
+ self._dtype = dtype
101
+ self._coeff_space = coeff_space
102
+ self._symbolic_space = None
103
+ self._refined_space = {}
104
+ self._interfaces = {}
105
+ self._interfaces_readonly = MappingProxyType(self._interfaces)
106
+
107
+ # If process does not own space, stop here
108
+ if coeff_space.parallel and cart.is_comm_null:
109
+ return
110
+
111
+ # Determine portion of logical domain local to process.
112
+ # This corresponds to the indices of the first and last elements
113
+ # owned by the current process, along each direction.
114
+ self._element_starts = self._coeff_space.cart.domain_decomposition.starts
115
+ self._element_ends = self._coeff_space.cart.domain_decomposition.ends
116
+
117
+ # Compute limits of eta_0, eta_1, eta_2, etc... in subdomain local to process
118
+ self._eta_limits = tuple((space.breaks[s], space.breaks[e+1])
119
+ for s, e, space in zip(self._element_starts, self._element_ends, self._spaces))
120
+
121
+ # Local domains for every process
122
+ self._global_element_starts = domain_decomposition.global_element_starts
123
+ self._global_element_ends = domain_decomposition.global_element_ends
124
+
125
+ # Extended 1D assembly grids (local to process) along each direction
126
+ self._assembly_grids = [{} for _ in range(self.ldim)]
127
+
128
+ # Flag: object NOT YET prepared for interpolation
129
+ self._interpolation_ready = False
130
+
131
+ # Store information about nested grids
132
+ self.set_refined_space(self.ncells, self)
133
+
134
+ #--------------------------------------------------------------------------
135
+ # Abstract interface: read-only attributes
136
+ #--------------------------------------------------------------------------
137
+
138
+ # @property
139
+ # def nquads( self ):
140
+ # assert self._nquads, "nquads has to be set with self._nquads = nquads"
141
+ # return self._nquads
142
+
143
+ # @nquads.setter
144
+ # def nquads(self, value):
145
+ # self._nquads = value
146
+
147
+ # @property
148
+ # def quad_grids( self ):
149
+ # assert self._nquads, "nquads has to be set with self._nquads = nquads"
150
+ # return tuple({q: gag} for q, gag in zip(self.nquads, self.get_assembly_grids(*self.nquads)))
151
+
152
+ @property
153
+ def ldim(self):
154
+ """ Parametric dimension.
155
+ """
156
+ return sum([V.ldim for V in self.spaces])
157
+
158
+ @property
159
+ def periodic(self):
160
+ """
161
+ Tuple of booleans: along each logical dimension,
162
+ say if domain is periodic.
163
+ :rtype: tuple[bool]
164
+ """
165
+ # [YG, 27.03.2025]: according to the abstract interface of FemSpace,
166
+ # this property should return a tuple of `ldim` booleans. However, the
167
+ # spaces in self.spaces seem to be returning a single scalar value.
168
+ return tuple(V.periodic for V in self.spaces)
169
+
170
+ @property
171
+ def domain_decomposition(self):
172
+ return self._domain_decomposition
173
+
174
+ @property
175
+ def mapping(self):
176
+ # [YG, 28.03.2025]: not clear why there should be no mapping here...
177
+ # Clearly this property is never used in feectools.
178
+ return None
179
+
180
+ @property
181
+ def coeff_space(self):
182
+ """
183
+ Vector space of the coefficients (mapping invariant).
184
+ :rtype: feectools.linalg.stencil.StencilVectorSpace
185
+ """
186
+ return self._coeff_space
187
+
188
+ @property
189
+ def symbolic_space( self ):
190
+ return self._symbolic_space
191
+
192
+ @property
193
+ def interfaces( self ):
194
+ return self._interfaces_readonly
195
+
196
+ @symbolic_space.setter
197
+ def symbolic_space( self, symbolic_space ):
198
+ #assert isinstance(symbolic_space, BasicFunctionSpace)
199
+ self._symbolic_space = symbolic_space
200
+
201
+ @property
202
+ def patch_spaces(self):
203
+ return (self,)
204
+
205
+ @property
206
+ def component_spaces(self):
207
+ return (self,)
208
+
209
+ @property
210
+ def axis_spaces(self):
211
+ return self._spaces
212
+
213
+ @property
214
+ def is_multipatch(self):
215
+ return False
216
+
217
+ @property
218
+ def is_vector_valued(self):
219
+ return False
220
+
221
+ #--------------------------------------------------------------------------
222
+ # Abstract interface: evaluation methods
223
+ #--------------------------------------------------------------------------
224
+ def eval_field( self, field, *eta, weights=None):
225
+
226
+ assert isinstance( field, FemField )
227
+ assert field.space is self
228
+ assert len( eta ) == self.ldim
229
+ if weights:
230
+ assert weights.space == field.coeffs.space
231
+
232
+ bases = []
233
+ index = []
234
+
235
+ # Necessary if vector coeffs is distributed across processes
236
+ if not field.coeffs.ghost_regions_in_sync:
237
+ field.coeffs.update_ghost_regions()
238
+
239
+ # Check if `x` is iterable and loop over elements
240
+ if isinstance(eta[0], (list, np.ndarray)) and np.ndim(eta[0]) > 0:
241
+ for dim in range(1, self.ldim):
242
+ assert len(eta[0]) == len(eta[dim])
243
+ res_list = []
244
+ for i in range(len(eta[0])):
245
+ x = [eta[j][i] for j in range(self.ldim)]
246
+ res_list.append(self.eval_field(field, *x, weights=weights))
247
+ return np.array(res_list)
248
+
249
+ for (x, xlim, space) in zip( eta, self.eta_lims, self.spaces ):
250
+
251
+ knots = space.knots
252
+ degree = space.degree
253
+ span = find_span( knots, degree, x )
254
+
255
+ #-------------------------------------------------#
256
+ # Fix span for boundaries between subdomains #
257
+ #-------------------------------------------------#
258
+ # TODO: Use local knot sequence instead of global #
259
+ # one to get correct span in all situations #
260
+ #-------------------------------------------------#
261
+ if x == xlim[1] and x != knots[-1-degree]:
262
+ span -= 1
263
+ #-------------------------------------------------#
264
+ basis = basis_funs( knots, degree, x, span)
265
+
266
+ # If needed, rescale B-splines to get M-splines
267
+ if space.basis == 'M':
268
+ basis *= space.scaling_array[span-degree : span+1]
269
+
270
+ # Determine local span
271
+ wrap_x = space.periodic and x > xlim[1]
272
+ loc_span = span - space.nbasis if wrap_x else span
273
+
274
+ bases.append( basis )
275
+ index.append( slice( loc_span-degree, loc_span+1 ) )
276
+ # Get contiguous copy of the spline coefficients required for evaluation
277
+ index = tuple( index )
278
+ coeffs = field.coeffs[index].copy()
279
+ if weights:
280
+ coeffs *= weights[index]
281
+
282
+ # Evaluation of multi-dimensional spline
283
+ # TODO: optimize
284
+
285
+ # Option 1: contract indices one by one and store intermediate results
286
+ # - Pros: small number of Python iterations = ldim
287
+ # - Cons: we create ldim-1 temporary objects of decreasing size
288
+ #
289
+ res = coeffs
290
+ for basis in bases[::-1]:
291
+ res = np.dot( res, basis )
292
+
293
+ # # Option 2: cycle over each element of 'coeffs' (touched only once)
294
+ # # - Pros: no temporary objects are created
295
+ # # - Cons: large number of Python iterations = number of elements in 'coeffs'
296
+ # #
297
+ # res = 0.0
298
+ # for idx,c in np.ndenumerate( coeffs ):
299
+ # ndbasis = np.prod( [b[i] for i,b in zip( idx, bases )] )
300
+ # res += c * ndbasis
301
+
302
+ return res
303
+
304
+ # ...
305
+ def preprocess_regular_tensor_grid(self, grid, der=0, overlap=0):
306
+ """Returns all the quantities needed to evaluate fields on a regular tensor-grid.
307
+
308
+ Parameters
309
+ ----------
310
+ grid : List of ndarray
311
+ List of 2D arrays representing each direction of the grid.
312
+ Each of these arrays should have shape (ne_xi, nv_xi) where ne_xi is the
313
+ number of cells in the domain in the direction xi and nv_xi is the number of
314
+ evaluation points in the same direction.
315
+
316
+ der : int, default=0
317
+ Number of derivatives of the basis functions to pre-compute.
318
+
319
+ overlap : int
320
+ How much to overlap. Only used in the distributed context.
321
+
322
+ Returns
323
+ -------
324
+ degree : tuple of int
325
+ Degree in each direction
326
+ global_basis : List of ndarray
327
+ List of 4D arrays, one per direction, containing the values of the p+1 non-vanishing
328
+ basis functions (and their derivatives) at each grid point.
329
+ The array for direction xi has shape (ne_xi, der + 1, p+1, nv_xi).
330
+
331
+ global_spans : List of ndarray
332
+ List of 1D arrays, one per direction, containing the index of the last non-vanishing
333
+ basis function in each cell. The array for direction xi has shape (ne_xi,).
334
+
335
+ local_shape : List of tuple
336
+ Shape of what is local to this instance.
337
+ """
338
+ # Check the grid
339
+ assert len(grid) == self.ldim
340
+
341
+ # Get the local domain
342
+ v = self.coeff_space
343
+ starts, ends = self.local_domain
344
+
345
+ # Add the overlap if we are in parallel
346
+ if v.parallel:
347
+ starts = tuple(s - overlap if s!=0 else s for s in starts)
348
+ ends = tuple(e + overlap for e in ends)
349
+
350
+ # Compute the basis functions and spans.
351
+ local_shape = []
352
+ global_basis = []
353
+ global_spans = []
354
+ for i in range(self.ldim):
355
+ # We only care about the local grid
356
+ grid_local = grid[i][slice(starts[i], ends[i] + 1)]
357
+
358
+ # Compute basis functions and spans
359
+ global_basis_i = basis_ders_on_quad_grid(self.knots[i], self.degree[i], grid_local, der, self.spaces[i].basis, offset=starts[i])
360
+ global_spans_i = elements_spans(self.knots[i], self.degree[i])[slice(starts[i], ends[i] + 1)] - v.starts[i] + v.shifts[i] * v.pads[i]
361
+
362
+ local_shape.append(grid_local.shape)
363
+ global_basis.append(global_basis_i)
364
+ global_spans.append(global_spans_i)
365
+ return self.degree, global_basis, global_spans, local_shape
366
+
367
+ #...
368
+ def preprocess_irregular_tensor_grid(self, grid, der=0, overlap=0):
369
+ """Returns all the quantities needed to evaluate fields on a regular tensor-grid.
370
+
371
+ Parameters
372
+ ----------
373
+ grid : List of ndarray
374
+ List of 1D arrays representing each direction of the grid.
375
+
376
+ der : int, default=0
377
+ Number of derivatives of the basis functions to pre-compute.
378
+
379
+ overlap : int
380
+ How much to overlap. Only used in the distributed context.
381
+
382
+ Returns
383
+ -------
384
+ pads : tuple of int
385
+ Padding in each direction
386
+ degree : tuple of int
387
+ Degree in each direction
388
+ global_basis : List of ndarray
389
+ List of 4D arrays, one per direction, containing the values of the p+1 non-vanishing
390
+ basis functions (and their derivatives) at each grid point.
391
+ The array for direction xi has shape (n_xi, p+1, der + 1).
392
+
393
+ global_spans : List of ndarray
394
+ List of 1D arrays, one per direction, containing the index of the last non-vanishing
395
+ basis function in each cell. The array for direction xi has shape (n_xi,).
396
+
397
+ cell_indexes : list of ndarray
398
+ List of 1D arrays, one per direction, containing the index of the cell in which
399
+ the corresponding point in grid is.
400
+
401
+ local_shape : List of tuple
402
+ Shape of what is local to this instance.
403
+ """
404
+ # Check the grid
405
+ assert len(grid) == self.ldim
406
+
407
+ # Get the local domain
408
+ v = self.coeff_space
409
+ starts, ends = self.local_domain
410
+
411
+ # Add the overlap if we are in parallel
412
+ if v.parallel:
413
+ starts = tuple(s - overlap if s!=0 else s for s in starts)
414
+ ends = tuple(e + overlap for e in ends)
415
+
416
+ # Compute the basis functions and spans and cell indexes.
417
+ global_basis = []
418
+ global_spans = []
419
+ cell_indexes = []
420
+ local_shape = []
421
+ for i in range(self.ldim):
422
+ # Check the that the grid is sorted.
423
+ grid_i = grid[i]
424
+ assert all(grid_i[j] <= grid_i[j + 1] for j in range(len(grid_i) - 1))
425
+
426
+ # Get the cell indexes
427
+ cell_index_i = cell_index(self.breaks[i], grid_i)
428
+ min_idx = np.searchsorted(cell_index_i, starts[i], side='left')
429
+ max_idx = np.searchsorted(cell_index_i, ends[i], side='right')
430
+ # We only care about the local cells.
431
+ cell_index_i = cell_index_i[min_idx:max_idx]
432
+ grid_local_i = grid_i[min_idx:max_idx]
433
+
434
+ # basis functions and spans
435
+ global_basis_i = basis_ders_on_irregular_grid(self.knots[i], self.degree[i], grid_local_i, cell_index_i, der, self.spaces[i].basis)
436
+ global_spans_i = elements_spans(self.knots[i], self.degree[i])[slice(starts[i], ends[i] + 1)] - v.starts[i] + v.shifts[i] * v.pads[i]
437
+
438
+ local_shape.append(len(grid_local_i))
439
+ global_basis.append(global_basis_i)
440
+ global_spans.append(global_spans_i)
441
+
442
+ # starts[i] is cell 0 of the local domain
443
+ cell_indexes.append(cell_index_i - starts[i])
444
+
445
+ return self.degree, global_basis, global_spans, cell_indexes, local_shape
446
+
447
+ # ...
448
+ def eval_fields(self, grid, *fields, weights=None, npts_per_cell=None, overlap=0):
449
+ """Evaluate one or several fields at the given location(s) grid.
450
+
451
+ Parameters
452
+ ----------
453
+ grid : List of ndarray
454
+ Grid on which to evaluate the fields
455
+
456
+ *fields : tuple of feectools.fem.basic.FemField
457
+ Fields to evaluate
458
+
459
+ weights : feectools.fem.basic.FemField or None, optional
460
+ Weights field.
461
+
462
+ npts_per_cell: int or tuple of int or None, optional
463
+ number of evaluation points in each cell.
464
+ If an integer is given, then assume that it is the same in every direction.
465
+
466
+ overlap : int
467
+ How much to overlap. Only used in the distributed context.
468
+
469
+ Returns
470
+ -------
471
+ List of ndarray of floats
472
+ List of the evaluated fields.
473
+ """
474
+ assert all(f.space is self for f in fields)
475
+ for f in fields:
476
+ # Necessary if vector coeffs is distributed across processes
477
+ if not f.coeffs.ghost_regions_in_sync:
478
+ f.coeffs.update_ghost_regions()
479
+
480
+ if weights is not None:
481
+ assert weights.space is self
482
+ assert all(f.coeffs.space is weights.coeffs.space for f in fields)
483
+ if not weights.coeffs.ghost_regions_in_sync:
484
+ weights.coeffs.update_ghost_regions()
485
+
486
+ assert len(grid) == self.ldim
487
+ grid = [np.asarray(grid[i]) for i in range(self.ldim)]
488
+ assert all(grid[i].ndim == grid[i + 1].ndim for i in range(self.ldim - 1))
489
+
490
+ # --------------------------
491
+ # Case 1. Scalar coordinates
492
+ if (grid[0].size == 1) or grid[0].ndim == 0:
493
+ if weights is not None:
494
+ return [self.eval_field(f, *grid, weights=weights.coeffs) for f in fields]
495
+ else:
496
+ return [self.eval_field(f, *grid) for f in fields]
497
+
498
+ # Case 2. 1D array of coordinates and no npts_per_cell is given
499
+ # -> grid is tensor-product, but npts_per_cell is not the same in each cell
500
+ elif grid[0].ndim == 1 and npts_per_cell is None:
501
+ out_fields = self.eval_fields_irregular_tensor_grid(grid, *fields, weights=weights, overlap=overlap)
502
+ return [np.ascontiguousarray(out_fields[..., i]) for i in range(len(fields))]
503
+
504
+ # Case 3. 1D arrays of coordinates and npts_per_cell is a tuple or an integer
505
+ # -> grid is tensor-product, and each cell has the same number of evaluation points
506
+ elif grid[0].ndim == 1 and npts_per_cell is not None:
507
+ if isinstance(npts_per_cell, int):
508
+ npts_per_cell = (npts_per_cell,) * self.ldim
509
+ for i in range(self.ldim):
510
+ ncells_i = len(self.breaks[i]) - 1
511
+ grid[i] = np.reshape(grid[i], (ncells_i, npts_per_cell[i]))
512
+ out_fields = self.eval_fields_regular_tensor_grid(grid, *fields, weights=weights, overlap=overlap)
513
+ # return a list
514
+ return [np.ascontiguousarray(out_fields[..., i]) for i in range(len(fields))]
515
+
516
+ # Case 4. (self.ldim)D arrays of coordinates and no npts_per_cell
517
+ # -> unstructured grid
518
+ elif grid[0].ndim == self.ldim and npts_per_cell is None:
519
+ raise NotImplementedError("Unstructured grids are not supported yet.")
520
+
521
+ # Case 5. Nonsensical input
522
+ else:
523
+ raise ValueError("This combination of argument isn't understood. The 4 cases understood are :\n"
524
+ "Case 1. Scalar coordinates\n"
525
+ "Case 2. 1D array of coordinates and no npts_per_cell is given\n"
526
+ "Case 3. 1D arrays of coordinates and npts_per_cell is a tuple or an integer\n"
527
+ "Case 4. {0}D arrays of coordinates and no npts_per_cell".format(self.ldim))
528
+
529
+ # ...
530
+ def eval_fields_regular_tensor_grid(self, grid, *fields, weights=None, overlap=0):
531
+ """Evaluate fields on a regular tensor grid
532
+
533
+ Parameters
534
+ ----------
535
+ grid : List of ndarray
536
+ List of 2D arrays representing each direction of the grid.
537
+ Each of these arrays should have shape (ne_xi, nv_xi) where ne is the
538
+ number of cells in the domain in the direction xi and nv_xi is the number of
539
+ evaluation points in the same direction.
540
+
541
+ *fields : tuple of feectools.fem.basic.FemField
542
+ Fields to evaluate on `grid`.
543
+
544
+ weights : feectools.fem.basic.FemField or None, optional
545
+ Weights to apply to our fields.
546
+
547
+ overlap : int
548
+ How much to overlap. Only used in the distributed context.
549
+
550
+ Returns
551
+ -------
552
+ List of ndarray of float
553
+ Values of the fields on the regular tensor grid
554
+ """
555
+ degree, global_basis, global_spans, local_shape = self.preprocess_regular_tensor_grid(grid, der=0, overlap=overlap)
556
+ ncells = [local_shape[i][0] for i in range(self.ldim)]
557
+ n_eval_points = [local_shape[i][1] for i in range(self.ldim)]
558
+ out_fields = np.zeros((*(tuple(ncells[i] * n_eval_points[i] for i in range(self.ldim))), len(fields)), dtype=self.dtype)
559
+
560
+ global_arr_coeffs = np.zeros(shape=(*fields[0].coeffs._data.shape, len(fields)), dtype=self.dtype)
561
+
562
+ for i in range(len(fields)):
563
+ global_arr_coeffs[..., i] = fields[i].coeffs._data
564
+
565
+ if weights is None:
566
+ args = (*ncells, *degree, *n_eval_points, *global_basis, *global_spans, global_arr_coeffs, out_fields)
567
+ if self.ldim == 1: eval_fields_1d_no_weights(*args)
568
+ elif self.ldim == 2: eval_fields_2d_no_weights(*args)
569
+ elif self.ldim == 3: eval_fields_3d_no_weights(*args)
570
+ else:
571
+ raise NotImplementedError(f"eval_fields_{self.ldim}d_no_weights not implemented")
572
+ else:
573
+ global_weight_coeffs = weights.coeffs._data
574
+ args = (*ncells, *degree, *n_eval_points, *global_basis, *global_spans, global_arr_coeffs, global_weight_coeffs, out_fields)
575
+ if self.ldim == 1: eval_fields_1d_weighted(*args)
576
+ elif self.ldim == 2: eval_fields_2d_weighted(*args)
577
+ elif self.ldim == 3: eval_fields_3d_weighted(*args)
578
+ else:
579
+ raise NotImplementedError(f"eval_fields_{self.ldim}d_weighted not implemented")
580
+
581
+ return out_fields
582
+
583
+ # ...
584
+ def eval_fields_irregular_tensor_grid(self, grid, *fields, weights=None, overlap=0):
585
+ """Evaluate fields on a regular tensor grid
586
+
587
+ Parameters
588
+ ----------
589
+ grid : List of ndarray
590
+ List of 2D arrays representing each direction of the grid.
591
+ Each of these arrays should have shape (ne_xi, nv_xi) where ne is the
592
+ number of cells in the domain in the direction xi and nv_xi is the number of
593
+ evaluation points in the same direction.
594
+
595
+ *fields : tuple of feectools.fem.basic.FemField
596
+ Fields to evaluate on `grid`.
597
+
598
+ weights : feectools.fem.basic.FemField or None, optional
599
+ Weights to apply to our fields.
600
+
601
+ overlap : int
602
+ How much to overlap. Only used in the distributed context.
603
+
604
+ Returns
605
+ -------
606
+ List of ndarray of float
607
+ Values of the fields on the regular tensor grid
608
+ """
609
+ degree, global_basis, global_spans, cell_indexes, local_shape = \
610
+ self.preprocess_irregular_tensor_grid(grid, overlap=overlap)
611
+ out_fields = np.zeros(tuple(local_shape) + (len(fields),), dtype=self.dtype)
612
+
613
+ global_arr_coeffs = np.zeros(shape=(*fields[0].coeffs._data.shape, len(fields)), dtype=self.dtype)
614
+
615
+ npoints = local_shape
616
+
617
+ for i in range(len(fields)):
618
+ global_arr_coeffs[..., i] = fields[i].coeffs._data
619
+
620
+ if weights is None:
621
+ args = (*npoints, *degree, *cell_indexes, *global_basis, *global_spans, global_arr_coeffs, out_fields)
622
+ if self.ldim == 1: eval_fields_1d_irregular_no_weights(*args)
623
+ elif self.ldim == 2: eval_fields_2d_irregular_no_weights(*args)
624
+ elif self.ldim == 3: eval_fields_3d_irregular_no_weights(*args)
625
+ else:
626
+ raise NotImplementedError(f"eval_fields_{self.ldim}d_irregular_no_weights not implemented")
627
+ else:
628
+ global_weight_coeffs = weights.coeffs._data
629
+ args = (*npoints, *degree, *cell_indexes, *global_basis, *global_spans, global_arr_coeffs, global_weight_coeffs, out_fields)
630
+ if self.ldim == 1: eval_fields_1d_irregular_weighted(*args)
631
+ elif self.ldim == 2: eval_fields_2d_irregular_weighted(*args)
632
+ elif self.ldim == 3: eval_fields_3d_irregular_weighted(*args)
633
+ else:
634
+ raise NotImplementedError(f"eval_fields_{self.ldim}d_irregular_weighted not implemented")
635
+
636
+ return out_fields
637
+
638
+ # ...
639
+ def eval_field_gradient(self, field, *eta, weights=None):
640
+
641
+ assert isinstance(field, FemField)
642
+ assert field.space is self
643
+ assert len(eta) == self.ldim
644
+
645
+ bases_0 = []
646
+ bases_1 = []
647
+ index = []
648
+
649
+ # Check if `x` is iterable and loop over elements
650
+ if isinstance(eta[0], (list, np.ndarray)) and np.ndim(eta[0]) > 0:
651
+ for dim in range(1, self.ldim):
652
+ assert len(eta[0]) == len(eta[dim])
653
+ res_list = []
654
+ for i in range(len(eta[0])):
655
+ x = [eta[j][i] for j in range(self.ldim)]
656
+ res_list.append(self.eval_field_gradient(field, *x, weights=weights))
657
+ return np.array(res_list)
658
+
659
+ for (x, xlim, space) in zip( eta, self.eta_lims, self.spaces ):
660
+
661
+ knots = space.knots
662
+ degree = space.degree
663
+ span = find_span( knots, degree, x )
664
+ #-------------------------------------------------#
665
+ # Fix span for boundaries between subdomains #
666
+ #-------------------------------------------------#
667
+ # TODO: Use local knot sequence instead of global #
668
+ # one to get correct span in all situations #
669
+ #-------------------------------------------------#
670
+ if x == xlim[1] and x != knots[-1-degree]:
671
+ span -= 1
672
+ #-------------------------------------------------#
673
+ basis_0 = basis_funs(knots, degree, x, span)
674
+ basis_1 = basis_funs_1st_der(knots, degree, x, span)
675
+
676
+ # If needed, rescale B-splines to get M-splines
677
+ if space.basis == 'M':
678
+ scaling = space.scaling_array[span-degree : span+1]
679
+ basis_0 *= scaling
680
+ basis_1 *= scaling
681
+
682
+ # Determine local span
683
+ wrap_x = space.periodic and x > xlim[1]
684
+ loc_span = span - space.nbasis if wrap_x else span
685
+
686
+ bases_0.append( basis_0 )
687
+ bases_1.append( basis_1 )
688
+ index.append( slice( loc_span-degree, loc_span+1 ) )
689
+
690
+ # Get contiguous copy of the spline coefficients required for evaluation
691
+ index = tuple( index )
692
+ coeffs = field.coeffs[index].copy()
693
+ if weights:
694
+ coeffs *= weights[index]
695
+
696
+ # Evaluate each component of the gradient using algorithm described in "Option 1" above
697
+ grad = []
698
+ for d in range( self.ldim ):
699
+ bases = [(bases_1[d] if i==d else bases_0[i]) for i in range( self.ldim )]
700
+ res = coeffs
701
+ for basis in bases[::-1]:
702
+ res = np.dot( res, basis )
703
+ grad.append( res )
704
+
705
+ return grad
706
+
707
+ # ...
708
+ def integral(self, f, *, nquads=None):
709
+
710
+ assert hasattr(f, '__call__')
711
+
712
+ if nquads is None:
713
+ nquads = [p + 1 for p in self.degree]
714
+ elif isinstance(nquads, int):
715
+ nquads = [nquads] * self.ldim
716
+ else:
717
+ nquads = list(nquads)
718
+
719
+ assert all(isinstance(nq, int) for nq in nquads)
720
+ assert all(nq >= 1 for nq in nquads)
721
+
722
+ # Extract and store quadrature data
723
+ assembly_grids = self.get_assembly_grids(*nquads)
724
+ nq = [g.num_quad_pts for g in assembly_grids]
725
+ points = [g.points for g in assembly_grids]
726
+ weights = [g.weights for g in assembly_grids]
727
+
728
+ # Get local element range
729
+ sk = [g.local_element_start for g in assembly_grids]
730
+ ek = [g.local_element_end for g in assembly_grids]
731
+
732
+ # Iterator over multi-index k (equivalent to nested loops over each dimension)
733
+ multi_range = lambda starts, ends: \
734
+ itertools.product(*[range(s, e+1) for s, e in zip(starts, ends)])
735
+
736
+ # Shortcut: Numpy product of all elements in a list
737
+ np_prod = np.prod
738
+
739
+ # Perform Gaussian quadrature in multiple dimensions
740
+ c = 0.0
741
+ for k in multi_range(sk, ek):
742
+
743
+ x = [ points_i[k_i, :] for points_i, k_i in zip( points, k)]
744
+ w = [weights_i[k_i, :] for weights_i, k_i in zip(weights, k)]
745
+
746
+ for q in np.ndindex(*nq):
747
+
748
+ y = [x_i[q_i] for x_i, q_i in zip(x, q)]
749
+ v = [w_i[q_i] for w_i, q_i in zip(w, q)]
750
+
751
+ c += f(*y) * np_prod(v)
752
+
753
+ # All reduce (MPI_SUM)
754
+ if self.coeff_space.parallel:
755
+ mpi_comm = self.coeff_space.cart.comm
756
+ c = mpi_comm.allreduce(c)
757
+
758
+ # convert to native python type if numpy to avoid errors with sympify
759
+ if isinstance(c, np.generic):
760
+ c = c.item()
761
+
762
+ return c
763
+
764
+ #--------------------------------------------------------------------------
765
+ # Other properties and methods
766
+ #--------------------------------------------------------------------------
767
+ @property
768
+ def dtype(self):
769
+ return self._dtype
770
+
771
+ #TODO: return tuple instead of product?
772
+ @property
773
+ def nbasis(self):
774
+ dims = [V.nbasis for V in self.spaces]
775
+ dim = 1
776
+ for d in dims:
777
+ dim *= d
778
+ return dim
779
+
780
+ @property
781
+ def knots(self):
782
+ return [V.knots for V in self.spaces]
783
+
784
+ @property
785
+ def breaks(self):
786
+ return [V.breaks for V in self.spaces]
787
+
788
+ @property
789
+ def degree(self):
790
+ return [V.degree for V in self.spaces]
791
+
792
+ @property
793
+ def multiplicity(self):
794
+ return [V.multiplicity for V in self.spaces]
795
+
796
+ @property
797
+ def pads(self):
798
+ return self.coeff_space.pads
799
+
800
+ @property
801
+ def ncells(self):
802
+ return [V.ncells for V in self.spaces]
803
+
804
+ @property
805
+ def spaces(self):
806
+ return self._spaces
807
+
808
+ def get_assembly_grids(self, *nquads):
809
+ """
810
+ Return a tuple of `FemAssemblyGrid` objects (one for each direction).
811
+
812
+ These objects are local to the process, and contain all 1D information
813
+ which is necessary for the correct assembly of the l.h.s. matrix and
814
+ r.h.s. vector in a finite element method. This information includes
815
+ the coordinates and weights of all quadrature points, as well as the
816
+ values of the non-zero basis functions, and their derivatives, at such
817
+ points.
818
+
819
+ The computed `FemAssemblyGrid` objects are stored in a dictionary in
820
+ `self` with `nquads` as key, and are reused if a match is found.
821
+
822
+ Parameters
823
+ ----------
824
+ *nquads : int
825
+ Number of quadrature points per cell, along each direction.
826
+
827
+ Returns
828
+ -------
829
+ tuple of FemAssemblyGrid
830
+ The 1D assembly grids along each direction.
831
+
832
+ """
833
+
834
+ assert len(nquads) == self.ldim
835
+ assert all(isinstance(nq, int) for nq in nquads)
836
+ assert all(nq >= 1 for nq in nquads)
837
+
838
+ assembly_grids = [None] * len(nquads)
839
+
840
+ for i, nq in enumerate(nquads):
841
+ # Get a reference to the local dictionary of FemAssemblyGrid along direction i
842
+ assembly_grids_dict_i = self._assembly_grids[i]
843
+ # If there is no FemAssemblyGrid for the required number of quadrature points,
844
+ # create a new FemAssemblyGrid and store it in the local dictionary.
845
+ if nq not in assembly_grids_dict_i:
846
+ V = self.spaces[i]
847
+ s = int(self._element_starts[i])
848
+ e = int(self._element_ends [i])
849
+ assembly_grids_dict_i[nq] = FemAssemblyGrid(V, s, e, nderiv=V.degree, nquads=nq)
850
+ # Store the required FemAssemblyGrid in the list
851
+ assembly_grids[i] = assembly_grids_dict_i[nq]
852
+
853
+ # Return a tuple with the FemAssemblyGrid objects
854
+ return tuple(assembly_grids)
855
+
856
+ @property
857
+ def local_domain(self):
858
+ """
859
+ Logical domain local to the process, assuming the global domain is
860
+ decomposed across processes without any overlapping.
861
+
862
+ This information is fundamental for avoiding double-counting when computing
863
+ integrals over the global domain.
864
+
865
+ Returns
866
+ -------
867
+ element_starts : tuple of int
868
+ Start element index along each direction.
869
+
870
+ element_ends : tuple of int
871
+ End element index along each direction.
872
+
873
+ """
874
+ return self._element_starts, self._element_ends
875
+
876
+ @property
877
+ def global_element_starts(self):
878
+ return self._global_element_starts
879
+
880
+ @property
881
+ def global_element_ends(self):
882
+ return self._global_element_ends
883
+
884
+ @property
885
+ def eta_lims(self):
886
+ """
887
+ Eta limits of domain local to the process (for field evaluation).
888
+
889
+ Returns
890
+ -------
891
+ eta_limits: tuple of (2-tuple of float)
892
+ Along each dimension i, limits are given as (eta^i_{min}, eta^i_{max}).
893
+
894
+ """
895
+ return self._eta_limits
896
+
897
+ # ...
898
+ def init_interpolation(self):
899
+ for V in self.spaces:
900
+ # TODO: check if OK to access private attribute...
901
+ if not V._interpolation_ready:
902
+ V.init_interpolation(dtype=self.dtype)
903
+
904
+ # ...
905
+ def init_histopolation(self):
906
+ for V in self.spaces:
907
+ # TODO: check if OK to access private attribute...
908
+ if not V._histopolation_ready:
909
+ V.init_histopolation(dtype=self.dtype)
910
+
911
+ # ...
912
+ def compute_interpolant(self, values, field):
913
+ """
914
+ Compute field (i.e. update its spline coefficients) such that it
915
+ interpolates a certain function $f(x1,x2,..)$ at the Greville points.
916
+
917
+ Parameters
918
+ ----------
919
+ values : StencilVector
920
+ Function values $f(x_i)$ at the n-dimensional tensor grid of
921
+ Greville points $x_i$, to be interpolated.
922
+
923
+ field : FemField
924
+ Input/output argument: tensor spline that has to interpolate the given
925
+ values.
926
+
927
+ """
928
+ assert values.space is self.coeff_space
929
+ assert isinstance( field, FemField )
930
+ assert field.space is self
931
+
932
+ if not self._interpolation_ready:
933
+ self.init_interpolation()
934
+
935
+ # TODO: check if OK to access private attribute '_interpolator' in self.spaces[i]
936
+ kronecker_solve(
937
+ solvers = [V._interpolator for V in self.spaces],
938
+ rhs = values,
939
+ out = field.coeffs,
940
+ )
941
+
942
+ # ...
943
+ def reduce_grid(self, axes=(), knots=()):
944
+ """
945
+ Create a new TensorFemSpace object with a coarser grid than the original one
946
+ we do that by giving a new knot sequence in the desired dimension.
947
+
948
+ Parameters
949
+ ----------
950
+ axes : List of int
951
+ Dimensions where we want to coarsen the grid.
952
+
953
+ knots : List or tuple
954
+ New knot sequences in each dimension.
955
+
956
+ Returns
957
+ -------
958
+ V : TensorFemSpace
959
+ New space with a coarser grid.
960
+
961
+ """
962
+ assert len(axes) == len(knots)
963
+
964
+ comm = MPI.COMM_WORLD
965
+ rank = comm.Get_rank()
966
+
967
+ v = self._coeff_space
968
+ spaces = list(self.spaces)
969
+
970
+ global_starts = list(v._cart._global_starts).copy()
971
+ global_ends = list(v._cart._global_ends).copy()
972
+ global_domains_ends = self._global_element_ends
973
+
974
+ for i, axis in enumerate(axes):
975
+ space = spaces[axis]
976
+ degree = space.degree
977
+ periodic = space.periodic
978
+ breaks = space.breaks
979
+ T = list(knots[i]).copy()
980
+ elements_ends = global_domains_ends[axis]
981
+ boundaries = breaks[elements_ends+1].tolist()
982
+
983
+ for b in boundaries:
984
+ if b not in T:
985
+ T.append(b)
986
+ T.sort()
987
+
988
+ new_space = SplineSpace(degree, knots=T, periodic=periodic,
989
+ dirichlet=space.dirichlet, basis=space.basis)
990
+ spaces[axis] = new_space
991
+ breaks = new_space.breaks.tolist()
992
+ elements_ends = np.array([breaks.index(bd) for bd in boundaries])-1
993
+ elements_starts = np.array([0] + (elements_ends[:-1]+1).tolist())
994
+
995
+ if periodic:
996
+ global_starts[axis] = elements_starts
997
+ global_ends[axis] = elements_ends
998
+ else:
999
+ global_starts[axis] = elements_starts + degree - 1
1000
+ global_ends[axis] = elements_ends + degree - 1
1001
+ global_ends[axis][-1] += 1
1002
+ global_starts[axis][0] = 0
1003
+
1004
+ cart = v._cart.reduce_grid(tuple(global_starts), tuple(global_ends))
1005
+ V = TensorFemSpace(cart.domain_decomposition, *spaces, cart=cart, dtype=v.dtype)
1006
+
1007
+ return V
1008
+
1009
+ # ...
1010
+ def export_fields(self, filename, **fields):
1011
+ """ Write spline coefficients of given fields to HDF5 file.
1012
+ """
1013
+ assert isinstance(filename, str)
1014
+ assert all(field.space is self for field in fields.values())
1015
+
1016
+ V = self.coeff_space
1017
+ comm = V.cart.comm if V.parallel else None
1018
+
1019
+ # Multi-dimensional index range local to process
1020
+ index = tuple(slice(s, e+1) for s,e in zip(V.starts, V.ends))
1021
+
1022
+ # Create HDF5 file (in parallel mode if MPI communicator size > 1)
1023
+ kwargs = {}
1024
+ if comm is not None:
1025
+ if comm.size > 1:
1026
+ kwargs.update(driver='mpio', comm=comm)
1027
+ h5 = h5py.File(filename, mode='w', **kwargs)
1028
+
1029
+ # Add field coefficients as named datasets
1030
+ for name,field in fields.items():
1031
+ dset = h5.create_dataset(name, shape=V.npts, dtype=V.dtype)
1032
+ dset[index] = field.coeffs[index]
1033
+
1034
+ # Close HDF5 file
1035
+ h5.close()
1036
+
1037
+ # ...
1038
+ def import_fields(self, filename, *field_names):
1039
+ """
1040
+ Load fields from HDF5 file containing spline coefficients.
1041
+
1042
+ Parameters
1043
+ ----------
1044
+ filename : str
1045
+ Name of HDF5 input file.
1046
+
1047
+ field_names : list of str
1048
+ Names of the datasets with the required spline coefficients.
1049
+
1050
+ Returns
1051
+ -------
1052
+ fields : list of FemSpace objects
1053
+ Distributed fields, given in the same order of the names.
1054
+
1055
+ """
1056
+ assert isinstance(filename, str)
1057
+ assert all(isinstance(name, str) for name in field_names)
1058
+
1059
+ V = self.coeff_space
1060
+ comm = V.cart.comm if V.parallel else None
1061
+
1062
+ # Multi-dimensional index range local to process
1063
+ index = tuple(slice(s, e+1) for s,e in zip(V.starts, V.ends))
1064
+
1065
+ # Open HDF5 file (in parallel mode if MPI communicator size > 1)
1066
+ kwargs = {}
1067
+ if comm is not None:
1068
+ if comm.size > 1:
1069
+ kwargs.update(driver='mpio', comm=comm)
1070
+ h5 = h5py.File(filename, mode='r', **kwargs)
1071
+
1072
+ # Create fields and load their coefficients from HDF5 datasets
1073
+ fields = []
1074
+ for name in field_names:
1075
+ dset = h5[name]
1076
+ if dset.shape != V.npts:
1077
+ h5.close()
1078
+ raise TypeError('Dataset not compatible with spline space.')
1079
+ field = FemField(self)
1080
+ field.coeffs[index] = dset[index]
1081
+ field.coeffs.update_ghost_regions()
1082
+ fields.append(field)
1083
+
1084
+ # Close HDF5 file
1085
+ h5.close()
1086
+
1087
+ return fields
1088
+
1089
+ # ...
1090
+ def reduce_degree(self, axes, multiplicity=None, basis='B'):
1091
+
1092
+ if isinstance(axes, int):
1093
+ axes = [axes]
1094
+
1095
+ if isinstance(multiplicity, int):
1096
+ multiplicity = [multiplicity]
1097
+
1098
+ if multiplicity is None:
1099
+ multiplicity = [self.multiplicity[i] for i in axes]
1100
+
1101
+ v = self._coeff_space
1102
+
1103
+ spaces = list(self.spaces)
1104
+
1105
+ for m, axis in zip(multiplicity, axes):
1106
+ space = spaces[axis]
1107
+
1108
+ reduced_space = SplineSpace(
1109
+ degree = space.degree - 1,
1110
+ pads = space.pads,
1111
+ grid = space.breaks,
1112
+ multiplicity= m,
1113
+ parent_multiplicity=space.multiplicity,
1114
+ periodic = space.periodic,
1115
+ dirichlet = space.dirichlet,
1116
+ basis = basis
1117
+ )
1118
+ spaces[axis] = reduced_space
1119
+
1120
+ npts = [s.nbasis for s in spaces]
1121
+ multiplicity = [s.multiplicity for s in spaces]
1122
+
1123
+ global_starts, global_ends = partition_coefficients(v.cart.domain_decomposition, spaces)
1124
+
1125
+ # create new CartDecomposition
1126
+ red_cart = v.cart.reduce_npts(npts, global_starts, global_ends, shifts=multiplicity)
1127
+
1128
+ # create new TensorFemSpace
1129
+
1130
+ tensor_vec = TensorFemSpace(self._domain_decomposition, *spaces, cart=red_cart, dtype=v.dtype)
1131
+ tensor_vec._interpolation_ready = False
1132
+
1133
+ for key in self._refined_space:
1134
+ if key == tuple(self.ncells):
1135
+ tensor_vec.set_refined_space(key, tensor_vec)
1136
+ else:
1137
+ tensor_vec.set_refined_space(key, self._refined_space[key].reduce_degree(axes, multiplicity, basis))
1138
+ return tensor_vec
1139
+
1140
+ # ...
1141
+ def add_refined_space(self, ncells):
1142
+ """ refine the space with new ncells and add it to the dictionary of refined_space"""
1143
+
1144
+ ncells = tuple(ncells)
1145
+ if ncells in self._refined_space: return
1146
+ if ncells == tuple(self.ncells):
1147
+ self.set_refined_space(ncells, self)
1148
+ return
1149
+
1150
+ spaces = [s.refine(n) for s,n in zip(self.spaces, ncells)]
1151
+ npts = [s.nbasis for s in spaces]
1152
+ domain = self.domain_decomposition
1153
+ new_global_starts = []
1154
+ new_global_ends = []
1155
+ for i in range(domain.ndim):
1156
+ gs = domain.global_element_starts[i]
1157
+ ge = domain.global_element_ends [i]
1158
+ new_global_starts.append([])
1159
+ new_global_ends .append([])
1160
+ for s,e in zip(gs, ge):
1161
+ bs = self.spaces[i].breaks[s]
1162
+ be = self.spaces[i].breaks[e+1]
1163
+ s = spaces[i].breaks.tolist().index(bs)
1164
+ e = spaces[i].breaks.tolist().index(be)
1165
+ new_global_starts[-1].append(s)
1166
+ new_global_ends [-1].append(e-1)
1167
+
1168
+ new_global_starts[-1] = np.array(new_global_starts[-1])
1169
+ new_global_ends [-1] = np.array(new_global_ends [-1])
1170
+
1171
+ new_domain = domain.refine(ncells, new_global_starts, new_global_ends)
1172
+ new_space = TensorFemSpace(new_domain, *spaces, dtype=self._coeff_space.dtype)
1173
+
1174
+ self.set_refined_space(ncells, new_space)
1175
+
1176
+ # ...
1177
+ def create_interface_space(self, axis, ext, cart):
1178
+ """ Create a new interface fem space along a given axis and extremity.
1179
+
1180
+ Parameters
1181
+ ----------
1182
+ axis : int
1183
+ The axis of the new Interface space.
1184
+
1185
+ ext: int
1186
+ The extremity of the new Interface space.
1187
+ the values must be 1 or -1.
1188
+
1189
+ cart: CartDecomposition
1190
+ The cart of the new space, needed in the parallel case.
1191
+ """
1192
+ axis = int(axis)
1193
+ ext = int(ext)
1194
+
1195
+ assert axis < self.ldim
1196
+ assert ext in [-1, 1]
1197
+
1198
+ if cart.is_comm_null or self._interfaces.get((axis, ext), None):
1199
+ return
1200
+
1201
+ spaces = self.spaces
1202
+ coeff_space = self.coeff_space
1203
+
1204
+ coeff_space.set_interface(axis, ext, cart)
1205
+
1206
+ space = TensorFemSpace(self._domain_decomposition, *spaces,
1207
+ coeff_space=coeff_space.interfaces[axis, ext],
1208
+ dtype=coeff_space.dtype)
1209
+
1210
+ self._interfaces[axis, ext] = space
1211
+
1212
+ def get_refined_space(self, ncells):
1213
+ return self._refined_space[tuple(ncells)]
1214
+
1215
+ def set_refined_space(self, ncells, new_space):
1216
+ assert all(nc1==nc2 for nc1,nc2 in zip(ncells, new_space.ncells))
1217
+ self._refined_space[tuple(ncells)] = new_space
1218
+
1219
+ # ...
1220
+ def plot_2d_decomposition(self, mapping=None, *, refine=10, fig=None, ax=None, mpi_root=0):
1221
+ """
1222
+ Plot decomposition of 2D TensorFemSpace w/ mapping to 2D physical space
1223
+
1224
+ Plot the domain decomposition across MPI processes of a 2D
1225
+ TensorFemSpace with a mapping between 2D logical and 2D physical spaces.
1226
+ This function must be called collectively, and only the root process will make
1227
+ the plot. On non-root processes the arguments `fig` and `ax` must be None.
1228
+
1229
+ Parameters
1230
+ ----------
1231
+ mapping : BasicCallableMapping
1232
+ Mapping from (eta1, eta2) to (x1, x2).
1233
+
1234
+ refine : int, default=10
1235
+ Cell refinement along the logical dimensions eta1 and eta2.
1236
+
1237
+ fig : plt.Figure, optional
1238
+ Figure where the plot should be made. Must be None on non-root processes.
1239
+
1240
+ ax : plt.Axes, optional
1241
+ Axes where the plot should be made. Must be None on non-root processes.
1242
+
1243
+ mpi_root: int, default=0
1244
+ The rank of the MPI root process which should create the plot.
1245
+
1246
+ Returns
1247
+ -------
1248
+ plt.Figure
1249
+ Figure where the plot was made. Coincides with `fig` if provided.
1250
+ """
1251
+ import matplotlib.pyplot as plt
1252
+ from matplotlib.patches import Polygon, Patch
1253
+ from sympde.topology.mapping import BasicCallableMapping
1254
+ from feectools.utilities.utils import refine_array_1d
1255
+
1256
+ # Sanity check
1257
+ assert self.ldim == 2, "Function only works in 2D"
1258
+
1259
+ # Check mapping
1260
+ if mapping is None:
1261
+ mapping = lambda eta: eta
1262
+ else:
1263
+ assert isinstance(mapping, BasicCallableMapping)
1264
+ assert mapping.ldim == 2, "Domain of argument `mapping` must be 2D"
1265
+ assert mapping.pdim == 2, "Codomain of argument `mapping` must be 2D"
1266
+
1267
+ # Check refine argument
1268
+ assert isinstance(refine, int), f"Argument `refine` must be int, got {type(refine)} instead"
1269
+ assert refine >= 1, f"Argument `refine` must be >= 1, got {refine} instead"
1270
+
1271
+ # Extract information about MPI communicator
1272
+ mpi_comm = self.coeff_space.cart.comm
1273
+ mpi_rank = mpi_comm.rank
1274
+ mpi_size = mpi_comm.size
1275
+
1276
+ # Check mpi_root argument
1277
+ assert isinstance(mpi_root, int), f"Argument `mpi_root` must be int, got {type(mpi_root)} instead"
1278
+ assert mpi_root >= 0, f"Argument `mpi_root` must be >= 0, got {mpi_root} instead"
1279
+ assert mpi_root < mpi_size, f"Argument `mpi_root` must be smaller than communicator size ({mpi_size}), got {mpi_root} instead"
1280
+
1281
+ # Check fig and ax arguments
1282
+ if mpi_rank == mpi_root:
1283
+ assert isinstance(fig, plt.Figure) or fig is None, f"Argument `fig` must be matplotlib Figure, got {type(fig)} instead"
1284
+ assert isinstance(ax, plt.Axes) or ax is None, f"Argument `ax` must be matplotlib Axes, got {type(ax)} instead"
1285
+ else:
1286
+ assert fig is None, f"Argument `fig` must be None on non-root process with rank {mpi_rank}"
1287
+ assert ax is None, f"Argument `ax` must be None on non-root process with rank {mpi_rank}"
1288
+
1289
+ N = refine
1290
+ V1, V2 = self.spaces
1291
+
1292
+ # Local grid, refined
1293
+ [sk1, sk2], [ek1, ek2] = self.local_domain
1294
+ eta1 = refine_array_1d(V1.breaks[sk1:ek1+2], N)
1295
+ eta2 = refine_array_1d(V2.breaks[sk2:ek2+2], N)
1296
+ pcoords = np.array([[mapping(e1, e2) for e2 in eta2] for e1 in eta1])
1297
+
1298
+ # Local domain as Matplotlib polygonal patch
1299
+ AB = pcoords[ :, 0, :] # eta2 = min
1300
+ BC = pcoords[ -1, :, :] # eta1 = max
1301
+ CD = pcoords[::-1, -1, :] # eta2 = max (points must be reversed)
1302
+ DA = pcoords[ 0, ::-1, :] # eta1 = min (points must be reversed)
1303
+ xy = np.concatenate([AB, BC, CD, DA], axis=0)
1304
+ poly = Polygon(xy, edgecolor='None')
1305
+
1306
+ # Gather polygons on master process
1307
+ polys = mpi_comm.gather(poly, root=mpi_root)
1308
+
1309
+ # Gather (s1, s2, e1, e2) on root
1310
+ if mpi_rank == mpi_root:
1311
+ s1_all = np.empty(mpi_size, dtype=int)
1312
+ s2_all = np.empty(mpi_size, dtype=int)
1313
+ e1_all = np.empty(mpi_size, dtype=int)
1314
+ e2_all = np.empty(mpi_size, dtype=int)
1315
+ else:
1316
+ s1_all = None
1317
+ s2_all = None
1318
+ e1_all = None
1319
+ e2_all = None
1320
+
1321
+ mpi_comm.Gather(sk1 * N, s1_all, root=mpi_root)
1322
+ mpi_comm.Gather(sk2 * N, s2_all, root=mpi_root)
1323
+ mpi_comm.Gather((ek1 + 1) * N, e1_all, root=mpi_root)
1324
+ mpi_comm.Gather((ek2 + 1) * N, e2_all, root=mpi_root)
1325
+
1326
+ # Gather pcoords on root
1327
+ # TODO: use Gatherv, and NumPy arrays as buffers
1328
+ gathered_pcoords = mpi_comm.gather(pcoords, root=mpi_root)
1329
+
1330
+ #-------------------------------
1331
+ # Non-master processes stop here
1332
+ if mpi_rank != mpi_root:
1333
+ return
1334
+ #-------------------------------
1335
+
1336
+ # Reconstruct global grid (refined) on root process
1337
+ global_shape = ((V1.breaks.size - 1) * N + 1,
1338
+ (V2.breaks.size - 1) * N + 1,
1339
+ 2)
1340
+ pcoords_global = np.empty(global_shape)
1341
+
1342
+ for rank in range(mpi_comm.size):
1343
+ s1 = s1_all[rank]
1344
+ e1 = e1_all[rank]
1345
+ s2 = s2_all[rank]
1346
+ e2 = e2_all[rank]
1347
+ pcoords_global[s1:e1+1, s2:e2+1, :] = gathered_pcoords[rank]
1348
+
1349
+ xx = pcoords_global[:, :, 0]
1350
+ yy = pcoords_global[:, :, 1]
1351
+
1352
+ # If fig or ax are given, get one from the other. Otherwise create new ones
1353
+ if fig and ax:
1354
+ assert ax in fig.axes, "Argument `ax` must be in `fig.axes`"
1355
+ elif fig:
1356
+ ax = fig.gca()
1357
+ elif ax:
1358
+ fig = ax.figure
1359
+ else:
1360
+ fig, ax = plt.subplots(1, 1)
1361
+
1362
+ # Plot decomposed domain
1363
+ colors = itertools.cycle(plt.rcParams['axes.prop_cycle'].by_key()['color'])
1364
+ handles = []
1365
+ for i, (poly, color) in enumerate(zip(polys, colors)):
1366
+ # Add patch
1367
+ poly.set_facecolor(color)
1368
+ ax.add_patch(poly)
1369
+ # Create legend entry
1370
+ handle = Patch(color=color, label='Rank {}'.format(i))
1371
+ handles.append(handle)
1372
+
1373
+ ax.set_xlabel(r'$x$', rotation='horizontal')
1374
+ ax.set_ylabel(r'$y$', rotation='horizontal')
1375
+ ax.set_title ('Domain decomposition')
1376
+ ax.plot(xx[:,::N] , yy[:,::N] , 'k')
1377
+ ax.plot(xx[::N,:].T, yy[::N,:].T, 'k')
1378
+ ax.set_aspect('equal')
1379
+ ax.legend(handles=handles, bbox_to_anchor=(1.05, 1), loc=2)
1380
+ fig.tight_layout()
1381
+
1382
+ return fig
1383
+
1384
+ # ...
1385
+ def __str__(self):
1386
+ """Pretty printing"""
1387
+ txt = '\n'
1388
+ txt += '> ldim :: {ldim}\n'.format(ldim=self.ldim)
1389
+ txt += '> total nbasis :: {dim}\n'.format(dim=self.nbasis)
1390
+
1391
+ dims = ', '.join(str(V.nbasis) for V in self.spaces)
1392
+ txt += '> nbasis :: ({dims})\n'.format(dims=dims)
1393
+ return txt