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
27 changes: 27 additions & 0 deletions autoarray/fit/fit_interferometer.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import functools
import numpy as np
from typing import Optional

from autoarray.dataset.interferometer.dataset import Interferometer

Expand Down Expand Up @@ -206,11 +207,37 @@ def signal_to_noise_map(self) -> np.ndarray:

return signal_to_noise_map_real + 1.0j * signal_to_noise_map_imag

@property
def sparse_chi_squared(self) -> Optional[float]:
"""
The chi-squared of this fit computed from the dataset's `sparse_operator` without any visibility-sized
array, used by `chi_squared` when the dataset is array-free (built by `from_stream` /
`from_sparse_terms`, so `data` is `None`).

`None` by default: a bare `FitInterferometer` has no real-space model image to evaluate it from.
Subclasses whose model visibilities are the Fourier transform `p = F i_p` of a real-space image `i_p`
(e.g. the light-profile fits of PyAutoGalaxy and PyAutoLens without an inversion) override it with
`data_term - 2 i_p^T d~ + i_p^T W~ i_p`
(`inversion_interferometer_util.sparse_profile_terms_from`), so `log_likelihood` and
`figure_of_merit` work array-free. The residual and chi-squared *maps* still need the visibilities and
raise on such a dataset.
"""
return None

@property
def chi_squared(self) -> float:
"""
Returns the chi-squared terms of the model data's fit to an dataset, by summing the chi-squared-map.

On an array-free dataset (no `data`) this is `sparse_chi_squared` when a subclass provides it, and
otherwise raises an `exc.DatasetException`.
"""
if self.data is None:
sparse_chi_squared = self.sparse_chi_squared

if sparse_chi_squared is not None:
return sparse_chi_squared

self._require("chi_squared", "data", "noise_map")
return fit_util.chi_squared_complex_from(
chi_squared_map=self.chi_squared_map.array,
Expand Down
12 changes: 12 additions & 0 deletions autoarray/inversion/inversion/dataset_interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ def __init__(
sparse_operator=None,
noise_covariance_matrix=None,
sparse_dirty_image=None,
data_term=None,
):
"""
Generic class which acts as an interface between a dataset and an inversion.
Expand Down Expand Up @@ -65,6 +66,16 @@ def __init__(
visibilities of ordinary light profiles have been subtracted). If `None`, the operator's cached dirty
image is used. This is distinct from `Interferometer.dirty_image`, the unweighted dirty image of the
data used for visualization.
data_term
The chi-squared data term `sum(d_r^2/sigma_r^2) + sum(d_i^2/sigma_i^2)` of *this interface's*
(possibly profile-subtracted) visibilities, read by the sparse interferometer inversion's
`fast_chi_squared` when `data` is `None` in preference to the `sparse_operator`'s cached scalar (which
is the data term of the raw, unsubtracted visibilities). It is how a fit with ordinary light profiles
on an array-free dataset passes `data=None`: the light profiles' visibilities `F i_p` are never
formed, and the subtracted data term `data_term - 2 i_p^T d~ + i_p^T W~ i_p` is computed by
`inversion_interferometer_util.sparse_profile_terms_from` alongside the subtracted
`sparse_dirty_image`. If `None`, the operator's cached scalar is used. A scalar (traced under
`jax.jit` when the profile image is).
"""
self.data = data
self.noise_map = noise_map
Expand All @@ -74,6 +85,7 @@ def __init__(
self.sparse_operator = sparse_operator
self.noise_covariance_matrix = noise_covariance_matrix
self.sparse_dirty_image = sparse_dirty_image
self.data_term = data_term

@property
def mask(self):
Expand Down
28 changes: 18 additions & 10 deletions autoarray/inversion/inversion/interferometer/abstract.py
Original file line number Diff line number Diff line change
Expand Up @@ -191,10 +191,14 @@ def fast_chi_squared(self):
where `s` is the reconstruction vector, `F` is the curvature matrix, `D` is the data vector,
and `d_r`/`d_i` are the real/imaginary parts of the observed visibilities.

When the dataset interface's `data` is `None` the third term is read from the scalar
`sparse_operator.data_term` cached when the operator was built, so no visibility array
is reduced over. That is only correct when the data fitted is the raw data the operator
was built from (nothing subtracted), which is the contract of passing `data=None`.
When the dataset interface's `data` is `None` the third term is a scalar, so no visibility
array is reduced over. It is the interface's own `data_term` when one is given (the data
term of profile-subtracted visibilities, computed via the identity in
`inversion_interferometer_util.sparse_profile_terms_from`), and otherwise the
`sparse_operator.data_term` cached when the operator was built -- which is only correct when
the data fitted is the raw data the operator was built from (nothing subtracted). The
interface's `data_term` always takes precedence, so a fit that subtracted light profiles can
never silently fall back to the unsubtracted scalar.

This avoids computing the full mapped reconstructed visibilities and is faster than computing
`chi_squared` via the residual visibilities when many source pixels are used.
Expand All @@ -217,12 +221,16 @@ def fast_chi_squared(self):
)

if self.dataset.data is None:
# The interface carries no visibilities (e.g. a pixelization-only fit on the sparse
# path, where nothing was subtracted from the data), so term 3 is the scalar
# `d^T N^-1 d` the sparse operator cached when it was built from the raw data.
chi_squared_term_3 = getattr(
self.dataset.sparse_operator, "data_term", None
)
# The interface carries no visibilities, so term 3 is a precomputed scalar: the
# interface's own `data_term` (that of profile-subtracted visibilities) when given,
# else the `d^T N^-1 d` the sparse operator cached when it was built from the raw data
# (e.g. a pixelization-only fit, where nothing was subtracted).
chi_squared_term_3 = getattr(self.dataset, "data_term", None)

if chi_squared_term_3 is None:
chi_squared_term_3 = getattr(
self.dataset.sparse_operator, "data_term", None
)

if chi_squared_term_3 is None:
raise exc.InversionException(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1804,6 +1804,78 @@ def curvature_matrix_func_list_from(
return curvature_weights_0.T @ operated


def sparse_profile_terms_from(
sparse_operator: "InterferometerSparseOperator",
image,
extent_index_for_masked_pixel,
xp=np,
):
"""
Returns the sparse-path quantities of fitting visibilities from which the visibilities `p = F i_p` of a
real-space image `i_p` (e.g. the summed image of a fit's ordinary light profiles) are subtracted, computed
without forming `p` or touching any visibility-sized array.

With `d` the visibilities the `sparse_operator` was built from, `W = 1 / sigma^2` their (equal real and
imaginary) inverse variances, `d~ = Re(F^H W d)` the operator's cached `dirty_image`, `W~ = Re(F^H W F)` the
operator itself and `data_term = sum(|d|^2 / sigma^2)` its cached scalar, linearity gives:

Re(F^H W (d - F i_p)) = d~ - W~ i_p
sum(|d - F i_p|^2 / sigma^2) = data_term - 2 i_p^T d~ + i_p^T W~ i_p

The first is the dirty image a sparse inversion of the profile-subtracted visibilities forms its data vector
from; the second is their chi-squared data term (term 3 of `fast_chi_squared`), and on its own the
chi-squared of a fit whose model is only `p`. Both reuse the single product `W~ i_p` (one FFT convolution on
the real-space grid), so the data term costs two dot products on top of the dirty image.

The data term is a difference of large numbers at high signal-to-noise (`data_term >> chi^2`); in float64
this is accurate to ~1e-15 relative to `data_term`, ample for likelihood comparisons.

Every operation is an `xp` array operation, so under `jax.jit` the returned scalar is traced with `i_p`.

Parameters
----------
sparse_operator
The `InterferometerSparseOperator` of the dataset, carrying the `dirty_image` and `data_term` of the
visibilities it was built from.
image
The image `i_p` on the slim masked real-space grid (an `Array2D` or a plain array).
extent_index_for_masked_pixel
The `real_space_mask.extent_index_for_masked_pixel` mapping slim masked pixels to the operator's
rectangular extent grid.
xp
The array module (`numpy` or `jax.numpy`).

Returns
-------
operated_image
`W~ i_p` on the slim masked grid.
sparse_dirty_image
`d~ - W~ i_p`.
data_term
`data_term - 2 i_p^T d~ + i_p^T W~ i_p`, or `None` if the operator carries no `data_term`.
"""
image = getattr(image, "array", image)

operated_image = sparse_operator.operated_matrix_slim_from(
matrix_slim=image[:, None],
extent_index_for_masked_pixel=extent_index_for_masked_pixel,
xp=xp,
)[:, 0]

dirty_image = xp.asarray(sparse_operator.dirty_image)

sparse_dirty_image = dirty_image - operated_image

data_term = getattr(sparse_operator, "data_term", None)

if data_term is not None:
data_term = (
data_term - 2.0 * xp.dot(image, dirty_image) + xp.dot(image, operated_image)
)

return operated_image, sparse_dirty_image, data_term


@dataclass(frozen=True)
class SparseTerms:
"""
Expand Down
64 changes: 64 additions & 0 deletions test_autoarray/fit/test_fit_interferometer.py
Original file line number Diff line number Diff line change
Expand Up @@ -568,3 +568,67 @@ def test__dirty_model_image_natural_from__no_sparse_operator__raises(

with pytest.raises(aa.exc.DatasetException, match="sparse_operator"):
dirty_model_image_natural_from(dataset=interferometer_7, image=image)


class _ProfileImageFit(aa.m.MockFitInterferometer):
"""
A fit whose model visibilities are the Fourier transform of a real-space `image`, overriding the
`sparse_chi_squared` hook with the data-term identity as PyAutoGalaxy / PyAutoLens light-profile fits do.
"""

def __init__(self, dataset, image):
super().__init__(dataset=dataset)
self.image = image

@property
def sparse_chi_squared(self):
return aa.util.inversion_interferometer.sparse_profile_terms_from(
sparse_operator=self.dataset.sparse_operator,
image=self.image,
extent_index_for_masked_pixel=self.dataset.real_space_mask.extent_index_for_masked_pixel,
)[2]


def test__fit_interferometer__array_free_dataset__sparse_chi_squared_hook():
"""
On an array-free dataset `chi_squared` (hence `log_likelihood`) is read from the `sparse_chi_squared` hook
when a subclass provides it, matching the in-memory fit of the same model visibilities; the maps still
raise. The hook is not consulted when the fit has data.
"""
pytest.importorskip("nufftax")

dataset_memory, dataset_stream, _ = _array_free_fit_setup()

mask = dataset_memory.real_space_mask

image = aa.Array2D(
values=np.random.default_rng(seed=2).normal(size=mask.pixels_in_mask),
mask=mask,
)

fit_memory = aa.m.MockFitInterferometer(
dataset=dataset_memory,
model_data=dataset_memory.transformer.visibilities_from(image=image),
)
fit_stream = _ProfileImageFit(dataset=dataset_stream, image=image)

assert fit_stream.chi_squared == pytest.approx(fit_memory.chi_squared, rel=1.0e-8)
assert fit_stream.log_likelihood == pytest.approx(
fit_memory.log_likelihood, rel=1.0e-8
)
assert fit_stream.figure_of_merit == pytest.approx(
fit_memory.figure_of_merit, rel=1.0e-8
)

for name in ("residual_map", "chi_squared_map", "normalized_residual_map"):
with pytest.raises(aa.exc.DatasetException, match="array-free"):
getattr(fit_stream, name)

# With data present the hook is ignored and the chi-squared-map is summed.
fit_memory_hook = _ProfileImageFit(dataset=dataset_memory, image=None)
fit_memory_hook._model_data = fit_memory.model_data

assert fit_memory_hook.chi_squared == fit_memory.chi_squared

# The base class provides no hook.
assert aa.m.MockFitInterferometer(dataset=dataset_stream).sparse_chi_squared is None
Loading
Loading