Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 35 additions & 0 deletions test/utils/test_parallel.py
Original file line number Diff line number Diff line change
@@ -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
3 changes: 2 additions & 1 deletion uxarray/core/gradient.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand Down
5 changes: 3 additions & 2 deletions uxarray/grid/angles.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion uxarray/grid/area.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
_numba_dot3,
_numba_mul3_scalar,
)
from uxarray.utils.parallel import parallel_njit


@njit(cache=True)
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion uxarray/grid/bounds.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
all_elements_nan,
any_close_lat,
)
from uxarray.utils.parallel import parallel_njit


def _populate_face_bounds(
Expand Down Expand Up @@ -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,
Expand Down
5 changes: 3 additions & 2 deletions uxarray/grid/connectivity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
):
Expand Down
5 changes: 3 additions & 2 deletions uxarray/grid/coordinates.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
_numba_div3_scalar,
_numba_norm3,
)
from uxarray.utils.parallel import parallel_njit


@njit(cache=True, nogil=True)
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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.

Expand Down
3 changes: 2 additions & 1 deletion uxarray/grid/dual.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion uxarray/grid/geometry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]),
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion uxarray/grid/integrate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]

Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion uxarray/grid/neighbors.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
INT_FILL_VALUE,
)
from uxarray.errors import DimensionError
from uxarray.utils.parallel import parallel_njit


class KDTree:
Expand Down Expand Up @@ -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)
Expand Down
3 changes: 2 additions & 1 deletion uxarray/grid/point_in_face.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion uxarray/grid/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
_numba_norm3,
_numba_sub3,
)
from uxarray.utils.parallel import parallel_njit


@njit(cache=True)
Expand Down Expand Up @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion uxarray/remap/bilinear.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
56 changes: 56 additions & 0 deletions uxarray/utils/parallel.py
Original file line number Diff line number Diff line change
@@ -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
Loading