Skip to content

Commit 90f6f0c

Browse files
authored
Merge pull request #460 from PyAutoLabs/feature/interferometer-operated-override
Honor operated_mapping_matrix_override in interferometer inversions
2 parents 25da365 + cff9e5f commit 90f6f0c

4 files changed

Lines changed: 232 additions & 8 deletions

File tree

autoarray/inversion/inversion/interferometer/abstract.py

Lines changed: 35 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import numpy as np
22
from typing import Dict, List, Optional, Union
33

4+
from autoarray import exc
45
from autoarray.dataset.interferometer.dataset import Interferometer
56
from autoarray.inversion.inversion.dataset_interface import DatasetInterface
67
from autoarray.inversion.inversion.abstract import AbstractInversion
@@ -69,13 +70,41 @@ def operated_mapping_matrix_list(self) -> List[np.ndarray]:
6970
This is used to construct the simultaneous linear equations which reconstruct the data.
7071
7172
This property returns the a list of each linear object's transformed mapping matrix.
73+
74+
A linear object may have a `operated_mapping_matrix_override` property, which bypasses the `mapping_matrix`
75+
computation and transformer operation and is directly placed in the `operated_mapping_matrix_list`. Because
76+
the override bypasses the transformer it must already be in the data's visibility space, with (complex)
77+
shape [total_visibilities, params] (e.g. computed via an analytic Fourier transform).
7278
"""
73-
return [
74-
self.transformer.transform_mapping_matrix(
75-
mapping_matrix=linear_obj.mapping_matrix, xp=self._xp
76-
)
77-
for linear_obj in self.linear_obj_list
78-
]
79+
operated_mapping_matrix_list = []
80+
81+
for linear_obj in self.linear_obj_list:
82+
operated_mapping_matrix_override = linear_obj.operated_mapping_matrix_override
83+
84+
if operated_mapping_matrix_override is not None:
85+
expected_shape = (self.data.shape[0], linear_obj.params)
86+
87+
if tuple(operated_mapping_matrix_override.shape) != expected_shape:
88+
raise exc.InversionException(
89+
f"The `operated_mapping_matrix_override` of a linear object input to an interferometer "
90+
f"inversion has shape {tuple(operated_mapping_matrix_override.shape)} but shape "
91+
f"{expected_shape} ([total_visibilities, params]) is required.\n\n"
92+
f"For an interferometer dataset the override bypasses the transformer entirely and is "
93+
f"placed directly in the `operated_mapping_matrix_list`, therefore it must be in the "
94+
f"data's visibility space (unlike the real-space `mapping_matrix`, which the transformer "
95+
f"maps to visibilities)."
96+
)
97+
98+
operated_mapping_matrix_list.append(operated_mapping_matrix_override)
99+
100+
else:
101+
operated_mapping_matrix_list.append(
102+
self.transformer.transform_mapping_matrix(
103+
mapping_matrix=linear_obj.mapping_matrix, xp=self._xp
104+
)
105+
)
106+
107+
return operated_mapping_matrix_list
79108

80109
@property
81110
def mapped_reconstructed_data_dict(

autoarray/inversion/inversion/interferometer/sparse.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import numpy as np
22
from typing import Dict, List, Union
33

4+
from autoarray import exc
45
from autoarray.dataset.interferometer.dataset import Interferometer
56
from autoarray.inversion.inversion.dataset_interface import DatasetInterface
67
from autoarray.inversion.inversion.interferometer.abstract import (
@@ -45,6 +46,16 @@ def __init__(
4546
The linear objects used to reconstruct the data's observed values. If multiple linear objects are passed
4647
the simultaneous linear equations are combined and solved simultaneously.
4748
"""
49+
for linear_obj in linear_obj_list:
50+
if linear_obj.operated_mapping_matrix_override is not None:
51+
raise exc.InversionException(
52+
"A linear object with an `operated_mapping_matrix_override` was passed to the sparse "
53+
"(w-tilde) interferometer inversion, which constructs its linear algebra without an "
54+
"explicit operated mapping matrix and therefore cannot apply the override.\n\n"
55+
"Use the mapping formalism instead (e.g. do not call `apply_sparse_operator` on the "
56+
"interferometer dataset)."
57+
)
58+
4859
super().__init__(
4960
dataset=dataset,
5061
linear_obj_list=linear_obj_list,

autoarray/inversion/linear_obj/linear_obj.py

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -122,7 +122,8 @@ def pixel_signals_from(self, signal_scale) -> np.ndarray:
122122
def operated_mapping_matrix_override(self) -> Optional[np.ndarray]:
123123
"""
124124
An `Inversion` takes the `mapping_matrix` of each linear object and combines it with the data's operators
125-
(e.g. a PSF for `Imaging` data) to compute the `operated_mapping_matrix`.
125+
(e.g. a PSF for `Imaging` data, the transformer for `Interferometer` data) to compute the
126+
`operated_mapping_matrix`.
126127
127128
If this property is overwritten this operation is not performed, with the `operated_mapping_matrix` output
128129
by this property automatically used instead.
@@ -132,9 +133,19 @@ def operated_mapping_matrix_override(self) -> Optional[np.ndarray]:
132133
region which is blurred into the masked region which is linear solved for. This flux is outside the region
133134
that defines the `mapping_matrix` and thus this override is required to properly incorporate it.
134135
136+
Because the override bypasses the data's operators entirely, it must be in the data's space, which depends
137+
on the dataset type being fitted:
138+
139+
- `Imaging`: a real matrix of dimensions (total_mask_pixels, total_parameters), e.g. the PSF-convolved
140+
image of each linear object's parameter.
141+
142+
- `Interferometer`: a complex matrix of dimensions (total_visibilities, total_parameters), e.g. the
143+
visibilities of each linear object's parameter computed via an analytic Fourier transform. The
144+
transformer (NUFFT / DFT) is not applied to the override.
145+
135146
Returns
136147
-------
137-
An operated mapping matrix of dimensions (total_mask_pixels, total_parameters) which overrides the mapping
148+
An operated mapping matrix in the data's space (see above) which overrides the mapping
138149
matrix calculations performed in the linear equation solvers.
139150
"""
140151
return None

test_autoarray/inversion/inversion/interferometer/test_interferometer.py

Lines changed: 173 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,179 @@ def test__fast_chi_squared(
6464
assert inversion.fast_chi_squared == pytest.approx(chi_squared, 1.0e-4)
6565

6666

67+
def test__operated_mapping_matrix_list__override_is_honored():
68+
mask = aa.Mask2D(
69+
mask=[
70+
[True, True, True, True, True, True, True],
71+
[True, True, True, True, True, True, True],
72+
[True, True, True, False, True, True, True],
73+
[True, True, False, False, False, True, True],
74+
[True, True, True, False, True, True, True],
75+
[True, True, True, True, True, True, True],
76+
[True, True, True, True, True, True, True],
77+
],
78+
pixel_scales=2.0,
79+
)
80+
81+
n_visibilities = 5
82+
rng = np.random.default_rng(seed=0)
83+
data = aa.Visibilities(
84+
visibilities=rng.normal(size=(n_visibilities, 2)).astype(np.float64)
85+
)
86+
noise_map = aa.VisibilitiesNoiseMap(
87+
visibilities=np.ones((n_visibilities, 2), dtype=np.float64)
88+
)
89+
uv_wavelengths = rng.normal(size=(n_visibilities, 2)).astype(np.float64)
90+
91+
dataset = aa.Interferometer(
92+
data=data,
93+
noise_map=noise_map,
94+
uv_wavelengths=uv_wavelengths,
95+
real_space_mask=mask,
96+
transformer_class=aa.TransformerDFT,
97+
)
98+
99+
mapping_matrix = np.ones((mask.pixels_in_mask, 1))
100+
override = (999.0 + 1.0j) * np.ones((n_visibilities, 1))
101+
102+
linear_obj_override = aa.m.MockLinearObjFuncList(
103+
parameters=1,
104+
mapping_matrix=mapping_matrix,
105+
operated_mapping_matrix_override=override,
106+
)
107+
linear_obj_no_override = aa.m.MockLinearObjFuncList(
108+
parameters=1,
109+
mapping_matrix=mapping_matrix,
110+
)
111+
112+
inversion = aa.Inversion(
113+
dataset=dataset,
114+
linear_obj_list=[linear_obj_override, linear_obj_no_override],
115+
)
116+
117+
operated_mapping_matrix_list = inversion.operated_mapping_matrix_list
118+
119+
assert operated_mapping_matrix_list[0] == pytest.approx(override, 1.0e-8)
120+
121+
transformed_mapping_matrix = dataset.transformer.transform_mapping_matrix(
122+
mapping_matrix=mapping_matrix
123+
)
124+
125+
assert operated_mapping_matrix_list[1] == pytest.approx(
126+
transformed_mapping_matrix, 1.0e-8
127+
)
128+
129+
assert inversion.operated_mapping_matrix[:, 0] == pytest.approx(
130+
override[:, 0], 1.0e-8
131+
)
132+
assert inversion.curvature_matrix.shape == (2, 2)
133+
assert inversion.data_vector.shape == (2,)
134+
135+
136+
def test__operated_mapping_matrix_override__wrong_shape_raises():
137+
mask = aa.Mask2D(
138+
mask=[
139+
[True, True, True, True, True, True, True],
140+
[True, True, True, True, True, True, True],
141+
[True, True, True, False, True, True, True],
142+
[True, True, False, False, False, True, True],
143+
[True, True, True, False, True, True, True],
144+
[True, True, True, True, True, True, True],
145+
[True, True, True, True, True, True, True],
146+
],
147+
pixel_scales=2.0,
148+
)
149+
150+
n_visibilities = 7
151+
rng = np.random.default_rng(seed=0)
152+
data = aa.Visibilities(
153+
visibilities=rng.normal(size=(n_visibilities, 2)).astype(np.float64)
154+
)
155+
noise_map = aa.VisibilitiesNoiseMap(
156+
visibilities=np.ones((n_visibilities, 2), dtype=np.float64)
157+
)
158+
uv_wavelengths = rng.normal(size=(n_visibilities, 2)).astype(np.float64)
159+
160+
dataset = aa.Interferometer(
161+
data=data,
162+
noise_map=noise_map,
163+
uv_wavelengths=uv_wavelengths,
164+
real_space_mask=mask,
165+
transformer_class=aa.TransformerDFT,
166+
)
167+
168+
# A real-space shaped override (e.g. [total_mask_pixels, params]) is not valid for an
169+
# interferometer inversion, whose override must be in visibility space.
170+
linear_obj = aa.m.MockLinearObjFuncList(
171+
parameters=1,
172+
mapping_matrix=np.ones((mask.pixels_in_mask, 1)),
173+
operated_mapping_matrix_override=np.ones((mask.pixels_in_mask, 1)),
174+
)
175+
176+
inversion = aa.Inversion(dataset=dataset, linear_obj_list=[linear_obj])
177+
178+
with pytest.raises(aa.exc.InversionException):
179+
inversion.operated_mapping_matrix_list
180+
181+
182+
def test__operated_mapping_matrix_override__sparse_operator_raises():
183+
mask = aa.Mask2D(
184+
mask=[
185+
[True, True, True, True, True, True, True],
186+
[True, True, True, True, True, True, True],
187+
[True, True, True, False, True, True, True],
188+
[True, True, False, False, False, True, True],
189+
[True, True, True, False, True, True, True],
190+
[True, True, True, True, True, True, True],
191+
[True, True, True, True, True, True, True],
192+
],
193+
pixel_scales=2.0,
194+
)
195+
196+
grid = aa.Grid2D.from_mask(mask=mask, over_sample_size=1)
197+
198+
mesh = aa.mesh.Delaunay(pixels=9)
199+
image_mesh = aa.image_mesh.Overlay(shape=(3, 3))
200+
image_mesh_grid = image_mesh.image_plane_mesh_grid_from(mask=mask, adapt_data=None)
201+
202+
interpolator = mesh.interpolator_from(
203+
source_plane_data_grid=grid,
204+
source_plane_mesh_grid=image_mesh_grid,
205+
)
206+
mapper = aa.Mapper(interpolator=interpolator)
207+
208+
n_visibilities = 5
209+
rng = np.random.default_rng(seed=0)
210+
data = aa.Visibilities(
211+
visibilities=rng.normal(size=(n_visibilities, 2)).astype(np.float64)
212+
)
213+
noise_map = aa.VisibilitiesNoiseMap(
214+
visibilities=np.ones((n_visibilities, 2), dtype=np.float64)
215+
)
216+
uv_wavelengths = rng.normal(size=(n_visibilities, 2)).astype(np.float64)
217+
218+
dataset_sparse = aa.Interferometer(
219+
data=data,
220+
noise_map=noise_map,
221+
uv_wavelengths=uv_wavelengths,
222+
real_space_mask=mask,
223+
transformer_class=aa.TransformerDFT,
224+
).apply_sparse_operator(use_jax=False)
225+
226+
linear_obj = aa.m.MockLinearObjFuncList(
227+
parameters=1,
228+
mapping_matrix=np.ones((mask.pixels_in_mask, 1)),
229+
operated_mapping_matrix_override=(999.0 + 1.0j)
230+
* np.ones((n_visibilities, 1)),
231+
)
232+
233+
with pytest.raises(aa.exc.InversionException):
234+
aa.Inversion(
235+
dataset=dataset_sparse,
236+
linear_obj_list=[mapper, linear_obj],
237+
)
238+
239+
67240
def test__curvature_matrix__interferometer_sparse_operator__delaunay__identical_to_mapping():
68241
mask = aa.Mask2D(
69242
mask=[

0 commit comments

Comments
 (0)