feectools 0.3.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.3.0/feectools.egg-info → feectools-0.3.1}/PKG-INFO +2 -2
  2. {feectools-0.3.0 → feectools-0.3.1}/feectools/core/bsplines.py +50 -32
  3. {feectools-0.3.0 → feectools-0.3.1}/feectools/core/tests/test_bsplines.py +1 -1
  4. {feectools-0.3.0 → feectools-0.3.1}/feectools/core/tests/test_bsplines_kernel.py +9 -8
  5. {feectools-0.3.0 → feectools-0.3.1}/feectools/core/tests/test_bsplines_pyccel.py +13 -4
  6. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/blocking_data_exchanger.py +9 -0
  7. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/cart.py +30 -18
  8. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/interface_data_exchanger.py +5 -0
  9. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/mpi.py +3 -5
  10. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/nonblocking_data_exchanger.py +8 -0
  11. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/partition.py +1 -2
  12. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/petsc.py +1 -1
  13. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/tests/test_cart_1d.py +2 -2
  14. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/tests/test_cart_2d.py +4 -3
  15. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/tests/test_cart_3d.py +6 -4
  16. feectools-0.3.1/feectools/ddm/tests/test_device_binding.py +18 -0
  17. {feectools-0.3.0 → feectools-0.3.1}/feectools/feec/derivatives.py +17 -5
  18. {feectools-0.3.0 → feectools-0.3.1}/feectools/feec/global_geometric_projectors.py +75 -59
  19. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/partitioning.py +7 -6
  20. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/splines.py +3 -12
  21. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tensor.py +19 -2
  22. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/analytical_profiles_1d.py +2 -0
  23. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/utilities.py +4 -0
  24. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/direct_solvers.py +23 -24
  25. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/fft.py +8 -1
  26. feectools-0.3.1/feectools/linalg/kernels/__init__.py +8 -0
  27. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_1d/__init__.py +6 -0
  28. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_1d/stencil_axpy_1d_cuda.cu +20 -0
  29. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_1d/stencil_axpy_1d_kernels.py +26 -0
  30. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_2d/__init__.py +6 -0
  31. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_2d/stencil_axpy_2d_cuda.cu +23 -0
  32. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_2d/stencil_axpy_2d_kernels.py +27 -0
  33. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_3d/__init__.py +6 -0
  34. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_3d/stencil_axpy_3d_cuda.cu +25 -0
  35. feectools-0.3.1/feectools/linalg/kernels/stencil_axpy_3d/stencil_axpy_3d_kernels.py +28 -0
  36. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_1d/__init__.py +6 -0
  37. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_1d/stencil_dot_1d_cuda.cu +40 -0
  38. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_1d/stencil_dot_1d_kernels.py +35 -0
  39. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_2d/__init__.py +6 -0
  40. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_2d/stencil_dot_2d_cuda.cu +42 -0
  41. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_2d/stencil_dot_2d_kernels.py +99 -0
  42. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_3d/__init__.py +6 -0
  43. feectools-0.3.1/feectools/linalg/kernels/stencil_dot_3d/stencil_dot_3d_cuda.cu +64 -0
  44. feectools-0.3.0/feectools/linalg/stencil_dot_kernels.py → feectools-0.3.1/feectools/linalg/kernels/stencil_dot_3d/stencil_dot_3d_kernels.py +14 -125
  45. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_1d/__init__.py +6 -0
  46. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_1d/stencil_inner_1d_cuda.cu +41 -0
  47. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_1d/stencil_inner_1d_kernels.py +36 -0
  48. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_2d/__init__.py +6 -0
  49. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_2d/stencil_inner_2d_cuda.cu +44 -0
  50. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_2d/stencil_inner_2d_kernels.py +40 -0
  51. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_3d/__init__.py +6 -0
  52. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_3d/stencil_inner_3d_cuda.cu +46 -0
  53. feectools-0.3.1/feectools/linalg/kernels/stencil_inner_3d/stencil_inner_3d_kernels.py +44 -0
  54. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_1d/__init__.py +6 -0
  55. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_1d/stencil_transpose_1d_cuda.cu +40 -0
  56. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_1d/stencil_transpose_1d_kernels.py +37 -0
  57. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_2d/__init__.py +6 -0
  58. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_2d/stencil_transpose_2d_cuda.cu +47 -0
  59. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_2d/stencil_transpose_2d_kernels.py +116 -0
  60. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_3d/__init__.py +6 -0
  61. feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_3d/stencil_transpose_3d_cuda.cu +76 -0
  62. feectools-0.3.0/feectools/linalg/stencil_transpose_kernels.py → feectools-0.3.1/feectools/linalg/kernels/stencil_transpose_3d/stencil_transpose_3d_kernels.py +20 -140
  63. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/kron.py +31 -16
  64. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/solvers.py +8 -4
  65. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/sparse.py +12 -1
  66. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/stencil.py +141 -134
  67. feectools-0.3.1/feectools/linalg/tests/cuda_parity_cases.py +54 -0
  68. feectools-0.3.1/feectools/linalg/tests/kernel_test_args.py +100 -0
  69. feectools-0.3.1/feectools/linalg/tests/test_cuda_emulation.py +38 -0
  70. feectools-0.3.1/feectools/linalg/tests/test_cuda_parity.py +99 -0
  71. feectools-0.3.1/feectools/linalg/tests/test_device_matvec.py +206 -0
  72. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_fft.py +2 -1
  73. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_kron_stencil_matrix.py +2 -1
  74. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_linalg.py +31 -26
  75. feectools-0.3.1/feectools/linalg/tests/test_mpi_device.py +251 -0
  76. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_stencil_interface_matrix.py +8 -6
  77. {feectools-0.3.0 → feectools-0.3.1}/feectools/utilities/utils.py +9 -15
  78. {feectools-0.3.0 → feectools-0.3.1/feectools.egg-info}/PKG-INFO +2 -2
  79. {feectools-0.3.0 → feectools-0.3.1}/feectools.egg-info/SOURCES.txt +43 -4
  80. {feectools-0.3.0 → feectools-0.3.1}/feectools.egg-info/requires.txt +1 -1
  81. {feectools-0.3.0 → feectools-0.3.1}/pyproject.toml +3 -3
  82. feectools-0.3.0/feectools/linalg/kernels/axpy_kernels.py +0 -57
  83. feectools-0.3.0/feectools/linalg/kernels/inner_kernels.py +0 -100
  84. feectools-0.3.0/feectools/utilities/__init__.py +0 -0
  85. {feectools-0.3.0 → feectools-0.3.1}/AUTHORS +0 -0
  86. {feectools-0.3.0 → feectools-0.3.1}/LICENSE +0 -0
  87. {feectools-0.3.0 → feectools-0.3.1}/README.md +0 -0
  88. {feectools-0.3.0 → feectools-0.3.1}/feectools/__init__.py +0 -0
  89. {feectools-0.3.0 → feectools-0.3.1}/feectools/accelerate/__init__.py +0 -0
  90. {feectools-0.3.0 → feectools-0.3.1}/feectools/accelerate/accelerate.py +0 -0
  91. {feectools-0.3.0 → feectools-0.3.1}/feectools/accelerate/compile_psydac.mk +0 -0
  92. {feectools-0.3.0 → feectools-0.3.1}/feectools/api/__init__.py +0 -0
  93. {feectools-0.3.0 → feectools-0.3.1}/feectools/api/essential_bc.py +0 -0
  94. {feectools-0.3.0 → feectools-0.3.1}/feectools/api/fem_bilinear_form.py +0 -0
  95. {feectools-0.3.0 → feectools-0.3.1}/feectools/api/fem_common.py +0 -0
  96. {feectools-0.3.0 → feectools-0.3.1}/feectools/api/fem_sum_form.py +0 -0
  97. {feectools-0.3.0 → feectools-0.3.1}/feectools/api/settings.py +0 -0
  98. {feectools-0.3.0 → feectools-0.3.1}/feectools/core/__init__.py +0 -0
  99. {feectools-0.3.0 → feectools-0.3.1}/feectools/core/bsplines_kernels.py +0 -0
  100. {feectools-0.3.0 → feectools-0.3.1}/feectools/core/field_evaluation_kernels.py +0 -0
  101. {feectools-0.3.0 → feectools-0.3.1}/feectools/core/tests/__init__.py +0 -0
  102. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/__init__.py +0 -0
  103. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/basic.py +0 -0
  104. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/tests/__init__.py +0 -0
  105. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/tests/test_coarsen.py +0 -0
  106. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/tests/test_multicart_2d.py +0 -0
  107. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/tests/test_partition.py +0 -0
  108. {feectools-0.3.0 → feectools-0.3.1}/feectools/ddm/utilities.py +0 -0
  109. {feectools-0.3.0 → feectools-0.3.1}/feectools/feec/__init__.py +0 -0
  110. {feectools-0.3.0 → feectools-0.3.1}/feectools/feec/dof_kernels.py +0 -0
  111. {feectools-0.3.0 → feectools-0.3.1}/feectools/feec/hodge.py +0 -0
  112. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/__init__.py +0 -0
  113. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/basic.py +0 -0
  114. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/grid.py +0 -0
  115. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/lst_preconditioner.py +0 -0
  116. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/projectors.py +0 -0
  117. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/__init__.py +0 -0
  118. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/analytical_profiles_base.py +0 -0
  119. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/splines_error_bounds.py +0 -0
  120. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/test_dirichlet_projectors.py +0 -0
  121. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/test_spline_histopolation.py +0 -0
  122. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/test_spline_interpolation.py +0 -0
  123. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/test_splines.py +0 -0
  124. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/test_splines_par.py +0 -0
  125. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/test_tensor.py +0 -0
  126. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/tests/test_vector_spaces.py +0 -0
  127. {feectools-0.3.0 → feectools-0.3.1}/feectools/fem/vector.py +0 -0
  128. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/__init__.py +0 -0
  129. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/basic.py +0 -0
  130. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/block.py +0 -0
  131. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/kernels/matvec_kernels.py +0 -0
  132. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/kernels/stencil2IJV_kernels.py +0 -0
  133. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/kernels/stencil2coo_kernels.py +0 -0
  134. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/kernels/transpose_kernels.py +0 -0
  135. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/memory.py +0 -0
  136. {feectools-0.3.0/feectools/linalg/kernels → feectools-0.3.1/feectools/linalg/tests}/__init__.py +0 -0
  137. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_axpy_ghost_sync.py +0 -0
  138. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_block.py +0 -0
  139. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_matrix_free.py +0 -0
  140. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_solvers.py +0 -0
  141. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_stencil_vector.py +0 -0
  142. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_stencil_vector_space.py +0 -0
  143. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/test_toarray.py +0 -0
  144. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/tests/utilities.py +0 -0
  145. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/topetsc.py +0 -0
  146. {feectools-0.3.0 → feectools-0.3.1}/feectools/linalg/utilities.py +0 -0
  147. {feectools-0.3.0/feectools/linalg/tests → feectools-0.3.1/feectools/utilities}/__init__.py +0 -0
  148. {feectools-0.3.0 → feectools-0.3.1}/feectools/utilities/quadratures.py +0 -0
  149. {feectools-0.3.0 → feectools-0.3.1}/feectools/version.py +0 -0
  150. {feectools-0.3.0 → feectools-0.3.1}/feectools.egg-info/dependency_links.txt +0 -0
  151. {feectools-0.3.0 → feectools-0.3.1}/feectools.egg-info/entry_points.txt +0 -0
  152. {feectools-0.3.0 → feectools-0.3.1}/feectools.egg-info/top_level.txt +0 -0
  153. {feectools-0.3.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.3.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
@@ -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
@@ -2,19 +2,21 @@
2
2
 
3
3
  import copy
4
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.
5
7
  import numpy as np
6
- import cunumpy as xp
7
- from cunumpy.xp import array_backend
8
+ import numpy as xp # this module is host-only MPI/index bookkeeping, never device data
8
9
  from itertools import product
9
10
 
10
- # Initialize CUDA context before MPI if using CuPy backend
11
- if array_backend.backend == "cupy":
12
- try:
13
- import cupy as cp
14
- cp.cuda.Device(0).use()
15
- cp.cuda.Stream.null.synchronize()
16
- except Exception:
17
- 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()
18
20
 
19
21
  from feectools.ddm.mpi import mpi as MPI
20
22
  from feectools.ddm.mpi import MockMPI
@@ -541,6 +543,12 @@ class CartDecomposition():
541
543
  """
542
544
  def __init__( self, domain_decomposition, npts, global_starts, global_ends, pads, shifts ):
543
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
+
544
552
  # Check input arguments
545
553
  # TODO: check that arguments are identical across all processes
546
554
  assert len( npts ) == len( global_starts ) == len( global_ends ) == len( pads ) == len(shifts)
@@ -553,8 +561,8 @@ class CartDecomposition():
553
561
  self._domain_decomposition = domain_decomposition
554
562
  self._npts = tuple( npts )
555
563
  # Convert to NumPy arrays for MPI compatibility (MPI can't handle CuPy arrays)
556
- self._global_starts = tuple( [ np.asarray(gs.get() if hasattr(gs, 'get') else gs) for gs in global_starts] )
557
- 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] )
558
566
  self._pads = tuple( pads )
559
567
  self._shifts = tuple( shifts )
560
568
  self._periods = domain_decomposition.periods
@@ -569,6 +577,11 @@ class CartDecomposition():
569
577
  self._shape = (0,)*self._ndims
570
578
  self._parent_starts = (None,)*self._ndims
571
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 = {}
572
585
 
573
586
  if self._comm == MPI.COMM_NULL:
574
587
  return
@@ -581,7 +594,11 @@ class CartDecomposition():
581
594
  # Know my coordinates in the topology
582
595
  self._coords = domain_decomposition.coords
583
596
  # Convert coords to NumPy for indexing (MPI coords should be on CPU)
584
- 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]
585
602
 
586
603
  # Start/end values of global indices (without ghost regions)
587
604
  self._starts = tuple( self._global_starts[axis][c] for axis,c in zip(range(self._ndims), coords_np) )
@@ -605,11 +622,6 @@ class CartDecomposition():
605
622
  # Create (N-1)-dimensional communicators within the Cartesian topology
606
623
  self._subcomm = domain_decomposition.subcomm
607
624
 
608
- # dict to store information for communicating with neighbors
609
- self._shift_info = {}
610
-
611
- # # dict to store information for communicating with neighbors using non blocking communications
612
- self._shift_info_non_blocking = {}
613
625
 
614
626
  #---------------------------------------------------------------------------
615
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
@@ -89,7 +89,7 @@ def run_cart_1d( data_exchanger_type, verbose=False ):
89
89
  #---------------------------------------------------------------------------
90
90
 
91
91
  # Fill in true domain with u[i1_loc]=i1_glob
92
- u[p1:-p1] = [i1 for i1 in range(s1,e1+1)]
92
+ u[p1:-p1] = xp.asarray([i1 for i1 in range(s1,e1+1)])
93
93
 
94
94
  request = synchronizer.prepare_communications(u)
95
95
 
@@ -101,7 +101,7 @@ def run_cart_1d( data_exchanger_type, verbose=False ):
101
101
  # CHECK RESULTS
102
102
  #---------------------------------------------------------------------------
103
103
  # Verify that ghost cells contain correct data (note periodic domain!)
104
- success = all( u[:] == [i1%n1 for i1 in range(s1-p1,e1+p1+1)] )
104
+ success = bool( (u[:] == xp.asarray([i1%n1 for i1 in range(s1-p1,e1+p1+1)])).all() )
105
105
 
106
106
  # MASTER only: collect information from all processes
107
107
  success_global = comm.reduce( success, op=MPI.LAND, root=0 )
@@ -92,7 +92,9 @@ def run_cart_2d( data_exchanger_type, verbose=False , nprocs=None, reverse_axis=
92
92
  #---------------------------------------------------------------------------
93
93
 
94
94
  # Fill in true domain with u[i1_loc,i2_loc,:]=[i1_glob,i2_glob]
95
- u[p1:-p1,p2:-p2,:] = [[(i1,i2) for i2 in range(s2,e2+1)] for i1 in range(s1,e1+1)]
95
+ u[p1:-p1,p2:-p2,:] = xp.asarray(
96
+ [[(i1, i2) for i2 in range(s2, e2 + 1)] for i1 in range(s1, e1 + 1)]
97
+ )
96
98
 
97
99
 
98
100
  request = synchronizer.prepare_communications(u)
@@ -109,7 +111,7 @@ def run_cart_2d( data_exchanger_type, verbose=False , nprocs=None, reverse_axis=
109
111
  val = lambda i1,i2: (i1%n1,i2) if 0<=i2<n2 else (0,0)
110
112
  uex = [[val(i1,i2) for i2 in range(s2-p2,e2+p2+1)] for i1 in range(s1-p1,e1+p1+1)]
111
113
 
112
- success = (u == uex).all()
114
+ success = (u == xp.asarray(uex)).all()
113
115
 
114
116
  # MASTER only: collect information from all processes
115
117
  success_global = comm.reduce( success, op=MPI.LAND, root=0 )
@@ -161,4 +163,3 @@ if __name__=='__main__':
161
163
  print( "PASSED", end='\n\n', flush=True )
162
164
  else:
163
165
  print( "FAILED", end='\n\n', flush=True )
164
-
@@ -98,9 +98,11 @@ def run_cart_3d( data_exchanger_type, verbose=False ):
98
98
  #---------------------------------------------------------------------------
99
99
 
100
100
  # Fill in true domain with u[i1_loc,i2_loc,i3_loc,:]=[i1_glob,i2_glob,i3_glob]
101
- u[p1:-p1,p2:-p2,p3:-p3,:] = [[[(i1,i2,i3) for i3 in range(s3,e3+1)] \
102
- for i2 in range(s2,e2+1)] \
103
- for i1 in range(s1,e1+1)]
101
+ u[p1:-p1,p2:-p2,p3:-p3,:] = xp.asarray(
102
+ [[[(i1, i2, i3) for i3 in range(s3, e3 + 1)]
103
+ for i2 in range(s2, e2 + 1)]
104
+ for i1 in range(s1, e1 + 1)]
105
+ )
104
106
 
105
107
  request = synchronizer.prepare_communications(u)
106
108
  # Update ghost regions
@@ -118,7 +120,7 @@ def run_cart_3d( data_exchanger_type, verbose=False ):
118
120
  for i2 in range(s2-p2,e2+p2+1)] \
119
121
  for i1 in range(s1-p1,e1+p1+1)]
120
122
 
121
- success = (u == uex).all()
123
+ success = (u == xp.asarray(uex)).all()
122
124
 
123
125
  # MASTER only: collect information from all processes
124
126
  success_global = comm.reduce( success, op=MPI.LAND, root=0 )
@@ -0,0 +1,18 @@
1
+ """Each process is bound to its own GPU when feectools.ddm.cart is imported."""
2
+ import cunumpy as xp
3
+ import pytest
4
+ from cunumpy.cuda import device_count
5
+ from cunumpy.kernel_testing import requires_cupy
6
+ from cunumpy.mpi import local_rank
7
+
8
+
9
+ @requires_cupy
10
+ def test_rank_is_bound_to_its_local_device():
11
+ import cupy as cp
12
+
13
+ import feectools.ddm.cart # noqa: F401 -- binds the device on import
14
+
15
+ if xp.get_backend() != "cupy":
16
+ pytest.skip("device binding only happens on the CuPy backend")
17
+ expected = local_rank() % device_count()
18
+ assert cp.cuda.runtime.getDevice() == expected