diff --git a/autolens/interferometer/fit_interferometer.py b/autolens/interferometer/fit_interferometer.py index 94a3492e9..2cb50cca9 100644 --- a/autolens/interferometer/fit_interferometer.py +++ b/autolens/interferometer/fit_interferometer.py @@ -127,7 +127,7 @@ def profile_image(self) -> aa.Array2D: return self.tracer.image_2d_from(grid=self.grids.lp, xp=self._xp) @cached_property - def profile_visibilities(self) -> aa.Visibilities: + def profile_visibilities(self) -> Optional[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. @@ -135,7 +135,26 @@ def profile_visibilities(self) -> aa.Visibilities: 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`. + + On an array-free dataset (built by `Interferometer.from_stream` / `from_sparse_terms`, which has no + `uv_wavelengths` and so no transformer) there are no visibilities to compute: this returns `None` when + the tracer has no ordinary light profile, and raises an `exc.DatasetException` when it does, because + subtracting a light profile's visibilities needs the visibility arrays. """ + if self.dataset.transformer is None: + if _has_light_profile_non_linear(galaxies=self.tracer.galaxies): + raise aa.exc.DatasetException( + "This FitInterferometer's dataset is array-free (built by from_stream / " + "from_sparse_terms) and has no visibilities or transformer, so the tracer's " + "ordinary (non-linear) light profiles cannot be Fourier transformed and " + "subtracted. An array-free dataset supports pixelization-only and linear-light " + "fits; non-linear light profiles arrive in a later phase. Use the in-memory " + "constructor (`Interferometer(data=..., noise_map=..., uv_wavelengths=..., ...)`) " + "to fit them." + ) + + return None + if _has_light_profile_non_linear(galaxies=self.tracer.galaxies): return self.dataset.transformer.visibilities_from( image=self.profile_image, xp=self._xp @@ -146,12 +165,21 @@ def profile_visibilities(self) -> aa.Visibilities: ) @cached_property - def profile_subtracted_visibilities(self) -> aa.Visibilities: + def profile_subtracted_visibilities(self) -> Optional[aa.Visibilities]: """ Returns the interferometer dataset's visibilities with all transformed light profile images in the fit's tracer subtracted. + + On an array-free dataset there are no visibilities, so this is `None`. `profile_visibilities` is + evaluated first, so a fit with ordinary light profiles on an array-free dataset raises rather than + silently fitting the unsubtracted sparse terms. """ - return self.data - self.profile_visibilities + profile_visibilities = self.profile_visibilities + + if self.data is None: + return None + + return self.data - profile_visibilities @property def _uses_precomputed_data_term(self) -> bool: @@ -250,10 +278,14 @@ def inversion_with_data(self) -> Optional[aa.AbstractInversion]: 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. + + On an array-free dataset (built by `Interferometer.from_stream` / `from_sparse_terms`) `fit.data` is + `None`, so there are no visibilities to carry and `inversion` itself is returned; output quantities + that read the visibilities are unavailable on such a fit. """ inversion = self.inversion - if inversion is None or inversion.dataset.data is not None: + if inversion is None or inversion.dataset.data is not None or self.data is None: return inversion # Solve first, so the copy shares the reconstruction (and everything it cached) rather than repeating it. diff --git a/autolens/interferometer/model/analysis.py b/autolens/interferometer/model/analysis.py index 3fbbcb593..db51dc2c7 100644 --- a/autolens/interferometer/model/analysis.py +++ b/autolens/interferometer/model/analysis.py @@ -17,11 +17,11 @@ from typing import Optional from autonerves.dictable import to_dict -from autonerves.fitsable import hdu_list_for_output_from import autofit as af import autoarray as aa import autogalaxy as ag +from autogalaxy.interferometer.model.analysis import interferometer_hdu_list_from from autolens.analysis.analysis.dataset import AnalysisDataset from autolens.analysis.exceptions import raise_fit_exception @@ -351,16 +351,9 @@ def save_attributes(self, paths: af.DirectoryPaths): """ super().save_attributes(paths=paths) - hdu_list = hdu_list_for_output_from( - values_list=[ - self.dataset.real_space_mask.astype("float"), - self.dataset.data.in_array, - self.dataset.noise_map.in_array, - self.dataset.uv_wavelengths, - ], - ext_name_list=["mask", "data", "noise_map", "uv_wavelengths"], - header_dict=self.dataset.real_space_mask.header_dict, - ) + # An array-free dataset writes its `SparseTerms` instead of the visibility arrays (see + # `autogalaxy.interferometer.model.analysis.interferometer_hdu_list_from`). + hdu_list = interferometer_hdu_list_from(dataset=self.dataset) # `dataset.fits` is written once per search, to the `image` folder, and is written # unconditionally (it is not gated on any visualization setting). The write is skipped @@ -374,7 +367,10 @@ def save_attributes(self, paths: af.DirectoryPaths): if not dataset_path.exists(): hdu_list.writeto(dataset_path, overwrite=True) - paths.save_json( - "transformer_class", - to_dict(self.dataset.transformer.__class__), - ) + # An array-free dataset has no transformer; its transformer class name is recorded in + # the `dataset.fits` header instead. + if self.dataset.transformer is not None: + paths.save_json( + "transformer_class", + to_dict(self.dataset.transformer.__class__), + ) diff --git a/test_autolens/interferometer/model/test_analysis_interferometer.py b/test_autolens/interferometer/model/test_analysis_interferometer.py index 8bfa06748..faadcf301 100644 --- a/test_autolens/interferometer/model/test_analysis_interferometer.py +++ b/test_autolens/interferometer/model/test_analysis_interferometer.py @@ -187,3 +187,104 @@ def test__shared_state_from__populates_mesh_geometry_fields(interferometer_7): shared = analysis.shared_state_from(instance=instance) assert shared.source_plane_mesh_grid is not None assert shared.image_plane_mesh_grid is not None + + +class _FitStub: + """ + The part of a `PyAutoFit` aggregator `Fit` the interferometer loader reads: `value(name)` returning the + `dataset.fits` HDU list written by `save_attributes` and the saved `transformer_class` json (if any). + """ + + def __init__(self, paths): + self.paths = paths + self.children = [] + + def value(self, name): + from astropy.io import fits + from autonerves.dictable import from_dict + + if name == "dataset": + return fits.open(self.paths.image_path / "dataset.fits") + + if name == "transformer_class": + path = self.paths._files_path / "transformer_class.json" + + if not path.exists(): + return None + + return from_dict(self.paths.load_json("transformer_class")) + + return None + + +def test__save_attributes__array_free_dataset__aggregator_round_trip( + interferometer_7, tmp_path +): + """ + `save_attributes` writes an array-free dataset's `SparseTerms` to `dataset.fits` and the aggregator loader + (shared with autogalaxy) rebuilds it via `Interferometer.from_sparse_terms`, so a fit of the reloaded + dataset with a tracer reproduces the original `log_evidence`. + """ + from astropy.io import fits + + from autolens.aggregator import _interferometer_from + + dataset = aa.Interferometer.from_stream( + [ + ( + interferometer_7.uv_wavelengths, + interferometer_7.data, + interferometer_7.noise_map, + ) + ], + real_space_mask=interferometer_7.real_space_mask, + transformer_class=type(interferometer_7.transformer), + ) + + paths = af.DirectoryPaths(name="array_free_round_trip", path_prefix=str(tmp_path)) + + analysis = al.AnalysisInterferometer(dataset=dataset, use_jax=False) + analysis.save_attributes(paths=paths) + + with fits.open(paths.image_path / "dataset.fits") as hdu_list: + assert [hdu.name for hdu in hdu_list] == [ + "MASK", + "NUFFT_PRECISION_OPERATOR", + "DIRTY_IMAGE", + "DIRTY_BEAM", + "SPARSE_TERMS_SCALARS", + ] + + assert not (paths._files_path / "transformer_class.json").exists() + + dataset_reloaded = _interferometer_from(fit=_FitStub(paths=paths))[0] + + assert dataset_reloaded.is_array_free + assert ( + dataset_reloaded.sparse_operator.data_term + == dataset.sparse_operator.data_term + ) + assert ( + dataset_reloaded.sparse_operator.noise_normalization + == dataset.sparse_operator.noise_normalization + ) + + lens = al.Galaxy( + redshift=0.5, + mass=al.mp.Isothermal(centre=(0.0, 0.0), einstein_radius=1.0), + ) + source = al.Galaxy( + redshift=1.0, + pixelization=al.Pixelization( + mesh=al.mesh.RectangularUniform(shape=(3, 3)), + regularization=al.reg.Constant(coefficient=1.0), + ), + ) + tracer = al.Tracer(galaxies=[lens, source]) + + log_evidence = al.FitInterferometer(dataset=dataset, tracer=tracer).log_evidence + log_evidence_reloaded = al.FitInterferometer( + dataset=dataset_reloaded, tracer=tracer + ).log_evidence + + assert log_evidence_reloaded == pytest.approx(log_evidence, rel=1.0e-8) diff --git a/test_autolens/interferometer/test_fit_interferometer.py b/test_autolens/interferometer/test_fit_interferometer.py index bda5b7486..130fa1b18 100644 --- a/test_autolens/interferometer/test_fit_interferometer.py +++ b/test_autolens/interferometer/test_fit_interferometer.py @@ -733,3 +733,99 @@ def figure_of_merit_from(coefficient, xp): assert float(figure_of_merit_jax) == pytest.approx( figure_of_merit_numpy, rel=1.0e-8 ) + + +def _array_free_dataset_from(dataset): + """ + The array-free counterpart of `dataset` (no visibilities, uv-wavelengths or transformer), built by + streaming its visibilities through `Interferometer.from_stream` with the same transformer class. + """ + return aa.Interferometer.from_stream( + [(dataset.uv_wavelengths, dataset.data, dataset.noise_map)], + real_space_mask=dataset.real_space_mask, + transformer_class=type(dataset.transformer), + ) + + +def test__fit_figure_of_merit__array_free_dataset__pixelization_only__matches_in_memory_sparse( + interferometer_7, +): + dataset_sparse = interferometer_7.apply_sparse_operator(use_jax=False) + dataset_array_free = _array_free_dataset_from(interferometer_7) + + assert dataset_array_free.is_array_free + + tracer = _pixelized_source_tracer() + + fit_sparse = al.FitInterferometer(dataset=dataset_sparse, tracer=tracer) + fit_array_free = al.FitInterferometer(dataset=dataset_array_free, tracer=tracer) + + assert fit_array_free._uses_precomputed_data_term + assert fit_array_free.inversion.dataset.data is None + assert isinstance(fit_array_free.inversion, aa.InversionInterferometerSparse) + + assert fit_array_free.figure_of_merit == pytest.approx( + fit_sparse.figure_of_merit, rel=1.0e-8 + ) + assert fit_array_free.log_evidence == pytest.approx( + fit_sparse.log_evidence, rel=1.0e-8 + ) + + assert fit_array_free.profile_visibilities is None + assert fit_array_free.profile_subtracted_visibilities is None + assert fit_array_free.inversion_with_data is fit_array_free.inversion + + +def test__fit_figure_of_merit__array_free_dataset__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) + dataset_array_free = _array_free_dataset_from(interferometer_7) + + def figure_of_merit_from(coefficient, dataset, xp): + fit = al.FitInterferometer( + dataset=dataset, + tracer=_pixelized_source_tracer(coefficient=coefficient), + xp=xp, + ) + + assert fit.inversion.dataset.data is None + + return fit.figure_of_merit + + figure_of_merit_sparse = figure_of_merit_from( + coefficient=1.0, dataset=dataset_sparse, xp=np + ) + figure_of_merit_numpy = figure_of_merit_from( + coefficient=1.0, dataset=dataset_array_free, xp=np + ) + + figure_of_merit_jax = jax.jit( + lambda c: figure_of_merit_from(c, dataset=dataset_array_free, xp=jnp) + )(1.0) + + assert figure_of_merit_numpy == pytest.approx(figure_of_merit_sparse, rel=1.0e-8) + assert float(figure_of_merit_jax) == pytest.approx( + figure_of_merit_numpy, rel=1.0e-8 + ) + + +def test__fit_figure_of_merit__array_free_dataset__lens_light_profile__raises( + interferometer_7, +): + dataset_array_free = _array_free_dataset_from(interferometer_7) + + tracer = _pixelized_source_tracer(lens_light=True) + + fit = al.FitInterferometer(dataset=dataset_array_free, tracer=tracer) + + assert not fit._uses_precomputed_data_term + + with pytest.raises(aa.exc.DatasetException): + fit.profile_visibilities + + with pytest.raises(aa.exc.DatasetException): + al.FitInterferometer(dataset=dataset_array_free, tracer=tracer).figure_of_merit