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
59 changes: 56 additions & 3 deletions autoarray/dataset/imaging/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,36 @@ def _validate_convolve_over_sample_size(
)


_SPARSE_OPERATOR_CONVOLVE_OVER_SAMPLE_SIZE_ERROR = (
"The sparse linear algebra formalism precomputes PSF products at "
"image resolution and is incompatible with an oversampled PSF "
"(convolve_over_sample_size > 1)."
)


def _warn_sparse_operator_discarded(sparse_operator, method_name: str) -> None:
"""
Log a warning when a dataset rebuild discards an attached sparse operator.

The sparse operator precomputes PSF products of every pair of masked noise-map
values, so it is invalidated by any change to the mask or noise-map. Rebuilds
which make such changes (e.g. `apply_mask`, `apply_noise_scaling`) therefore do
not carry it, and the returned dataset falls back to the (slower) mapping matrix
formalism unless the operator is re-applied. This warning stops that fallback
from happening silently.
"""
if sparse_operator is None:
return

logger.warning(
f"IMAGING - {method_name} changes the mask or noise-map, which invalidates the "
f"sparse operator attached to this dataset, so it has been discarded. Pixelized "
f"reconstructions will use the slower mapping matrix formalism unless the sparse "
f"operator is re-applied (via apply_sparse_operator() or "
f"apply_sparse_operator_cpu()) as the last dataset operation."
)


class Imaging(AbstractDataset):
def __init__(
self,
Expand Down Expand Up @@ -439,6 +469,10 @@ def apply_mask(self, mask: Mask2D) -> "Imaging":
values=self.over_sample_size_pixelization.native, mask=mask
)

_warn_sparse_operator_discarded(
sparse_operator=self.sparse_operator, method_name="apply_mask"
)

dataset = Imaging(
data=data,
noise_map=noise_map,
Expand Down Expand Up @@ -520,13 +554,19 @@ def apply_noise_scaling(

noise_map = Array2D(values=noise_map, mask=self.data.mask)

_warn_sparse_operator_discarded(
sparse_operator=self.sparse_operator, method_name="apply_noise_scaling"
)

dataset = Imaging(
data=data,
noise_map=noise_map,
psf=self.psf,
noise_covariance_matrix=self.noise_covariance_matrix,
over_sample_size_lp=self.over_sample_size_lp,
over_sample_size_pixelization=self.over_sample_size_pixelization,
convolve_over_sample_size_lp=self.convolve_over_sample_size_lp,
convolve_over_sample_size_pixelization=self.convolve_over_sample_size_pixelization,
check_noise_map=False,
)

Expand All @@ -551,6 +591,10 @@ def apply_over_sampling(
This function resets the cached properties so that the new over sampling is used in the grid and grid
pixelization.

The `noise_covariance_matrix` and any attached `sparse_operator` are carried over, as neither depends on
the over sampling (the sparse operator depends only on the noise-map, PSF kernel and mask). This means
`apply_sparse_operator()` / `apply_sparse_operator_cpu()` may be called before or after this method.

Parameters
----------
over_sample_size_lp
Expand All @@ -566,12 +610,14 @@ def apply_over_sampling(
data=self.data,
noise_map=self.noise_map,
psf=self.psf,
noise_covariance_matrix=self.noise_covariance_matrix,
over_sample_size_lp=over_sample_size_lp or self.over_sample_size_lp,
over_sample_size_pixelization=over_sample_size_pixelization
or self.over_sample_size_pixelization,
convolve_over_sample_size_lp=self.convolve_over_sample_size_lp,
convolve_over_sample_size_pixelization=self.convolve_over_sample_size_pixelization,
check_noise_map=False,
sparse_operator=self.sparse_operator,
)

return dataset
Expand Down Expand Up @@ -620,9 +666,7 @@ def apply_sparse_operator(

if self.psf is not None and self.psf.convolve_over_sample_size > 1:
raise exc.DatasetException(
"The sparse linear algebra formalism precomputes PSF products at "
"image resolution and is incompatible with an oversampled PSF "
"(convolve_over_sample_size > 1)."
_SPARSE_OPERATOR_CONVOLVE_OVER_SAMPLE_SIZE_ERROR
)

logger.info(
Expand All @@ -645,6 +689,8 @@ def apply_sparse_operator(
noise_covariance_matrix=self.noise_covariance_matrix,
over_sample_size_lp=self.over_sample_size_lp,
over_sample_size_pixelization=self.over_sample_size_pixelization,
convolve_over_sample_size_lp=self.convolve_over_sample_size_lp,
convolve_over_sample_size_pixelization=self.convolve_over_sample_size_pixelization,
check_noise_map=False,
sparse_operator=sparse_operator,
)
Expand All @@ -668,6 +714,11 @@ def apply_sparse_operator_cpu(
A new `Imaging` dataset with a precomputed Numba-based sparse operator attached,
enabling efficient pixelized source reconstruction on CPU hardware.
"""
if self.psf is not None and self.psf.convolve_over_sample_size > 1:
raise exc.DatasetException(
_SPARSE_OPERATOR_CONVOLVE_OVER_SAMPLE_SIZE_ERROR
)

try:
import numba
except ModuleNotFoundError:
Expand Down Expand Up @@ -711,6 +762,8 @@ def apply_sparse_operator_cpu(
noise_covariance_matrix=self.noise_covariance_matrix,
over_sample_size_lp=self.over_sample_size_lp,
over_sample_size_pixelization=self.over_sample_size_pixelization,
convolve_over_sample_size_lp=self.convolve_over_sample_size_lp,
convolve_over_sample_size_pixelization=self.convolve_over_sample_size_pixelization,
check_noise_map=False,
sparse_operator=sparse_operator,
)
Expand Down
11 changes: 10 additions & 1 deletion autoarray/inversion/inversion/abstract.py
Original file line number Diff line number Diff line change
Expand Up @@ -299,8 +299,17 @@ def operated_mapping_matrix(self) -> np.ndarray:

If there are multiple linear objects, the blurred mapping matrices are stacked such that their simultaneous
linear equations are solved simultaneously.

The dataset-specific `operated_mapping_matrix_list` is itself cached, so for a single linear object its
matrix is returned directly rather than copied via `hstack`, avoiding holding two identical copies of
the (large) operated mapping matrix in memory.
"""
return self._xp.hstack(self.operated_mapping_matrix_list)
operated_mapping_matrix_list = self.operated_mapping_matrix_list

if len(operated_mapping_matrix_list) == 1:
return operated_mapping_matrix_list[0]

return self._xp.hstack(operated_mapping_matrix_list)

@property
def data_vector(self) -> np.ndarray:
Expand Down
2 changes: 1 addition & 1 deletion autoarray/inversion/inversion/imaging/abstract.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,7 @@ def mapping_matrix_list(self) -> List[np.ndarray]:
"""
return [linear_obj.mapping_matrix for linear_obj in self.linear_obj_list]

@property
@cached_property
def operated_mapping_matrix_list(self) -> List[np.ndarray]:
"""
The `operated_mapping_matrix` of a linear object describes the mappings between the observed data's values and
Expand Down
4 changes: 3 additions & 1 deletion autoarray/inversion/inversion/interferometer/abstract.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
import numpy as np
from typing import Dict, List, Optional, Union

from autonerves import cached_property

from autoarray import exc
from autoarray.dataset.interferometer.dataset import Interferometer
from autoarray.inversion.inversion.dataset_interface import DatasetInterface
Expand Down Expand Up @@ -61,7 +63,7 @@ def transformer(self):
def mask(self) -> Mask2D:
return self.transformer.real_space_mask

@property
@cached_property
def operated_mapping_matrix_list(self) -> List[np.ndarray]:
"""
The `operated_mapping_matrix` of a linear object describes the mappings between the observed data's values
Expand Down
96 changes: 96 additions & 0 deletions test_autoarray/dataset/imaging/test_dataset.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,15 @@
import copy
import logging

import numpy as np
import pytest

import autoarray as aa

from autoarray import exc
from autoarray.inversion.inversion.imaging_numba.sparse import (
InversionImagingSparseNumba,
)
from pathlib import Path

test_data_path = Path(Path(__file__).resolve().parent) / "files"
Expand Down Expand Up @@ -504,3 +508,95 @@ def test__convolve_over_sample_size__sparse_operator_guard():

with pytest.raises(aa.exc.DatasetException):
dataset.apply_sparse_operator()


def test__apply_over_sampling__keeps_sparse_operator_and_noise_covariance(
masked_imaging_7x7, masked_imaging_covariance_7x7, delaunay_mapper_9_3x3
):
# The sparse operator depends only on the noise-map, PSF kernel and mask, so it
# stays valid when only the over sampling changes and must not be dropped.
dataset_operator_first = (
masked_imaging_7x7.apply_sparse_operator_cpu().apply_over_sampling(
over_sample_size_lp=2
)
)

assert dataset_operator_first.sparse_operator is not None

dataset_operator_last = masked_imaging_7x7.apply_over_sampling(
over_sample_size_lp=2
).apply_sparse_operator_cpu()

inversion_operator_first = aa.Inversion(
dataset=dataset_operator_first,
linear_obj_list=[delaunay_mapper_9_3x3],
)
inversion_operator_last = aa.Inversion(
dataset=dataset_operator_last,
linear_obj_list=[delaunay_mapper_9_3x3],
)

assert isinstance(inversion_operator_first, InversionImagingSparseNumba)
assert inversion_operator_first.reconstruction == pytest.approx(
inversion_operator_last.reconstruction, 1.0e-8
)
assert inversion_operator_first.log_det_curvature_reg_matrix_term == pytest.approx(
inversion_operator_last.log_det_curvature_reg_matrix_term, 1.0e-8
)

# The noise covariance matrix is independent of over sampling and is also kept.
dataset_covariance = masked_imaging_covariance_7x7.apply_over_sampling(
over_sample_size_lp=2
)

assert dataset_covariance.noise_covariance_matrix == pytest.approx(
masked_imaging_covariance_7x7.noise_covariance_matrix, 1.0e-8
)


def test__apply_mask_and_noise_scaling__warn_when_sparse_operator_discarded(
masked_imaging_7x7, caplog
):
dataset = masked_imaging_7x7.apply_sparse_operator_cpu()

with caplog.at_level(logging.WARNING, logger="autoarray.dataset.imaging.dataset"):
masked = dataset.apply_mask(mask=masked_imaging_7x7.mask)

assert masked.sparse_operator is None
assert "sparse operator" in caplog.text
assert "apply_mask" in caplog.text

caplog.clear()

with caplog.at_level(logging.WARNING, logger="autoarray.dataset.imaging.dataset"):
scaled = dataset.apply_noise_scaling(mask=masked_imaging_7x7.mask)

assert scaled.sparse_operator is None
assert "sparse operator" in caplog.text
assert "apply_noise_scaling" in caplog.text

# No warning when there is no operator to discard.
caplog.clear()

with caplog.at_level(logging.WARNING, logger="autoarray.dataset.imaging.dataset"):
masked_imaging_7x7.apply_mask(mask=masked_imaging_7x7.mask)

assert "sparse operator" not in caplog.text


def test__convolve_over_sample_size__sparse_operator_cpu_guard():
data = aa.Array2D.no_mask(values=np.ones((11, 11)), pixel_scales=1.0)
noise_map = aa.Array2D.no_mask(values=np.ones((11, 11)), pixel_scales=1.0)
kernel_fine = aa.Array2D.no_mask(values=np.ones((9, 9)), pixel_scales=0.5)
psf = aa.Convolver(kernel=kernel_fine)

dataset = aa.Imaging(
data=data,
noise_map=noise_map,
psf=psf,
over_sample_size_pixelization=2,
convolve_over_sample_size_pixelization=2,
)

with pytest.raises(aa.exc.DatasetException):
dataset.apply_sparse_operator_cpu()
32 changes: 32 additions & 0 deletions test_autoarray/inversion/inversion/imaging/test_imaging.py
Original file line number Diff line number Diff line change
Expand Up @@ -267,3 +267,35 @@ def test__mapping_matrix_over_sampled_for__kxs__full_bin_reproduces_mapping_matr
# s=2 rows mean-binned to image resolution also reproduce mapping_matrix.
binned = m_s2.reshape(n_pix, 4, mapping_matrix.shape[1]).mean(axis=1)
assert binned == pytest.approx(mapping_matrix, abs=1.0e-14)


def test__operated_mapping_matrix_list__psf_convolution_performed_once_per_linear_obj(
masked_imaging_7x7, delaunay_mapper_9_3x3, monkeypatch
):
# The dense (mapping) route reaches `operated_mapping_matrix_list` from the
# curvature matrix / data vector (via `operated_mapping_matrix`) and again from
# `mapped_reconstructed_data_dict`. It must be cached so the PSF convolution of
# each linear object's mapping matrix happens once per inversion.
calls = []

convolved_mapping_matrix_from = aa.Convolver.convolved_mapping_matrix_from

def counted(self, *args, **kwargs):
calls.append(1)
return convolved_mapping_matrix_from(self, *args, **kwargs)

monkeypatch.setattr(aa.Convolver, "convolved_mapping_matrix_from", counted)

inversion = aa.Inversion(
dataset=masked_imaging_7x7,
linear_obj_list=[delaunay_mapper_9_3x3],
)

assert isinstance(inversion, aa.InversionImagingMapping)

inversion.log_det_curvature_reg_matrix_term
inversion.reconstruction
inversion.mapped_reconstructed_operated_data
inversion.mapped_reconstructed_data

assert len(calls) == 1
Original file line number Diff line number Diff line change
Expand Up @@ -1178,6 +1178,40 @@ def test__interferometer_mapping__curvature_matrix_and_data_vector_evaluated_onc
assert log_evidence_terms == reference


def test__operated_mapping_matrix_list__transform_performed_once_per_linear_obj(
interferometer_7_no_fft, rectangular_mapper_7x7_3x3, monkeypatch
):
# The dense (mapping) route reaches `operated_mapping_matrix_list` from the
# curvature matrix / data vector (via `operated_mapping_matrix`) and again from
# `mapped_reconstructed_data_dict`. It must be cached so the Fourier transform of
# each linear object's mapping matrix happens once per inversion.
calls = []

transformer_cls = type(interferometer_7_no_fft.transformer)
transform_mapping_matrix = transformer_cls.transform_mapping_matrix

def counted(self, *args, **kwargs):
calls.append(1)
return transform_mapping_matrix(self, *args, **kwargs)

monkeypatch.setattr(transformer_cls, "transform_mapping_matrix", counted)

inversion = aa.Inversion(
dataset=interferometer_7_no_fft,
linear_obj_list=[rectangular_mapper_7x7_3x3],
settings=aa.Settings(),
)

assert isinstance(inversion, aa.InversionInterferometerMapping)

inversion.log_det_curvature_reg_matrix_term
inversion.reconstruction
inversion.mapped_reconstructed_operated_data
inversion.mapped_reconstructed_data

assert len(calls) == 1


def _sparse_interface_setup():
mask = aa.Mask2D.circular(shape_native=(10, 10), pixel_scales=1.0, radius=3.0)

Expand Down
Loading