From b5ef2e89eaeb7ab4b240adb0e3f035f8104323ec Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Sat, 26 Sep 2026 11:53:14 +0100 Subject: [PATCH] feat: route func-list-only interferometer inversions through the sparse operator (#575) - factory: an interferometer dataset with a sparse operator now selects InversionInterferometerSparse for every linear object list, including MGE-only (no mapper); imaging routing unchanged (commented why). - DatasetInterface gains sparse_dirty_image; the sparse data vector uses it when set, else the operator's cached dirty image. - TransformerDFT.image_from / image_direct_from take xp so the adjoint DFT traces under jax.jit. - Tests: func-list-only sparse vs dense parity (1e-10), sparse_dirty_image override with a cached-image control, JAX-vs-NumPy func-list-only case, factory routing, DFT image_from under jit. Refs PyAutoLabs/PyAutoArray#575 Co-Authored-By: Claude Opus 5.5 --- .../inversion/inversion/dataset_interface.py | 9 + autoarray/inversion/inversion/factory.py | 19 +- .../inversion/interferometer/sparse.py | 14 +- autoarray/operators/transformer.py | 5 +- autoarray/operators/transformer_util.py | 17 +- .../interferometer/test_interferometer.py | 178 ++++++++++++++++++ .../inversion/inversion/test_factory.py | 37 ++++ test_autoarray/operators/test_transformer.py | 28 +++ 8 files changed, 288 insertions(+), 19 deletions(-) diff --git a/autoarray/inversion/inversion/dataset_interface.py b/autoarray/inversion/inversion/dataset_interface.py index 5cc3b2163..d5abf0101 100644 --- a/autoarray/inversion/inversion/dataset_interface.py +++ b/autoarray/inversion/inversion/dataset_interface.py @@ -8,6 +8,7 @@ def __init__( transformer=None, sparse_operator=None, noise_covariance_matrix=None, + sparse_dirty_image=None, ): """ Generic class which acts as an interface between a dataset and an inversion. @@ -49,6 +50,13 @@ def __init__( noise_covariance_matrix A noise-map covariance matrix representing the covariance between noise in every `data` value, which can be used via a bespoke fit to account for correlated noise in the data. + sparse_dirty_image + The noise-weighted dirty image `Re(Fᴴ W d)` of this interface's `data`, used by the sparse (w-tilde) + interferometer inversion to form its data vector. The `sparse_operator` caches the dirty image of the + visibilities it was built from, so this is only needed when `data` differs from them (e.g. when the + visibilities of ordinary light profiles have been subtracted). If `None`, the operator's cached dirty + image is used. This is distinct from `Interferometer.dirty_image`, the unweighted dirty image of the + data used for visualization. """ self.data = data self.noise_map = noise_map @@ -57,6 +65,7 @@ def __init__( self.transformer = transformer self.sparse_operator = sparse_operator self.noise_covariance_matrix = noise_covariance_matrix + self.sparse_dirty_image = sparse_dirty_image @property def mask(self): diff --git a/autoarray/inversion/inversion/factory.py b/autoarray/inversion/inversion/factory.py index 745a60447..f6fd4555e 100644 --- a/autoarray/inversion/inversion/factory.py +++ b/autoarray/inversion/inversion/factory.py @@ -125,6 +125,10 @@ def inversion_imaging_from( An `Inversion` whose type is determined by the input `dataset` and `settings`. """ + # Unlike `inversion_interferometer_from`, imaging keeps func-list-only inversions (e.g. an MGE with no + # pixelization) on the mapping formalism. A func-list-only imaging mapping inversion is cheap (one PSF + # convolution per basis function), and the imaging sparse data vector's handling of profile-subtracted + # data is tracked separately, so the interferometer routing change is deliberately not mirrored here. use_sparse_operator = True if all( @@ -199,15 +203,12 @@ def inversion_interferometer_from( ------- An `Inversion` whose type is determined by the input `dataset` and `settings`. """ - use_sparse_operator = True - - if all( - isinstance(linear_obj, AbstractLinearObjFuncList) - for linear_obj in linear_obj_list - ): - use_sparse_operator = False - - if dataset.sparse_operator is not None and use_sparse_operator: + # A sparse operator selects the w-tilde formalism for every linear object list, including func-list-only + # ones (e.g. an MGE with no pixelization). For interferometer data this avoids the dense + # `transform_mapping_matrix` of shape (N_vis, S), which dominates run time and memory at large N_vis, and + # `InversionInterferometerSparse` handles zero mappers (the func-func curvature blocks and the data vector + # need no mapper). `_use_interferometer_numba` returns False whenever a func-list is present. + if dataset.sparse_operator is not None: if _use_interferometer_numba( linear_obj_list=linear_obj_list, diff --git a/autoarray/inversion/inversion/interferometer/sparse.py b/autoarray/inversion/inversion/interferometer/sparse.py index 8d7185e14..4da505ae1 100644 --- a/autoarray/inversion/inversion/interferometer/sparse.py +++ b/autoarray/inversion/inversion/interferometer/sparse.py @@ -88,10 +88,18 @@ def data_vector(self) -> np.ndarray: which is exactly the entry the mapping (dense) formalism computes via the linear function's transformed mapping matrix. Linear function lists therefore require no separate branch here. + + The cached dirty image is that of the visibilities the `sparse_operator` was built from. When + the inversion's data differs from them (e.g. a `DatasetInterface` whose data has the visibilities + of ordinary light profiles subtracted), the dataset's `sparse_dirty_image` of that data is used + instead, so that `D` is consistent with the data the chi-squared is computed from. """ - return self._xp.dot( - self.mapping_matrix.T, self.dataset.sparse_operator.dirty_image - ) + dirty_image = getattr(self.dataset, "sparse_dirty_image", None) + + if dirty_image is None: + dirty_image = self.dataset.sparse_operator.dirty_image + + return self._xp.dot(self.mapping_matrix.T, dirty_image) def _sparse_triplets_curvature_from(self, mapper: Mapper): """ diff --git a/autoarray/operators/transformer.py b/autoarray/operators/transformer.py index ddea2b7df..6ff5a878c 100644 --- a/autoarray/operators/transformer.py +++ b/autoarray/operators/transformer.py @@ -199,9 +199,12 @@ def image_from(self, visibilities: Visibilities, xp=np) -> Array2D: mask as this transformer's `real_space_mask`. """ image_slim = transformer_util.image_direct_from( - visibilities=visibilities.in_array, + visibilities=xp.stack( + (xp.real(visibilities.array), xp.imag(visibilities.array)), axis=-1 + ), grid_radians=self.grid.array, uv_wavelengths=self.uv_wavelengths, + xp=xp, ) image_native = array_2d_util.array_2d_native_from( diff --git a/autoarray/operators/transformer_util.py b/autoarray/operators/transformer_util.py index 56b1614e6..4e35e29e0 100644 --- a/autoarray/operators/transformer_util.py +++ b/autoarray/operators/transformer_util.py @@ -50,7 +50,10 @@ def visibilities_from( def image_direct_from( - visibilities: np.ndarray, grid_radians: np.ndarray, uv_wavelengths: np.ndarray + visibilities: np.ndarray, + grid_radians: np.ndarray, + uv_wavelengths: np.ndarray, + xp=np, ) -> np.ndarray: """ Reconstruct a real-valued sky image from complex interferometric visibilities @@ -69,6 +72,8 @@ def image_direct_from( uv_wavelengths The (u, v) spatial frequencies in units of wavelengths for each baseline. + xp + The array module (`numpy` or `jax.numpy`) the transform is computed with. Returns ------- @@ -78,15 +83,15 @@ def image_direct_from( # Compute the phase term for each (pixel, visibility) pair phase = ( 2.0 - * np.pi + * xp.pi * ( - np.outer(grid_radians[:, 1], uv_wavelengths[:, 0]) - + np.outer(grid_radians[:, 0], uv_wavelengths[:, 1]) + xp.outer(grid_radians[:, 1], uv_wavelengths[:, 0]) + + xp.outer(grid_radians[:, 0], uv_wavelengths[:, 1]) ) ) - real_part = np.dot(np.cos(phase), visibilities[:, 0]) - imag_part = np.dot(np.sin(phase), visibilities[:, 1]) + real_part = xp.dot(xp.cos(phase), visibilities[:, 0]) + imag_part = xp.dot(xp.sin(phase), visibilities[:, 1]) image_1d = real_part - imag_part diff --git a/test_autoarray/inversion/inversion/interferometer/test_interferometer.py b/test_autoarray/inversion/inversion/interferometer/test_interferometer.py index 7cf1bf1cc..e879c7a96 100644 --- a/test_autoarray/inversion/inversion/interferometer/test_interferometer.py +++ b/test_autoarray/inversion/inversion/interferometer/test_interferometer.py @@ -563,6 +563,133 @@ def test__interferometer_sparse_operator__func_list_and_mapper__identical_to_map ) +def test__interferometer_sparse_operator__func_list_only__identical_to_mapping(): + """ + A linear function list with no mapper (e.g. an MGE with no pixelization) is routed to the sparse + path whenever the dataset has a sparse operator, where its curvature matrix is the func-func block + `Bᵀ W~ B` alone and its data vector is `Bᵀ d~`, reproducing the dense (mapping formalism) inversion. + """ + mask, rng, dataset, dataset_sparse = _sparse_parity_setup() + + linear_obj = aa.m.MockLinearObjFuncList( + parameters=3, + mapping_matrix=rng.normal(size=(mask.pixels_in_mask, 3)), + ) + + inversion_sparse = aa.Inversion( + dataset=dataset_sparse, linear_obj_list=[linear_obj] + ) + inversion_mapping = aa.Inversion(dataset=dataset, linear_obj_list=[linear_obj]) + + assert type(inversion_sparse) is aa.InversionInterferometerSparse + assert isinstance(inversion_mapping, aa.InversionInterferometerMapping) + + # The reconstruction comes out of a linear solve of this small, poorly conditioned system and is + # compared (to its looser tolerance) by `_assert_sparse_matches_mapping` below. + for name in ("curvature_matrix", "data_vector"): + reference = np.asarray(getattr(inversion_mapping, name)) + + np.testing.assert_allclose( + np.asarray(getattr(inversion_sparse, name)), + reference, + rtol=1.0e-10, + atol=1.0e-10 * np.abs(reference).max(), + err_msg=name, + ) + + _assert_sparse_matches_mapping( + dataset=dataset, + dataset_sparse=dataset_sparse, + linear_obj_list=[linear_obj], + ) + + +def test__interferometer_sparse_operator__sparse_dirty_image_override__used_by_data_vector(): + """ + The sparse operator caches the dirty image of the visibilities it was built from. When an inversion + fits different visibilities (e.g. with the visibilities of ordinary light profiles subtracted), the + `DatasetInterface` supplies their dirty image via `sparse_dirty_image`, which the data vector must use + so that the sparse inversion reproduces the dense inversion of the subtracted visibilities. + """ + mask, rng, dataset, dataset_sparse = _sparse_parity_setup() + + linear_obj = aa.m.MockLinearObjFuncList( + parameters=2, + mapping_matrix=rng.normal(size=(mask.pixels_in_mask, 2)), + ) + + mapper = _mapper_from( + mask=mask, + pixels=9, + shape=(3, 3), + regularization=aa.reg.Constant(coefficient=1.0), + ) + + subtracted = aa.Visibilities( + visibilities=dataset.data.array + - rng.normal(size=dataset.data.shape) + - 1j * rng.normal(size=dataset.data.shape) + ) + + sparse_dirty_image = dataset.transformer.image_from( + visibilities=aa.Visibilities( + visibilities=subtracted.array.real * dataset.noise_map.array.real**-2.0 + + 1j * subtracted.array.imag * dataset.noise_map.array.imag**-2.0 + ) + ).array + + dataset_interface_mapping = aa.DatasetInterface( + data=subtracted, + noise_map=dataset.noise_map, + grids=dataset.grids, + transformer=dataset.transformer, + ) + + for linear_obj_list in ([linear_obj], [linear_obj, mapper]): + inversion_mapping = aa.Inversion( + dataset=dataset_interface_mapping, linear_obj_list=linear_obj_list + ) + + inversion_sparse = aa.Inversion( + dataset=aa.DatasetInterface( + data=subtracted, + noise_map=dataset.noise_map, + grids=dataset.grids, + transformer=dataset.transformer, + sparse_operator=dataset_sparse.sparse_operator, + sparse_dirty_image=sparse_dirty_image, + ), + linear_obj_list=linear_obj_list, + ) + + assert isinstance(inversion_sparse, aa.InversionInterferometerSparse) + + reference = np.asarray(inversion_mapping.data_vector) + + np.testing.assert_allclose( + np.asarray(inversion_sparse.data_vector), + reference, + rtol=1.0e-10, + atol=1.0e-10 * np.abs(reference).max(), + ) + + # The control: without the override the data vector is that of the unsubtracted visibilities. + inversion_sparse_cached = aa.Inversion( + dataset=aa.DatasetInterface( + data=subtracted, + noise_map=dataset.noise_map, + grids=dataset.grids, + transformer=dataset.transformer, + sparse_operator=dataset_sparse.sparse_operator, + ), + linear_obj_list=linear_obj_list, + ) + + assert np.abs( + np.asarray(inversion_sparse_cached.data_vector) - reference + ).max() > 1.0e-4 * np.abs(reference).max() + + def test__interferometer_sparse_operator__x2_mappers__identical_to_mapping(): """ Two `Mapper` objects fitted simultaneously require the mapper-mapper off-diagonal block @@ -866,3 +993,54 @@ def test__interferometer_sparse_operator__numpy_inversion_matches_jax_inversion( ).max() assert difference_sparse <= max(10.0 * difference_dense, 1.0e-14) + + # A linear function list with no mapper (e.g. an MGE with no pixelization) takes the same + # sparse path, with only the func-func curvature block and `Bᵀ d~` data vector. + linear_obj = aa.m.MockLinearObjFuncList( + parameters=3, + mapping_matrix=np.random.default_rng(seed=seed).normal( + size=(mask.pixels_in_mask, 3) + ), + ) + + inversion_np = aa.Inversion( + dataset=dataset_sparse, linear_obj_list=[linear_obj], xp=np + ) + inversion_jax = aa.Inversion( + dataset=dataset_sparse, linear_obj_list=[linear_obj], xp=jnp + ) + + assert type(inversion_np) is aa.InversionInterferometerSparse + assert type(inversion_jax) is aa.InversionInterferometerSparse + + for name in ("curvature_matrix", "data_vector"): + reference = np.asarray(getattr(inversion_jax, name)) + + np.testing.assert_allclose( + np.asarray(getattr(inversion_np, name)), + reference, + rtol=1.0e-10, + atol=1.0e-10 * np.abs(reference).max(), + err_msg=name, + ) + + # The unregularized func-list system is poorly conditioned, so the reconstruction is held to + # the same control as above: no worse than the dense inversion's NumPy / JAX spread. + difference_dense = np.abs( + np.asarray( + aa.Inversion( + dataset=dataset, linear_obj_list=[linear_obj], xp=np + ).reconstruction + ) + - np.asarray( + aa.Inversion( + dataset=dataset, linear_obj_list=[linear_obj], xp=jnp + ).reconstruction + ) + ).max() + difference_sparse = np.abs( + np.asarray(inversion_np.reconstruction) + - np.asarray(inversion_jax.reconstruction) + ).max() + + assert difference_sparse <= max(10.0 * difference_dense, 1.0e-14) diff --git a/test_autoarray/inversion/inversion/test_factory.py b/test_autoarray/inversion/inversion/test_factory.py index aded1b5b1..870fac628 100644 --- a/test_autoarray/inversion/inversion/test_factory.py +++ b/test_autoarray/inversion/inversion/test_factory.py @@ -490,6 +490,43 @@ def test__inversion_imaging__linear_obj_func_with_sparse_operator( ) +def test__inversion_interferometer__via_linear_obj_func_list__sparse_operator( + interferometer_7_no_fft, +): + mask = interferometer_7_no_fft.real_space_mask + + grid = aa.Grid2D.from_mask(mask=mask) + + linear_obj = aa.m.MockLinearObjFuncList( + parameters=2, + grid=grid, + mapping_matrix=np.random.default_rng(seed=1).normal( + size=(mask.pixels_in_mask, 2) + ), + ) + + inversion = aa.Inversion( + dataset=interferometer_7_no_fft, + linear_obj_list=[linear_obj], + ) + + assert isinstance(inversion, aa.InversionInterferometerMapping) + + # Unlike imaging, a func-list-only interferometer inversion (e.g. an MGE with no pixelization) uses the + # sparse operator when the dataset has one. + + inversion_sparse = aa.Inversion( + dataset=interferometer_7_no_fft.apply_sparse_operator(use_jax=False), + linear_obj_list=[linear_obj], + ) + + assert type(inversion_sparse) is aa.InversionInterferometerSparse + assert inversion_sparse.data_vector == pytest.approx(inversion.data_vector, 1.0e-8) + assert inversion_sparse.curvature_matrix == pytest.approx( + inversion.curvature_matrix, 1.0e-8 + ) + + def test__inversion_interferometer__via_mapper( interferometer_7_no_fft, rectangular_mapper_7x7_3x3, diff --git a/test_autoarray/operators/test_transformer.py b/test_autoarray/operators/test_transformer.py index 2802a961c..42fd8e674 100644 --- a/test_autoarray/operators/test_transformer.py +++ b/test_autoarray/operators/test_transformer.py @@ -54,6 +54,34 @@ def test__dft__image_from__visibilities_7__first_three_image_pixels_match_expect assert image[0:3] == pytest.approx([-1.49022481, -0.22395855, -0.45588535], 1.0e-4) +def test__dft__image_from__jax_jit_matches_numpy( + visibilities_7, uv_wavelengths_7x2, mask_2d_7x7 +): + """ + The adjoint DFT is used inside `jax.jit` fits (the dirty image of profile-subtracted visibilities for the + sparse operator), so it must trace with `xp=jnp` and match the NumPy result. + """ + jax = pytest.importorskip("jax") + import jax.numpy as jnp + + transformer = aa.TransformerDFT( + uv_wavelengths=uv_wavelengths_7x2, + real_space_mask=mask_2d_7x7, + ) + + @jax.jit + def f(visibilities): + return transformer.image_from( + visibilities=aa.Visibilities(visibilities=visibilities), xp=jnp + ).array + + image = transformer.image_from(visibilities=visibilities_7) + + assert np.asarray(f(jnp.asarray(visibilities_7.array))) == pytest.approx( + image.array, rel=1.0e-10, abs=1.0e-12 + ) + + def test__nufft__visibilities_from__all_ones_image__first_visibility_matches_expected(): uv_wavelengths = np.array([[0.2, 1.0], [0.5, 1.1], [0.8, 1.2]])