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
28 changes: 25 additions & 3 deletions autoarray/dataset/interferometer/dataset.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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
------
Expand All @@ -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(
Expand Down Expand Up @@ -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()`.
Expand All @@ -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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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.
Expand All @@ -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)
Expand All @@ -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(
Expand All @@ -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):
"""
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
-------
Expand All @@ -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:
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading