From 4c468a6830bd3fadb21abcd6d8b59f6aa947cc38 Mon Sep 17 00:00:00 2001 From: cmdupuis3 Date: Tue, 29 Sep 2026 17:45:06 -0500 Subject: [PATCH] Guard numba's thread pool against concurrent entry from thread pools 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; nothing to do. * omp starts a full team per calling thread: 179 OS threads from a 12-thread pool. * workqueue aborts the process on concurrent entry. That is live today: a chunked neighborhood reduction (the target="parallel" gufunc runs inside apply_ufunc(dask="parallelized")) kills the interpreter. It predates the nogil change; it aborts with the GIL held too. numba_pool() in uxarray/utils/parallel.py guards each call by layer: a lock under workqueue (on every thread, since a main-thread call racing a dask task aborts too), numba.set_num_threads(1) off the main thread under omp, restored afterwards because the setting is thread-local. The 14 parallel=True kernels move to @parallel_njit, which is njit plus that guard; the neighborhood gufunc wraps its call in it. The two remedies are per layer on purpose. Capping does not stop the workqueue abort, and capping under the lock would run every pool-thread call one at a time on one thread. Chunked neighborhood percentile, 64 steps, best of 3: layer before after tbb 139ms 139ms omp 179 OS threads 47 OS threads workqueue abort ~1.1s The workqueue time is the layer, not the lock: main takes 1176ms under dask's sync scheduler. Test suite: 983 passed, 1 skipped. test_plot_with_features fails identically before and after (matplotlib figure size, unrelated). Co-Authored-By: Claude Opus 5.5 --- test/utils/test_parallel.py | 47 +++++++++++++++++++++++++++++ uxarray/core/gradient.py | 3 +- uxarray/grid/angles.py | 5 ++-- uxarray/grid/area.py | 3 +- uxarray/grid/bounds.py | 3 +- uxarray/grid/connectivity.py | 5 ++-- uxarray/grid/coordinates.py | 5 ++-- uxarray/grid/dual.py | 3 +- uxarray/grid/geometry.py | 3 +- uxarray/grid/integrate.py | 3 +- uxarray/grid/neighbors.py | 4 ++- uxarray/grid/point_in_face.py | 3 +- uxarray/grid/utils.py | 3 +- uxarray/remap/bilinear.py | 3 +- uxarray/utils/parallel.py | 56 +++++++++++++++++++++++++++++++++++ 15 files changed, 133 insertions(+), 16 deletions(-) create mode 100644 test/utils/test_parallel.py create mode 100644 uxarray/utils/parallel.py diff --git a/test/utils/test_parallel.py b/test/utils/test_parallel.py new file mode 100644 index 000000000..6f53eaace --- /dev/null +++ b/test/utils/test_parallel.py @@ -0,0 +1,47 @@ +import os +import subprocess +import sys + +_CODE = """ +from concurrent.futures import ThreadPoolExecutor + +import numpy as np + +from uxarray.grid.coordinates import _construct_face_centroids +from uxarray.grid.neighbors import Neighborhood + +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)) +n_nodes_per_face = np.full(n_face, 4) +centroids = (*nodes, face_nodes, n_nodes_per_face) + +data = rng.random((8, n_node)) +counts = np.full(n_node, 16) +starts = np.arange(n_node) * 16 +flat = rng.integers(0, n_node, n_node * 16) +mean = (data, flat, starts, counts, 0.0) + +for kernel, args in [ + (_construct_face_centroids, centroids), + (Neighborhood._mean_kernel, mean), +]: + expected = kernel(*args) + with ThreadPoolExecutor(8) as pool: + results = list(pool.map(lambda _: kernel(*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 6b1823a7c..782949fd5 100644 --- a/uxarray/grid/neighbors.py +++ b/uxarray/grid/neighbors.py @@ -14,6 +14,7 @@ INT_FILL_VALUE, ) from uxarray.errors import DimensionError +from uxarray.utils.parallel import numba_pool class KDTree: @@ -1243,7 +1244,8 @@ def kernel(data, flat, starts, counts, param, out): return kernel def kernel(*args): - return build()(*args) + with numba_pool(): + return build()(*args) return kernel 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