From 3ea8fd5f9c1834b51dea4d9ef36c4108a667aeb4 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Wed, 30 Sep 2026 09:40:33 +0100 Subject: [PATCH] feat(interferometer): pass data=None on the sparse precomputed-data-term path (#756) Mirror PyAutoGalaxy#637 in autolens: gate FitInterferometer.tracer_to_inversion on uses_precomputed_data_term_from so sparse pixelization-only lens fits build their inversion with data=None and never allocate profile_visibilities / profile_subtracted_visibilities per likelihood call. Add inversion_with_data for output paths, use it in the interferometer visualizer, and cache the two profile visibility properties. Tests mirror autogalaxy's phase-1 additions. Phase 1 of https://github.com/orgs/PyAutoLabs/discussions/13. Co-Authored-By: Claude Fable 5.1 --- autolens/interferometer/fit_interferometer.py | 65 ++++- autolens/interferometer/model/visualizer.py | 2 +- .../interferometer/test_fit_interferometer.py | 229 ++++++++++++++++++ 3 files changed, 292 insertions(+), 4 deletions(-) diff --git a/autolens/interferometer/fit_interferometer.py b/autolens/interferometer/fit_interferometer.py index 65e10adde..94a3492e9 100644 --- a/autolens/interferometer/fit_interferometer.py +++ b/autolens/interferometer/fit_interferometer.py @@ -15,6 +15,7 @@ The ``TracerToInversion`` helper is used to assemble the linear system in step 4. """ +import copy import numpy as np from typing import Dict, List, Optional @@ -27,6 +28,7 @@ from autogalaxy.interferometer.fit_interferometer import ( _has_light_profile_non_linear, sparse_dirty_image_from, + uses_precomputed_data_term_from, ) from autolens.lens.tracer import Tracer @@ -124,7 +126,7 @@ def profile_image(self) -> aa.Array2D: """ return self.tracer.image_2d_from(grid=self.grids.lp, xp=self._xp) - @property + @cached_property def profile_visibilities(self) -> aa.Visibilities: """ Returns the visibilities of every light profile in the tracer, which are computed by performing a Fourier @@ -143,7 +145,7 @@ def profile_visibilities(self) -> aa.Visibilities: shape_slim=(self.dataset.transformer.uv_wavelengths.shape[0],) ) - @property + @cached_property def profile_subtracted_visibilities(self) -> aa.Visibilities: """ Returns the interferometer dataset's visibilities with all transformed light profile images in the fit's @@ -151,10 +153,39 @@ def profile_subtracted_visibilities(self) -> aa.Visibilities: """ return self.data - self.profile_visibilities + @property + def _uses_precomputed_data_term(self) -> bool: + """ + Whether this fit's inversion reads its data term from the scalar cached on the dataset's + `sparse_operator` (see `uses_precomputed_data_term_from`), in which case `tracer_to_inversion` passes + `data=None` and the likelihood never evaluates `profile_visibilities` or + `profile_subtracted_visibilities`. + """ + return uses_precomputed_data_term_from( + dataset=self.dataset, + galaxies=self.tracer.galaxies, + data=self.data, + noise_map=self.noise_map, + ) + @property def tracer_to_inversion(self) -> TracerToInversion: + """ + Returns the object which builds this fit's inversion from its tracer's linear objects. + + The inversion fits the `profile_subtracted_visibilities`, except on the sparse path when no galaxy has an + ordinary light profile (`_uses_precomputed_data_term`): nothing is then subtracted, and `data=None` is passed + so the sparse inversion takes its data vector from the operator's cached dirty image and the data term of + its `fast_chi_squared` from the operator's cached scalar, touching no visibility-sized array. The + visibilities remain available to outputs via `fit.data`. + """ + if self._uses_precomputed_data_term: + data = None + else: + data = self.profile_subtracted_visibilities + dataset = aa.DatasetInterface( - data=self.profile_subtracted_visibilities, + data=data, noise_map=self.noise_map, grids=self.grids, transformer=self.dataset.transformer, @@ -208,6 +239,34 @@ def inversion(self) -> Optional[aa.AbstractInversion]: if self.perform_inversion: return self.tracer_to_inversion.inversion + @property + def inversion_with_data(self) -> Optional[aa.AbstractInversion]: + """ + The fit's `inversion`, guaranteed to carry the visibilities it fitted as its dataset's `data`, for + output quantities that read them (e.g. `data_subtracted_dict`, plotted by `subplot_of_mapper`). + + On the sparse path with no ordinary light profile (`_uses_precomputed_data_term`) the likelihood's + inversion is built with `data=None`, so that it touches no visibility-sized array. Nothing was subtracted + from the visibilities in that case, so the data it fitted are `fit.data`: this returns a shallow copy of + the inversion (solved first, so it shares the reconstruction and every other cached quantity) whose dataset interface carries + `fit.data`. In every other case it returns `inversion` itself. + """ + inversion = self.inversion + + if inversion is None or inversion.dataset.data is not None: + return inversion + + # Solve first, so the copy shares the reconstruction (and everything it cached) rather than repeating it. + inversion.reconstruction + + dataset = copy.copy(inversion.dataset) + dataset.data = self.data + + inversion_with_data = copy.copy(inversion) + inversion_with_data.dataset = dataset + + return inversion_with_data + @property def model_data(self) -> aa.Visibilities: """ diff --git a/autolens/interferometer/model/visualizer.py b/autolens/interferometer/model/visualizer.py index 39ecf3a2f..893c9d06b 100644 --- a/autolens/interferometer/model/visualizer.py +++ b/autolens/interferometer/model/visualizer.py @@ -157,7 +157,7 @@ def visualize( if fit.inversion is not None: try: plotter.inversion( - inversion=fit.inversion, + inversion=fit.inversion_with_data, ) except IndexError: pass diff --git a/test_autolens/interferometer/test_fit_interferometer.py b/test_autolens/interferometer/test_fit_interferometer.py index bd784777e..bda5b7486 100644 --- a/test_autolens/interferometer/test_fit_interferometer.py +++ b/test_autolens/interferometer/test_fit_interferometer.py @@ -504,3 +504,232 @@ def spy(*args, **kwargs): assert calls == [] assert profile_visibilities.shape == interferometer_7.data.shape assert np.all(profile_visibilities.array == 0.0) + + # Sparse: nothing is subtracted from the visibilities, so the likelihood must not even build the (zero) + # profile visibilities: the inversion is passed `data=None` and reads its data term from the sparse + # operator's cached scalar. + dataset_sparse = interferometer_7.apply_sparse_operator(use_jax=False) + + # The sparse dataset reuses the dense dataset's transformer, so one spy covers both. + assert dataset_sparse.transformer is interferometer_7.transformer + + zeros_calls = [] + + zeros = aa.Visibilities.zeros + + def zeros_spy(*args, **kwargs): + zeros_calls.append(1) + return zeros(*args, **kwargs) + + monkeypatch.setattr(aa.Visibilities, "zeros", zeros_spy) + + fit_sparse = al.FitInterferometer(dataset=dataset_sparse, tracer=tracer) + + figure_of_merit = fit_sparse.figure_of_merit + + assert fit_sparse._uses_precomputed_data_term + assert fit_sparse.inversion.dataset.data is None + assert calls == [] + assert zeros_calls == [] + assert "profile_visibilities" not in fit_sparse.__dict__ + assert "profile_subtracted_visibilities" not in fit_sparse.__dict__ + + # The value is the one given by passing the visibilities explicitly (the array path). + with monkeypatch.context() as m: + m.setattr( + al.FitInterferometer, + "_uses_precomputed_data_term", + property(lambda self: False), + ) + + fit_array = al.FitInterferometer(dataset=dataset_sparse, tracer=tracer) + + assert fit_array.inversion.dataset.data is not None + assert figure_of_merit == pytest.approx(fit_array.figure_of_merit, rel=1.0e-12) + + # Output paths still see real (zero) profile visibilities. + assert np.all(fit_sparse.profile_visibilities.array == 0.0) + + +def _pixelized_source_tracer(coefficient=1.0, lens_light=False): + lens = al.Galaxy( + redshift=0.5, + mass=al.mp.Isothermal(centre=(0.0, 0.0), einstein_radius=1.0), + ) + + if lens_light: + lens.bulge = al.lp.Sersic(intensity=0.1, centre=(0.05, 0.05)) + + source = al.Galaxy( + redshift=1.0, + pixelization=al.Pixelization( + mesh=al.mesh.RectangularUniform(shape=(3, 3)), + regularization=al.reg.Constant(coefficient=coefficient), + ), + ) + + return al.Tracer(galaxies=[lens, source]) + + +def test__fit_figure_of_merit__sparse_operator__pixelization_only__data_term_scalar_matches_dense( + interferometer_7, monkeypatch +): + """ + A lens with only mass and a pixelized source on the sparse path passes `data=None` to its inversion, whose + `fast_chi_squared` then reads the data term cached on the sparse operator; the log evidence must match the + dense fit, and be exactly the value obtained when the visibilities are passed explicitly. + """ + dataset_sparse = interferometer_7.apply_sparse_operator(use_jax=False) + + tracer = _pixelized_source_tracer() + + fit = al.FitInterferometer(dataset=interferometer_7, tracer=tracer) + fit_sparse = al.FitInterferometer(dataset=dataset_sparse, tracer=tracer) + + assert isinstance(fit.inversion, aa.InversionInterferometerMapping) + assert isinstance(fit_sparse.inversion, aa.InversionInterferometerSparse) + + assert fit_sparse._uses_precomputed_data_term + assert fit_sparse.inversion.dataset.data is None + + assert fit_sparse.log_likelihood == pytest.approx(fit.log_likelihood, rel=1.0e-8) + assert fit_sparse.log_evidence == pytest.approx(fit.log_evidence, rel=1.0e-8) + assert fit_sparse.noise_normalization == fit.noise_normalization + + # Output quantities that read the visibilities get them via `inversion_with_data`. + inversion_with_data = fit_sparse.inversion_with_data + + assert inversion_with_data.dataset.data is fit_sparse.data + assert inversion_with_data.reconstruction is fit_sparse.inversion.reconstruction + assert inversion_with_data.fast_chi_squared == pytest.approx( + fit_sparse.inversion.fast_chi_squared, rel=1.0e-12 + ) + + mapper = inversion_with_data.cls_list_from(cls=aa.Mapper)[0] + + np.testing.assert_array_equal( + inversion_with_data.data_subtracted_dict[mapper].array, + interferometer_7.data.array, + ) + + # The dense fit's inversion already carries its data, so it is returned unchanged. + assert fit.inversion_with_data is fit.inversion + + # Control: with the gate forced off the visibilities are passed explicitly, and the figure of merit is + # bit-identical. + figure_of_merit = fit_sparse.figure_of_merit + + monkeypatch.setattr( + "autolens.interferometer.fit_interferometer.uses_precomputed_data_term_from", + lambda **kwargs: False, + ) + + fit_forced = al.FitInterferometer(dataset=dataset_sparse, tracer=tracer) + + assert not fit_forced._uses_precomputed_data_term + assert fit_forced.inversion.dataset.data is not None + assert fit_forced.figure_of_merit == figure_of_merit + + +def test__fit_figure_of_merit__sparse_operator__pixelization_only__no_visibility_arrays( + interferometer_7, monkeypatch +): + """ + On the gated sparse path the likelihood must perform no Fourier transform and build no visibility-sized + zeros, and must never evaluate `profile_visibilities` or `profile_subtracted_visibilities`. + """ + dataset_sparse = interferometer_7.apply_sparse_operator(use_jax=False) + + calls = [] + + visibilities_from = dataset_sparse.transformer.visibilities_from + + def spy(*args, **kwargs): + calls.append(1) + return visibilities_from(*args, **kwargs) + + monkeypatch.setattr(dataset_sparse.transformer, "visibilities_from", spy) + + zeros_calls = [] + + zeros = aa.Visibilities.zeros + + def zeros_spy(*args, **kwargs): + zeros_calls.append(1) + return zeros(*args, **kwargs) + + monkeypatch.setattr(aa.Visibilities, "zeros", zeros_spy) + + fit = al.FitInterferometer( + dataset=dataset_sparse, tracer=_pixelized_source_tracer() + ) + + fit.figure_of_merit + + assert calls == [] + assert zeros_calls == [] + assert "profile_visibilities" not in fit.__dict__ + assert "profile_subtracted_visibilities" not in fit.__dict__ + + # Output paths still see real (zero) profile visibilities. + assert np.all(fit.profile_visibilities.array == 0.0) + + +def test__fit_figure_of_merit__sparse_operator__light_profile__unchanged_vs_data_passed( + interferometer_7, monkeypatch +): + """ + With a lens ordinary light profile the sparse fit must keep passing the profile-subtracted visibilities, so + its figure of merit is exactly the one computed with the data passed explicitly. + """ + dataset_sparse = interferometer_7.apply_sparse_operator(use_jax=False) + + tracer = _pixelized_source_tracer(lens_light=True) + + fit_sparse = al.FitInterferometer(dataset=dataset_sparse, tracer=tracer) + + assert not fit_sparse._uses_precomputed_data_term + + figure_of_merit = fit_sparse.figure_of_merit + + monkeypatch.setattr( + al.FitInterferometer, + "_uses_precomputed_data_term", + property(lambda self: False), + ) + + fit_forced = al.FitInterferometer(dataset=dataset_sparse, tracer=tracer) + + assert fit_forced.figure_of_merit == figure_of_merit + np.testing.assert_array_equal( + fit_sparse.inversion.dataset.data.array, + (interferometer_7.data - fit_sparse.profile_visibilities).array, + ) + + +def test__fit_figure_of_merit__sparse_operator__pixelization_only__jax_jit_matches_numpy( + interferometer_7, +): + jax = pytest.importorskip("jax") + import jax.numpy as jnp + + dataset_sparse = interferometer_7.apply_sparse_operator(use_jax=False) + + def figure_of_merit_from(coefficient, xp): + fit = al.FitInterferometer( + dataset=dataset_sparse, + tracer=_pixelized_source_tracer(coefficient=coefficient), + xp=xp, + ) + + assert fit.inversion.dataset.data is None + + return fit.figure_of_merit + + figure_of_merit_numpy = figure_of_merit_from(coefficient=1.0, xp=np) + + figure_of_merit_jax = jax.jit(lambda c: figure_of_merit_from(c, xp=jnp))(1.0) + + assert float(figure_of_merit_jax) == pytest.approx( + figure_of_merit_numpy, rel=1.0e-8 + )