feectools 0.2.0__tar.gz → 0.3.1__tar.gz

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 (153) hide show
  1. {feectools-0.2.0/feectools.egg-info → feectools-0.3.1}/PKG-INFO +2 -2
  2. {feectools-0.2.0 → feectools-0.3.1}/feectools/api/essential_bc.py +3 -0
  3. {feectools-0.2.0 → feectools-0.3.1}/feectools/core/bsplines.py +50 -32
  4. {feectools-0.2.0 → feectools-0.3.1}/feectools/core/tests/test_bsplines.py +1 -1
  5. {feectools-0.2.0 → feectools-0.3.1}/feectools/core/tests/test_bsplines_kernel.py +9 -8
  6. {feectools-0.2.0 → feectools-0.3.1}/feectools/core/tests/test_bsplines_pyccel.py +13 -4
  7. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/blocking_data_exchanger.py +9 -0
  8. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/cart.py +102 -31
  9. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/interface_data_exchanger.py +5 -0
  10. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/mpi.py +3 -5
  11. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/nonblocking_data_exchanger.py +8 -0
  12. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/partition.py +1 -2
  13. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/petsc.py +1 -1
  14. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/tests/test_cart_1d.py +2 -2
  15. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/tests/test_cart_2d.py +4 -3
  16. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/tests/test_cart_3d.py +6 -4
  17. feectools-0.3.1/feectools/ddm/tests/test_coarsen.py +67 -0
  18. feectools-0.3.1/feectools/ddm/tests/test_device_binding.py +18 -0
  19. {feectools-0.2.0 → feectools-0.3.1}/feectools/feec/derivatives.py +17 -5
  20. {feectools-0.2.0 → feectools-0.3.1}/feectools/feec/global_geometric_projectors.py +75 -59
  21. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/partitioning.py +7 -6
  22. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/splines.py +3 -12
  23. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tensor.py +19 -2
  24. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/analytical_profiles_1d.py +2 -0
  25. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/utilities.py +4 -0
  26. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/direct_solvers.py +23 -24
  27. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/fft.py +8 -1
  28. feectools-0.3.1/feectools/linalg/kernels/__init__.py +8 -0
  29. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_1d/__init__.py +6 -0
  30. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_1d/stencil_axpy_1d_cuda.cu +20 -0
  31. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_1d/stencil_axpy_1d_kernels.py +26 -0
  32. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_2d/__init__.py +6 -0
  33. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_2d/stencil_axpy_2d_cuda.cu +23 -0
  34. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_2d/stencil_axpy_2d_kernels.py +27 -0
  35. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_3d/__init__.py +6 -0
  36. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_3d/stencil_axpy_3d_cuda.cu +25 -0
  37. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_3d/stencil_axpy_3d_kernels.py +28 -0
  38. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_1d/__init__.py +6 -0
  39. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_1d/stencil_dot_1d_cuda.cu +40 -0
  40. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_1d/stencil_dot_1d_kernels.py +35 -0
  41. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_2d/__init__.py +6 -0
  42. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_2d/stencil_dot_2d_cuda.cu +42 -0
  43. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_2d/stencil_dot_2d_kernels.py +99 -0
  44. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_3d/__init__.py +6 -0
  45. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_3d/stencil_dot_3d_cuda.cu +64 -0
  46. feectools-0.2.0/feectools/linalg/stencil_dot_kernels.py → feectools-0.3.1/feectools/linalg/kernels/stencil_dot_3d/stencil_dot_3d_kernels.py +14 -125
  47. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_1d/__init__.py +6 -0
  48. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_1d/stencil_inner_1d_cuda.cu +41 -0
  49. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_1d/stencil_inner_1d_kernels.py +36 -0
  50. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_2d/__init__.py +6 -0
  51. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_2d/stencil_inner_2d_cuda.cu +44 -0
  52. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_2d/stencil_inner_2d_kernels.py +40 -0
  53. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_3d/__init__.py +6 -0
  54. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_3d/stencil_inner_3d_cuda.cu +46 -0
  55. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_3d/stencil_inner_3d_kernels.py +44 -0
  56. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_1d/__init__.py +6 -0
  57. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_1d/stencil_transpose_1d_cuda.cu +40 -0
  58. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_1d/stencil_transpose_1d_kernels.py +37 -0
  59. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_2d/__init__.py +6 -0
  60. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_2d/stencil_transpose_2d_cuda.cu +47 -0
  61. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_2d/stencil_transpose_2d_kernels.py +116 -0
  62. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_3d/__init__.py +6 -0
  63. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_3d/stencil_transpose_3d_cuda.cu +76 -0
  64. feectools-0.2.0/feectools/linalg/stencil_transpose_kernels.py → feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_3d/stencil_transpose_3d_kernels.py +20 -140
  65. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/kron.py +31 -16
  66. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/solvers.py +8 -4
  67. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/sparse.py +12 -1
  68. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/stencil.py +142 -135
  69. feectools-0.3.1/feectools/linalg/tests/cuda_parity_cases.py +54 -0
  70. feectools-0.3.1/feectools/linalg/tests/kernel_test_args.py +100 -0
  71. feectools-0.3.1/feectools/linalg/tests/test_axpy_ghost_sync.py +31 -0
  72. feectools-0.3.1/feectools/linalg/tests/test_cuda_emulation.py +38 -0
  73. feectools-0.3.1/feectools/linalg/tests/test_cuda_parity.py +99 -0
  74. feectools-0.3.1/feectools/linalg/tests/test_device_matvec.py +206 -0
  75. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_fft.py +2 -1
  76. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_kron_stencil_matrix.py +2 -1
  77. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_linalg.py +31 -26
  78. feectools-0.3.1/feectools/linalg/tests/test_mpi_device.py +251 -0
  79. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_stencil_interface_matrix.py +8 -6
  80. {feectools-0.2.0 → feectools-0.3.1}/feectools/utilities/utils.py +9 -15
  81. {feectools-0.2.0 → feectools-0.3.1/feectools.egg-info}/PKG-INFO +2 -2
  82. {feectools-0.2.0 → feectools-0.3.1}/feectools.egg-info/SOURCES.txt +45 -4
  83. {feectools-0.2.0 → feectools-0.3.1}/feectools.egg-info/requires.txt +1 -1
  84. {feectools-0.2.0 → feectools-0.3.1}/pyproject.toml +3 -3
  85. feectools-0.2.0/feectools/linalg/kernels/axpy_kernels.py +0 -57
  86. feectools-0.2.0/feectools/linalg/kernels/inner_kernels.py +0 -100
  87. feectools-0.2.0/feectools/utilities/__init__.py +0 -0
  88. {feectools-0.2.0 → feectools-0.3.1}/AUTHORS +0 -0
  89. {feectools-0.2.0 → feectools-0.3.1}/LICENSE +0 -0
  90. {feectools-0.2.0 → feectools-0.3.1}/README.md +0 -0
  91. {feectools-0.2.0 → feectools-0.3.1}/feectools/__init__.py +0 -0
  92. {feectools-0.2.0 → feectools-0.3.1}/feectools/accelerate/__init__.py +0 -0
  93. {feectools-0.2.0 → feectools-0.3.1}/feectools/accelerate/accelerate.py +0 -0
  94. {feectools-0.2.0 → feectools-0.3.1}/feectools/accelerate/compile_psydac.mk +0 -0
  95. {feectools-0.2.0 → feectools-0.3.1}/feectools/api/__init__.py +0 -0
  96. {feectools-0.2.0 → feectools-0.3.1}/feectools/api/fem_bilinear_form.py +0 -0
  97. {feectools-0.2.0 → feectools-0.3.1}/feectools/api/fem_common.py +0 -0
  98. {feectools-0.2.0 → feectools-0.3.1}/feectools/api/fem_sum_form.py +0 -0
  99. {feectools-0.2.0 → feectools-0.3.1}/feectools/api/settings.py +0 -0
  100. {feectools-0.2.0 → feectools-0.3.1}/feectools/core/__init__.py +0 -0
  101. {feectools-0.2.0 → feectools-0.3.1}/feectools/core/bsplines_kernels.py +0 -0
  102. {feectools-0.2.0 → feectools-0.3.1}/feectools/core/field_evaluation_kernels.py +0 -0
  103. {feectools-0.2.0 → feectools-0.3.1}/feectools/core/tests/__init__.py +0 -0
  104. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/__init__.py +0 -0
  105. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/basic.py +0 -0
  106. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/tests/__init__.py +0 -0
  107. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/tests/test_multicart_2d.py +0 -0
  108. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/tests/test_partition.py +0 -0
  109. {feectools-0.2.0 → feectools-0.3.1}/feectools/ddm/utilities.py +0 -0
  110. {feectools-0.2.0 → feectools-0.3.1}/feectools/feec/__init__.py +0 -0
  111. {feectools-0.2.0 → feectools-0.3.1}/feectools/feec/dof_kernels.py +0 -0
  112. {feectools-0.2.0 → feectools-0.3.1}/feectools/feec/hodge.py +0 -0
  113. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/__init__.py +0 -0
  114. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/basic.py +0 -0
  115. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/grid.py +0 -0
  116. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/lst_preconditioner.py +0 -0
  117. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/projectors.py +0 -0
  118. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/__init__.py +0 -0
  119. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/analytical_profiles_base.py +0 -0
  120. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/splines_error_bounds.py +0 -0
  121. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/test_dirichlet_projectors.py +0 -0
  122. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/test_spline_histopolation.py +0 -0
  123. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/test_spline_interpolation.py +0 -0
  124. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/test_splines.py +0 -0
  125. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/test_splines_par.py +0 -0
  126. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/test_tensor.py +0 -0
  127. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/tests/test_vector_spaces.py +0 -0
  128. {feectools-0.2.0 → feectools-0.3.1}/feectools/fem/vector.py +0 -0
  129. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/__init__.py +0 -0
  130. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/basic.py +0 -0
  131. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/block.py +0 -0
  132. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/kernels/matvec_kernels.py +0 -0
  133. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/kernels/stencil2IJV_kernels.py +0 -0
  134. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/kernels/stencil2coo_kernels.py +0 -0
  135. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/kernels/transpose_kernels.py +0 -0
  136. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/memory.py +0 -0
  137. {feectools-0.2.0/feectools/linalg/kernels → feectools-0.3.1/feectools/linalg/tests}/__init__.py +0 -0
  138. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_block.py +0 -0
  139. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_matrix_free.py +0 -0
  140. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_solvers.py +0 -0
  141. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_stencil_vector.py +0 -0
  142. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_stencil_vector_space.py +0 -0
  143. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/test_toarray.py +0 -0
  144. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/tests/utilities.py +0 -0
  145. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/topetsc.py +0 -0
  146. {feectools-0.2.0 → feectools-0.3.1}/feectools/linalg/utilities.py +0 -0
  147. {feectools-0.2.0/feectools/linalg/tests → feectools-0.3.1/feectools/utilities}/__init__.py +0 -0
  148. {feectools-0.2.0 → feectools-0.3.1}/feectools/utilities/quadratures.py +0 -0
  149. {feectools-0.2.0 → feectools-0.3.1}/feectools/version.py +0 -0
  150. {feectools-0.2.0 → feectools-0.3.1}/feectools.egg-info/dependency_links.txt +0 -0
  151. {feectools-0.2.0 → feectools-0.3.1}/feectools.egg-info/entry_points.txt +0 -0
  152. {feectools-0.2.0 → feectools-0.3.1}/feectools.egg-info/top_level.txt +0 -0
  153. {feectools-0.2.0 → feectools-0.3.1}/setup.cfg +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: feectools
3
- Version: 0.2.0
3
+ Version: 0.3.1
4
4
  Summary: Slimmed-down fork of Psydac (https://github.com/pyccel/psydac) with less functionality and fewer dependencies.
5
5
  Author-email: Psydac development team <psydac@googlegroups.com>
6
6
  Maintainer-email: Stefan Possanner <stefan.possanner@ipp.mpg.de>, Max Lindqvist <max.lindqvist@ipp.mpg.de>, Yaman Güçlü <yaman.guclu@gmail.com>, Martin Campos Pinto <martin.campos-pinto@ipp.mpg.de>, Ahmed Ratnani <ratnaniahmed@gmail.com>
@@ -42,7 +42,7 @@ Requires-Dist: matplotlib
42
42
  Requires-Dist: pyyaml>=5.1
43
43
  Requires-Dist: packaging
44
44
  Requires-Dist: pyevtk
45
- Requires-Dist: cunumpy
45
+ Requires-Dist: cunumpy<0.6,>=0.5.0
46
46
  Requires-Dist: pyccel>=2.1.0
47
47
  Requires-Dist: h5py
48
48
  Requires-Dist: tblib
@@ -72,6 +72,9 @@ def apply_essential_bc_stencil(a, *, axis, ext, order, identity=False):
72
72
  if isinstance(a, StencilVector):
73
73
  V = a.space
74
74
  n = V.ndim
75
+ # Boundary entries may be ghost entries of neighbouring processes, which
76
+ # all call this function: their ghost regions are no longer up to date.
77
+ a.ghost_regions_in_sync = False
75
78
  elif isinstance(a, StencilMatrix):
76
79
  V = a.codomain
77
80
  n = V.ndim * 2
@@ -16,6 +16,7 @@ References:
16
16
 
17
17
  """
18
18
  import cunumpy as xp
19
+ from cunumpy.kernels import PyccelKernel
19
20
  from cunumpy.xp import array_backend
20
21
  import numpy as np
21
22
 
@@ -38,6 +39,27 @@ from feectools.core.bsplines_kernels import (find_span_p,
38
39
  cell_index_p,
39
40
  basis_ders_on_irregular_grid_p)
40
41
 
42
+ # Kernels generated by Pyccel only understand NumPy arrays; wrap them so they
43
+ # can also be called with CuPy arrays (see cunumpy.kernels.PyccelKernel).
44
+ find_span_p = PyccelKernel(find_span_p)
45
+ find_spans_p = PyccelKernel(find_spans_p)
46
+ basis_funs_p = PyccelKernel(basis_funs_p)
47
+ basis_funs_array_p = PyccelKernel(basis_funs_array_p)
48
+ basis_funs_1st_der_p = PyccelKernel(basis_funs_1st_der_p)
49
+ basis_funs_all_ders_p = PyccelKernel(basis_funs_all_ders_p)
50
+ collocation_matrix_p = PyccelKernel(collocation_matrix_p)
51
+ histopolation_matrix_p = PyccelKernel(histopolation_matrix_p)
52
+ greville_p = PyccelKernel(greville_p)
53
+ breakpoints_p = PyccelKernel(breakpoints_p)
54
+ elements_spans_p = PyccelKernel(elements_spans_p)
55
+ make_knots_p = PyccelKernel(make_knots_p)
56
+ elevate_knots_p = PyccelKernel(elevate_knots_p)
57
+ quadrature_grid_p = PyccelKernel(quadrature_grid_p)
58
+ basis_ders_on_quad_grid_p = PyccelKernel(basis_ders_on_quad_grid_p)
59
+ basis_integrals_p = PyccelKernel(basis_integrals_p)
60
+ cell_index_p = PyccelKernel(cell_index_p)
61
+ basis_ders_on_irregular_grid_p = PyccelKernel(basis_ders_on_irregular_grid_p)
62
+
41
63
  __all__ = ('find_span',
42
64
  'find_spans',
43
65
  'basis_funs',
@@ -84,7 +106,7 @@ def find_span(knots, degree, x):
84
106
  Knot span index.
85
107
  """
86
108
  x = float(x)
87
- knots = xp.ascontiguousarray(knots, dtype=float)
109
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
88
110
  return find_span_p(knots, degree, x)
89
111
 
90
112
  #==============================================================================
@@ -116,8 +138,8 @@ def find_spans(knots, degree, x, out=None):
116
138
  spans : array of ints
117
139
  Knots span indexes.
118
140
  """
119
- knots = xp.ascontiguousarray(knots, dtype=float)
120
- x = xp.ascontiguousarray(x, dtype=float)
141
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
142
+ x = xp.ascontiguousarray(xp.asarray(x), dtype=float)
121
143
  if out is None:
122
144
  out = xp.zeros_like(x, dtype=int)
123
145
  else:
@@ -155,7 +177,7 @@ def basis_funs(knots, degree, x, span, out=None):
155
177
  1D array containing the values of ``degree + 1`` non-zero
156
178
  Bsplines at location ``x``.
157
179
  """
158
- knots = xp.ascontiguousarray(knots, dtype=float)
180
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
159
181
  # Get native float
160
182
  x = float(x)
161
183
  if out is None:
@@ -193,8 +215,8 @@ def basis_funs_array(knots, degree, span, x, out=None):
193
215
  2D array of shape ``(len(x), degree + 1)`` containing the values of ``degree + 1`` non-zero
194
216
  Bsplines at each location in ``x``.
195
217
  """
196
- knots = xp.ascontiguousarray(knots, dtype=float)
197
- x = xp.ascontiguousarray(x, dtype=float)
218
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
219
+ x = xp.ascontiguousarray(xp.asarray(x), dtype=float)
198
220
  if out is None:
199
221
  out = xp.zeros(x.shape + (degree + 1,), dtype=float)
200
222
  else:
@@ -240,7 +262,7 @@ def basis_funs_1st_der(knots, degree, x, span, out=None):
240
262
  ----------
241
263
  .. [2] SELALIB, Semi-Lagrangian Library. http://selalib.gforge.inria.fr
242
264
  """
243
- knots = xp.ascontiguousarray(knots, dtype=float)
265
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
244
266
  # Get native float to work on windows
245
267
  x = float(x)
246
268
  if out is None:
@@ -291,7 +313,7 @@ def basis_funs_all_ders(knots, degree, x, span, n, normalization='B', out=None):
291
313
  ders[i,j] = (d/dx)^i B_k(x) with k=(span-degree+j),
292
314
  for 0 <= i <= n and 0 <= j <= degree+1.
293
315
  """
294
- knots = xp.ascontiguousarray(knots, dtype=float)
316
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
295
317
  # Get native float to work on windows
296
318
  x = float(x)
297
319
  if out is None:
@@ -346,8 +368,8 @@ def collocation_matrix(knots, degree, periodic, normalization, xgrid, out=None,
346
368
  if xgrid.size == 1:
347
369
  return xp.ones((1, 1), dtype=float)
348
370
 
349
- knots = xp.ascontiguousarray(knots, dtype=float)
350
- xgrid = xp.ascontiguousarray(xgrid, dtype=float)
371
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
372
+ xgrid = xp.ascontiguousarray(xp.asarray(xgrid), dtype=float)
351
373
  if out is None:
352
374
  nb = len(knots) - degree - 1
353
375
  if periodic:
@@ -430,8 +452,8 @@ def histopolation_matrix(knots, degree, periodic, normalization, xgrid, multipli
430
452
  if not xp.all(xp.diff(xgrid) > 0):
431
453
  raise ValueError("Grid points must be ordered, with no repetitions: {}".format(xgrid))
432
454
 
433
- knots = xp.ascontiguousarray(knots, dtype=float)
434
- xgrid = xp.ascontiguousarray(xgrid, dtype=float)
455
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
456
+ xgrid = xp.ascontiguousarray(xp.asarray(xgrid), dtype=float)
435
457
  elevated_knots = elevate_knots(knots, degree, periodic, multiplicity=multiplicity)
436
458
 
437
459
  normalization = normalization == "M"
@@ -477,7 +499,7 @@ def breakpoints(knots, degree, tol=1e-15, out=None):
477
499
  breaks : numpy.ndarray (1D)
478
500
  Abscissas of all breakpoints.
479
501
  """
480
- knots = xp.ascontiguousarray(knots, dtype=float)
502
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
481
503
  if out is None:
482
504
  out = xp.zeros(len(knots), dtype=float)
483
505
  else:
@@ -518,8 +540,7 @@ def greville(knots, degree, periodic, out=None, multiplicity=1):
518
540
  # Greville points are index arrays, keep on NumPy
519
541
  if isinstance(knots, (list, tuple)):
520
542
  knots = np.asarray(knots, dtype=float)
521
- if hasattr(knots, 'get'):
522
- knots = knots.get() # Convert CuPy to NumPy
543
+ knots = xp.to_numpy(knots)
523
544
  knots = np.ascontiguousarray(knots, dtype=float)
524
545
  if out is None:
525
546
  n = len(knots) - 2 * degree - 2 + multiplicity if periodic else len(knots) - degree - 1
@@ -572,7 +593,7 @@ def elements_spans(knots, degree, out=None):
572
593
  spans = xp.searchsorted( knots, breaks[:-1], side='right' ) - 1
573
594
 
574
595
  """
575
- knots = xp.ascontiguousarray(knots, dtype=float)
596
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
576
597
  if out is None:
577
598
  out = np.zeros(len(knots), dtype=xp.int64)
578
599
  else:
@@ -624,7 +645,7 @@ def make_knots(breaks, degree, periodic, multiplicity=1, out=None):
624
645
  # Consistency checks
625
646
  assert len(breaks) > 1
626
647
  # Convert to numpy for comparison since assertion needs Python bool
627
- breaks_np = breaks.get() if hasattr(breaks, 'get') else breaks
648
+ breaks_np = xp.to_numpy(breaks)
628
649
  if isinstance(breaks_np, (list, tuple)):
629
650
  breaks_np = np.asarray(breaks_np)
630
651
  assert all( np.diff(breaks_np) > 0 )
@@ -638,8 +659,7 @@ def make_knots(breaks, degree, periodic, multiplicity=1, out=None):
638
659
 
639
660
  # Keep breaks on NumPy for initialization - knots are index arrays needed for CPU operations
640
661
  breaks = np.asarray(breaks, dtype=float) if isinstance(breaks, (list, tuple)) else breaks
641
- if hasattr(breaks, 'get'):
642
- breaks = breaks.get() # Convert CuPy to NumPy
662
+ breaks = xp.to_numpy(breaks)
643
663
  breaks = np.ascontiguousarray(breaks, dtype=float)
644
664
  if out is None:
645
665
  # Knots are index arrays, keep them on NumPy
@@ -693,8 +713,7 @@ def elevate_knots(knots, degree, periodic, multiplicity=1, tol=1e-15, out=None):
693
713
  multiplicity = int(multiplicity)
694
714
  if isinstance(knots, (list, tuple)):
695
715
  knots = np.asarray(knots, dtype=float)
696
- if hasattr(knots, 'get'):
697
- knots = knots.get() # Convert CuPy to NumPy
716
+ knots = xp.to_numpy(knots)
698
717
  knots = np.ascontiguousarray(knots, dtype=float)
699
718
  if out is None:
700
719
  if periodic:
@@ -771,14 +790,13 @@ def quadrature_grid(breaks, quad_rule_x, quad_rule_w):
771
790
  assert max(quad_rule_x) <= +1
772
791
 
773
792
  # Convert breaks to numpy if CuPy (breaks/grids should stay on CPU)
774
- if hasattr(breaks, 'get'):
775
- breaks = breaks.get()
793
+ breaks = xp.to_numpy(breaks)
776
794
  breaks = np.ascontiguousarray(breaks, dtype=float)
777
795
 
778
796
  if array_backend.backend == "cupy":
779
797
  # Convert CuPy arrays to NumPy
780
- quad_rule_x = quad_rule_x.get() if hasattr(quad_rule_x, 'get') else quad_rule_x
781
- quad_rule_w = quad_rule_w.get() if hasattr(quad_rule_w, 'get') else quad_rule_w
798
+ quad_rule_x = xp.to_numpy(quad_rule_x)
799
+ quad_rule_w = xp.to_numpy(quad_rule_w)
782
800
 
783
801
  quad_rule_x = np.ascontiguousarray(quad_rule_x, dtype=float)
784
802
  quad_rule_w = np.ascontiguousarray(quad_rule_w, dtype=float)
@@ -848,8 +866,8 @@ def basis_ders_on_quad_grid(knots, degree, quad_grid, nders, normalization, offs
848
866
  """
849
867
  offset = int(offset)
850
868
  ne, nq = quad_grid.shape
851
- knots = xp.ascontiguousarray(knots, dtype=float)
852
- quad_grid = xp.ascontiguousarray(quad_grid, dtype=float)
869
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
870
+ quad_grid = xp.ascontiguousarray(xp.asarray(quad_grid), dtype=float)
853
871
  if out is None:
854
872
  out = xp.zeros((ne, degree + 1, nders + 1, nq), dtype=float)
855
873
  else:
@@ -892,7 +910,7 @@ def basis_integrals(knots, degree, out=None):
892
910
  to (len(knots)-degree-1). In the periodic case the last (degree) values in
893
911
  the array are redundant, as they are a copy of the first (degree) values.
894
912
  """
895
- knots = xp.ascontiguousarray(knots, dtype=float)
913
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
896
914
  if out is None:
897
915
  out = xp.zeros(len(knots) - degree - 1, dtype=float)
898
916
  else:
@@ -934,8 +952,8 @@ def cell_index(breaks, i_grid, tol=1e-15, out=None):
934
952
  ``cell_index[i]`` is the index of the cell in which
935
953
  ``i_grid[i]`` belong.
936
954
  """
937
- breaks = xp.ascontiguousarray(breaks, dtype=float)
938
- i_grid = xp.ascontiguousarray(i_grid, dtype=float)
955
+ breaks = xp.ascontiguousarray(xp.asarray(breaks), dtype=float)
956
+ i_grid = xp.ascontiguousarray(xp.asarray(i_grid), dtype=float)
939
957
  if out is None:
940
958
  out = np.zeros_like(i_grid, dtype=xp.int64)
941
959
  else:
@@ -990,8 +1008,8 @@ def basis_ders_on_irregular_grid(knots, degree, i_grid, cell_index, nders, norma
990
1008
  . il: local basis function (0 <= il <= degree)
991
1009
  . id: derivative (0 <= id <= nders )
992
1010
  """
993
- knots = xp.ascontiguousarray(knots, dtype=float)
994
- i_grid = xp.ascontiguousarray(i_grid, dtype=float)
1011
+ knots = xp.ascontiguousarray(xp.asarray(knots), dtype=float)
1012
+ i_grid = xp.ascontiguousarray(xp.asarray(i_grid), dtype=float)
995
1013
  if out is None:
996
1014
  nx = i_grid.shape[0]
997
1015
  out = xp.zeros((nx, degree + 1, nders + 1), dtype=float)
@@ -159,7 +159,7 @@ def test_histopolation_matrix(lims, nc, p, periodic, tol=1e-13):
159
159
  def test_cell_index(i_grid, expected):
160
160
  breaks = xp.array([0. , 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.])
161
161
  out = cell_index(breaks, xp.asarray(i_grid))
162
- assert xp.array_equal(expected, out)
162
+ assert xp.array_equal(xp.asarray(expected), out)
163
163
 
164
164
  #==============================================================================
165
165
  # SCRIPT FUNCTIONALITY: PLOT BASIS FUNCTIONS
@@ -2,14 +2,16 @@
2
2
 
3
3
  import pytest
4
4
  import cunumpy as xp
5
+ import numpy as np
5
6
 
6
7
 
7
8
  from feectools.core.bsplines_kernels import cell_index_p
8
9
 
9
10
  def test_cell_index_p():
10
- breaks = xp.array([0. , 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.])
11
- breaks = xp.ascontiguousarray(breaks, dtype=float)
12
- out = xp.zeros_like(breaks, dtype=xp.int64)
11
+ # This directly tests the raw Pyccel kernel, which intentionally accepts
12
+ # NumPy host arrays only; CuPy coverage belongs to the public wrapper.
13
+ breaks = np.ascontiguousarray([0. , 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.], dtype=float)
14
+ out = np.zeros_like(breaks, dtype=np.int64)
13
15
  tol = 1e-15
14
16
 
15
17
  # limit case: code should decide wether point is in or out, not fall in infinite loop
@@ -26,15 +28,14 @@ def test_cell_index_p():
26
28
  assert status == expected_status
27
29
 
28
30
  # checking that the values match those of searchsorted (-1) for arbitrary grid points
29
- i_grid = xp.array([0.14320482, 0.86569833, 0.77775327, 0.00895956, 0.074629 ,
31
+ i_grid = np.array([0.14320482, 0.86569833, 0.77775327, 0.00895956, 0.074629 ,
30
32
  0.45682646, 0.5384352 , 0.20915311, 0.73121977, 0.01057414,
31
33
  0.33756086, 0.17839759, 0.14023414, 0.09846206, 0.79970392,
32
34
  0.65330406, 0.82716552, 0.24185731, 0.24054685, 0.72466651,
33
35
  0.69125033, 0.3136558 , 0.64794089, 0.47975527, 0.99802844,
34
36
  0.64402598, 0.41263526, 0.28178414, 0.57274384, 0.73218562])
35
- out = xp.zeros_like(i_grid, dtype=xp.int64)
37
+ out = np.zeros_like(i_grid, dtype=np.int64)
36
38
  status = cell_index_p(breaks, i_grid, tol, out)
37
39
  assert status == 0
38
- nps = xp.searchsorted(breaks, i_grid)-1
39
- assert xp.allclose(out, nps)
40
-
40
+ nps = np.searchsorted(breaks, i_grid)-1
41
+ assert np.allclose(out, nps)
@@ -187,8 +187,9 @@ def basis_funs_all_ders_true(knots, degree, x, span, n, normalization='B'):
187
187
 
188
188
  # Normalization to get M-Splines
189
189
  if normalization == 'M':
190
- ders *= [(degree + 1) / (knots[i + degree + 1] - knots[i]) \
191
- for i in range(span - degree, span + 1)]
190
+ scaling = xp.asarray([(degree + 1) / (knots[i + degree + 1] - knots[i])
191
+ for i in range(span - degree, span + 1)])
192
+ ders *= scaling
192
193
  return ders
193
194
 
194
195
  #==============================================================================
@@ -221,7 +222,15 @@ def collocation_matrix_true(knots, degree, periodic, normalization, xgrid):
221
222
  for i,x in enumerate( xgrid ):
222
223
  span = find_span_true( knots, degree, x )
223
224
  basis = basis_funs_true( knots, degree, x, span )
224
- mat[i,js(span)] = normalize(basis, span)
225
+ values = normalize(basis, span)
226
+ if periodic:
227
+ # NumPy and CuPy differ for indexed assignment with repeated
228
+ # indices (which occurs when nb <= degree). The production
229
+ # kernel assigns in loop order, so make the reference explicit.
230
+ for j, value in zip(js(span), values):
231
+ mat[i, j] = value
232
+ else:
233
+ mat[i, js(span)] = values
225
234
 
226
235
  # Mitigate round-off errors
227
236
  mat[abs(mat) < 1e-14] = 0.0
@@ -293,7 +302,7 @@ def histopolation_matrix_true(knots, degree, periodic, normalization, xgrid):
293
302
  # Compute span for each row (index of last non-zero basis function)
294
303
  # TODO: would be better to have this ready beforehand
295
304
  # TODO: use tolerance instead of comparing against zero
296
- spans = [(row != 0).argmax() + (degree+1) for row in C]
305
+ spans = [int((row != 0).argmax()) + (degree+1) for row in C]
297
306
 
298
307
  # Compute histopolation matrix from collocation matrix of higher degree
299
308
  m = C.shape[0] - 1
@@ -1,6 +1,7 @@
1
1
  # coding: utf-8
2
2
 
3
3
  import cunumpy as xp
4
+ from cunumpy.mpi import synchronize_for_mpi
4
5
  import numpy as np
5
6
  from feectools.ddm.mpi import mpi as MPI
6
7
 
@@ -82,6 +83,10 @@ class BlockingCartDataExchanger(CartDataExchanger):
82
83
 
83
84
  assert isinstance( array, xp.ndarray )
84
85
 
86
+ # MPI reads/writes `array` directly; on a device backend the
87
+ # kernels that produced it must have finished first.
88
+ synchronize_for_mpi( array )
89
+
85
90
  # Shortcuts
86
91
  cart = self._cart
87
92
  comm = self._comm
@@ -123,6 +128,10 @@ class BlockingCartDataExchanger(CartDataExchanger):
123
128
 
124
129
  assert isinstance( array, xp.ndarray )
125
130
 
131
+ # MPI reads/writes `array` directly; on a device backend the
132
+ # kernels that produced it must have finished first.
133
+ synchronize_for_mpi( array )
134
+
126
135
  # Shortcuts
127
136
  cart = self._cart
128
137
  comm = self._comm
@@ -1,19 +1,22 @@
1
1
  # coding: utf-8
2
2
 
3
+ import copy
3
4
  import os
5
+ import cunumpy # only for its backend-agnostic to_numpy(), see below -- not aliased to
6
+ # xp here, since that alias is reserved for plain NumPy in this module.
4
7
  import numpy as np
5
- import cunumpy as xp
6
- from cunumpy.xp import array_backend
8
+ import numpy as xp # this module is host-only MPI/index bookkeeping, never device data
7
9
  from itertools import product
8
10
 
9
- # Initialize CUDA context before MPI if using CuPy backend
10
- if array_backend.backend == "cupy":
11
- try:
12
- import cupy as cp
13
- cp.cuda.Device(0).use()
14
- cp.cuda.Stream.null.synchronize()
15
- except Exception:
16
- pass
11
+ from cunumpy.xp import array_backend, to_numpy
12
+
13
+ # Bind this rank to its own GPU (by its rank within the node) and create the
14
+ # CUDA context before MPI is initialized, as CUDA-aware MPI requires. Must stay
15
+ # above the feectools.ddm.mpi import, which initializes MPI as a side effect.
16
+ # A no-op on the NumPy backend.
17
+ from cunumpy.cuda import bind_local_device
18
+
19
+ bind_local_device()
17
20
 
18
21
  from feectools.ddm.mpi import mpi as MPI
19
22
  from feectools.ddm.mpi import MockMPI
@@ -409,42 +412,100 @@ class DomainDecomposition:
409
412
  def refine(self, ncells, global_element_starts, global_element_ends):
410
413
  """ Create the new Cartesian decomposition of the refined domain.
411
414
 
415
+ The process topology (and its communicators) is shared with ``self``.
416
+
412
417
  Parameters
413
418
  ----------
414
419
  ncells : list or tuple of int
415
420
  Number of cells of refined space.
416
421
 
417
- global_starts: list of list of int
418
- The starts of the coefficients for every process along each direction.
422
+ global_element_starts : list of list of int
423
+ The element starts for every process along each direction.
419
424
 
420
- global_ends: list of list of int
421
- The ends of the coefficients for every process along each direction.
425
+ global_element_ends : list of list of int
426
+ The element ends for every process along each direction.
422
427
 
423
428
  Returns
424
429
  -------
425
- domain : CartDecomposition
426
- Cartesian decomposition of the refined domain.
430
+ domain : DomainDecomposition
431
+ Domain decomposition of the refined domain.
427
432
  """
428
433
 
429
434
  # Check input arguments
430
435
  assert len( ncells ) == len( self.ncells )
431
436
  assert all(nc>=snc for nc, snc in zip(ncells, self.ncells))
432
437
 
433
- domain = DomainDecomposition(self.ncells, self.periods, comm=self.comm,
434
- global_comm=self.global_comm, num_threads=self.num_threads,
435
- size=self.size)
436
- domain._ncells = tuple ( ncells )
438
+ return self._with_element_partition(ncells, global_element_starts, global_element_ends)
439
+
440
+ def coarsen(self, factors):
441
+ """ Create the Cartesian decomposition of a coarsened domain, aligned with ``self``.
442
+
443
+ Along axis ``i`` every ``factors[i]`` consecutive cells are merged into one coarse cell.
444
+ Each process owns the coarse cells covering exactly its fine cells, so the process
445
+ topology (and its communicators) is shared with ``self``. This requires that the
446
+ element starts and ends+1 of every process are divisible by ``factors[i]``.
447
+
448
+ Parameters
449
+ ----------
450
+ factors : list or tuple of int
451
+ Coarsening factor (>= 1) along each direction.
452
+
453
+ Returns
454
+ -------
455
+ domain : DomainDecomposition
456
+ Domain decomposition of the coarse domain.
457
+ """
458
+
459
+ assert len( factors ) == self.ndim
460
+ assert all( isinstance(f, (int, np.integer)) and f >= 1 for f in factors )
461
+
462
+ ncells = []
463
+ global_element_starts = []
464
+ global_element_ends = []
465
+ for axis, f in enumerate(factors):
466
+ gs = xp.asarray(self._global_element_starts[axis])
467
+ ge = xp.asarray(self._global_element_ends [axis])
468
+ if self._ncells[axis] % f != 0 or xp.any(gs % f != 0) or xp.any((ge + 1) % f != 0):
469
+ raise ValueError(
470
+ f"Cannot coarsen axis {axis} by a factor {f}: ncells={self._ncells[axis]}, "
471
+ f"element starts={gs.tolist()}, ends={ge.tolist()} are not all aligned."
472
+ )
473
+ ncells.append(self._ncells[axis] // f)
474
+ global_element_starts.append(xp.array(gs // f))
475
+ global_element_ends .append(xp.array((ge + 1) // f - 1))
476
+
477
+ return self._with_element_partition(ncells, global_element_starts, global_element_ends)
478
+
479
+ def _with_element_partition(self, ncells, global_element_starts, global_element_ends):
480
+ """ Return a copy of ``self`` with the same process topology but a new element partition.
481
+
482
+ Communicators are shared (not duplicated), hence this method is not collective.
483
+ """
484
+
485
+ assert len( ncells ) == self.ndim
486
+ for axis in range(self.ndim):
487
+ gs = xp.asarray(global_element_starts[axis])
488
+ ge = xp.asarray(global_element_ends [axis])
489
+ assert len(gs) == len(ge) == self._nprocs[axis], \
490
+ f"Axis {axis}: need one block per process ({self._nprocs[axis]}), got {len(gs)}."
491
+ assert gs[0] == 0 and ge[-1] == ncells[axis] - 1, \
492
+ f"Axis {axis}: blocks must cover [0, {ncells[axis] - 1}]."
493
+ assert xp.all(ge >= gs), f"Axis {axis}: empty blocks are not allowed."
494
+ assert xp.all(gs[1:] == ge[:-1] + 1), f"Axis {axis}: blocks must be contiguous."
495
+
496
+ domain = copy.copy(self)
497
+ domain._ncells = tuple( int(n) for n in ncells )
437
498
 
438
499
  # Store arrays with all the starts and ends along each direction for every process
439
- domain._global_element_starts = tuple(global_element_starts)
440
- domain._global_element_ends = tuple(global_element_ends)
500
+ domain._global_element_starts = list(global_element_starts)
501
+ domain._global_element_ends = list(global_element_ends)
441
502
  if self.is_comm_null:return domain
442
503
 
443
504
  # Start/end values of global indices (without ghost regions)
444
505
  domain._starts = tuple( domain._global_element_starts[axis][c] for axis,c in zip(range(self._ndims), self._coords) )
445
506
  domain._ends = tuple( domain._global_element_ends [axis][c] for axis,c in zip(range(self._ndims), self._coords) )
446
507
 
447
- domain._local_ncells = tuple(e-s+1 for s,e in zip(self._starts, self._ends))
508
+ domain._local_ncells = tuple(e-s+1 for s,e in zip(domain._starts, domain._ends))
448
509
  return domain
449
510
 
450
511
  #==================================================================================
@@ -482,6 +543,12 @@ class CartDecomposition():
482
543
  """
483
544
  def __init__( self, domain_decomposition, npts, global_starts, global_ends, pads, shifts ):
484
545
 
546
+ # global_starts/global_ends are host-side decomposition metadata; callers
547
+ # may hand them in as CuPy arrays (e.g. built with cunumpy under the CuPy
548
+ # backend), so coerce them to NumPy up front.
549
+ global_starts = [ to_numpy(gs) for gs in global_starts ]
550
+ global_ends = [ to_numpy(ge) for ge in global_ends ]
551
+
485
552
  # Check input arguments
486
553
  # TODO: check that arguments are identical across all processes
487
554
  assert len( npts ) == len( global_starts ) == len( global_ends ) == len( pads ) == len(shifts)
@@ -494,8 +561,8 @@ class CartDecomposition():
494
561
  self._domain_decomposition = domain_decomposition
495
562
  self._npts = tuple( npts )
496
563
  # Convert to NumPy arrays for MPI compatibility (MPI can't handle CuPy arrays)
497
- self._global_starts = tuple( [ np.asarray(gs.get() if hasattr(gs, 'get') else gs) for gs in global_starts] )
498
- self._global_ends = tuple( [ np.asarray(ge.get() if hasattr(ge, 'get') else ge) for ge in global_ends] )
564
+ self._global_starts = tuple( [ to_numpy(gs) for gs in global_starts] )
565
+ self._global_ends = tuple( [ to_numpy(ge) for ge in global_ends] )
499
566
  self._pads = tuple( pads )
500
567
  self._shifts = tuple( shifts )
501
568
  self._periods = domain_decomposition.periods
@@ -510,6 +577,11 @@ class CartDecomposition():
510
577
  self._shape = (0,)*self._ndims
511
578
  self._parent_starts = (None,)*self._ndims
512
579
  self._parent_ends = (None,)*self._ndims
580
+ # Serial decompositions have no neighbour exchanges, but exchange
581
+ # helpers still inspect these caches. Define them before the early
582
+ # communicator exits so those helpers are backend-independent.
583
+ self._shift_info = {}
584
+ self._shift_info_non_blocking = {}
513
585
 
514
586
  if self._comm == MPI.COMM_NULL:
515
587
  return
@@ -522,7 +594,11 @@ class CartDecomposition():
522
594
  # Know my coordinates in the topology
523
595
  self._coords = domain_decomposition.coords
524
596
  # Convert coords to NumPy for indexing (MPI coords should be on CPU)
525
- coords_np = [c.get() if hasattr(c, 'get') else c for c in self._coords]
597
+ # cunumpy.to_numpy, not used here: self._coords may hold plain Python ints
598
+ # (mpi4py's Get_coords returns a plain list), and to_numpy would wrap those
599
+ # into 0-d NumPy arrays via np.asarray -- the wrong type to index a tuple of
600
+ # global_starts/ends with below. is_gpu leaves non-CuPy values untouched.
601
+ coords_np = [c.get() if cunumpy.is_gpu(c) else c for c in self._coords]
526
602
 
527
603
  # Start/end values of global indices (without ghost regions)
528
604
  self._starts = tuple( self._global_starts[axis][c] for axis,c in zip(range(self._ndims), coords_np) )
@@ -546,11 +622,6 @@ class CartDecomposition():
546
622
  # Create (N-1)-dimensional communicators within the Cartesian topology
547
623
  self._subcomm = domain_decomposition.subcomm
548
624
 
549
- # dict to store information for communicating with neighbors
550
- self._shift_info = {}
551
-
552
- # # dict to store information for communicating with neighbors using non blocking communications
553
- self._shift_info_non_blocking = {}
554
625
 
555
626
  #---------------------------------------------------------------------------
556
627
  # Global properties (same for each process)
@@ -1,5 +1,6 @@
1
1
  # coding: utf-8
2
2
 
3
+ from cunumpy.mpi import synchronize_for_mpi
3
4
  from feectools.ddm.mpi import mpi as MPI
4
5
 
5
6
  from .cart import InterfaceCartDecomposition, find_mpi_type
@@ -48,6 +49,10 @@ class InterfaceCartDataExchanger:
48
49
 
49
50
  # ...
50
51
  def start_update_ghost_regions( self, array_minus=None, array_plus=None ):
52
+ # MPI reads/writes these buffers directly; on a device backend the
53
+ # kernels that produced them must have finished first.
54
+ synchronize_for_mpi( array_minus, array_plus )
55
+
51
56
  send_req = []
52
57
  recv_req = []
53
58
  cart = self._cart
@@ -187,11 +187,9 @@ def _mpi_disabled():
187
187
 
188
188
  if launched_under_mpi():
189
189
  try:
190
- # Disable MPI when using CuPy due to known segfault issues with OpenMPI + CUDA
191
- import os
192
- if os.environ.get('ARRAY_BACKEND') == 'cupy':
193
- raise ImportError("MPI disabled when using CuPy backend")
194
-
190
+ # MPI with the CuPy backend needs a CUDA-aware MPI library: device
191
+ # buffers are passed to MPI directly (after synchronize_for_mpi, see
192
+ # the data exchangers). A non-CUDA-aware MPI segfaults on them.
195
193
  if _mpi_disabled():
196
194
  raise ImportError("MPI disabled (feectools.use_mpi = False or FEECTOOLS_MPI=0)")
197
195
 
@@ -1,6 +1,7 @@
1
1
  # coding: utf-8
2
2
 
3
3
  import cunumpy as xp
4
+ from cunumpy.mpi import synchronize_for_mpi
4
5
  import numpy as np
5
6
  from itertools import product
6
7
 
@@ -98,6 +99,9 @@ class NonBlockingCartDataExchanger(CartDataExchanger):
98
99
  return tuple(requests)
99
100
 
100
101
  def start_update_ghost_regions(self, array, requests ):
102
+ # The persistent requests read/write `array` directly; on a device
103
+ # backend the kernels that produced it must have finished first.
104
+ synchronize_for_mpi( array )
101
105
  MPI.Prequest.Startall( requests )
102
106
 
103
107
  def end_update_ghost_regions(self, array, requests):
@@ -108,6 +112,10 @@ class NonBlockingCartDataExchanger(CartDataExchanger):
108
112
 
109
113
  assert isinstance( array, xp.ndarray )
110
114
 
115
+ # MPI reads/writes `array` directly; on a device backend the
116
+ # kernels that produced it must have finished first.
117
+ synchronize_for_mpi( array )
118
+
111
119
  # Shortcuts
112
120
  cart = self._cart
113
121
  comm = self._comm
@@ -1,5 +1,4 @@
1
- import cunumpy as xp
2
- import numpy as np
1
+ import numpy as xp # this module is host-only MPI/index bookkeeping, never device data
3
2
  import numpy.ma as ma
4
3
 
5
4
  __all__ = ('compute_dims', 'partition_procs_per_patch')
@@ -1,6 +1,6 @@
1
1
  # coding: utf-8
2
2
 
3
- import cunumpy as xp
3
+ import numpy as xp # this module is host-only MPI/index bookkeeping, never device data
4
4
  from itertools import product
5
5
 
6
6
  import cunumpy as xp