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
65 changes: 62 additions & 3 deletions autolens/interferometer/fit_interferometer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -143,18 +145,47 @@ 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
tracer subtracted.
"""
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,
Expand Down Expand Up @@ -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:
"""
Expand Down
2 changes: 1 addition & 1 deletion autolens/interferometer/model/visualizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
229 changes: 229 additions & 0 deletions test_autolens/interferometer/test_fit_interferometer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Loading