Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 36 additions & 4 deletions autolens/interferometer/fit_interferometer.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,15 +127,34 @@ 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.

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
Expand All @@ -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:
Expand Down Expand Up @@ -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.
Expand Down
26 changes: 11 additions & 15 deletions autolens/interferometer/model/analysis.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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__),
)
101 changes: 101 additions & 0 deletions test_autolens/interferometer/model/test_analysis_interferometer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
96 changes: 96 additions & 0 deletions test_autolens/interferometer/test_fit_interferometer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading