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
6 changes: 4 additions & 2 deletions autoarray/inversion/inversion/interferometer/mapping.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
8 changes: 5 additions & 3 deletions autoarray/inversion/inversion/interferometer/sparse.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 @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading