From 1b066ccd26c2be1dd04e1e239818fa99da7c2818 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Thu, 1 Oct 2026 09:16:53 +0100 Subject: [PATCH] feat: per-interface data_term override and sparse profile-term identity (streaming P4) Array-free interferometer fits with ordinary light profiles need the chi-squared data term of the profile-subtracted visibilities without forming them. - DatasetInterface gains `data_term`; InversionInterferometer.fast_chi_squared prefers it over `sparse_operator.data_term` when `data is None`, so a fit that subtracted light profiles can never silently fall back to the unsubtracted scalar. - inversion_interferometer_util.sparse_profile_terms_from(operator, image, ...) returns (W~ i_p, d~ - W~ i_p, data_term - 2 i_p.d~ + i_p.W~ i_p) from a single W~ product; jit-safe, traced with the image. - FitInterferometer gains an overridable `sparse_chi_squared` hook (None by default); `chi_squared` returns it on an array-free dataset before raising, so subclasses can provide log_likelihood array-free. Maps still raise. - Tests: identity vs the dense subtracted visibilities (rel 1e-10), override used / operator scalar when absent / ignored when data given, numpy == jax.jit, and the fit hook vs the in-memory fit. Refs PyAutoLabs/PyAutoArray#598 Co-Authored-By: Claude Fable 5.1 --- autoarray/fit/fit_interferometer.py | 27 +++ .../inversion/inversion/dataset_interface.py | 12 ++ .../inversion/interferometer/abstract.py | 28 ++- .../inversion_interferometer_util.py | 72 +++++++ test_autoarray/fit/test_fit_interferometer.py | 64 ++++++ .../interferometer/test_interferometer.py | 204 ++++++++++++++++++ 6 files changed, 397 insertions(+), 10 deletions(-) diff --git a/autoarray/fit/fit_interferometer.py b/autoarray/fit/fit_interferometer.py index 802d3c934..66c46c9f1 100644 --- a/autoarray/fit/fit_interferometer.py +++ b/autoarray/fit/fit_interferometer.py @@ -1,5 +1,6 @@ import functools import numpy as np +from typing import Optional from autoarray.dataset.interferometer.dataset import Interferometer @@ -206,11 +207,37 @@ def signal_to_noise_map(self) -> np.ndarray: return signal_to_noise_map_real + 1.0j * signal_to_noise_map_imag + @property + def sparse_chi_squared(self) -> Optional[float]: + """ + The chi-squared of this fit computed from the dataset's `sparse_operator` without any visibility-sized + array, used by `chi_squared` when the dataset is array-free (built by `from_stream` / + `from_sparse_terms`, so `data` is `None`). + + `None` by default: a bare `FitInterferometer` has no real-space model image to evaluate it from. + Subclasses whose model visibilities are the Fourier transform `p = F i_p` of a real-space image `i_p` + (e.g. the light-profile fits of PyAutoGalaxy and PyAutoLens without an inversion) override it with + `data_term - 2 i_p^T d~ + i_p^T W~ i_p` + (`inversion_interferometer_util.sparse_profile_terms_from`), so `log_likelihood` and + `figure_of_merit` work array-free. The residual and chi-squared *maps* still need the visibilities and + raise on such a dataset. + """ + return None + @property def chi_squared(self) -> float: """ Returns the chi-squared terms of the model data's fit to an dataset, by summing the chi-squared-map. + + On an array-free dataset (no `data`) this is `sparse_chi_squared` when a subclass provides it, and + otherwise raises an `exc.DatasetException`. """ + if self.data is None: + sparse_chi_squared = self.sparse_chi_squared + + if sparse_chi_squared is not None: + return sparse_chi_squared + self._require("chi_squared", "data", "noise_map") return fit_util.chi_squared_complex_from( chi_squared_map=self.chi_squared_map.array, diff --git a/autoarray/inversion/inversion/dataset_interface.py b/autoarray/inversion/inversion/dataset_interface.py index 032bcb2b5..782e0de56 100644 --- a/autoarray/inversion/inversion/dataset_interface.py +++ b/autoarray/inversion/inversion/dataset_interface.py @@ -9,6 +9,7 @@ def __init__( sparse_operator=None, noise_covariance_matrix=None, sparse_dirty_image=None, + data_term=None, ): """ Generic class which acts as an interface between a dataset and an inversion. @@ -65,6 +66,16 @@ def __init__( 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. + data_term + The chi-squared data term `sum(d_r^2/sigma_r^2) + sum(d_i^2/sigma_i^2)` of *this interface's* + (possibly profile-subtracted) visibilities, read by the sparse interferometer inversion's + `fast_chi_squared` when `data` is `None` in preference to the `sparse_operator`'s cached scalar (which + is the data term of the raw, unsubtracted visibilities). It is how a fit with ordinary light profiles + on an array-free dataset passes `data=None`: the light profiles' visibilities `F i_p` are never + formed, and the subtracted data term `data_term - 2 i_p^T d~ + i_p^T W~ i_p` is computed by + `inversion_interferometer_util.sparse_profile_terms_from` alongside the subtracted + `sparse_dirty_image`. If `None`, the operator's cached scalar is used. A scalar (traced under + `jax.jit` when the profile image is). """ self.data = data self.noise_map = noise_map @@ -74,6 +85,7 @@ def __init__( self.sparse_operator = sparse_operator self.noise_covariance_matrix = noise_covariance_matrix self.sparse_dirty_image = sparse_dirty_image + self.data_term = data_term @property def mask(self): diff --git a/autoarray/inversion/inversion/interferometer/abstract.py b/autoarray/inversion/inversion/interferometer/abstract.py index 8915c6532..9db642643 100644 --- a/autoarray/inversion/inversion/interferometer/abstract.py +++ b/autoarray/inversion/inversion/interferometer/abstract.py @@ -191,10 +191,14 @@ def fast_chi_squared(self): where `s` is the reconstruction vector, `F` is the curvature matrix, `D` is the data vector, and `d_r`/`d_i` are the real/imaginary parts of the observed visibilities. - When the dataset interface's `data` is `None` the third term is read from the scalar - `sparse_operator.data_term` cached when the operator was built, so no visibility array - is reduced over. That is only correct when the data fitted is the raw data the operator - was built from (nothing subtracted), which is the contract of passing `data=None`. + When the dataset interface's `data` is `None` the third term is a scalar, so no visibility + array is reduced over. It is the interface's own `data_term` when one is given (the data + term of profile-subtracted visibilities, computed via the identity in + `inversion_interferometer_util.sparse_profile_terms_from`), and otherwise the + `sparse_operator.data_term` cached when the operator was built -- which is only correct when + the data fitted is the raw data the operator was built from (nothing subtracted). The + interface's `data_term` always takes precedence, so a fit that subtracted light profiles can + never silently fall back to the unsubtracted scalar. This avoids computing the full mapped reconstructed visibilities and is faster than computing `chi_squared` via the residual visibilities when many source pixels are used. @@ -217,12 +221,16 @@ def fast_chi_squared(self): ) if self.dataset.data is None: - # The interface carries no visibilities (e.g. a pixelization-only fit on the sparse - # path, where nothing was subtracted from the data), so term 3 is the scalar - # `d^T N^-1 d` the sparse operator cached when it was built from the raw data. - chi_squared_term_3 = getattr( - self.dataset.sparse_operator, "data_term", None - ) + # The interface carries no visibilities, so term 3 is a precomputed scalar: the + # interface's own `data_term` (that of profile-subtracted visibilities) when given, + # else the `d^T N^-1 d` the sparse operator cached when it was built from the raw data + # (e.g. a pixelization-only fit, where nothing was subtracted). + chi_squared_term_3 = getattr(self.dataset, "data_term", None) + + if chi_squared_term_3 is None: + chi_squared_term_3 = getattr( + self.dataset.sparse_operator, "data_term", None + ) if chi_squared_term_3 is None: raise exc.InversionException( diff --git a/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py b/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py index 5bf8fd0e0..25360e114 100644 --- a/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py +++ b/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py @@ -1804,6 +1804,78 @@ def curvature_matrix_func_list_from( return curvature_weights_0.T @ operated +def sparse_profile_terms_from( + sparse_operator: "InterferometerSparseOperator", + image, + extent_index_for_masked_pixel, + xp=np, +): + """ + Returns the sparse-path quantities of fitting visibilities from which the visibilities `p = F i_p` of a + real-space image `i_p` (e.g. the summed image of a fit's ordinary light profiles) are subtracted, computed + without forming `p` or touching any visibility-sized array. + + With `d` the visibilities the `sparse_operator` was built from, `W = 1 / sigma^2` their (equal real and + imaginary) inverse variances, `d~ = Re(F^H W d)` the operator's cached `dirty_image`, `W~ = Re(F^H W F)` the + operator itself and `data_term = sum(|d|^2 / sigma^2)` its cached scalar, linearity gives: + + Re(F^H W (d - F i_p)) = d~ - W~ i_p + sum(|d - F i_p|^2 / sigma^2) = data_term - 2 i_p^T d~ + i_p^T W~ i_p + + The first is the dirty image a sparse inversion of the profile-subtracted visibilities forms its data vector + from; the second is their chi-squared data term (term 3 of `fast_chi_squared`), and on its own the + chi-squared of a fit whose model is only `p`. Both reuse the single product `W~ i_p` (one FFT convolution on + the real-space grid), so the data term costs two dot products on top of the dirty image. + + The data term is a difference of large numbers at high signal-to-noise (`data_term >> chi^2`); in float64 + this is accurate to ~1e-15 relative to `data_term`, ample for likelihood comparisons. + + Every operation is an `xp` array operation, so under `jax.jit` the returned scalar is traced with `i_p`. + + Parameters + ---------- + sparse_operator + The `InterferometerSparseOperator` of the dataset, carrying the `dirty_image` and `data_term` of the + visibilities it was built from. + image + The image `i_p` on the slim masked real-space grid (an `Array2D` or a plain array). + extent_index_for_masked_pixel + The `real_space_mask.extent_index_for_masked_pixel` mapping slim masked pixels to the operator's + rectangular extent grid. + xp + The array module (`numpy` or `jax.numpy`). + + Returns + ------- + operated_image + `W~ i_p` on the slim masked grid. + sparse_dirty_image + `d~ - W~ i_p`. + data_term + `data_term - 2 i_p^T d~ + i_p^T W~ i_p`, or `None` if the operator carries no `data_term`. + """ + image = getattr(image, "array", image) + + operated_image = sparse_operator.operated_matrix_slim_from( + matrix_slim=image[:, None], + extent_index_for_masked_pixel=extent_index_for_masked_pixel, + xp=xp, + )[:, 0] + + dirty_image = xp.asarray(sparse_operator.dirty_image) + + sparse_dirty_image = dirty_image - operated_image + + data_term = getattr(sparse_operator, "data_term", None) + + if data_term is not None: + data_term = ( + data_term - 2.0 * xp.dot(image, dirty_image) + xp.dot(image, operated_image) + ) + + return operated_image, sparse_dirty_image, data_term + + @dataclass(frozen=True) class SparseTerms: """ diff --git a/test_autoarray/fit/test_fit_interferometer.py b/test_autoarray/fit/test_fit_interferometer.py index cb9cf0715..91878451b 100644 --- a/test_autoarray/fit/test_fit_interferometer.py +++ b/test_autoarray/fit/test_fit_interferometer.py @@ -568,3 +568,67 @@ def test__dirty_model_image_natural_from__no_sparse_operator__raises( with pytest.raises(aa.exc.DatasetException, match="sparse_operator"): dirty_model_image_natural_from(dataset=interferometer_7, image=image) + + +class _ProfileImageFit(aa.m.MockFitInterferometer): + """ + A fit whose model visibilities are the Fourier transform of a real-space `image`, overriding the + `sparse_chi_squared` hook with the data-term identity as PyAutoGalaxy / PyAutoLens light-profile fits do. + """ + + def __init__(self, dataset, image): + super().__init__(dataset=dataset) + self.image = image + + @property + def sparse_chi_squared(self): + return aa.util.inversion_interferometer.sparse_profile_terms_from( + sparse_operator=self.dataset.sparse_operator, + image=self.image, + extent_index_for_masked_pixel=self.dataset.real_space_mask.extent_index_for_masked_pixel, + )[2] + + +def test__fit_interferometer__array_free_dataset__sparse_chi_squared_hook(): + """ + On an array-free dataset `chi_squared` (hence `log_likelihood`) is read from the `sparse_chi_squared` hook + when a subclass provides it, matching the in-memory fit of the same model visibilities; the maps still + raise. The hook is not consulted when the fit has data. + """ + pytest.importorskip("nufftax") + + dataset_memory, dataset_stream, _ = _array_free_fit_setup() + + mask = dataset_memory.real_space_mask + + image = aa.Array2D( + values=np.random.default_rng(seed=2).normal(size=mask.pixels_in_mask), + mask=mask, + ) + + fit_memory = aa.m.MockFitInterferometer( + dataset=dataset_memory, + model_data=dataset_memory.transformer.visibilities_from(image=image), + ) + fit_stream = _ProfileImageFit(dataset=dataset_stream, image=image) + + assert fit_stream.chi_squared == pytest.approx(fit_memory.chi_squared, rel=1.0e-8) + assert fit_stream.log_likelihood == pytest.approx( + fit_memory.log_likelihood, rel=1.0e-8 + ) + assert fit_stream.figure_of_merit == pytest.approx( + fit_memory.figure_of_merit, rel=1.0e-8 + ) + + for name in ("residual_map", "chi_squared_map", "normalized_residual_map"): + with pytest.raises(aa.exc.DatasetException, match="array-free"): + getattr(fit_stream, name) + + # With data present the hook is ignored and the chi-squared-map is summed. + fit_memory_hook = _ProfileImageFit(dataset=dataset_memory, image=None) + fit_memory_hook._model_data = fit_memory.model_data + + assert fit_memory_hook.chi_squared == fit_memory.chi_squared + + # The base class provides no hook. + assert aa.m.MockFitInterferometer(dataset=dataset_stream).sparse_chi_squared is None diff --git a/test_autoarray/inversion/inversion/interferometer/test_interferometer.py b/test_autoarray/inversion/inversion/interferometer/test_interferometer.py index 262b9b218..df33d8ae3 100644 --- a/test_autoarray/inversion/inversion/interferometer/test_interferometer.py +++ b/test_autoarray/inversion/inversion/interferometer/test_interferometer.py @@ -1369,6 +1369,210 @@ def test__fast_chi_squared__data_none__jax_matches_numpy(): ) +def _profile_subtracted_from(dataset_sparse, seed=5): + """ + A random real-space image `i_p` on the dataset's masked grid, the visibilities `d - F i_p` with its + Fourier transform subtracted, and the sparse dirty image / data term of those visibilities computed by + `sparse_profile_terms_from` without forming `F i_p`. + """ + mask = dataset_sparse.real_space_mask + + rng = np.random.default_rng(seed=seed) + + image = aa.Array2D(values=rng.normal(size=mask.pixels_in_mask), mask=mask) + + subtracted = aa.Visibilities( + visibilities=dataset_sparse.data.array + - dataset_sparse.transformer.visibilities_from(image=image).array + ) + + operated_image, sparse_dirty_image, data_term = ( + aa.util.inversion_interferometer.sparse_profile_terms_from( + sparse_operator=dataset_sparse.sparse_operator, + image=image, + extent_index_for_masked_pixel=mask.extent_index_for_masked_pixel, + ) + ) + + return image, subtracted, operated_image, sparse_dirty_image, data_term + + +def test__sparse_profile_terms_from__matches_the_dense_subtracted_visibilities(): + """ + The identity `sum(|d - F i_p|^2 / sigma^2) = data_term - 2 i_p^T d~ + i_p^T W~ i_p` and the dirty image + `d~ - W~ i_p` of the subtracted visibilities, against the same quantities reduced from `d - F i_p`. + """ + dataset_sparse, _ = _sparse_interface_setup() + + image, subtracted, operated_image, sparse_dirty_image, data_term = ( + _profile_subtracted_from(dataset_sparse) + ) + + noise_map = dataset_sparse.noise_map.array + + data_term_dense = np.sum( + subtracted.array.real**2.0 / noise_map.real**2.0 + ) + np.sum(subtracted.array.imag**2.0 / noise_map.imag**2.0) + + assert data_term == pytest.approx(data_term_dense, rel=1.0e-10) + + sparse_dirty_image_dense = dataset_sparse.transformer.image_from( + visibilities=aa.Visibilities( + visibilities=subtracted.array.real * noise_map.real**-2.0 + + 1j * subtracted.array.imag * noise_map.imag**-2.0 + ) + ).array + + np.testing.assert_allclose( + sparse_dirty_image, sparse_dirty_image_dense, rtol=1.0e-10, atol=1.0e-10 + ) + np.testing.assert_allclose( + operated_image, + np.asarray(dataset_sparse.sparse_operator.dirty_image) - sparse_dirty_image, + rtol=1.0e-12, + atol=1.0e-12, + ) + + +def test__sparse_profile_terms_from__no_operator_data_term__returns_none(): + dataset_sparse, _ = _sparse_interface_setup() + + operator = dataset_sparse.sparse_operator + + operator_without_scalars = ( + aa.InterferometerSparseOperator.from_nufft_precision_operator( + nufft_precision_operator=operator.nufft_precision_operator, + dirty_image=operator.dirty_image, + ) + ) + + mask = dataset_sparse.real_space_mask + + _, sparse_dirty_image, data_term = ( + aa.util.inversion_interferometer.sparse_profile_terms_from( + sparse_operator=operator_without_scalars, + image=np.ones(mask.pixels_in_mask), + extent_index_for_masked_pixel=mask.extent_index_for_masked_pixel, + ) + ) + + assert data_term is None + assert sparse_dirty_image.shape == (mask.pixels_in_mask,) + + +def test__fast_chi_squared__data_none__interface_data_term_overrides_the_operator(): + """ + With `data=None` the interface's own `data_term` (that of profile-subtracted visibilities) takes precedence + over the operator's scalar (that of the raw visibilities), so the sparse inversion of the subtracted + visibilities passed as scalars equals the one passed the visibilities themselves; without it the + operator's scalar is used. + """ + dataset_sparse, mapper = _sparse_interface_setup() + + _, subtracted, _, sparse_dirty_image, data_term = _profile_subtracted_from( + dataset_sparse + ) + + def interface_from(data, data_term): + return aa.DatasetInterface( + data=data, + noise_map=dataset_sparse.noise_map, + grids=dataset_sparse.grids, + transformer=dataset_sparse.transformer, + sparse_operator=dataset_sparse.sparse_operator, + sparse_dirty_image=sparse_dirty_image, + data_term=data_term, + ) + + inversion_array = aa.Inversion( + dataset=interface_from(data=subtracted, data_term=None), + linear_obj_list=[mapper], + ) + inversion_override = aa.Inversion( + dataset=interface_from(data=None, data_term=data_term), + linear_obj_list=[mapper], + ) + inversion_operator = aa.Inversion( + dataset=interface_from(data=None, data_term=None), + linear_obj_list=[mapper], + ) + + assert isinstance(inversion_override, aa.InversionInterferometerSparse) + + assert inversion_override.fast_chi_squared == pytest.approx( + inversion_array.fast_chi_squared, rel=1.0e-10 + ) + + # Absent the override, term 3 is the operator's (unsubtracted) data term. + difference = ( + inversion_operator.fast_chi_squared - inversion_override.fast_chi_squared + ) + + assert difference == pytest.approx( + dataset_sparse.sparse_operator.data_term - data_term, rel=1.0e-10 + ) + assert abs(difference) > 1.0e-3 + + # When the visibilities are passed, they are reduced over and the override is not read. + inversion_array_with_override = aa.Inversion( + dataset=interface_from(data=subtracted, data_term=0.0), + linear_obj_list=[mapper], + ) + + assert ( + inversion_array_with_override.fast_chi_squared + == inversion_array.fast_chi_squared + ) + + +def test__fast_chi_squared__data_none__interface_data_term__jax_jit_matches_numpy(): + jax = pytest.importorskip("jax") + + import jax.numpy as jnp + + dataset_sparse, mapper = _sparse_interface_setup() + + mask = dataset_sparse.real_space_mask + + image = np.random.default_rng(seed=5).normal(size=mask.pixels_in_mask) + + def fast_chi_squared_from(image, xp): + _, sparse_dirty_image, data_term = ( + aa.util.inversion_interferometer.sparse_profile_terms_from( + sparse_operator=dataset_sparse.sparse_operator, + image=image, + extent_index_for_masked_pixel=mask.extent_index_for_masked_pixel, + xp=xp, + ) + ) + + inversion = aa.Inversion( + dataset=aa.DatasetInterface( + data=None, + noise_map=dataset_sparse.noise_map, + grids=dataset_sparse.grids, + transformer=dataset_sparse.transformer, + sparse_operator=dataset_sparse.sparse_operator, + sparse_dirty_image=sparse_dirty_image, + data_term=data_term, + ), + linear_obj_list=[mapper], + xp=xp, + ) + + return inversion.fast_chi_squared + + fast_chi_squared_numpy = fast_chi_squared_from(image, xp=np) + + fast_chi_squared_jax = jax.jit(lambda i: fast_chi_squared_from(i, xp=jnp))( + jnp.asarray(image) + ) + + assert float(fast_chi_squared_jax) == pytest.approx( + float(fast_chi_squared_numpy), rel=1.0e-8 + ) + + def _array_free_setup(n_visibilities=60, seed=3): """ A NUFFT dataset with random data and non-uniform (equal real/imaginary) noise, its