diff --git a/autoarray/config/general.yaml b/autoarray/config/general.yaml index 6ac374c1b..b79ea4bb9 100644 --- a/autoarray/config/general.yaml +++ b/autoarray/config/general.yaml @@ -13,6 +13,7 @@ inversion: reconstruction_vmax_factor: 0.5 # Plots of an Inversion's reconstruction use the reconstructed data's bright value multiplied by this factor. log_det_method: cholesky # How the Bayesian-evidence log-determinant terms are computed. "cholesky" (default) is the historical 2*sum(log(diag(cholesky(M)))); "slogdet" uses logabsdet of slogdet(M), which is identical where M is positive-definite but finite (not NaN) where the Cholesky fails, for gradient-based searches (opt-in, non-default; does not change the default evidence). Under "slogdet" the kernel regularization schemes (Matern/Gaussian/Exponential) also compute the regularization log-det analytically from a Cholesky of their covariance instead of factorizing the formed inverse. See PyAutoArray#391. regularization_term_method: matmul # How the Bayesian-evidence regularization term s^T H s is computed. "matmul" (default) is the historical s @ (H @ s) against the explicitly formed regularization matrix; "cho_solve" evaluates coefficient * s^T C^-1 s for the kernel schemes (Matern/Gaussian/Exponential/MaternAdapt) via one Cholesky solve of their covariance C, avoiding the explicit inverse whose round-off is amplified by cond(C) (~1e9 on clustered traced mesh vertices). Opt-in, non-default; does not change the default evidence. Schemes with no such factorization fall back to the formed matrix. + interferometer_numba_nnz_per_source_max: 60.0 # Geometry gate for the numba `direct_conv` interferometer curvature path, in mean non-zeros per source column (mapper.pix_sizes_for_sub_slim_index.sum() / mapper.params). At or below this the factory routes a NumPy (xp=np) single-mapper interferometer inversion to InversionInterferometerSparseNumba, which is 2-7x faster than the JAX/FFT route while the mapping operator stays sparse; above it the FFT route wins and is used. Measured crossovers are ~60 (Delaunay) and ~77 (rectangular) on the reference CPU (autolens_profiling#226 verdict section 2) -- this is a machine-dependent constant, so re-measure before tuning. Set 0 to disable the numba path entirely. numba: use_numba: true cache: true diff --git a/autoarray/inversion/inversion/factory.py b/autoarray/inversion/inversion/factory.py index de1fe8a10..745a60447 100644 --- a/autoarray/inversion/inversion/factory.py +++ b/autoarray/inversion/inversion/factory.py @@ -11,6 +11,13 @@ from autoarray.inversion.inversion.interferometer.sparse import ( InversionInterferometerSparse, ) +from autoarray.inversion.inversion.interferometer_numba.sparse import ( + InversionInterferometerSparseNumba, +) +from autoarray.inversion.inversion.interferometer_numba import ( + inversion_interferometer_numba_util, +) +from autoarray.inversion.mappers.abstract import Mapper from autoarray.inversion.inversion.dataset_interface import DatasetInterface from autoarray.inversion.linear_obj.linear_obj import LinearObj from autoarray.inversion.linear_obj.func_list import AbstractLinearObjFuncList @@ -202,6 +209,19 @@ def inversion_interferometer_from( if dataset.sparse_operator is not None and use_sparse_operator: + if _use_interferometer_numba( + linear_obj_list=linear_obj_list, + settings=settings, + xp=xp, + ): + return InversionInterferometerSparseNumba( + dataset=dataset, + linear_obj_list=linear_obj_list, + settings=settings, + xp=xp, + preloads=preloads, + ) + return InversionInterferometerSparse( dataset=dataset, linear_obj_list=linear_obj_list, @@ -216,3 +236,85 @@ def inversion_interferometer_from( settings=settings, xp=xp, ) + + +def _use_interferometer_numba( + linear_obj_list: List[LinearObj], + settings: Settings = None, + xp=np, +) -> bool: + """ + Whether an interferometer inversion is routed to the numba `direct_conv` curvature + path (`InversionInterferometerSparseNumba`) rather than the FFT one + (`InversionInterferometerSparse`). + + Every condition below is a routing decision, not an error: a model the kernel cannot + represent, or a geometry where the FFT route is faster, simply falls through to the + sparse path silently. Constructing `InversionInterferometerSparseNumba` directly with + such inputs still raises -- the class checks the same preconditions itself, so the two + cannot drift apart in meaning, only in whether they are fatal. + + The conditions are, in order of cost to evaluate: + + - `xp is np` -- the kernel is numba, with no JAX path. + - `settings.interferometer_numba_nnz_per_source_max > 0` -- `0` is the kill switch. + - exactly one `Mapper` and no `AbstractLinearObjFuncList` -- the kernel builds a single + mapper-mapper block and has no off-diagonal or function blocks. + - no over-sampling (`sub_fraction == 1`) -- the kernel uses the mapper's weights as-is. + - the mapper's mean non-zeros per source column is at or below the gate -- above it the + FFT route is faster (see `Settings.interferometer_numba_nnz_per_source_max`). + - `import numba` succeeds. + + Parameters + ---------- + linear_obj_list + The linear objects reconstructing the data. + settings + The inversion settings, whose `interferometer_numba_nnz_per_source_max` is the + geometry gate. `None` uses the packaged defaults. + xp + The array module the inversion runs on. + """ + if xp is not np: + return False + + settings = settings if settings is not None else Settings() + + nnz_max = settings.interferometer_numba_nnz_per_source_max + + if nnz_max is None or nnz_max <= 0: + return False + + if any( + isinstance(linear_obj, AbstractLinearObjFuncList) + for linear_obj in linear_obj_list + ): + return False + + mapper_list = [ + linear_obj for linear_obj in linear_obj_list if isinstance(linear_obj, Mapper) + ] + + if len(mapper_list) != 1 or len(mapper_list) != len(linear_obj_list): + return False + + mapper = mapper_list[0] + + sub_fraction = np.asarray(mapper.over_sampler.sub_fraction.array) + + if not np.all(sub_fraction == 1.0): + return False + + nnz_per_source_column = ( + inversion_interferometer_numba_util.nnz_per_source_column_from(mapper=mapper) + ) + + if nnz_per_source_column > nnz_max: + return False + + try: + import numba # noqa: F401 + except ModuleNotFoundError: + return False + + return True diff --git a/autoarray/inversion/inversion/interferometer_numba/__init__.py b/autoarray/inversion/inversion/interferometer_numba/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/autoarray/inversion/inversion/interferometer_numba/inversion_interferometer_numba_util.py b/autoarray/inversion/inversion/interferometer_numba/inversion_interferometer_numba_util.py new file mode 100644 index 000000000..034909ae9 --- /dev/null +++ b/autoarray/inversion/inversion/interferometer_numba/inversion_interferometer_numba_util.py @@ -0,0 +1,302 @@ +import numpy as np + +from autoarray import numba_util + + +@numba_util.jit() +def curvature_direct_conv( + preload, + iy, + ix, + flat, + indptr, + col, + val, + cscptr, + csc_row, + csc_val, + ny, + nx, + pix_pixels, +): + """ + `F = Aᵀ W~ A` by convolving each source column of `A` over the extent rectangle. + + For source column `s`: + + 1. `u = W~ A[:, s]` on the `(ny, nx)` extent grid -- each of the `nnz_s` non-zeros of + the column scatters a shifted copy of the `W~` kernel onto `u`. Splitting the row + into the two contiguous halves of the wrapped `preload` row makes the inner loop a + pure contiguous AXPY. + 2. `F[s, :] = Aᵀ u` -- one gather per non-zero of the whole mapping operator. + + Cost `O(nnz·M + S·nnz)` with `M = ny·nx`, versus the historic pair loop's `O(N² P²)`: + the convolution replaces the `N²` pixel-pair space with the `M`-cell extent rectangle, + which is the whole point of the extent-grid form. It beats the FFT route (the JAX and + NumPy `InterferometerSparseOperator` paths) while the source columns stay sparse -- + measured crossover ~60 non-zeros per source column on Delaunay meshes, ~77 on + rectangular ones (autolens_profiling issue #226). + + The result is a complete, symmetric `[pix_pixels, pix_pixels]` matrix: the kernel + loops the full column x row space rather than halving on symmetry, so no mirroring + pass is required afterwards. + + Parameters + ---------- + preload + The real `(2 * ny, 2 * nx)` `W~` operator as a function of pixel offsets, i.e. + `InterferometerSparseOperator.nufft_precision_operator`. Indexed with wrapped + (negative) offsets, exactly as the FFT paths convolve with it. + iy, ix + The extent-grid row / column of every masked (sub-slim) pixel. + flat + `iy * nx + ix`, the extent-flat index of every masked pixel. + indptr, col, val + The mapping operator `A` in CSR form (rows = masked pixels, columns = source + pixels). + cscptr, csc_row, csc_val + The same triplets in CSC form (source-pixel major), which the convolution sweeps. + ny, nx + The extent rectangle's shape. + pix_pixels + The number of source pixels, i.e. `mapper.params`. + + Returns + ------- + The curvature matrix `F`, of shape `[pix_pixels, pix_pixels]`. + """ + n_pix = indptr.shape[0] - 1 + m_cells = ny * nx + nx2 = 2 * nx + + curvature_matrix = np.zeros((pix_pixels, pix_pixels)) + u = np.zeros(m_cells) + + for sp in range(pix_pixels): + u[:] = 0.0 + + for t in range(cscptr[sp], cscptr[sp + 1]): + i0 = csc_row[t] + wi = csc_val[t] + i_y = iy[i0] + i_x = ix[i0] + off = nx2 - i_x + + for jy in range(ny): + dy = jy - i_y + base = jy * nx + + for jx in range(i_x): + u[base + jx] += wi * preload[dy, off + jx] + + for jx in range(i_x, nx): + u[base + jx] += wi * preload[dy, jx - i_x] + + row = curvature_matrix[sp] + + for i1 in range(n_pix): + ui = u[flat[i1]] + + for t in range(indptr[i1], indptr[i1 + 1]): + row[col[t]] += val[t] * ui + + return curvature_matrix + + +_PARALLEL_CACHE: dict = {} + + +def direct_conv_parallel_kernel(): + """ + Compile (once) :func:`curvature_direct_conv` with `prange` over source columns. + + Built lazily rather than decorated at import time for three reasons: `numba.prange` + has to be resolvable in the function's own scope under `nopython`; + `numba_util.jit` cannot express `parallel=True` (it takes the shared library-wide + `general.yaml` numba options); and the thread count numba bakes in is read from + `NUMBA_NUM_THREADS` at *its* import, so a thread-scaling arm must set that variable + before this is first called. + + Each source column owns its accumulator `u` and writes only its own row of `F`, so + the parallel loop needs neither a reduction nor a lock -- the parallel kernel returns + exactly the serial kernel's matrix, which the tests pin. + + Returns + ------- + The compiled parallel kernel, with the same signature as + :func:`curvature_direct_conv`. + """ + if "kernel" in _PARALLEL_CACHE: + return _PARALLEL_CACHE["kernel"] + + import numba + from numba import prange + + @numba.njit(cache=True, parallel=True, nogil=True) + def _curvature_direct_conv_parallel( + preload, + iy, + ix, + flat, + indptr, + col, + val, + cscptr, + csc_row, + csc_val, + ny, + nx, + pix_pixels, + ): + n_pix = indptr.shape[0] - 1 + m_cells = ny * nx + nx2 = 2 * nx + + curvature_matrix = np.zeros((pix_pixels, pix_pixels)) + + for sp in prange(pix_pixels): + u = np.zeros(m_cells) + + for t in range(cscptr[sp], cscptr[sp + 1]): + i0 = csc_row[t] + wi = csc_val[t] + i_y = iy[i0] + i_x = ix[i0] + off = nx2 - i_x + + for jy in range(ny): + dy = jy - i_y + base = jy * nx + + for jx in range(i_x): + u[base + jx] += wi * preload[dy, off + jx] + + for jx in range(i_x, nx): + u[base + jx] += wi * preload[dy, jx - i_x] + + row = curvature_matrix[sp] + + for i1 in range(n_pix): + ui = u[flat[i1]] + + for t in range(indptr[i1], indptr[i1 + 1]): + row[col[t]] += val[t] * ui + + return curvature_matrix + + _PARALLEL_CACHE["kernel"] = _curvature_direct_conv_parallel + return _curvature_direct_conv_parallel + + +def kernel_inputs_from( + pix_indexes_for_sub_slim_index: np.ndarray, + pix_sizes_for_sub_slim_index: np.ndarray, + pix_weights_for_sub_slim_index: np.ndarray, + extent_index_for_masked_pixel: np.ndarray, + extent_shape, + pix_pixels: int, +) -> dict: + """ + Flatten a mapper's `[N_pix, P_max]` triplet arrays into the flat CSR / CSC / extent + layout :func:`curvature_direct_conv` consumes. + + The rows are indexed on the *unmasked extent* rectangle, which is the grid the `W~` + operator lives on -- the same grid `InversionInterferometerSparse` builds its COO + triplets against. The extent row / column of each masked pixel is recovered from the + mask's own `extent_index_for_masked_pixel` (`iy = flat // nx`, `ix = flat % nx`) + rather than from a re-origined `native_index_for_slim_index`, so nothing new has to + be stored on the dataset and the two paths cannot drift. + + A real `Mapper` pads its `[N_pix, P_max]` triplet rows to the longest row, so the + valid entries are the first `pix_sizes_for_sub_slim_index[i]` of each. Selecting them + with a boolean mask keeps C (row-major) order, which is exactly CSR order, and keeps + this out of Python: the marshalling runs once per likelihood evaluation and a + per-pixel loop here would show up as kernel cost. + + Parameters + ---------- + pix_indexes_for_sub_slim_index, pix_sizes_for_sub_slim_index, pix_weights_for_sub_slim_index + The mapper's dense triplet arrays. Weights are used as-is, which is only correct + without over-sampling (`sub_fraction == 1`) -- the caller enforces that. + extent_index_for_masked_pixel + The mask's flat extent-grid index of every masked pixel, ordered to match the + triplet rows. + extent_shape + The `(ny, nx)` shape of the unmasked extent, i.e. + `mask.shape_native_masked_pixels`. + pix_pixels + The number of source pixels, i.e. `mapper.params`. + + Returns + ------- + A dict of the flat arrays and scalars the kernel takes, keyed by its argument names + (`preload` excepted, which the caller supplies). + """ + flat = np.asarray(extent_index_for_masked_pixel, dtype=np.int64) + sizes = np.asarray(pix_sizes_for_sub_slim_index, dtype=np.int64) + indexes = np.asarray(pix_indexes_for_sub_slim_index, dtype=np.int64) + weights = np.asarray(pix_weights_for_sub_slim_index, dtype=np.float64) + + ny, nx = int(extent_shape[0]), int(extent_shape[1]) + + n_pix = flat.shape[0] + + iy = flat // nx + ix = flat % nx + + indptr = np.zeros(n_pix + 1, dtype=np.int64) + np.cumsum(sizes, out=indptr[1:]) + nnz = int(indptr[-1]) + + valid = np.arange(indexes.shape[1], dtype=np.int64)[None, :] < sizes[:, None] + + col = np.ascontiguousarray(indexes[valid]) + val = np.ascontiguousarray(weights[valid]) + + # Source-major (CSC) view of the same triplets. + row_of_nnz = np.repeat(np.arange(n_pix, dtype=np.int64), sizes) + order = np.argsort(col, kind="stable") + csc_row = row_of_nnz[order] + csc_val = val[order] + counts = np.bincount(col, minlength=int(pix_pixels)).astype(np.int64) + cscptr = np.zeros(int(pix_pixels) + 1, dtype=np.int64) + np.cumsum(counts, out=cscptr[1:]) + + return { + "iy": np.ascontiguousarray(iy), + "ix": np.ascontiguousarray(ix), + "flat": np.ascontiguousarray(flat), + "indptr": indptr, + "col": col, + "val": val, + "cscptr": cscptr, + "csc_row": np.ascontiguousarray(csc_row), + "csc_val": np.ascontiguousarray(csc_val), + "n_pix": n_pix, + "nnz": nnz, + "ny": ny, + "nx": nx, + "pix_pixels": int(pix_pixels), + } + + +def nnz_per_source_column_from(mapper) -> float: + """ + The mean number of non-zero mapping weights per source column of `A`. + + This is the geometry the `direct_conv` kernel's cost scales with, and therefore the + quantity the factory gates on: the kernel's convolution step costs `O(nnz · M)`, + while the FFT route's cost is set by the number of source *columns* rather than their + density, so the two cross at a roughly fixed non-zeros-per-column value (~60 on + Delaunay meshes, ~77 on rectangular ones). + + `nnz = pix_sizes_for_sub_slim_index.sum()` is the total number of valid triplets and + `mapper.params` the number of source pixels, so this is simply their ratio. + """ + nnz = float(np.asarray(mapper.pix_sizes_for_sub_slim_index).sum()) + params = int(mapper.params) + + if params == 0: + return 0.0 + + return nnz / params diff --git a/autoarray/inversion/inversion/interferometer_numba/sparse.py b/autoarray/inversion/inversion/interferometer_numba/sparse.py new file mode 100644 index 000000000..a18c066fa --- /dev/null +++ b/autoarray/inversion/inversion/interferometer_numba/sparse.py @@ -0,0 +1,233 @@ +import numpy as np +from typing import List, Union + +from autonerves import cached_property, conf + +from autoarray import exc +from autoarray.dataset.interferometer.dataset import Interferometer +from autoarray.inversion.inversion.dataset_interface import DatasetInterface +from autoarray.inversion.inversion.interferometer.sparse import ( + InversionInterferometerSparse, +) +from autoarray.inversion.linear_obj.linear_obj import LinearObj +from autoarray.inversion.linear_obj.func_list import AbstractLinearObjFuncList +from autoarray.inversion.mappers.abstract import Mapper +from autoarray.settings import Settings + +from autoarray.inversion.inversion.interferometer_numba import ( + inversion_interferometer_numba_util, +) + + +class InversionInterferometerSparseNumba(InversionInterferometerSparse): + def __init__( + self, + dataset: Union[Interferometer, DatasetInterface], + linear_obj_list: List[LinearObj], + settings: Settings = None, + xp=np, + preloads=None, + ): + """ + The single-mapper interferometer inversion with its curvature matrix `F` assembled + by a numba CPU kernel instead of the FFT route. + + `InversionInterferometerSparse` forms `F = Aᵀ W~ A` by applying `W~` as an FFT + convolution to blocks of `A`'s columns. That cost is set by the number of source + columns, whatever their density. The `direct_conv` kernel instead convolves each + source column over the `(ny, nx)` extent rectangle directly, at a cost that scales + with the column's non-zeros -- so it wins while the mapping operator stays sparse + (2-7x below ~60 non-zeros per source column on Delaunay meshes, ~77 on rectangular + ones; autolens_profiling issue #226) and loses above it. + + Everything except `curvature_matrix_diag` is inherited: the data vector, + regularization, the reconstruction and the evidence terms are the parent's, and + `curvature_matrix` keeps the parent's no-regularization diagonal handling. The + kernel returns a complete symmetric `F`, so no mirroring pass is needed (and the + parent's single-mapper branch does not apply one). + + The class raises on every configuration the kernel cannot represent rather than + silently working around it -- see `_check_preconditions`. The factory + (`inversion_interferometer_from`) checks the same conditions before it routes + here and falls through to `InversionInterferometerSparse` when any fails; a + direct construction with bad inputs still raises. + + Parameters + ---------- + dataset + The interferometer dataset (or `DatasetInterface`) being reconstructed. It + must carry a `sparse_operator`, whose `nufft_precision_operator` is the `W~` + preload the kernel convolves with. + linear_obj_list + The linear objects reconstructing the data. Exactly one `Mapper`, and nothing + else. + settings + The inversion settings (`autoarray.settings.Settings`). + xp + The array module. Must be `numpy`: the kernel is numba, with no JAX path. + preloads + Optional `AbstractPreloads`, forwarded to the parent unchanged. + """ + if xp is not np: + raise exc.InversionException( + "`InversionInterferometerSparseNumba` was passed a non-NumPy array module " + f"({xp!r}). The `direct_conv` curvature kernel is a numba `@jit` function " + "with no JAX path; use `InversionInterferometerSparse` for the JAX/FFT " + "route." + ) + + try: + import numba # noqa: F401 + except ModuleNotFoundError as error: + raise exc.InversionException( + "The numba interferometer inversion requires numba, which is not " + "installed. Install it, or use `InversionInterferometerSparse`." + ) from error + + super().__init__( + dataset=dataset, + linear_obj_list=linear_obj_list, + settings=settings, + xp=xp, + preloads=preloads, + ) + + self._check_preconditions() + + def _check_preconditions(self) -> None: + """ + Raise on every configuration the `direct_conv` kernel cannot represent. + + These are raised, not worked around: each one would otherwise compute a different + (and silently wrong) `F`, or none at all. The factory checks the same conditions + up front and simply does not route here when one fails. + """ + if self.has(cls=AbstractLinearObjFuncList): + raise exc.InversionException( + "A linear-function list (e.g. a linear light profile or MGE basis) was " + "passed to `InversionInterferometerSparseNumba`. The `direct_conv` kernel " + "assembles F only from a mapper's `pix_indexes/sizes/weights_for_sub_slim_index` " + "triplets and has no mapper x function or function x function block.\n\n" + "Use `InversionInterferometerSparse` (which does support mixed linear " + "objects) for such a model." + ) + + total_mappers = self.total(cls=Mapper) + + if total_mappers != 1: + raise exc.InversionException( + f"`InversionInterferometerSparseNumba` was passed {total_mappers} mappers. " + "The `direct_conv` kernel builds a single [pix_pixels, pix_pixels] " + "curvature matrix from one mapper and has no off-diagonal mapper x mapper " + "block.\n\n" + "Use `InversionInterferometerSparse` for a multi-mapper model." + ) + + mapper = self.cls_list_from(cls=Mapper)[0] + + sub_fraction = np.asarray(mapper.over_sampler.sub_fraction.array) + + if not np.all(sub_fraction == 1.0): + raise exc.InversionException( + "`InversionInterferometerSparseNumba` was passed a mapper whose " + f"over-sampler has `sub_fraction != 1` (minimum {float(sub_fraction.min())}, " + "i.e. `over_sample_size` up to " + f"{int(round(1.0 / float(sub_fraction.min())))}).\n\n" + "The sparse triplets fold `over_sampler.sub_fraction` into the mapping " + "weights (`interferometer/sparse.py::_sparse_triplets_curvature_from`); " + "the `direct_conv` kernel uses `pix_weights_for_sub_slim_index` as-is and " + "its rows are indexed on the slim grid, so with over-sampling it would " + "silently compute a different F.\n\n" + "Apply `over_sample_size_pixelization=1` to the dataset, or use " + "`InversionInterferometerSparse`." + ) + + @cached_property + def kernel_index_arrays(self) -> dict: + """ + The mapper's triplets in the flat CSR / CSC / extent-grid layout the kernel takes. + + The extent implied by the mask is checked against the preload's `(2ny, 2nx)` + shape, because the kernel indexes the preload by extent offsets with wrapped + (negative) indices -- a mismatch would wrap silently rather than fail. + """ + mapper = self.cls_list_from(cls=Mapper)[0] + + extent_index_for_masked_pixel = np.asarray( + self.mask.extent_index_for_masked_pixel + )[np.asarray(mapper.slim_index_for_sub_slim_index)] + + inputs = inversion_interferometer_numba_util.kernel_inputs_from( + pix_indexes_for_sub_slim_index=mapper.pix_indexes_for_sub_slim_index, + pix_sizes_for_sub_slim_index=mapper.pix_sizes_for_sub_slim_index, + pix_weights_for_sub_slim_index=mapper.pix_weights_for_sub_slim_index, + extent_index_for_masked_pixel=extent_index_for_masked_pixel, + extent_shape=self.mask.shape_native_masked_pixels, + pix_pixels=int(mapper.params), + ) + + preload_shape = np.asarray( + self.dataset.sparse_operator.nufft_precision_operator + ).shape + + if (2 * inputs["ny"], 2 * inputs["nx"]) != preload_shape: + raise exc.InversionException( + "The unmasked extent implied by the mask " + f"({inputs['ny']} x {inputs['nx']}) does not match the sparse operator's " + f"`nufft_precision_operator` shape {preload_shape}, which must be " + "(2ny, 2nx). The `direct_conv` kernel indexes the preload by extent " + "offsets, so a mismatch would wrap silently instead of failing." + ) + + return inputs + + @property + def curvature_matrix_diag(self) -> np.ndarray: + """ + `F = Aᵀ W~ A` for the inversion's single mapper, from the `direct_conv` numba + kernel. + + The returned matrix is complete and symmetric (the kernel loops the full source + column x masked pixel space rather than halving on symmetry), so it drops straight + into the parent's single-mapper branch with no mirroring pass. + """ + inputs = self.kernel_index_arrays + + preload = np.ascontiguousarray( + np.asarray( + self.dataset.sparse_operator.nufft_precision_operator, dtype=np.float64 + ) + ) + + if _numba_parallel(): + kernel = inversion_interferometer_numba_util.direct_conv_parallel_kernel() + else: + kernel = inversion_interferometer_numba_util.curvature_direct_conv + + return kernel( + preload, + inputs["iy"], + inputs["ix"], + inputs["flat"], + inputs["indptr"], + inputs["col"], + inputs["val"], + inputs["cscptr"], + inputs["csc_row"], + inputs["csc_val"], + inputs["ny"], + inputs["nx"], + inputs["pix_pixels"], + ) + + +def _numba_parallel() -> bool: + """ + Whether the parallel (`prange`) kernel is used, read from the same + `general.yaml -> numba -> parallel` flag the shared `numba_util.jit` decorator reads, + with the same fallback when no config supplies it. + """ + try: + return bool(conf.instance["general"]["numba"]["parallel"]) + except Exception: + return False diff --git a/autoarray/settings.py b/autoarray/settings.py index 29fa8c0ed..e72c9e07b 100644 --- a/autoarray/settings.py +++ b/autoarray/settings.py @@ -21,6 +21,7 @@ def __init__( nnls_warm_start_error_tolerance: Optional[float] = None, log_det_method: Optional[str] = None, regularization_term_method: Optional[str] = None, + interferometer_numba_nnz_per_source_max: Optional[float] = None, ): """ The settings of an Inversion, customizing how a linear set of equations are solved for. @@ -174,6 +175,13 @@ def __init__( Note neither option removes the explicit inverse from the inversion as a whole — ``curvature_reg_matrix`` is a dense ``F + H`` feeding the dense solve for the reconstruction, so ``H`` is still formed there regardless. + interferometer_numba_nnz_per_source_max + The geometry gate above which the numba `direct_conv` interferometer curvature + path is not used, in mean non-zeros per source column + (``mapper.pix_sizes_for_sub_slim_index.sum() / mapper.params``). `None` + (default) reads the packaged value (`60.0`); `0` disables the numba path. See + the property of the same name for the measured crossovers and why the constant + is machine-dependent. """ self.use_mixed_precision = use_mixed_precision self.nnls_solver_tol = nnls_solver_tol @@ -188,6 +196,9 @@ def __init__( ) self._log_det_method = log_det_method self._regularization_term_method = regularization_term_method + self._interferometer_numba_nnz_per_source_max = ( + interferometer_numba_nnz_per_source_max + ) @property def use_positive_only_solver(self): @@ -285,3 +296,41 @@ def regularization_term_method(self): return conf.instance["general"]["inversion"]["regularization_term_method"] return self._regularization_term_method + + @property + def interferometer_numba_nnz_per_source_max(self) -> float: + """ + The geometry gate above which the numba `direct_conv` interferometer curvature + path is not used. + + `InversionInterferometerSparseNumba` convolves each source column of the mapping + operator over the extent rectangle, at a cost that scales with the column's + non-zeros; the FFT route (`InversionInterferometerSparse`) costs the same whatever + the density. The two therefore cross at a roughly fixed number of non-zeros per + source column, `mapper.pix_sizes_for_sub_slim_index.sum() / mapper.params`, and + the factory routes to numba only at or below this value. + + The measured crossovers are **~60 non-zeros per source column on Delaunay meshes** + and **~77 on rectangular meshes** (autolens_profiling issue #226 verdict, section + 2), where the numba kernel is 2-7x faster than JAX-CPU well below the crossover. + The default is the conservative of the two. + + This is a **machine-dependent constant**: it is set by the ratio of scalar AXPY + throughput to FFT throughput on the CPU running the fit, so a machine with a very + different cache hierarchy or FFT library will cross somewhere else. Re-measure + before tuning it for a new machine; `0` disables the numba path entirely. + """ + if self._interferometer_numba_nnz_per_source_max is None: + try: + return conf.instance["general"]["inversion"][ + "interferometer_numba_nnz_per_source_max" + ] + except KeyError: + # A workspace `general.yaml` normally omits this key, so autoconf's + # config-path list falls through to autoarray's packaged value (`60.0`). + # This fallback fires only when the workspace config is the sole config + # path (isolated test configs push one dir) and returns that same value, + # so both routes resolve identically. + return 60.0 + + return self._interferometer_numba_nnz_per_source_max diff --git a/test_autoarray/inversion/inversion/interferometer_numba/test_interferometer_numba.py b/test_autoarray/inversion/inversion/interferometer_numba/test_interferometer_numba.py new file mode 100644 index 000000000..b8e1196dc --- /dev/null +++ b/test_autoarray/inversion/inversion/interferometer_numba/test_interferometer_numba.py @@ -0,0 +1,514 @@ +import numpy as np +import pytest +import types + +import autoarray as aa + +from autoarray import exc +from autoarray.inversion.inversion.factory import _use_interferometer_numba +from autoarray.inversion.inversion.interferometer.sparse import ( + InversionInterferometerSparse, +) +from autoarray.inversion.inversion.interferometer_numba.sparse import ( + InversionInterferometerSparseNumba, +) +from autoarray.inversion.inversion.interferometer_numba import ( + inversion_interferometer_numba_util as numba_util, +) + +pytest.importorskip("numba") + + +def _dataset_from(seed=0, n_visibilities=5): + """ + The 7x7 / `TransformerDFT` interferometer dataset the sparse-operator parity tests in + `interferometer/test_interferometer.py::_sparse_parity_setup` use, plus its + sparse-operator form. Duplicated here rather than imported so the numba tests do not + depend on a sibling test module's private helper. + """ + mask = aa.Mask2D( + mask=[ + [True, True, True, True, True, True, True], + [True, True, True, True, True, True, True], + [True, True, True, False, True, True, True], + [True, True, False, False, False, True, True], + [True, True, True, False, True, True, True], + [True, True, True, True, True, True, True], + [True, True, True, True, True, True, True], + ], + pixel_scales=2.0, + ) + + rng = np.random.default_rng(seed=seed) + + data = aa.Visibilities( + visibilities=rng.normal(size=(n_visibilities, 2)).astype(np.float64) + ) + noise_map = aa.VisibilitiesNoiseMap( + visibilities=np.ones((n_visibilities, 2), dtype=np.float64) + ) + uv_wavelengths = rng.normal(size=(n_visibilities, 2)).astype(np.float64) + + dataset = aa.Interferometer( + data=data, + noise_map=noise_map, + uv_wavelengths=uv_wavelengths, + real_space_mask=mask, + transformer_class=aa.TransformerDFT, + ) + + return mask, dataset, dataset.apply_sparse_operator(use_jax=False) + + +def _delaunay_mapper_from(mask, pixels=9, shape=(3, 3), over_sample_size=1): + grid = aa.Grid2D.from_mask(mask=mask, over_sample_size=over_sample_size) + + mesh = aa.mesh.Delaunay(pixels=pixels) + image_mesh = aa.image_mesh.Overlay(shape=shape) + image_mesh_grid = image_mesh.image_plane_mesh_grid_from(mask=mask, adapt_data=None) + + interpolator = mesh.interpolator_from( + source_plane_data_grid=grid, + source_plane_mesh_grid=image_mesh_grid, + ) + + return aa.Mapper( + interpolator=interpolator, + regularization=aa.reg.Constant(coefficient=1.0), + ) + + +def _rectangular_mapper_from(mask, shape=(3, 3), over_sample_size=1): + from autoarray.inversion.mesh.mesh.rectangular_rtu_adapt_density import ( + overlay_grid_from, + ) + + grid = aa.Grid2D.from_mask(mask=mask, over_sample_size=over_sample_size) + + source_plane_mesh_grid = overlay_grid_from( + shape_native=shape, grid=grid.over_sampled + ) + + mesh = aa.mesh.RectangularUniform(shape=shape) + + interpolator = mesh.interpolator_from( + source_plane_data_grid=grid, + source_plane_mesh_grid=aa.Grid2DIrregular(source_plane_mesh_grid), + adapt_data=aa.Array2D.ones(shape, pixel_scales=0.1), + ) + + return aa.Mapper( + interpolator=interpolator, + regularization=aa.reg.Constant(coefficient=1.0), + ) + + +def _assert_numba_matches_sparse(dataset_sparse, mapper): + """ + Pin the numba `direct_conv` inversion against the NumPy FFT (`xp=np`) inversion, which + is the path it replaces, at the repo's exact-parity idiom + `rtol=1e-10, atol=1e-10 * max|reference|`. + """ + inversion_sparse = InversionInterferometerSparse( + dataset=dataset_sparse, + linear_obj_list=[mapper], + xp=np, + ) + inversion_numba = InversionInterferometerSparseNumba( + dataset=dataset_sparse, + linear_obj_list=[mapper], + xp=np, + ) + + curvature_matrix = np.asarray(inversion_sparse.curvature_matrix) + atol = 1.0e-10 * np.abs(curvature_matrix).max() + + np.testing.assert_allclose( + np.asarray(inversion_numba.curvature_matrix), + curvature_matrix, + rtol=1.0e-10, + atol=atol, + ) + + data_vector = np.asarray(inversion_sparse.data_vector) + + np.testing.assert_allclose( + np.asarray(inversion_numba.data_vector), + data_vector, + rtol=1.0e-10, + atol=1.0e-10 * np.abs(data_vector).max(), + ) + + reconstruction = np.asarray(inversion_sparse.reconstruction) + + np.testing.assert_allclose( + np.asarray(inversion_numba.reconstruction), + reconstruction, + rtol=1.0e-10, + atol=1.0e-10 * np.abs(reconstruction).max(), + ) + + assert inversion_numba.log_det_curvature_reg_matrix_term == pytest.approx( + inversion_sparse.log_det_curvature_reg_matrix_term, rel=1.0e-10 + ) + assert inversion_numba.log_det_regularization_matrix_term == pytest.approx( + inversion_sparse.log_det_regularization_matrix_term, rel=1.0e-10 + ) + + return inversion_sparse, inversion_numba + + +def test__numba_inversion__delaunay__matches_sparse_numpy_inversion(): + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + _assert_numba_matches_sparse(dataset_sparse=dataset_sparse, mapper=mapper) + + +def test__numba_inversion__rectangular__matches_sparse_numpy_inversion(): + mask, _, dataset_sparse = _dataset_from(seed=1) + mapper = _rectangular_mapper_from(mask=mask) + + _assert_numba_matches_sparse(dataset_sparse=dataset_sparse, mapper=mapper) + + +def test__numba_inversion__control__one_percent_scale_of_curvature_matrix_fails_the_pin(): + """ + The parity pin above is only meaningful if it can fail: a 1% rescaling of `F` — far + smaller than any real kernel bug — must break it. + """ + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + inversion_sparse, inversion_numba = _assert_numba_matches_sparse( + dataset_sparse=dataset_sparse, mapper=mapper + ) + + curvature_matrix = np.asarray(inversion_sparse.curvature_matrix) + atol = 1.0e-10 * np.abs(curvature_matrix).max() + + with pytest.raises(AssertionError): + np.testing.assert_allclose( + 1.01 * np.asarray(inversion_numba.curvature_matrix), + curvature_matrix, + rtol=1.0e-10, + atol=atol, + ) + + +def _kernel_args_from(dataset_sparse, mapper): + inversion_numba = InversionInterferometerSparseNumba( + dataset=dataset_sparse, + linear_obj_list=[mapper], + xp=np, + ) + + inputs = inversion_numba.kernel_index_arrays + + preload = np.ascontiguousarray( + np.asarray( + dataset_sparse.sparse_operator.nufft_precision_operator, dtype=np.float64 + ) + ) + + return ( + preload, + inputs["iy"], + inputs["ix"], + inputs["flat"], + inputs["indptr"], + inputs["col"], + inputs["val"], + inputs["cscptr"], + inputs["csc_row"], + inputs["csc_val"], + inputs["ny"], + inputs["nx"], + inputs["pix_pixels"], + ) + + +def test__parallel_kernel_equals_serial_kernel(): + """ + Both kernels are called through the module functions directly rather than through the + `general.yaml -> numba -> parallel` flag, so the comparison does not depend on which + one the config would have selected. + """ + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + args = _kernel_args_from(dataset_sparse=dataset_sparse, mapper=mapper) + + curvature_matrix_serial = numba_util.curvature_direct_conv(*args) + curvature_matrix_parallel = numba_util.direct_conv_parallel_kernel()(*args) + + np.testing.assert_allclose( + curvature_matrix_parallel, + curvature_matrix_serial, + rtol=1.0e-10, + atol=1.0e-10 * np.abs(curvature_matrix_serial).max(), + ) + + +def test__kernel_inputs_from__matches_a_hand_built_expectation(): + """ + A three-pixel, two-source-pixel mapper whose CSR / CSC / extent arrays are small + enough to write out by hand. + + The extent rectangle is 2 x 3, so the mask's flat extent indices `[0, 2, 4]` are rows + `[0, 0, 1]` and columns `[0, 2, 1]`. The triplet rows are padded to width 2, and the + second pixel maps to a single source pixel, so its padding entry must be dropped. + """ + pix_indexes = np.array([[0, 1], [1, 0], [0, 1]], dtype=np.int64) + pix_sizes = np.array([2, 1, 2], dtype=np.int64) + pix_weights = np.array([[0.25, 0.75], [1.0, 0.0], [0.5, 0.5]], dtype=np.float64) + + inputs = numba_util.kernel_inputs_from( + pix_indexes_for_sub_slim_index=pix_indexes, + pix_sizes_for_sub_slim_index=pix_sizes, + pix_weights_for_sub_slim_index=pix_weights, + extent_index_for_masked_pixel=np.array([0, 2, 4], dtype=np.int64), + extent_shape=(2, 3), + pix_pixels=2, + ) + + assert inputs["ny"] == 2 + assert inputs["nx"] == 3 + assert inputs["n_pix"] == 3 + assert inputs["nnz"] == 5 + assert inputs["pix_pixels"] == 2 + + assert inputs["iy"].tolist() == [0, 0, 1] + assert inputs["ix"].tolist() == [0, 2, 1] + assert inputs["flat"].tolist() == [0, 2, 4] + + # CSR: the padding entry of row 1 is dropped, so `col`/`val` hold five entries. + assert inputs["indptr"].tolist() == [0, 2, 3, 5] + assert inputs["col"].tolist() == [0, 1, 1, 0, 1] + assert inputs["val"].tolist() == [0.25, 0.75, 1.0, 0.5, 0.5] + + # CSC: source 0 is hit by rows 0 and 2, source 1 by rows 0, 1 and 2. + assert inputs["cscptr"].tolist() == [0, 2, 5] + assert inputs["csc_row"].tolist() == [0, 2, 0, 1, 2] + assert inputs["csc_val"].tolist() == [0.25, 0.5, 0.75, 1.0, 0.5] + + +def test__nnz_per_source_column_from(): + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + expected = float(np.asarray(mapper.pix_sizes_for_sub_slim_index).sum()) / float( + mapper.params + ) + + assert numba_util.nnz_per_source_column_from(mapper=mapper) == pytest.approx( + expected, 1.0e-12 + ) + + +def test__factory__routes_to_numba_below_the_gate_and_sparse_above_it(): + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + nnz_per_source_column = numba_util.nnz_per_source_column_from(mapper=mapper) + + inversion = aa.Inversion( + dataset=dataset_sparse, + linear_obj_list=[mapper], + settings=aa.Settings( + interferometer_numba_nnz_per_source_max=nnz_per_source_column + 1.0 + ), + ) + + assert isinstance(inversion, InversionInterferometerSparseNumba) + + inversion = aa.Inversion( + dataset=dataset_sparse, + linear_obj_list=[mapper], + settings=aa.Settings( + interferometer_numba_nnz_per_source_max=nnz_per_source_column - 1.0 + ), + ) + + assert isinstance(inversion, InversionInterferometerSparse) + assert not isinstance(inversion, InversionInterferometerSparseNumba) + + +def test__factory__gate_of_zero_disables_the_numba_path(): + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + inversion = aa.Inversion( + dataset=dataset_sparse, + linear_obj_list=[mapper], + settings=aa.Settings(interferometer_numba_nnz_per_source_max=0), + ) + + assert isinstance(inversion, InversionInterferometerSparse) + assert not isinstance(inversion, InversionInterferometerSparseNumba) + + +def test__routing_predicate__falls_through_on_every_unsupported_configuration(): + """ + The routing predicate is exercised directly, because a routing miss is silent by + design: it returns the sparse inversion rather than raising, so a test that only + looked at the returned type could not tell *which* condition rejected it. + + `xp` is checked with a stand-in module rather than `jax.numpy`, so this test carries + no JAX dependency and runs on the no-JAX CI leg. + """ + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + settings = aa.Settings(interferometer_numba_nnz_per_source_max=1.0e6) + + assert _use_interferometer_numba(linear_obj_list=[mapper], settings=settings, xp=np) + + # A non-NumPy array module never routes to numba. + assert not _use_interferometer_numba( + linear_obj_list=[mapper], + settings=settings, + xp=types.SimpleNamespace(__name__="jax.numpy"), + ) + + # A linear function list has no block in the kernel. + func_list = aa.m.MockLinearObjFuncList( + parameters=1, + mapping_matrix=np.ones((mask.pixels_in_mask, 1)), + ) + + assert not _use_interferometer_numba( + linear_obj_list=[func_list, mapper], settings=settings, xp=np + ) + + # More than one mapper has no off-diagonal block in the kernel. + mapper_1 = _delaunay_mapper_from(mask=mask, pixels=4, shape=(2, 2)) + + assert not _use_interferometer_numba( + linear_obj_list=[mapper, mapper_1], settings=settings, xp=np + ) + + # Over-sampling is folded into the sparse weights but not the kernel's. + mapper_over_sampled = _delaunay_mapper_from(mask=mask, over_sample_size=2) + + assert not _use_interferometer_numba( + linear_obj_list=[mapper_over_sampled], settings=settings, xp=np + ) + + # The geometry gate. + assert not _use_interferometer_numba( + linear_obj_list=[mapper], + settings=aa.Settings(interferometer_numba_nnz_per_source_max=0.5), + xp=np, + ) + + +def test__precondition__non_numpy_array_module_raises(): + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + with pytest.raises(exc.InversionException, match="non-NumPy array module"): + InversionInterferometerSparseNumba( + dataset=dataset_sparse, + linear_obj_list=[mapper], + xp=types.SimpleNamespace(__name__="jax.numpy"), + ) + + +def test__precondition__linear_func_list_raises(): + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + func_list = aa.m.MockLinearObjFuncList( + parameters=1, + mapping_matrix=np.ones((mask.pixels_in_mask, 1)), + ) + + with pytest.raises(exc.InversionException, match="linear-function list"): + InversionInterferometerSparseNumba( + dataset=dataset_sparse, + linear_obj_list=[func_list, mapper], + xp=np, + ) + + +def test__precondition__multiple_mappers_raise(): + mask, _, dataset_sparse = _dataset_from() + mapper_0 = _delaunay_mapper_from(mask=mask) + mapper_1 = _delaunay_mapper_from(mask=mask, pixels=4, shape=(2, 2)) + + with pytest.raises(exc.InversionException, match="was passed 2 mappers"): + InversionInterferometerSparseNumba( + dataset=dataset_sparse, + linear_obj_list=[mapper_0, mapper_1], + xp=np, + ) + + +def test__precondition__over_sampling_raises(): + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask, over_sample_size=2) + + with pytest.raises(exc.InversionException, match="sub_fraction"): + InversionInterferometerSparseNumba( + dataset=dataset_sparse, + linear_obj_list=[mapper], + xp=np, + ) + + +def test__kernel_index_arrays__preload_shape_mismatch_raises(monkeypatch): + """ + The kernel indexes the preload with wrapped (negative) extent offsets, so a preload + whose shape is not `(2ny, 2nx)` would wrap silently rather than fail. The check names + both shapes. + """ + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + inversion_numba = InversionInterferometerSparseNumba( + dataset=dataset_sparse, + linear_obj_list=[mapper], + xp=np, + ) + + preload = np.asarray(dataset_sparse.sparse_operator.nufft_precision_operator) + + monkeypatch.setattr( + type(dataset_sparse.sparse_operator), + "nufft_precision_operator", + property(lambda self: preload[:-2, :]), + raising=False, + ) + + with pytest.raises(exc.InversionException, match="does not match the sparse"): + _ = inversion_numba.kernel_index_arrays + + +def test__existing_sparse_path_is_unchanged_when_the_gate_rejects_the_geometry(): + """ + The default interferometer route must be untouched by the new class: with the gate + closed, the inversion the factory builds is the sparse one and its matrices are + identical to those of a directly constructed `InversionInterferometerSparse`. + """ + mask, _, dataset_sparse = _dataset_from() + mapper = _delaunay_mapper_from(mask=mask) + + inversion_factory = aa.Inversion( + dataset=dataset_sparse, + linear_obj_list=[mapper], + settings=aa.Settings(interferometer_numba_nnz_per_source_max=0), + ) + inversion_direct = InversionInterferometerSparse( + dataset=dataset_sparse, + linear_obj_list=[mapper], + xp=np, + ) + + assert type(inversion_factory) is InversionInterferometerSparse + + np.testing.assert_allclose( + np.asarray(inversion_factory.curvature_matrix), + np.asarray(inversion_direct.curvature_matrix), + rtol=1.0e-12, + atol=0.0, + )