diff --git a/autoarray/dataset/interferometer/dataset.py b/autoarray/dataset/interferometer/dataset.py index df057bbab..ef1d1f1a6 100644 --- a/autoarray/dataset/interferometer/dataset.py +++ b/autoarray/dataset/interferometer/dataset.py @@ -1,6 +1,6 @@ import logging import numpy as np -from typing import Optional +from typing import Optional, Tuple from autonerves.fitsable import ndarray_via_fits_from from autonerves import cached_property @@ -338,6 +338,7 @@ def from_stream( use_jax: bool = False, show_progress: bool = False, batch_size: int = 128, + phase_centre: Optional[Tuple[float, float]] = None, ) -> "Interferometer": """ Build an array-free `Interferometer` by accumulating a stream of visibility chunks, @@ -360,6 +361,12 @@ def from_stream( Passed to `sparse_terms_from_chunks`. batch_size The number of source-pixel columns processed per batch by the sparse operator. + phase_centre + The `(y, x)` phase-centre shift in arcseconds applied to every chunk's + visibilities, `d' = d * exp(+2 pi i (u * x0 + v * y0))` (radians), so a source at + `(y0, x0)` lands at the image origin; recorded as `sparse_terms.phase_centre`. + `None` applies no shift (recorded as `(0.0, 0.0)`). See + `sparse_terms_from_chunks`. Raises ------ @@ -379,6 +386,7 @@ def from_stream( chunk_k=chunk_k, use_jax=use_jax, show_progress=show_progress, + phase_centre=phase_centre, ) return cls.from_sparse_terms( @@ -670,7 +678,10 @@ def apply_sparse_operator_from_chunks( The number of source-pixel columns processed per batch by the sparse operator. accumulator_kwargs Passed to `sparse_terms_from_chunks` (`transformer_class`, `method`, `eps`, - `chunk_size`, `chunk_k`, `use_jax`, `show_progress`). When not given, + `chunk_size`, `chunk_k`, `use_jax`, `show_progress`). `phase_centre` is rejected: + it would shift the operator's terms but not this dataset's retained `data`, so the + two would describe different phase centres; use `Interferometer.from_stream` + (array-free, no retained data) to stream with a phase-centre shift. When not given, `transformer_class`, `eps` and `chunk_size` follow this dataset's transformer (the same defaults `psf_precision_operator_from` takes), so the accumulated terms match `apply_sparse_operator()`. @@ -684,8 +695,19 @@ def apply_sparse_operator_from_chunks( Raises ------ exc.DatasetException - If any chunk has unequal real and imaginary noise sigma. + If any chunk has unequal real and imaginary noise sigma, or a `phase_centre` is + passed. """ + if accumulator_kwargs.get("phase_centre") is not None: + raise exc.DatasetException( + "Interferometer.apply_sparse_operator_from_chunks does not take a " + "`phase_centre`: it would shift the sparse operator's dirty image and " + "`data_term` but not this dataset's retained `data`, so the two would describe " + "different phase centres and give contradictory likelihoods. Use " + "`Interferometer.from_stream(..., phase_centre=...)`, which builds an " + "array-free dataset with no retained visibilities." + ) + if disable_jax() and accumulator_kwargs.get("use_jax", False): accumulator_kwargs["use_jax"] = False diff --git a/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py b/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py index 25360e114..cad4afdba 100644 --- a/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py +++ b/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py @@ -1926,6 +1926,12 @@ class SparseTerms: The NUFFT precision the precision operator was built with. transformer_class_name The class name of the transformer that formed the dirty image and beam. + phase_centre + The `(y, x)` phase-centre shift in arcseconds the visibilities were re-centred on + before forming the dirty image (`sparse_terms_from_chunks(phase_centre=...)`, which + records `(0.0, 0.0)` when no shift is applied). `None` means not recorded. Terms with + different phase centres have dirty images referred to different sky origins and + cannot be summed. """ nufft_precision_operator: np.ndarray @@ -1940,13 +1946,15 @@ class SparseTerms: origin: Optional[tuple] = None eps: Optional[float] = None transformer_class_name: Optional[str] = None + phase_centre: Optional[Tuple[float, float]] = None def __add__(self, other: "SparseTerms") -> "SparseTerms": """ The field-wise sum of two `SparseTerms`. Raises `exc.InversionException` if a provenance field (`shape_native`, - `pixel_scales`, `origin`, `eps`) is recorded on both sides with different values. + `pixel_scales`, `origin`, `eps`, `phase_centre`, `transformer_class_name`) is recorded on + both sides with different values. The result carries, for each provenance field, the recorded value from either side (the left one when both are recorded), so an unrecorded operand never erases the provenance of a recorded one. @@ -1965,14 +1973,23 @@ def __add__(self, other: "SparseTerms") -> "SparseTerms": f"{other.nufft_precision_operator.shape} differ." ) - for name in ("shape_native", "pixel_scales", "origin", "eps"): + for name in ( + "shape_native", + "pixel_scales", + "origin", + "eps", + "phase_centre", + "transformer_class_name", + ): value_self = getattr(self, name) value_other = getattr(other, name) if value_self is None or value_other is None: continue - if isinstance(value_self, float) or isinstance(value_other, float): + if isinstance(value_self, str) or isinstance(value_other, str): + differ = str(value_self) != str(value_other) + elif isinstance(value_self, float) or isinstance(value_other, float): differ = float(value_self) != float(value_other) else: differ = tuple(value_self) != tuple(value_other) @@ -1981,8 +1998,8 @@ def __add__(self, other: "SparseTerms") -> "SparseTerms": raise exc.InversionException( "SparseTerms can only be added when accumulated with the same provenance: " f"`{name}` is {value_self!r} on one and {value_other!r} on the other. " - "Terms accumulated on different real-space masks or NUFFT accuracies do " - "not describe the same operator." + "Terms accumulated on different real-space masks, NUFFT accuracies, " + "phase centres or transformers do not describe the same operator." ) return SparseTerms( @@ -2003,8 +2020,23 @@ def __add__(self, other: "SparseTerms") -> "SparseTerms": transformer_class_name=_recorded( self.transformer_class_name, other.transformer_class_name ), + phase_centre=_recorded(self.phase_centre, other.phase_centre), ) + def __radd__(self, other) -> "SparseTerms": + """ + Support `sum(list_of_terms)`, which starts from the integer `0`: `0 + terms` is + `terms`. Any other left operand is not supported. + """ + if ( + isinstance(other, (int, float)) + and not isinstance(other, bool) + and other == 0 + ): + return self + + return NotImplemented + def _recorded(value_left, value_right): """ @@ -2042,6 +2074,7 @@ def sparse_terms_from_chunks( chunk_k: int = 2048, use_jax: bool = False, show_progress: bool = False, + phase_centre: Optional[Tuple[float, float]] = None, ) -> SparseTerms: """ Accumulate the `SparseTerms` of an interferometer dataset one chunk of visibilities at a @@ -2068,6 +2101,35 @@ def sparse_terms_from_chunks( `K` may differ between chunks; empty chunks are skipped. Chunks are consumed once, in order, and only one is referenced at a time. + Multi-channel (MFS) terms + ------------------------- + Because every field is a sum over visibilities, the terms of several channels (or any + partition of the visibilities) accumulated separately on the same mask sum to the terms + of all of them accumulated together: `sum(per_channel_terms)` (via `SparseTerms.__add__` + / `__radd__`) is the multi-frequency-synthesis (MFS) terms, equal to one accumulation + over every channel's chunks to summation order. + + Phase-centre shift + ------------------ + With `phase_centre=(y0, x0)` (arcseconds, autoarray `(y, x)` order like a mask `origin`), + each chunk's visibilities are multiplied by the unit phase + + d' = d * exp(+2 pi i (u * l0 + v * m0)), l0 = x0, m0 = y0 in radians, + + before any data-dependent term is formed. The forward transform is + `V(u, v) = sum I exp(-2 pi i (u x + v y))`, so this re-centres the phase centre onto + `(y0, x0)`: a source at `(y0, x0)` lands at the image origin. The shifted visibilities + form both the dirty image and `data_term`, so every data-dependent term describes the same + visibilities. The precision operator, dirty beam, `sum_weights`, `noise_normalization` and + `n_vis` depend on the baselines and sigmas alone and are bit-identical to the unshifted + accumulation. `data_term` is invariant under a unit phase only when the real and imaginary + sigmas are exactly equal; `check_noise_map_real_imag_equal` accepts them equal to a relative + tolerance, so with slightly unequal sigmas it differs from the unshifted value at that + tolerance (it is still the correct `data_term` of the shifted data). The shift is recorded + as `SparseTerms.phase_centre` provenance (`(0.0, 0.0)` when no shift is applied), so terms + with different phase centres -- including shifted and unshifted ones -- refuse to be + summed. + Parameters ---------- chunks @@ -2082,6 +2144,10 @@ def sparse_terms_from_chunks( method, eps, chunk_size, chunk_k, use_jax, show_progress Passed to `nufft_precision_operator_from` for each chunk (`chunk_size` is the NUFFT builder's own inner chunk, a memory ceiling within a chunk). `eps=None` is `1e-12`. + phase_centre + The `(y, x)` phase-centre shift in arcseconds applied to every chunk's visibilities + before forming the dirty image (see "Phase-centre shift" above). `None` applies no + shift and records `(0.0, 0.0)`. Returns ------- @@ -2105,6 +2171,18 @@ def sparse_terms_from_chunks( if eps is None: eps = 1.0e-12 + shift = phase_centre is not None + + if not shift: + phase_centre = (0.0, 0.0) + else: + from astropy import units + + phase_centre = (float(phase_centre[0]), float(phase_centre[1])) + arcsec_to_rad = units.arcsec.to(units.rad) + m0 = phase_centre[0] * arcsec_to_rad + l0 = phase_centre[1] * arcsec_to_rad + terms = None for uv_wavelengths, data, noise_map in chunks: @@ -2135,6 +2213,15 @@ def sparse_terms_from_chunks( uv_wavelengths=uv_wavelengths, real_space_mask=real_space_mask ) + # The phase-centre shift is applied to the data before every data-dependent term (the + # dirty image and `data_term`), so all of them describe the same shifted visibilities. + # The precision operator, dirty beam, `sum_weights` and `noise_normalization` depend on + # the baselines and sigmas alone. + if shift: + data = data * np.exp( + 2j * np.pi * (uv_wavelengths[:, 0] * l0 + uv_wavelengths[:, 1] * m0) + ) + # The same arguments `Interferometer.psf_precision_operator_from` builds from the # dataset's own transformer, so the per-chunk operators sum to the dataset's. nufft_precision_operator = np.asarray( @@ -2192,6 +2279,7 @@ def sparse_terms_from_chunks( origin=tuple(real_space_mask.origin), eps=float(eps), transformer_class_name=type(transformer).__name__, + phase_centre=phase_centre, ) terms = chunk_terms if terms is None else terms + chunk_terms diff --git a/test_autoarray/dataset/interferometer/test_dataset.py b/test_autoarray/dataset/interferometer/test_dataset.py index ff0482c31..345f04eb3 100644 --- a/test_autoarray/dataset/interferometer/test_dataset.py +++ b/test_autoarray/dataset/interferometer/test_dataset.py @@ -670,6 +670,29 @@ def test__apply_sparse_operator_from_chunks__unequal_real_imag_noise__raises( dataset.apply_sparse_operator_from_chunks(chunks) +def test__apply_sparse_operator_from_chunks__phase_centre__raises(mask_2d_7x7): + """ + A `phase_centre` would shift the operator's terms but not the retained `data`. Codex's + example: a quarter-turn baseline with `d = -1j` makes the shifted sparse chi-squared of a + unit point at the origin 0 while the residual against the retained data is 2. + """ + dataset = _random_interferometer(mask_2d_7x7, transformer.TransformerDFT) + + chunks = [ + (dataset.uv_wavelengths, dataset.data.array, dataset.noise_map.array), + ] + + with pytest.raises(aa.exc.DatasetException, match="from_stream"): + dataset.apply_sparse_operator_from_chunks(chunks, phase_centre=(0.0, 1.0)) + + # An explicit `phase_centre=None` (no shift) is accepted. + dataset_sparse = dataset.apply_sparse_operator_from_chunks( + chunks, phase_centre=None, method="numpy" + ) + + assert dataset_sparse.sparse_terms.phase_centre == (0.0, 0.0) + + def _chunks_of(dataset, edges): return [ ( @@ -861,3 +884,175 @@ def test__apply_sparse_operator_from_chunks__result_carries_sparse_terms(mask_2d # The in-memory `apply_sparse_operator` path records no terms. assert dataset.apply_sparse_operator().sparse_terms is None + + +def _delaunay_mapper(mask): + grid = aa.Grid2D.from_mask(mask=mask, over_sample_size=1) + mesh = aa.mesh.Delaunay(pixels=9) + image_mesh_grid = aa.image_mesh.Overlay(shape=(3, 3)).image_plane_mesh_grid_from( + mask=mask, adapt_data=None + ) + return aa.Mapper( + interpolator=mesh.interpolator_from( + source_plane_data_grid=grid, source_plane_mesh_grid=image_mesh_grid + ), + regularization=aa.reg.Constant(coefficient=1.0), + ) + + +def _assert_sparse_fits_match(dataset_a, dataset_b, mapper, rel): + inversion_a = aa.Inversion(dataset=dataset_a, linear_obj_list=[mapper]) + inversion_b = aa.Inversion(dataset=dataset_b, linear_obj_list=[mapper]) + + assert isinstance(inversion_a, aa.InversionInterferometerSparse) + assert isinstance(inversion_b, aa.InversionInterferometerSparse) + + for name in ( + "fast_chi_squared", + "regularization_term", + "log_det_curvature_reg_matrix_term", + "log_det_regularization_matrix_term", + ): + assert float(getattr(inversion_a, name)) == pytest.approx( + float(getattr(inversion_b, name)), rel=rel + ), name + + log_evidence_a = aa.m.MockFitInterferometer( + dataset=dataset_a, inversion=inversion_a + ).log_evidence + log_evidence_b = aa.m.MockFitInterferometer( + dataset=dataset_b, inversion=inversion_b + ).log_evidence + + assert float(log_evidence_a) == pytest.approx(float(log_evidence_b), rel=rel) + + +def test__from_stream__phase_centre__matches_pre_shifted_in_memory_dataset( + mask_2d_7x7, +): + pytest.importorskip("nufftax") + + dataset = _random_interferometer(mask_2d_7x7, transformer.TransformerNUFFT) + + phase_centre = (0.4, -0.9) + + dataset_stream = aa.Interferometer.from_stream( + _chunks_of(dataset, [0, 13, 40]), mask_2d_7x7, phase_centre=phase_centre + ) + + assert dataset_stream.is_array_free + assert dataset_stream.sparse_terms.phase_centre == phase_centre + + uv_wavelengths = dataset.uv_wavelengths + m0, l0 = np.deg2rad(np.asarray(phase_centre) / 3600.0) + data_shifted = dataset.data.array * np.exp( + 2j * np.pi * (uv_wavelengths[:, 0] * l0 + uv_wavelengths[:, 1] * m0) + ) + + dataset_shifted = aa.Interferometer( + data=aa.Visibilities(visibilities=data_shifted), + noise_map=dataset.noise_map, + uv_wavelengths=uv_wavelengths, + real_space_mask=mask_2d_7x7, + transformer_class=transformer.TransformerNUFFT, + ).apply_sparse_operator() + + stream = dataset_stream.sparse_operator + shifted = dataset_shifted.sparse_operator + + np.testing.assert_allclose( + stream.dirty_image, + shifted.dirty_image, + rtol=1.0e-12, + atol=1.0e-12 * np.abs(shifted.dirty_image).max(), + ) + assert stream.data_term == pytest.approx(shifted.data_term, rel=1.0e-12) + assert stream.noise_normalization == pytest.approx( + shifted.noise_normalization, rel=1.0e-12 + ) + + _assert_sparse_fits_match( + dataset_stream, dataset_shifted, _delaunay_mapper(mask_2d_7x7), rel=1.0e-10 + ) + + # The unshifted stream records a zero phase centre, so the two cannot be summed. + dataset_unshifted = aa.Interferometer.from_stream( + _chunks_of(dataset, [0, 40]), mask_2d_7x7 + ) + + assert dataset_unshifted.sparse_terms.phase_centre == (0.0, 0.0) + + with pytest.raises(aa.exc.InversionException, match="phase_centre"): + dataset_stream.sparse_terms + dataset_unshifted.sparse_terms + + +def test__from_sparse_terms__sum_of_per_channel_terms_matches_in_memory_mfs( + mask_2d_7x7, +): + pytest.importorskip("nufftax") + + channels = [ + _random_interferometer( + mask_2d_7x7, + transformer.TransformerNUFFT, + n_visibilities=n_visibilities, + seed=seed, + ) + for n_visibilities, seed in ((40, 21), (23, 22), (31, 23)) + ] + + per_channel_terms = [ + aa.Interferometer.from_stream( + _chunks_of(channel, [0, 10, channel.uv_wavelengths.shape[0]]), + mask_2d_7x7, + ).sparse_terms + for channel in channels + ] + + dataset_mfs_stream = aa.Interferometer.from_sparse_terms( + sum(per_channel_terms), real_space_mask=mask_2d_7x7 + ) + + assert dataset_mfs_stream.is_array_free + assert dataset_mfs_stream.sparse_terms.n_vis == 94 + + dataset_mfs_memory = aa.Interferometer( + data=aa.Visibilities( + visibilities=np.concatenate([channel.data.array for channel in channels]) + ), + noise_map=aa.VisibilitiesNoiseMap( + visibilities=np.concatenate( + [channel.noise_map.array for channel in channels] + ) + ), + uv_wavelengths=np.concatenate([channel.uv_wavelengths for channel in channels]), + real_space_mask=mask_2d_7x7, + transformer_class=transformer.TransformerNUFFT, + ).apply_sparse_operator() + + stream = dataset_mfs_stream.sparse_operator + memory = dataset_mfs_memory.sparse_operator + + np.testing.assert_allclose( + stream.nufft_precision_operator, + memory.nufft_precision_operator, + rtol=1.0e-12, + atol=1.0e-12 * np.abs(memory.nufft_precision_operator).max(), + ) + np.testing.assert_allclose( + stream.dirty_image, + memory.dirty_image, + rtol=1.0e-12, + atol=1.0e-12 * np.abs(memory.dirty_image).max(), + ) + assert stream.data_term == pytest.approx(memory.data_term, rel=1.0e-12) + assert stream.noise_normalization == pytest.approx( + memory.noise_normalization, rel=1.0e-12 + ) + + _assert_sparse_fits_match( + dataset_mfs_stream, + dataset_mfs_memory, + _delaunay_mapper(mask_2d_7x7), + rel=1.0e-10, + ) diff --git a/test_autoarray/inversion/inversion/interferometer/test_inversion_interferometer_util.py b/test_autoarray/inversion/inversion/interferometer/test_inversion_interferometer_util.py index 8cc511715..545d7cc5b 100644 --- a/test_autoarray/inversion/inversion/interferometer/test_inversion_interferometer_util.py +++ b/test_autoarray/inversion/inversion/interferometer/test_inversion_interferometer_util.py @@ -1428,6 +1428,7 @@ def test__sparse_terms__add__unrecorded_plus_unrecorded_stays_unrecorded(): "origin", "eps", "transformer_class_name", + "phase_centre", ): assert getattr(total, name) is None @@ -1448,6 +1449,7 @@ def test__sparse_terms_from_chunks__records_provenance(): assert terms.origin == (0.0, 0.0) assert terms.eps == 1.0e-10 assert terms.transformer_class_name == "TransformerNUFFT" + assert terms.phase_centre == (0.0, 0.0) terms_dft = aa.util.inversion_interferometer.sparse_terms_from_chunks( _chunks_from(uv_wavelengths, data, noise_map, [0, 20]), @@ -1461,3 +1463,366 @@ def test__sparse_terms_from_chunks__records_provenance(): # Terms accumulated at different NUFFT accuracies cannot be summed. with pytest.raises(aa.exc.InversionException): terms + terms_dft + + +def _assert_terms_match_operator(terms, operator, mask, n_vis, sum_weights, rel): + np.testing.assert_allclose( + terms.nufft_precision_operator, + operator.nufft_precision_operator, + rtol=rel, + atol=rel * np.abs(operator.nufft_precision_operator).max(), + ) + + dirty_image = aa.Array2D(values=terms.dirty_image_native, mask=mask).slim.array + + np.testing.assert_allclose( + dirty_image, + operator.dirty_image, + rtol=rel, + atol=rel * np.abs(operator.dirty_image).max(), + ) + + assert terms.data_term == pytest.approx(operator.data_term, rel=rel) + assert terms.noise_normalization == pytest.approx( + operator.noise_normalization, rel=rel + ) + assert terms.sum_weights == pytest.approx(sum_weights, rel=rel) + assert terms.n_vis == n_vis + + +@pytest.mark.parametrize( + "method, use_jax", [("nufft", False), ("numpy", False), ("numpy", True)] +) +def test__sparse_terms__sum_of_channels_equals_mfs(method, use_jax, monkeypatch): + """ + `use_jax` only reaches the accumulator through `nufft_precision_operator_from`, which + ignores it under the default `method="nufft"` (already JAX). The JAX leg therefore uses + `method="numpy"`, which `use_jax=True` upgrades to the JAX brute force; a spy asserts that + builder actually runs. + """ + pytest.importorskip("nufftax") + + util = aa.util.inversion_interferometer + + jax_builder_calls = [] + + if use_jax: + pytest.importorskip("jax") + + if util.disable_jax(): + pytest.skip("PYAUTO_DISABLE_JAX demotes the JAX brute force to NumPy.") + + jax_builder = util.nufft_precision_operator_via_jax_from + + def spy(**kwargs): + jax_builder_calls.append(1) + return jax_builder(**kwargs) + + monkeypatch.setattr(util, "nufft_precision_operator_via_jax_from", spy) + + channels = [ + _streaming_inputs(n_visibilities=n_visibilities, seed=seed) + for n_visibilities, seed in ((40, 11), (25, 12), (57, 13)) + ] + + mask = channels[0][0] + + per_channel_terms = [ + aa.util.inversion_interferometer.sparse_terms_from_chunks( + _chunks_from( + uv_wavelengths, + data, + noise_map, + [0, uv_wavelengths.shape[0] // 2, uv_wavelengths.shape[0]], + ), + real_space_mask=mask, + method=method, + use_jax=use_jax, + ) + for _, uv_wavelengths, data, noise_map, _ in channels + ] + + # `sum` starts from `0`, exercising `SparseTerms.__radd__`. + mfs_terms = sum(per_channel_terms) + + assert isinstance(mfs_terms, aa.SparseTerms) + + uv_wavelengths = np.concatenate([channel[1] for channel in channels]) + data = np.concatenate([channel[2] for channel in channels]) + noise_map = np.concatenate([channel[3] for channel in channels]) + n_vis = uv_wavelengths.shape[0] + + one_accumulation = aa.util.inversion_interferometer.sparse_terms_from_chunks( + [ + chunk + for _, uv_c, data_c, noise_c, _ in channels + for chunk in _chunks_from(uv_c, data_c, noise_c, [0, uv_c.shape[0]]) + ], + real_space_mask=mask, + method=method, + use_jax=use_jax, + ) + + dataset_mfs = aa.Interferometer( + data=aa.Visibilities(visibilities=data), + noise_map=aa.VisibilitiesNoiseMap(visibilities=noise_map), + uv_wavelengths=uv_wavelengths, + real_space_mask=mask, + transformer_class=aa.TransformerNUFFT, + ) + operator_mfs = dataset_mfs.apply_sparse_operator( + method=method, use_jax=use_jax + ).sparse_operator + + # 3 channels x 2 chunks, 3 chunks of the one accumulation, 1 in-memory build. + assert len(jax_builder_calls) == (10 if use_jax else 0) + + sum_weights = np.sum(noise_map.real**-2.0) + + for terms in (mfs_terms, one_accumulation): + _assert_terms_match_operator( + terms, operator_mfs, mask, n_vis, sum_weights, rel=1.0e-12 + ) + + np.testing.assert_allclose( + mfs_terms.dirty_beam_native, + one_accumulation.dirty_beam_native, + rtol=1.0e-12, + atol=1.0e-12 * np.abs(one_accumulation.dirty_beam_native).max(), + ) + + +def test__sparse_terms__radd__supports_zero_only(): + terms = _terms_with_provenance(1.0, 3) + + assert (0 + terms) is terms + assert sum([terms]) is terms + + with pytest.raises(TypeError): + 1 + terms + + with pytest.raises(TypeError): + "a" + terms + + +def _phase_shifted(uv_wavelengths, data, phase_centre): + m0, l0 = np.deg2rad(np.asarray(phase_centre) / 3600.0) + + return data * np.exp( + 2j * np.pi * (uv_wavelengths[:, 0] * l0 + uv_wavelengths[:, 1] * m0) + ) + + +@pytest.mark.parametrize("transformer_class", [aa.TransformerNUFFT, aa.TransformerDFT]) +def test__sparse_terms_from_chunks__phase_centre__matches_in_memory_on_shifted_data( + transformer_class, +): + pytest.importorskip("nufftax") + + mask, uv_wavelengths, data, noise_map, _ = _streaming_inputs() + + phase_centre = (0.7, -1.3) + + chunks = _chunks_from(uv_wavelengths, data, noise_map, [0, 17, 60]) + + terms_shifted = aa.util.inversion_interferometer.sparse_terms_from_chunks( + chunks, + real_space_mask=mask, + transformer_class=transformer_class, + phase_centre=phase_centre, + ) + terms_unshifted = aa.util.inversion_interferometer.sparse_terms_from_chunks( + chunks, + real_space_mask=mask, + transformer_class=transformer_class, + ) + + assert terms_shifted.phase_centre == phase_centre + assert terms_unshifted.phase_centre == (0.0, 0.0) + + # Shifted and unshifted terms describe different sky origins and cannot be summed. + with pytest.raises(aa.exc.InversionException, match="phase_centre"): + terms_shifted + terms_unshifted + + dataset_shifted = aa.Interferometer( + data=aa.Visibilities( + visibilities=_phase_shifted(uv_wavelengths, data, phase_centre) + ), + noise_map=aa.VisibilitiesNoiseMap(visibilities=noise_map), + uv_wavelengths=uv_wavelengths, + real_space_mask=mask, + transformer_class=transformer_class, + ) + operator_shifted = dataset_shifted.apply_sparse_operator().sparse_operator + + _assert_terms_match_operator( + terms_shifted, + operator_shifted, + mask, + 60, + np.sum(noise_map.real**-2.0), + rel=1.0e-12, + ) + + # The shift is a real change to the dirty image ... + assert not np.allclose( + terms_shifted.dirty_image_native, terms_unshifted.dirty_image_native + ) + + # ... leaves every (uv, sigma)-only term bit-identical to the unshifted accumulation ... + np.testing.assert_array_equal( + terms_shifted.nufft_precision_operator, terms_unshifted.nufft_precision_operator + ) + np.testing.assert_array_equal( + terms_shifted.dirty_beam_native, terms_unshifted.dirty_beam_native + ) + for name in ("sum_weights", "noise_normalization", "n_vis"): + assert getattr(terms_shifted, name) == getattr(terms_unshifted, name) + + # ... and `data_term`, formed from the shifted data, is phase-invariant (to rounding) + # because these sigmas are exactly equal in real and imaginary parts. + assert terms_shifted.data_term == pytest.approx( + terms_unshifted.data_term, rel=1.0e-12 + ) + + +def test__sparse_terms_from_chunks__phase_centre__point_source_recentres(): + # An odd-shaped mask, so the image origin (0, 0) is the centre pixel (5, 5). + mask = aa.Mask2D.all_false(shape_native=(11, 11), pixel_scales=0.5) + + # A point source at (y, x) = (1.0, -1.5) arcsec: native pixel (3, 2). + phase_centre = (1.0, -1.5) + image_native = np.zeros((11, 11)) + image_native[3, 2] = 1.0 + + np.testing.assert_array_equal( + mask.derive_grid.all_false.native.array[3, 2], phase_centre + ) + + rng = np.random.default_rng(seed=1) + uv_wavelengths = rng.normal(size=(400, 2)) * 3.0e5 + + data = np.asarray( + aa.TransformerDFT(uv_wavelengths=uv_wavelengths, real_space_mask=mask) + .visibilities_from(image=aa.Array2D(values=image_native, mask=mask)) + .array + ) + noise_map = np.ones(400) + 1j * np.ones(400) + + def peak_pixel(phase_centre): + terms = aa.util.inversion_interferometer.sparse_terms_from_chunks( + _chunks_from(uv_wavelengths, data, noise_map, [0, 150, 400]), + real_space_mask=mask, + transformer_class=aa.TransformerDFT, + phase_centre=phase_centre, + ) + return np.unravel_index(np.argmax(terms.dirty_image_native), mask.shape_native) + + assert peak_pixel(None) == (3, 2) + assert peak_pixel(phase_centre) == (5, 5) + + # (x, y) order instead of (y, x) moves the source elsewhere, not to the centre. + assert peak_pixel((phase_centre[1], phase_centre[0])) != (5, 5) + + +def test__sparse_terms__add__phase_centre_mismatch_raises(): + with pytest.raises(aa.exc.InversionException, match="phase_centre"): + _terms_with_provenance(phase_centre=(0.0, 0.0)) + _terms_with_provenance( + phase_centre=(0.5, 0.0) + ) + + total = _terms_with_provenance(phase_centre=(0.5, 0.0)) + _terms_with_provenance( + phase_centre=(0.5, 0.0) + ) + + assert total.phase_centre == (0.5, 0.0) + + +def test__sparse_terms__add__unrecorded_left_operand_keeps_right_phase_centre(): + unknown = _terms_with_provenance() + a = _terms_with_provenance(phase_centre=(0.1, 0.2)) + b = _terms_with_provenance(phase_centre=(0.3, 0.2)) + + assert (unknown + a).phase_centre == (0.1, 0.2) + assert (a + unknown).phase_centre == (0.1, 0.2) + assert (unknown + unknown).phase_centre is None + + with pytest.raises(aa.exc.InversionException, match="phase_centre"): + (unknown + a) + b + + +def test__sparse_terms_from_chunks__phase_centre__data_term_uses_shifted_data(): + """ + `check_noise_map_real_imag_equal` accepts real and imaginary sigmas equal to a relative + tolerance, so `data_term` is not phase-invariant for slightly unequal sigmas. It must then + be formed from the shifted data, like the dirty image: a quarter-turn baseline moves + `d = 10000` wholly into the imaginary part, whose sigma is 1.000009. + """ + from astropy import units + + arcsec_to_rad = units.arcsec.to(units.rad) + + mask = aa.Mask2D.all_false(shape_native=(3, 3), pixel_scales=1.0) + uv_wavelengths = np.array([[1.0 / (4.0 * arcsec_to_rad), 0.0]]) + data = np.array([10000.0 + 0.0j]) + noise_map = np.array([1.0 + 1.000009j]) + + terms = aa.util.inversion_interferometer.sparse_terms_from_chunks( + [(uv_wavelengths, data, noise_map)], + real_space_mask=mask, + transformer_class=aa.TransformerDFT, + phase_centre=(0.0, 1.0), + ) + + data_shifted = _phase_shifted(uv_wavelengths, data, (0.0, 1.0)) + + data_term_shifted = float( + np.sum(data_shifted.real**2.0 / noise_map.real**2.0) + + np.sum(data_shifted.imag**2.0 / noise_map.imag**2.0) + ) + + assert data_term_shifted == pytest.approx(99998200.0243, rel=1.0e-10) + assert terms.data_term == pytest.approx(data_term_shifted, rel=1.0e-12) + + # The dirty image is formed from the same shifted data. + dataset_shifted = aa.Interferometer( + data=aa.Visibilities(visibilities=data_shifted), + noise_map=aa.VisibilitiesNoiseMap(visibilities=noise_map), + uv_wavelengths=uv_wavelengths, + real_space_mask=mask, + transformer_class=aa.TransformerDFT, + ) + operator_shifted = dataset_shifted.apply_sparse_operator( + method="numpy" + ).sparse_operator + + assert terms.data_term == pytest.approx(operator_shifted.data_term, rel=1.0e-12) + np.testing.assert_allclose( + aa.Array2D(values=terms.dirty_image_native, mask=mask).slim.array, + operator_shifted.dirty_image, + rtol=1.0e-12, + atol=1.0e-12 * np.abs(operator_shifted.dirty_image).max(), + ) + + +def test__sparse_terms__add__transformer_class_name_mismatch_raises(): + with pytest.raises(aa.exc.InversionException, match="transformer_class_name"): + _terms_with_provenance( + transformer_class_name="TransformerNUFFT" + ) + _terms_with_provenance(transformer_class_name="TransformerDFT") + + total = _terms_with_provenance( + transformer_class_name="TransformerDFT" + ) + _terms_with_provenance(transformer_class_name="TransformerDFT") + + assert total.transformer_class_name == "TransformerDFT" + + # Unrecorded on either side skips the check and carries the recorded value. + unknown = _terms_with_provenance() + + assert ( + unknown + _terms_with_provenance(transformer_class_name="TransformerDFT") + ).transformer_class_name == "TransformerDFT" + assert ( + _terms_with_provenance(transformer_class_name="TransformerDFT") + unknown + ).transformer_class_name == "TransformerDFT"