diff --git a/autoarray/inversion/inversion/interferometer/mapping.py b/autoarray/inversion/inversion/interferometer/mapping.py index 4651374a1..b313eed91 100644 --- a/autoarray/inversion/inversion/interferometer/mapping.py +++ b/autoarray/inversion/inversion/interferometer/mapping.py @@ -1,6 +1,8 @@ import numpy as np from typing import Dict, List, Union +from autonerves import cached_property + from autoarray.dataset.interferometer.dataset import Interferometer from autoarray.inversion.inversion.dataset_interface import DatasetInterface from autoarray.inversion.inversion.interferometer.abstract import ( @@ -50,7 +52,7 @@ def __init__( dataset=dataset, linear_obj_list=linear_obj_list, settings=settings, xp=xp ) - @property + @cached_property def data_vector(self) -> np.ndarray: """ The `data_vector` is a 1D vector whose values are solved for by the simultaneous linear equations constructed @@ -71,7 +73,7 @@ def data_vector(self) -> np.ndarray: noise_map=self.noise_map, ) - @property + @cached_property def curvature_matrix(self) -> np.ndarray: """ The `curvature_matrix` is a 2D matrix which uses the mappings between the data and the linear objects to diff --git a/autoarray/inversion/inversion/interferometer/sparse.py b/autoarray/inversion/inversion/interferometer/sparse.py index 4da505ae1..914905189 100644 --- a/autoarray/inversion/inversion/interferometer/sparse.py +++ b/autoarray/inversion/inversion/interferometer/sparse.py @@ -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 @@ -67,7 +69,7 @@ def __init__( preloads=preloads, ) - @property + @cached_property def data_vector(self) -> np.ndarray: """ The `data_vector` is a 1D vector whose values are solved for by the simultaneous linear equations constructed @@ -129,7 +131,7 @@ def _sparse_triplets_curvature_from(self, mapper: Mapper): xp=self._xp, ) - @property + @cached_property def curvature_matrix(self) -> np.ndarray: """ The `curvature_matrix` is a 2D matrix which uses the mappings between the data and the linear objects to @@ -187,7 +189,7 @@ def curvature_matrix(self) -> np.ndarray: return curvature_matrix - @property + @cached_property def curvature_matrix_diag(self) -> np.ndarray: """ The `curvature_matrix` is a 2D matrix which uses the mappings between the data and the linear objects to diff --git a/autoarray/inversion/inversion/interferometer_numba/sparse.py b/autoarray/inversion/inversion/interferometer_numba/sparse.py index a18c066fa..58cc8b4be 100644 --- a/autoarray/inversion/inversion/interferometer_numba/sparse.py +++ b/autoarray/inversion/inversion/interferometer_numba/sparse.py @@ -181,7 +181,7 @@ def kernel_index_arrays(self) -> dict: return inputs - @property + @cached_property def curvature_matrix_diag(self) -> np.ndarray: """ `F = Aᵀ W~ A` for the inversion's single mapper, from the `direct_conv` numba diff --git a/test_autoarray/inversion/inversion/interferometer/test_interferometer.py b/test_autoarray/inversion/inversion/interferometer/test_interferometer.py index e879c7a96..e9af578f4 100644 --- a/test_autoarray/inversion/inversion/interferometer/test_interferometer.py +++ b/test_autoarray/inversion/inversion/interferometer/test_interferometer.py @@ -1044,3 +1044,135 @@ def test__interferometer_sparse_operator__numpy_inversion_matches_jax_inversion( ).max() assert difference_sparse <= max(10.0 * difference_dense, 1.0e-14) + + +def _count_evaluations(monkeypatch, cls, name, counts): + """ + Wrap the body of the `cls.name` property with a counter while keeping its descriptor + type, so a `cached_property` stays cached (and a plain `property` stays uncached) and + the count is the number of times the body actually runs. + """ + import functools + + descriptor = cls.__dict__[name] + func = descriptor.fget if isinstance(descriptor, property) else descriptor.func + + @functools.wraps(func) + def counted(self): + counts[name] += 1 + return func(self) + + monkeypatch.setattr(cls, name, type(descriptor)(counted)) + + +def _log_evidence_terms_from(inversion): + """ + The inversion terms of `FitInterferometer.log_evidence`, which is what a figure of merit + evaluates: the fast chi-squared (reads F, D and the reconstruction), the regularization + term and both log-determinants (the cached `curvature_reg_matrix` reads F). + """ + return ( + float(inversion.fast_chi_squared) + + float(inversion.regularization_term) + + float(inversion.log_det_curvature_reg_matrix_term) + - float(inversion.log_det_regularization_matrix_term) + ) + + +def test__interferometer_sparse_operator__curvature_matrix_and_data_vector_evaluated_once_per_likelihood( + monkeypatch, +): + """ + `fast_chi_squared`, `curvature_reg_matrix` and `reconstruction` all read `curvature_matrix` + (F) and `data_vector` (D). Each must be built once per inversion, not once per reader: + on the NumPy path nothing merges the repeated builds, which were ~45 % of an alma call + (autolens_profiling #326). + """ + mask, _, _, dataset_sparse = _sparse_parity_setup() + + mapper = _mapper_from( + mask=mask, pixels=9, shape=(3, 3), regularization=aa.reg.Constant(coefficient=1.0) + ) + + # Built directly (not via `aa.Inversion`) so the NumPy FFT class runs even when numba is + # installed and the factory would route to `InversionInterferometerSparseNumba`. + def inversion_from(): + return aa.InversionInterferometerSparse( + dataset=dataset_sparse, linear_obj_list=[mapper], xp=np + ) + + reference = _log_evidence_terms_from(inversion_from()) + + counts = { + "curvature_matrix_diag": 0, + "data_vector": 0, + "curvature_matrix_diag_from": 0, + } + + _count_evaluations( + monkeypatch, aa.InversionInterferometerSparse, "curvature_matrix_diag", counts + ) + _count_evaluations( + monkeypatch, aa.InversionInterferometerSparse, "data_vector", counts + ) + + from autoarray.inversion.inversion.interferometer.inversion_interferometer_util import ( + InterferometerSparseOperator, + ) + + curvature_matrix_diag_from = InterferometerSparseOperator.curvature_matrix_diag_from + + def counted_curvature_matrix_diag_from(self, *args, **kwargs): + counts["curvature_matrix_diag_from"] += 1 + return curvature_matrix_diag_from(self, *args, **kwargs) + + monkeypatch.setattr( + InterferometerSparseOperator, + "curvature_matrix_diag_from", + counted_curvature_matrix_diag_from, + ) + + inversion = inversion_from() + + log_evidence_terms = _log_evidence_terms_from(inversion) + + assert counts == { + "curvature_matrix_diag": 1, + "data_vector": 1, + "curvature_matrix_diag_from": 1, + } + assert log_evidence_terms == reference + + +def test__interferometer_mapping__curvature_matrix_and_data_vector_evaluated_once_per_likelihood( + monkeypatch, +): + """ + The dense mapping inversion's `curvature_matrix` and `data_vector` are cached per inversion + too, as the imaging mapping inversion's are. + """ + mask, _, dataset, _ = _sparse_parity_setup() + + mapper = _mapper_from( + mask=mask, pixels=9, shape=(3, 3), regularization=aa.reg.Constant(coefficient=1.0) + ) + + reference = _log_evidence_terms_from( + aa.Inversion(dataset=dataset, linear_obj_list=[mapper]) + ) + + counts = {"curvature_matrix": 0, "data_vector": 0} + + for name in counts: + _count_evaluations( + monkeypatch, aa.InversionInterferometerMapping, name, counts + ) + + inversion = aa.Inversion(dataset=dataset, linear_obj_list=[mapper]) + + assert isinstance(inversion, aa.InversionInterferometerMapping) + + log_evidence_terms = _log_evidence_terms_from(inversion) + + assert counts == {"curvature_matrix": 1, "data_vector": 1} + assert log_evidence_terms == reference diff --git a/test_autoarray/inversion/inversion/interferometer_numba/test_interferometer_numba.py b/test_autoarray/inversion/inversion/interferometer_numba/test_interferometer_numba.py index b8e1196dc..def6e4636 100644 --- a/test_autoarray/inversion/inversion/interferometer_numba/test_interferometer_numba.py +++ b/test_autoarray/inversion/inversion/interferometer_numba/test_interferometer_numba.py @@ -512,3 +512,81 @@ def test__existing_sparse_path_is_unchanged_when_the_gate_rejects_the_geometry() rtol=1.0e-12, atol=0.0, ) + + +@pytest.mark.parametrize("parallel", [False, True]) +def test__numba_inversion__curvature_matrix_and_data_vector_evaluated_once_per_likelihood( + monkeypatch, parallel +): + """ + `fast_chi_squared`, `curvature_reg_matrix` and `reconstruction` all read `curvature_matrix` + (F, via the numba `curvature_matrix_diag`) and `data_vector` (D). The `direct_conv` kernel + must run once per inversion, not once per reader (autolens_profiling #326: F twice and + D four times on one figure of merit). + """ + import functools + + from autoarray.inversion.inversion.interferometer_numba import sparse as numba_sparse + + monkeypatch.setattr(numba_sparse, "_numba_parallel", lambda: parallel) + + mask, _, dataset_sparse = _dataset_from() + + mapper = _delaunay_mapper_from(mask=mask) + + def log_evidence_terms_from(inversion): + return ( + float(inversion.fast_chi_squared) + + float(inversion.regularization_term) + + float(inversion.log_det_curvature_reg_matrix_term) + - float(inversion.log_det_regularization_matrix_term) + ) + + reference = log_evidence_terms_from( + InversionInterferometerSparseNumba( + dataset=dataset_sparse, linear_obj_list=[mapper], xp=np + ) + ) + + counts = {"kernel": 0, "data_vector": 0} + + def counted(kernel): + @functools.wraps(kernel) + def wrapper(*args, **kwargs): + counts["kernel"] += 1 + return kernel(*args, **kwargs) + + return wrapper + + monkeypatch.setattr( + numba_util, "curvature_direct_conv", counted(numba_util.curvature_direct_conv) + ) + parallel_kernel = numba_util.direct_conv_parallel_kernel + monkeypatch.setattr( + numba_util, + "direct_conv_parallel_kernel", + lambda: counted(parallel_kernel()), + ) + + descriptor = InversionInterferometerSparse.__dict__["data_vector"] + func = descriptor.fget if isinstance(descriptor, property) else descriptor.func + + @functools.wraps(func) + def counted_data_vector(self): + counts["data_vector"] += 1 + return func(self) + + monkeypatch.setattr( + InversionInterferometerSparse, + "data_vector", + type(descriptor)(counted_data_vector), + ) + + inversion = InversionInterferometerSparseNumba( + dataset=dataset_sparse, linear_obj_list=[mapper], xp=np + ) + + log_evidence_terms = log_evidence_terms_from(inversion) + + assert counts == {"kernel": 1, "data_vector": 1} + assert log_evidence_terms == reference