diff --git a/test/utils/test_parallel.py b/test/utils/test_parallel.py new file mode 100644 index 000000000..73cce0049 --- /dev/null +++ b/test/utils/test_parallel.py @@ -0,0 +1,35 @@ +import os +import subprocess +import sys + +_CODE = """ +from concurrent.futures import ThreadPoolExecutor + +import numpy as np + +from uxarray.grid.coordinates import _construct_face_centroids + +rng = np.random.default_rng(0) +n_node, n_face = 100_000, 400_000 +nodes = [rng.random(n_node) for _ in range(3)] +face_nodes = rng.integers(0, n_node, (n_face, 4)) +args = (*nodes, face_nodes, np.full(n_face, 4)) + +expected = _construct_face_centroids(*args) +with ThreadPoolExecutor(8) as pool: + results = list(pool.map(lambda _: _construct_face_centroids(*args), range(32))) +for result in results: + np.testing.assert_allclose(result, expected) +""" + + +def test_parallel_kernels_called_from_a_thread_pool(): + """Numba's ``workqueue`` layer aborts the process when two threads enter a + parallel region at once, so this passes only if the kernels serialize.""" + result = subprocess.run( + [sys.executable, "-c", _CODE], + capture_output=True, + text=True, + env={**os.environ, "NUMBA_THREADING_LAYER": "workqueue"}, + ) + assert result.returncode == 0, result.stderr diff --git a/uxarray/core/gradient.py b/uxarray/core/gradient.py index 0b7d8b607..43b395c39 100644 --- a/uxarray/core/gradient.py +++ b/uxarray/core/gradient.py @@ -5,6 +5,7 @@ from uxarray.constants import INT_FILL_VALUE from uxarray.errors import DataCenteringError, DimensionError +from uxarray.utils.parallel import parallel_njit def _calculate_edge_face_difference(d_var, edge_faces, n_edge): @@ -350,7 +351,7 @@ def _dual_cell_area(sx, sy, sz, angles, n): return np.abs(area) -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _compute_gradients_on_faces( data, n_face, diff --git a/uxarray/grid/angles.py b/uxarray/grid/angles.py index 0e7df3d1b..83de9da6e 100644 --- a/uxarray/grid/angles.py +++ b/uxarray/grid/angles.py @@ -3,12 +3,13 @@ """ import numpy as np -from numba import njit, prange +from numba import prange from uxarray.grid.utils import _numba_norm3, _small_angle_of_2_vectors +from uxarray.utils.parallel import parallel_njit -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _compute_face_node_angles_convex( node_x, node_y, diff --git a/uxarray/grid/area.py b/uxarray/grid/area.py index f50d9c820..5279eae7c 100644 --- a/uxarray/grid/area.py +++ b/uxarray/grid/area.py @@ -7,6 +7,7 @@ _numba_dot3, _numba_mul3_scalar, ) +from uxarray.utils.parallel import parallel_njit @njit(cache=True) @@ -206,7 +207,7 @@ def _edge_passes_through_pole(node1, node2): ) -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _get_all_face_area_from_coords( x, y, diff --git a/uxarray/grid/bounds.py b/uxarray/grid/bounds.py index 1fcca8827..55e4997cc 100644 --- a/uxarray/grid/bounds.py +++ b/uxarray/grid/bounds.py @@ -15,6 +15,7 @@ all_elements_nan, any_close_lat, ) +from uxarray.utils.parallel import parallel_njit def _populate_face_bounds( @@ -134,7 +135,7 @@ def _populate_face_bounds( grid._ds["bounds"] = bounds_da -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _construct_face_bounds_array( face_node_connectivity, n_nodes_per_face, diff --git a/uxarray/grid/connectivity.py b/uxarray/grid/connectivity.py index 097e99cb7..19ae21cd2 100644 --- a/uxarray/grid/connectivity.py +++ b/uxarray/grid/connectivity.py @@ -12,6 +12,7 @@ _search_bucket, _sort_bucket, ) +from uxarray.utils.parallel import parallel_njit def close_face_nodes(face_node_connectivity, n_face, n_max_face_nodes): @@ -302,7 +303,7 @@ def _emit_bucket_edges( face_edge_flat[half_edge_slot[i]] = edge_idx -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _build_edge_node_connectivity(face_node_connectivity, n_nodes_per_face, n_node): """Constructs the ``edge_node_connectivity`` variable, which represents the indices of the two nodes that make up each edge. Additionally, the ``face_edge_connectivity`` is derived during construction, which represents the @@ -456,7 +457,7 @@ def _populate_face_edge_connectivity(grid): ) -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _build_face_edge_connectivity( face_node_connectivity, n_nodes_per_face, edge_node_connectivity, n_node ): diff --git a/uxarray/grid/coordinates.py b/uxarray/grid/coordinates.py index 73ad3aed6..9742ec79c 100644 --- a/uxarray/grid/coordinates.py +++ b/uxarray/grid/coordinates.py @@ -12,6 +12,7 @@ _numba_div3_scalar, _numba_norm3, ) +from uxarray.utils.parallel import parallel_njit @njit(cache=True, nogil=True) @@ -275,7 +276,7 @@ def _populate_face_centroids(grid, repopulate=False): ) -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _construct_face_centroids(node_x, node_y, node_z, face_nodes, n_nodes_per_face): """Constructs the xyz centroid coordinate for each face using Cartesian Averaging. @@ -541,7 +542,7 @@ def _populate_face_centerpoints(grid, repopulate=False): ) -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _construct_face_centerpoints(node_lon, node_lat, face_nodes, n_nodes_per_face): """Constructs the face centerpoint using Welzl's algorithm. diff --git a/uxarray/grid/dual.py b/uxarray/grid/dual.py index d33052d3b..4f7b7a622 100644 --- a/uxarray/grid/dual.py +++ b/uxarray/grid/dual.py @@ -2,6 +2,7 @@ from numba import njit, prange from uxarray.constants import INT_DTYPE, INT_FILL_VALUE +from uxarray.utils.parallel import parallel_njit def construct_dual(grid): @@ -61,7 +62,7 @@ def construct_dual(grid): return new_node_face_connectivity -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def construct_faces( valid_node_indices, n_edges, diff --git a/uxarray/grid/geometry.py b/uxarray/grid/geometry.py index 31481956d..0925020bd 100644 --- a/uxarray/grid/geometry.py +++ b/uxarray/grid/geometry.py @@ -16,6 +16,7 @@ from uxarray.grid.point_in_face import _face_contains_point from uxarray.grid.utils import _get_cartesian_face_edge_nodes from uxarray.utils.imports import _raise_hint_if_optional_deps_missing +from uxarray.utils.parallel import parallel_njit POLE_POINTS_XYZ = { "North": np.array([0.0, 0.0, 1.0]), @@ -1084,7 +1085,7 @@ def _populate_max_face_radius(grid): return max_distance -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def calculate_max_face_radius( face_node_connectivity: np.ndarray, node_x: np.ndarray, diff --git a/uxarray/grid/integrate.py b/uxarray/grid/integrate.py index c975929bc..91665d945 100644 --- a/uxarray/grid/integrate.py +++ b/uxarray/grid/integrate.py @@ -10,6 +10,7 @@ gca_const_lat_intersection, get_number_of_intersections, ) +from uxarray.utils.parallel import parallel_njit DUMMY_EDGE_VALUE = [INT_FILL_VALUE, INT_FILL_VALUE, INT_FILL_VALUE] @@ -598,7 +599,7 @@ def _compute_face_arc_length(face_edges_xyz, z): return total_length -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _zonal_face_weights_util_numba( face_edges_xyz: np.ndarray, n_edges_per_face: np.ndarray, diff --git a/uxarray/grid/neighbors.py b/uxarray/grid/neighbors.py index 4eb695f73..80a7b2c1c 100644 --- a/uxarray/grid/neighbors.py +++ b/uxarray/grid/neighbors.py @@ -15,6 +15,7 @@ INT_FILL_VALUE, ) from uxarray.errors import DimensionError +from uxarray.utils.parallel import parallel_njit class KDTree: @@ -1329,7 +1330,7 @@ def _reduce_rows(data, flat, starts, counts, op, param, out): _reduce_row(data[r], flat, starts, counts, op, param, out[r], buffer) -@njit(cache=True, nogil=True, parallel=True) +@parallel_njit def _reduce_rows_parallel(data, flat, starts, counts, op, param, out): """Reduces the rows of the 2-D ``data`` in parallel, for in-memory arrays.""" widest = _widest(counts) diff --git a/uxarray/grid/point_in_face.py b/uxarray/grid/point_in_face.py index ed6faf7fd..61e71a399 100644 --- a/uxarray/grid/point_in_face.py +++ b/uxarray/grid/point_in_face.py @@ -8,6 +8,7 @@ from uxarray.constants import ERROR_TOLERANCE, INT_DTYPE, INT_FILL_VALUE from uxarray.grid.arcs import point_within_gca from uxarray.grid.utils import _get_cartesian_face_edge_nodes, _small_angle_of_2_vectors +from uxarray.utils.parallel import parallel_njit if TYPE_CHECKING: from numpy.typing import ArrayLike @@ -120,7 +121,7 @@ def _get_faces_containing_point( return hit_buf[:count] -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _batch_point_in_face( points: np.ndarray, flat_candidate_indices: np.ndarray, diff --git a/uxarray/grid/utils.py b/uxarray/grid/utils.py index 4d01e5049..1b492efdd 100644 --- a/uxarray/grid/utils.py +++ b/uxarray/grid/utils.py @@ -9,6 +9,7 @@ _numba_norm3, _numba_sub3, ) +from uxarray.utils.parallel import parallel_njit @njit(cache=True) @@ -266,7 +267,7 @@ def _get_cartesian_face_edge_nodes_array( return face_edges_cartesian.reshape(n_face, n_max_face_edges, 2, 3) -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _get_cartesian_face_edge_nodes_array_subset( face_indices, face_node_connectivity, diff --git a/uxarray/remap/bilinear.py b/uxarray/remap/bilinear.py index 9c66edb24..7e8df9592 100644 --- a/uxarray/remap/bilinear.py +++ b/uxarray/remap/bilinear.py @@ -7,6 +7,7 @@ from numba import njit, prange from uxarray.errors import DataCenteringError +from uxarray.utils.parallel import parallel_njit if TYPE_CHECKING: from uxarray.core.dataarray import UxDataArray @@ -158,7 +159,7 @@ def _barycentric_weights(point_xyz, dual, data_size, source_grid): return all_weights, all_indices -@njit(cache=True, parallel=True, nogil=True) +@parallel_njit def _calculate_weights( valid_idxs, point_xyz, diff --git a/uxarray/utils/parallel.py b/uxarray/utils/parallel.py new file mode 100644 index 000000000..350215a14 --- /dev/null +++ b/uxarray/utils/parallel.py @@ -0,0 +1,56 @@ +"""Guarding numba's thread pool against being entered from several threads. + +A ``parallel=True`` kernel called from a thread pool -- dask's threaded +scheduler, typically -- nests numba's pool under it. What that does depends on +numba's threading layer: ``tbb`` composes, ``omp`` starts a full team per +calling thread, and ``workqueue`` kills the process on concurrent entry. +""" + +import contextlib +import functools +import threading + +import numba +from numba import njit + +_WORKQUEUE_LOCK = threading.Lock() + + +@functools.cache +def _threading_layer(): + # get_num_threads starts the layer, as a first parallel call would + numba.get_num_threads() + return numba.threading_layer() + + +@contextlib.contextmanager +def numba_pool(): + """Makes the enclosed parallel kernel call safe under the current layer: + one call at a time under ``workqueue``, one thread per call off the main + thread under ``omp``.""" + layer = _threading_layer() + if layer == "workqueue": + with _WORKQUEUE_LOCK: + yield + elif layer == "omp" and threading.current_thread() is not threading.main_thread(): + previous = numba.get_num_threads() + numba.set_num_threads(1) + try: + yield + finally: + numba.set_num_threads(previous) + else: + yield + + +def parallel_njit(func): + """``@njit(cache=True, parallel=True, nogil=True)``, called under + :func:`numba_pool`. Only for kernels Python calls directly.""" + kernel = njit(cache=True, parallel=True, nogil=True)(func) + + @functools.wraps(func) + def call(*args, **kwargs): + with numba_pool(): + return kernel(*args, **kwargs) + + return call