diff --git a/autolens/interferometer/fit_interferometer.py b/autolens/interferometer/fit_interferometer.py index 99060015e..65e10adde 100644 --- a/autolens/interferometer/fit_interferometer.py +++ b/autolens/interferometer/fit_interferometer.py @@ -24,6 +24,10 @@ import autogalaxy as ag from autogalaxy.abstract_fit import AbstractFitInversion +from autogalaxy.interferometer.fit_interferometer import ( + _has_light_profile_non_linear, + sparse_dirty_image_from, +) from autolens.lens.tracer import Tracer from autolens.lens.to_inversion import TracerToInversion @@ -112,14 +116,31 @@ def _xp(self): return jnp return np + @cached_property + def profile_image(self) -> aa.Array2D: + """ + Returns the summed image of every ordinary (non-linear) light profile in the tracer, which is Fourier + transformed to the `profile_visibilities`. + """ + return self.tracer.image_2d_from(grid=self.grids.lp, xp=self._xp) + @property def profile_visibilities(self) -> aa.Visibilities: """ Returns the visibilities of every light profile in the tracer, which are computed by performing a Fourier transform to the sum of light profile images. + + If the tracer has no ordinary (non-linear) light profile (e.g. its light is entirely an MGE of linear + Gaussians), the image is all zeros and the Fourier transform is skipped. This is decided structurally, + so it is safe under `jax.jit`. """ - return self.tracer.visibilities_from( - grid=self.grids.lp, transformer=self.dataset.transformer, xp=self._xp + if _has_light_profile_non_linear(galaxies=self.tracer.galaxies): + return self.dataset.transformer.visibilities_from( + image=self.profile_image, xp=self._xp + ) + + return aa.Visibilities.zeros( + shape_slim=(self.dataset.transformer.uv_wavelengths.shape[0],) ) @property @@ -138,6 +159,12 @@ def tracer_to_inversion(self) -> TracerToInversion: grids=self.grids, transformer=self.dataset.transformer, sparse_operator=self.dataset.sparse_operator, + sparse_dirty_image=sparse_dirty_image_from( + dataset=self.dataset, + galaxies=self.tracer.galaxies, + image=self.profile_image, + xp=self._xp, + ), ) return TracerToInversion( diff --git a/test_autolens/interferometer/test_fit_interferometer.py b/test_autolens/interferometer/test_fit_interferometer.py index 46cef27ce..bd784777e 100644 --- a/test_autolens/interferometer/test_fit_interferometer.py +++ b/test_autolens/interferometer/test_fit_interferometer.py @@ -1,6 +1,7 @@ import numpy as np import pytest +import autoarray as aa import autolens as al @@ -386,3 +387,120 @@ def test__model_visibilities_of_planes_list(interferometer_7): + fit.galaxy_model_visibilities_dict[galaxy_pix_1].array, 1.0e-4, ) + + +def test__fit_figure_of_merit__sparse_operator__lens_light_profile_and_source_mge__matches_dense( + interferometer_7, +): + """ + With the sparse operator applied, a lens ordinary light profile plus a source MGE (linear Gaussians) + must reproduce the dense fit: the lens light's visibilities are subtracted before the inversion, so the + sparse data vector must use the dirty image of the profile-subtracted visibilities. + """ + dataset_sparse = interferometer_7.apply_sparse_operator(use_jax=False) + + lens = al.Galaxy( + redshift=0.5, + bulge=al.lp.Sersic(intensity=0.1, centre=(0.05, 0.05)), + mass=al.mp.Isothermal(centre=(0.0, 0.0), einstein_radius=1.0), + ) + source = al.Galaxy( + redshift=1.0, + bulge=al.lp_basis.Basis( + profile_list=[ + al.lp_linear.Gaussian(sigma=sigma, centre=(0.1, 0.1)) + for sigma in (0.3, 1.0, 3.0) + ] + ), + ) + + # The second case, with no lens light, is the all-linear control which uses the cached dirty image. + for galaxies, has_lens_light in ( + ([lens, source], True), + ([al.Galaxy(redshift=0.5, mass=lens.mass), source], False), + ): + tracer = al.Tracer(galaxies=galaxies) + + fit = al.FitInterferometer(dataset=interferometer_7, tracer=tracer) + fit_sparse = al.FitInterferometer(dataset=dataset_sparse, tracer=tracer) + + + assert isinstance(fit_sparse.inversion, aa.InversionInterferometerSparse) + assert isinstance(fit.inversion, aa.InversionInterferometerMapping) + + 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.inversion.dataset.sparse_dirty_image is not None + ) is has_lens_light + + # The image `i_p` the sparse dirty image is corrected with (`d~ - W~ i_p`) must be exactly the image + # the fit's `profile_visibilities` are the Fourier transform of. + profile_visibilities = tracer.visibilities_from( + grid=dataset_sparse.grids.lp, transformer=dataset_sparse.transformer + ) + + np.testing.assert_allclose( + dataset_sparse.transformer.visibilities_from( + image=fit_sparse.profile_image + ).array, + profile_visibilities.array, + rtol=1.0e-12, + atol=1.0e-12, + ) + np.testing.assert_allclose( + fit_sparse.profile_visibilities.array, + profile_visibilities.array, + rtol=1.0e-12, + atol=1.0e-12, + ) + + +def test__profile_visibilities__linear_light_only__zeros_without_fourier_transform( + interferometer_7, monkeypatch +): + """ + A tracer whose light is entirely linear (a lens with only mass and a source MGE `Basis` of linear + Gaussians) has an all-zero ordinary light image, so `profile_visibilities` must be zeros without + performing a Fourier transform. + """ + calls = [] + + visibilities_from = interferometer_7.transformer.visibilities_from + + def spy(*args, **kwargs): + calls.append(1) + return visibilities_from(*args, **kwargs) + + monkeypatch.setattr(interferometer_7.transformer, "visibilities_from", spy) + + source = al.Galaxy( + redshift=1.0, + bulge=al.lp_basis.Basis( + profile_list=[ + al.lp_linear.Gaussian(sigma=sigma, centre=(0.1, 0.1)) + for sigma in (0.3, 1.0, 3.0) + ] + ), + ) + + tracer = al.Tracer( + galaxies=[ + al.Galaxy( + redshift=0.5, + mass=al.mp.Isothermal(centre=(0.0, 0.0), einstein_radius=1.0), + ), + source, + ] + ) + + fit = al.FitInterferometer(dataset=interferometer_7, tracer=tracer) + + profile_visibilities = fit.profile_visibilities + + assert calls == [] + assert profile_visibilities.shape == interferometer_7.data.shape + assert np.all(profile_visibilities.array == 0.0)