diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index c11e500cc..5cd15ce3a 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -25,3 +25,12 @@ repos: - id: ruff-format name: format with ruff types_or: [ python, pyi, jupyter ] + +- repo: local + hooks: + - id: clean-numba-cache + name: clear stale numba caches + entry: python ci/clean_numba_cache.py + language: python + files: ^uxarray/(utils/(computing|numba_math)|grid/(angles|arcs|area|coordinates|geometry|intersections|point_in_face|utils))\.py$ + pass_filenames: false diff --git a/ci/clean_numba_cache.py b/ci/clean_numba_cache.py new file mode 100644 index 000000000..c245f78b8 --- /dev/null +++ b/ci/clean_numba_cache.py @@ -0,0 +1,56 @@ +"""Delete numba's on-disk cache for ``uxarray``. + +Numba stamps each ``cache=True`` index (``*.nbi``) with the ``(mtime, size)`` of the +source file that *defines* the function, and nothing else. A kernel that calls a jitted +function from another module has that callee's code linked into its own cached object, +so editing the callee leaves the caller's cache entry "valid" with the old body baked in. + +Run this after editing any shared jitted module, or let the pre-commit hook do it:: + + python ci/clean_numba_cache.py [--dry-run] +""" + +import argparse +import os +from pathlib import Path + +PACKAGE_DIR = Path(__file__).resolve().parent.parent / "uxarray" +CACHE_SUFFIXES = {".nbi", ".nbc"} + + +def _cache_files(): + """Yield numba cache files stored next to the package's sources.""" + for path in PACKAGE_DIR.rglob("__pycache__/*"): + if path.suffix in CACHE_SUFFIXES: + yield path + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__.split("\n\n")[0]) + parser.add_argument( + "-n", + "--dry-run", + action="store_true", + help="list what would be removed without deleting anything", + ) + args = parser.parse_args(argv) + + files = sorted(set(_cache_files())) + for path in files: + if args.dry_run: + print(path) + else: + path.unlink(missing_ok=True) + + if os.environ.get("NUMBA_CACHE_DIR"): + # Numba files caches there under a hash of each source directory, so they cannot + # be told apart from other projects' entries; leave them alone rather than guess. + print("NUMBA_CACHE_DIR is set: clear it by hand, it is not touched here") + + verb = "would remove" if args.dry_run else "removed" + print(f"{verb} {len(files)} numba cache file(s)") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/docs/contributing.rst b/docs/contributing.rst index cffde8e8f..a07d30b70 100644 --- a/docs/contributing.rst +++ b/docs/contributing.rst @@ -360,6 +360,16 @@ following command:: $ pre-commit run --all-files +One of the hooks, ``clean-numba-cache``, deletes numba's on-disk kernel cache when you stage a +change to one of the shared jitted modules (``uxarray/utils/computing.py``, ``numba_math.py`` +and the ``uxarray/grid`` modules that call them). Numba only checks a cached kernel against +the file that defines it, so a kernel that calls a jitted function from another file would +otherwise keep running the old code after that function is edited. The hook only fires at +commit time, so after editing one of these files, run the same cleanup by hand before you +test the change:: + + $ python ci/clean_numba_cache.py + 3.5. Use Feature Branches ------------------------- diff --git a/test/grid/geometry/test_numba_flags.py b/test/grid/geometry/test_numba_flags.py new file mode 100644 index 000000000..d283e80f2 --- /dev/null +++ b/test/grid/geometry/test_numba_flags.py @@ -0,0 +1,112 @@ +"""Guard the numba flags that the compensated (EFT) kernels depend on. + +``fastmath`` lets LLVM reassociate ``e = (a - (s - bp)) + (b - bp)`` to 0, silently undoing +every compensated result from ``uxarray.utils.computing``. ``acc_sqrt_re``, ``_accux_gca`` +and ``_accux_constlat`` need ``error_model="numpy"`` so that 1/0 and sqrt(<0) give inf/nan +for the status layer to mask, instead of raising. + +An ``inline="always"`` callee is compiled with its caller's flags, so both checks apply to +every kernel this code is inlined into, not just where it is defined. +""" + +import ast +import importlib +import inspect +import pkgutil +import sys +import textwrap +from pathlib import Path + +import pytest +from numba.core.registry import CPUDispatcher + +import uxarray + +EFT_MODULE = "uxarray.utils.computing" +NUMPY_ERROR_MODEL_KERNELS = { + "uxarray.utils.computing.acc_sqrt_re", + "uxarray.grid.intersections._accux_gca", + "uxarray.grid.intersections._accux_constlat", +} + + +def _qualname(kernel): + return f"{kernel.py_func.__module__}.{kernel.py_func.__name__}" + + +def _callees(kernel): + """Yield the jitted functions called by name in ``kernel``'s body.""" + namespace = kernel.py_func.__globals__ + for node in ast.walk(ast.parse(textwrap.dedent(inspect.getsource(kernel.py_func)))): + if not isinstance(node, ast.Call): + continue + func = node.func + if isinstance(func, ast.Name): + target = namespace.get(func.id) + elif isinstance(func, ast.Attribute) and isinstance(func.value, ast.Name): + target = getattr(namespace.get(func.value.id), func.attr, None) + else: + continue + if isinstance(target, CPUDispatcher): + yield target + + +def _compiled_into(kernel): + """Names of ``kernel`` and every kernel inlined into it, all built with its flags.""" + names, stack = {_qualname(kernel)}, [kernel] + while stack: + for callee in _callees(stack.pop()): + name = _qualname(callee) + if callee.targetoptions.get("inline") == "always" and name not in names: + names.add(name) + stack.append(callee) + return names + + +@pytest.fixture(scope="module") +def kernels(): + """Map each module-level kernel in uxarray to ``(kernel, names compiled into it)``.""" + found = {} + for info in pkgutil.walk_packages(uxarray.__path__, "uxarray."): + try: + module = importlib.import_module(info.name) + except ImportError: # missing optional dependency: its kernels can't run either + continue + for obj in vars(module).values(): + if isinstance(obj, CPUDispatcher) and obj.py_func.__module__ == info.name: + found[_qualname(obj)] = (obj, _compiled_into(obj)) + assert NUMPY_ERROR_MODEL_KERNELS <= found.keys() # else the checks pass vacuously + return found + + +def test_numpy_error_model(kernels): + offenders = sorted( + name + for name, (kernel, compiled) in kernels.items() + if compiled & NUMPY_ERROR_MODEL_KERNELS + and kernel.targetoptions.get("error_model") != "numpy" + ) + assert not offenders, f"need error_model='numpy' to mask inf/nan: {offenders}" + + +def test_no_fastmath(kernels): + eft = { + name: kernel + for name, (kernel, compiled) in kernels.items() + if any(n.startswith(EFT_MODULE + ".") for n in compiled) + } + offenders = [ + name for name, kernel in eft.items() if kernel.targetoptions.get("fastmath") + ] + # Also scan the source, for kernels the import walk can't see: nested ones and the + # ``two_prod`` variant not selected on this machine. + for module in {kernel.py_func.__module__ for kernel in eft.values()}: + tree = ast.parse(Path(sys.modules[module].__file__).read_text()) + offenders += [ + f"{module}:{node.lineno}" + for node in ast.walk(tree) + if isinstance(node, ast.keyword) and node.arg == "fastmath" + ] + assert not offenders, ( + f"fastmath would cancel the EFT error terms: {sorted(offenders)}" + ) diff --git a/uxarray/grid/utils.py b/uxarray/grid/utils.py index bfa389371..458283303 100644 --- a/uxarray/grid/utils.py +++ b/uxarray/grid/utils.py @@ -529,7 +529,7 @@ def setter(self, value): # # NOTE: these are inlined into ``cache=True`` kernels in another module, and numba stamps its # cache against the defining file alone, so editing them does not invalidate a caller's cached -# object. Clear ``uxarray/grid/__pycache__/*.nbi *.nbc`` after changing anything here. +# object. Run ``python ci/clean_numba_cache.py`` after changing anything here. @njit(cache=True) diff --git a/uxarray/utils/computing.py b/uxarray/utils/computing.py index aff800372..098c39ee5 100644 --- a/uxarray/utils/computing.py +++ b/uxarray/utils/computing.py @@ -8,6 +8,12 @@ hardware FMA when one is available (validated bit-exact at import time) and falls back to the Veltkamp split otherwise, so there is no hard FMA dependency. +.. warning:: + Every primitive here is inlined into ``cache=True`` kernels in other modules. Numba + stamps a cache entry against the file that *defines* the kernel only, + so after editing this file those callers keep running the old body from their cache, + with no warning. Run ``python ci/clean_numba_cache.py`` before testing changes here. + Python/Numba port of the AccuSphGeom C++ library (EFT tier only; the adaptive Shewchuk predicate and exact-arithmetic fallback tiers are not ported):