Skip to content
Merged
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
3 changes: 2 additions & 1 deletion benchmarks/mpas_ocean.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,8 @@ def track_nbytes_gradient(self, resolution):

def track_peakmem_gradient(self, resolution):
"""Transient high-water allocation of taking a gradient."""
return peak_allocated(lambda: self.uxds[data_var].gradient())
with numba_threads(1):
return peak_allocated(lambda: self.uxds[data_var].gradient())

track_peakmem_gradient.unit = "bytes"

Expand Down
30 changes: 30 additions & 0 deletions test/test_dependencies.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,4 +38,34 @@ def test_hvplot_optional():
_assert_not_imported_after_import_uxarray("hvplot")


def test_no_numba_kernels_built_on_import():
"""Test that `import uxarray` does not build any numba kernel.

``guvectorize`` compiles at decoration time when it is given explicit
signatures, so a kernel assigned at module scope is built during the
import. This compilation can dominate the uxarray import, and building a
``target="parallel"`` kernel starts numba's threading layer, which
leaves a thread pool running, making forks unsafe.
"""
code = (
"import numba, uxarray\n"
"try:\n"
" layer = numba.threading_layer()\n"
"except ValueError:\n"
" pass\n"
"else:\n"
" raise AssertionError(\n"
" f'`import uxarray` started numba threading layer {layer!r}. '\n"
" 'Something it imports builds a parallel kernel at module '\n"
" 'scope; build it on first use instead.'\n"
" )\n"
)
result = subprocess.run(
[sys.executable, "-c", code],
capture_output=True,
text=True,
)
assert result.returncode == 0, result.stderr


# TODO: similar tests for cartopy, holoviews, and other optional deps.
118 changes: 69 additions & 49 deletions uxarray/grid/neighbors.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import functools
import warnings
from typing import Callable

Expand Down Expand Up @@ -1214,8 +1215,8 @@ def _get_element_coords(grid, data_mapping: str, coordinate_system: str):


def _make_kernel(reduce_fn):
"""Builds a kernel that gathers each neighborhood, then calls
``reduce_fn(window, param)`` on the 1-D result.
"""Returns a kernel, compiled on its first call, that gathers each
neighborhood, then calls ``reduce_fn(window, param)`` on the 1-D result.

``reduce_fn`` must be numba-compilable, and must be defined in a real
source file for ``cache=True`` to find it.
Expand All @@ -1225,23 +1226,30 @@ def _make_kernel(reduce_fn):
if not hasattr(reduce_fn, "py_func"):
reduce_fn = njit(cache=True)(reduce_fn)

@guvectorize(_GUFUNC_SIGNATURES, _GUFUNC_LAYOUT, **_GUFUNC_KWARGS)
def kernel(data, flat, starts, counts, param, out):
widest = 0
for i in range(counts.shape[0]):
if counts[i] > widest:
widest = counts[i]
buffer = np.empty(widest, dtype=np.float64)

for i in range(starts.shape[0]):
count = counts[i]
if count == 0:
out[i] = np.nan
continue
start = starts[i]
for j in range(count):
buffer[j] = data[flat[start + j]]
out[i] = reduce_fn(buffer[:count], param)
@functools.cache
def build():
@guvectorize(_GUFUNC_SIGNATURES, _GUFUNC_LAYOUT, **_GUFUNC_KWARGS)
def kernel(data, flat, starts, counts, param, out):
widest = 0
for i in range(counts.shape[0]):
if counts[i] > widest:
widest = counts[i]
buffer = np.empty(widest, dtype=np.float64)

for i in range(starts.shape[0]):
count = counts[i]
if count == 0:
out[i] = np.nan
continue
start = starts[i]
for j in range(count):
buffer[j] = data[flat[start + j]]
out[i] = reduce_fn(buffer[:count], param)

return kernel

def kernel(*args):
return build()(*args)

return kernel

Expand Down Expand Up @@ -1277,26 +1285,6 @@ def _median(window, _):
return np.median(window)


# One compiled kernel per reduction. The methods on ``Neighborhood`` below name
# these directly, so there is no dispatch table between the public API and the
# gufuncs: a reduction is reachable only if a method exists for it, and a method
# can only reach the kernel it names. ``Neighborhood`` is the only class that
# names them -- the data-bound classes reach a kernel by naming the
# ``Neighborhood`` method for it, so there is one place per reduction where its
# kernel and parameter are chosen.
_MEAN_KERNEL = _make_kernel(lambda window, _: np.mean(window))
_SUM_KERNEL = _make_kernel(lambda window, _: np.sum(window))
_MIN_KERNEL = _make_kernel(lambda window, _: np.min(window))
_MAX_KERNEL = _make_kernel(lambda window, _: np.max(window))
_PTP_KERNEL = _make_kernel(lambda window, _: np.max(window) - np.min(window))
_MEDIAN_KERNEL = _make_kernel(_median)
_VAR_KERNEL = _make_kernel(_variance)
_STD_KERNEL = _make_kernel(lambda window, ddof: np.sqrt(_variance(window, ddof)))
# ``percentile`` is ``quantile`` on a 0-100 scale, so both methods rescale onto
# this one kernel rather than compiling a near-duplicate.
_QUANTILE_KERNEL = _make_kernel(lambda window, q: np.quantile(window, q))


def _as_quantile(q, scale: float):
"""Validates ``q`` on a 0-``scale`` scale and returns it as a 0-1 fraction."""
value = float(q)
Expand Down Expand Up @@ -1507,47 +1495,79 @@ def __repr__(self) -> str:
f"neighbors_per_element=[{self._counts.min()}, {self._counts.max()}]>"
)

# One compiled kernel per reduction. The methods below call into these
# directly. Non-compiled functions are only provided hooks through
# ``reduce``. If new compiled reductions are desired, they should follow
# this pattern.
#
# ``_make_kernel`` defers each build to the kernel's first call. The
# deferred compilation ensures that these kernels will only be compiled
# individually and lazily. Further, the lazy compilation prevents gufuncs
# from spawning threadpools eagerly and disrupting threading and forking in
# other contexts. They are wrapped in ``staticmethod`` because a plain
# function in a class body would bind ``self`` as the kernel's first
# argument.

_mean_kernel = staticmethod(_make_kernel(lambda window, _: np.mean(window)))
_sum_kernel = staticmethod(_make_kernel(lambda window, _: np.sum(window)))
_min_kernel = staticmethod(_make_kernel(lambda window, _: np.min(window)))
_max_kernel = staticmethod(_make_kernel(lambda window, _: np.max(window)))
_ptp_kernel = staticmethod(
_make_kernel(lambda window, _: np.max(window) - np.min(window))
)
_median_kernel = staticmethod(_make_kernel(_median))
_var_kernel = staticmethod(_make_kernel(_variance))
_std_kernel = staticmethod(
_make_kernel(lambda window, ddof: np.sqrt(_variance(window, ddof)))
)

# ``percentile`` is ``quantile`` on a 0-100 scale, so both methods
# rescale onto this one kernel rather than compiling a near-duplicate.
_quantile_kernel = staticmethod(
_make_kernel(lambda window, q: np.quantile(window, q))
)

def mean(self, uxda):
"""Mean of each neighborhood."""
return self._apply_kernel(uxda, _MEAN_KERNEL, 0.0)
return self._apply_kernel(uxda, self._mean_kernel, 0.0)

def sum(self, uxda):
"""Sum of each neighborhood."""
return self._apply_kernel(uxda, _SUM_KERNEL, 0.0)
return self._apply_kernel(uxda, self._sum_kernel, 0.0)

def min(self, uxda):
"""Smallest value in each neighborhood."""
return self._apply_kernel(uxda, _MIN_KERNEL, 0.0)
return self._apply_kernel(uxda, self._min_kernel, 0.0)

def max(self, uxda):
"""Largest value in each neighborhood."""
return self._apply_kernel(uxda, _MAX_KERNEL, 0.0)
return self._apply_kernel(uxda, self._max_kernel, 0.0)

def ptp(self, uxda):
"""Peak-to-peak spread (``max - min``) of each neighborhood."""
return self._apply_kernel(uxda, _PTP_KERNEL, 0.0)
return self._apply_kernel(uxda, self._ptp_kernel, 0.0)

def median(self, uxda):
"""Median of each neighborhood."""
return self._apply_kernel(uxda, _MEDIAN_KERNEL, 0.0)
return self._apply_kernel(uxda, self._median_kernel, 0.0)

def var(self, uxda, ddof: int = 0):
"""Variance of each neighborhood, with ``ddof`` delta degrees of
freedom."""
return self._apply_kernel(uxda, _VAR_KERNEL, float(ddof))
return self._apply_kernel(uxda, self._var_kernel, float(ddof))

def std(self, uxda, ddof: int = 0):
"""Standard deviation of each neighborhood, with ``ddof`` delta degrees
of freedom."""
return self._apply_kernel(uxda, _STD_KERNEL, float(ddof))
return self._apply_kernel(uxda, self._std_kernel, float(ddof))

def quantile(self, uxda, q: float):
"""Quantile ``q`` (between 0 and 1) of each neighborhood."""
return self._apply_kernel(uxda, _QUANTILE_KERNEL, _as_quantile(q, 1.0))
return self._apply_kernel(uxda, self._quantile_kernel, _as_quantile(q, 1.0))

def percentile(self, uxda, q: float):
"""Percentile ``q`` (between 0 and 100) of each neighborhood."""
return self._apply_kernel(uxda, _QUANTILE_KERNEL, _as_quantile(q, 100.0))
return self._apply_kernel(uxda, self._quantile_kernel, _as_quantile(q, 100.0))

def reduce(self, uxda, func: Callable):
"""Reduces each neighborhood with an arbitrary callable.
Expand Down