From d3069dbdb4beebabdba7f8d798624729b80198a6 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Thu, 1 Oct 2026 10:20:35 +0100 Subject: [PATCH 1/2] feat: per-channel SparseTerms sum and phase-centre shifts in sparse_terms_from_chunks (streaming P5) - SparseTerms gains `phase_centre` provenance ((y, x) arcsec), checked in __add__ like origin/eps, and __radd__ for 0 so sum(list_of_terms) works (MFS terms = sum of per-channel terms). - sparse_terms_from_chunks(phase_centre=...) multiplies each chunk's visibilities by exp(+2 pi i (u l0 + v m0)) before forming the dirty image, re-centring a source at (y0, x0) onto the origin; every other term is built from the unshifted chunk and is bit-identical. Unshifted accumulations record (0.0, 0.0), so shifted and unshifted terms refuse to be summed. - Interferometer.from_stream forwards phase_centre. - Tests: channel sum == one accumulation == in-memory MFS (numpy + JAX), phase centre == in-memory on pre-shifted data (NUFFT + DFT), DFT point-source sign/order pin, provenance checks, array-free inversion parity. Refs PyAutoLabs/PyAutoArray#600 Co-Authored-By: Claude Fable 5.1 --- autoarray/dataset/interferometer/dataset.py | 14 +- .../inversion_interferometer_util.py | 87 +++++- .../dataset/interferometer/test_dataset.py | 172 ++++++++++++ .../test_inversion_interferometer_util.py | 252 ++++++++++++++++++ 4 files changed, 517 insertions(+), 8 deletions(-) diff --git a/autoarray/dataset/interferometer/dataset.py b/autoarray/dataset/interferometer/dataset.py index df057bbab..88be6f2f8 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,9 @@ 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`). A + `phase_centre` shifts only the operator's dirty image, not this dataset's retained + `data`, so the two then describe different phase centres. 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()`. diff --git a/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py b/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py index 25360e114..9e10d0934 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,14 @@ 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`) 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,7 +1972,7 @@ 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"): value_self = getattr(self, name) value_other = getattr(other, name) @@ -1981,8 +1988,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 or " + "phase centres do not describe the same operator." ) return SparseTerms( @@ -2003,8 +2010,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 +2064,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 +2091,32 @@ 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 the dirty image 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. Only `dirty_image_native` + changes. The precision operator, dirty beam, `sum_weights`, `noise_normalization` and + `n_vis` depend on the baselines and sigmas alone, and `data_term` is invariant under a unit + phase because the real and imaginary sigmas are equal; all are computed from the + unshifted chunk and are bit-identical to the unshifted accumulation. 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 +2131,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 +2158,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 +2200,15 @@ def sparse_terms_from_chunks( uv_wavelengths=uv_wavelengths, real_space_mask=real_space_mask ) + # The phase-centre shift only enters the dirty image; every other term is built from + # the unshifted chunk (they are invariant under a unit phase). + if not shift: + data_shifted = data + else: + data_shifted = 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( @@ -2157,8 +2231,8 @@ def sparse_terms_from_chunks( dirty_image_native = np.asarray( transformer.image_from( visibilities=Visibilities( - visibilities=data.real * noise_map_real**-2.0 - + 1j * data.imag * noise_map_imag**-2.0 + visibilities=data_shifted.real * noise_map_real**-2.0 + + 1j * data_shifted.imag * noise_map_imag**-2.0 ), ).native.array, dtype=np.float64, @@ -2192,6 +2266,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..d49b8e4e0 100644 --- a/test_autoarray/dataset/interferometer/test_dataset.py +++ b/test_autoarray/dataset/interferometer/test_dataset.py @@ -861,3 +861,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..172611719 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,253 @@ 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("use_jax", [False, True]) +def test__sparse_terms__sum_of_channels_equals_mfs(use_jax): + pytest.importorskip("nufftax") + + if use_jax: + pytest.importorskip("jax") + + 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, + 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, + 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(use_jax=use_jax).sparse_operator + + 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 + ) + + # ... and leaves every other 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", "data_term", "noise_normalization", "n_vis"): + assert getattr(terms_shifted, name) == getattr(terms_unshifted, name) + + +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 From 8a2332f763b905b982be46dcc0d38fc57b741069 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Thu, 1 Oct 2026 10:37:28 +0100 Subject: [PATCH 2/2] fix: phase shift feeds data_term, reject phase_centre on retained-data path, check transformer provenance (streaming P5 review) - sparse_terms_from_chunks applies the phase-centre shift to the data before every data-dependent term, so the dirty image and data_term describe the same shifted visibilities (with sigma_r ~= sigma_i within the check's rtol, the unshifted data_term was off: 1e8 vs 99998200.02 on a quarter-turn baseline). data_term is phase-invariant only for exactly equal sigmas; the P5 test now compares it at rel 1e-12. - Interferometer.apply_sparse_operator_from_chunks rejects phase_centre with a DatasetException pointing at from_stream: it shifted the operator but not the retained data (sparse chi2 0 vs 2 against the retained data). - SparseTerms.__add__ checks transformer_class_name with the other recorded provenance (mismatch raises, unrecorded skips). - The MFS channel-sum test's JAX leg now runs the JAX brute force (method="numpy" + use_jax=True, spied); method="nufft" ignores use_jax. Refs PyAutoLabs/PyAutoArray#600 Co-Authored-By: Claude Fable 5.1 --- autoarray/dataset/interferometer/dataset.py | 20 ++- .../inversion_interferometer_util.py | 51 +++++--- .../dataset/interferometer/test_dataset.py | 23 ++++ .../test_inversion_interferometer_util.py | 123 +++++++++++++++++- 4 files changed, 189 insertions(+), 28 deletions(-) diff --git a/autoarray/dataset/interferometer/dataset.py b/autoarray/dataset/interferometer/dataset.py index 88be6f2f8..ef1d1f1a6 100644 --- a/autoarray/dataset/interferometer/dataset.py +++ b/autoarray/dataset/interferometer/dataset.py @@ -678,9 +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`, `phase_centre`). A - `phase_centre` shifts only the operator's dirty image, not this dataset's retained - `data`, so the two then describe different phase centres. 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()`. @@ -694,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 9e10d0934..cad4afdba 100644 --- a/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py +++ b/autoarray/inversion/inversion/interferometer/inversion_interferometer_util.py @@ -1953,7 +1953,8 @@ 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`, `phase_centre`) 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. @@ -1972,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", "phase_centre"): + 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) @@ -1988,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, NUFFT accuracies or " - "phase centres 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( @@ -2106,13 +2116,16 @@ def sparse_terms_from_chunks( d' = d * exp(+2 pi i (u * l0 + v * m0)), l0 = x0, m0 = y0 in radians, - before the dirty image is formed. The forward transform is + 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. Only `dirty_image_native` - changes. The precision operator, dirty beam, `sum_weights`, `noise_normalization` and - `n_vis` depend on the baselines and sigmas alone, and `data_term` is invariant under a unit - phase because the real and imaginary sigmas are equal; all are computed from the - unshifted chunk and are bit-identical to the unshifted accumulation. The shift is recorded + `(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. @@ -2200,12 +2213,12 @@ def sparse_terms_from_chunks( uv_wavelengths=uv_wavelengths, real_space_mask=real_space_mask ) - # The phase-centre shift only enters the dirty image; every other term is built from - # the unshifted chunk (they are invariant under a unit phase). - if not shift: - data_shifted = data - else: - data_shifted = data * np.exp( + # 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) ) @@ -2231,8 +2244,8 @@ def sparse_terms_from_chunks( dirty_image_native = np.asarray( transformer.image_from( visibilities=Visibilities( - visibilities=data_shifted.real * noise_map_real**-2.0 - + 1j * data_shifted.imag * noise_map_imag**-2.0 + visibilities=data.real * noise_map_real**-2.0 + + 1j * data.imag * noise_map_imag**-2.0 ), ).native.array, dtype=np.float64, diff --git a/test_autoarray/dataset/interferometer/test_dataset.py b/test_autoarray/dataset/interferometer/test_dataset.py index d49b8e4e0..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 [ ( 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 172611719..545d7cc5b 100644 --- a/test_autoarray/inversion/inversion/interferometer/test_inversion_interferometer_util.py +++ b/test_autoarray/inversion/inversion/interferometer/test_inversion_interferometer_util.py @@ -1490,13 +1490,36 @@ def _assert_terms_match_operator(terms, operator, mask, n_vis, sum_weights, rel) assert terms.n_vis == n_vis -@pytest.mark.parametrize("use_jax", [False, True]) -def test__sparse_terms__sum_of_channels_equals_mfs(use_jax): +@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)) @@ -1513,6 +1536,7 @@ def test__sparse_terms__sum_of_channels_equals_mfs(use_jax): [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 @@ -1535,6 +1559,7 @@ def test__sparse_terms__sum_of_channels_equals_mfs(use_jax): 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, ) @@ -1545,7 +1570,12 @@ def test__sparse_terms__sum_of_channels_equals_mfs(use_jax): real_space_mask=mask, transformer_class=aa.TransformerNUFFT, ) - operator_mfs = dataset_mfs.apply_sparse_operator(use_jax=use_jax).sparse_operator + 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) @@ -1639,16 +1669,22 @@ def test__sparse_terms_from_chunks__phase_centre__matches_in_memory_on_shifted_d terms_shifted.dirty_image_native, terms_unshifted.dirty_image_native ) - # ... and leaves every other term bit-identical to the unshifted accumulation. + # ... 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", "data_term", "noise_normalization", "n_vis"): + 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). @@ -1713,3 +1749,80 @@ def test__sparse_terms__add__unrecorded_left_operand_keeps_right_phase_centre(): 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"