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
14 changes: 14 additions & 0 deletions autoarray/dataset/imaging/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -602,6 +602,20 @@ def apply_sparse_operator(
Imaging
A new `Imaging` dataset with the precomputed `ImagingSparseOperator` attached, enabling
efficient pixelized source reconstruction via the sparse linear algebra formalism.

Notes
-----
`PYAUTO_DISABLE_JAX=1` is *not* honoured here, unlike
`Interferometer.apply_sparse_operator`, and the asymmetry is deliberate. There the
variable overrides a `use_jax` argument that already selects between two backends
computing the same operator. Here there is no such argument: this method is the JAX
implementation, and the NumPy/CPU alternative is the separately named
`apply_sparse_operator_cpu`, which returns a different operator class
(`SparseLinAlgImagingNumba`) and requires numba. Silently returning that under an
environment variable would change the type of the returned object based on the
environment, which is a larger change than honouring a switch -- and an unmeasured
one: every JIT cost the phase-8 workspace timings attribute to this variable
(2.3-3.2 s per script) was on the interferometer path.
"""

if self.psf is not None and self.psf.convolve_over_sample_size > 1:
Expand Down
29 changes: 29 additions & 0 deletions autoarray/dataset/interferometer/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,24 @@
from autonerves.fitsable import ndarray_via_fits_from
from autonerves import cached_property

try:
from autonerves.test_mode import disable_jax
except ImportError:
# `disable_jax()` arrives in the autonerves release that closes
# PyAutoNerves#159. Importing it unconditionally would make an autonerves
# older than that release an `ImportError` at module load -- and a
# `--no-deps` install, an editable checkout or a hand-built virtualenv can
# all put one on the path regardless of the floor in `pyproject.toml`,
# which constrains resolution only. This is the same trade `dataset_util`
# records against `SMALL_DATASETS_HEADER_KEY`: degrade to the predicate's
# own one-line body rather than fail hard to avoid restating it. Delete the
# fallback when the floor names a release carrying the predicate.
import os

def disable_jax():
return os.environ.get("PYAUTO_DISABLE_JAX") == "1"


from autoarray.dataset.abstract.dataset import AbstractDataset
from autoarray.dataset.grids import GridsDataset
from autoarray.inversion.inversion.interferometer.inversion_interferometer_util import (
Expand Down Expand Up @@ -266,6 +284,14 @@ def apply_sparse_operator(
use_jax
If `True`, JAX is used to accelerate the NUFFT precision matrix computation.

`PYAUTO_DISABLE_JAX=1` overrides this to `False`. That variable is a
harness-level switch, not a preference: it is the documented way to force the
NumPy path (the workspace `start_here` guides name it beside `use_jax=False`),
and the smoke profiles set it so a fast run does not pay a JIT compile. An
explicit `use_jax=True` in a script -- which is the right thing for a script
demonstrating the production path to say -- must therefore not defeat it, or
the harness pays 2.3-3.2 s of compile for a backend it asked to disable.

Precondition
------------
Every visibility must have equal real and imaginary noise sigma
Expand All @@ -289,6 +315,9 @@ def apply_sparse_operator(
If any visibility has unequal real and imaginary noise sigma.
"""

if disable_jax():
use_jax = False

noise_map_real = np.asarray(self.noise_map.real)
noise_map_imag = np.asarray(self.noise_map.imag)

Expand Down
11 changes: 10 additions & 1 deletion autoarray/inversion/mesh/image_mesh/overlay.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

from autoarray.geometry import geometry_util
from autoarray.structures.grids import grid_2d_util
from autoarray.util.dataset_util import cap_mesh_shape_for_small_datasets
from autoarray import numba_util


Expand Down Expand Up @@ -158,7 +159,7 @@ def overlay_via_unmasked_overlaid_from(


class Overlay(AbstractImageMesh):
def __init__(self, shape=(3, 3)):
def __init__(self, shape=(3, 3), respect_small_datasets: bool = True):
"""
Computes an image-mesh by overlaying a uniform grid of (y,x) coordinates over the masked image that the
pixelization is fitting.
Expand All @@ -176,10 +177,18 @@ def __init__(self, shape=(3, 3)):
----------
shape
The 2D shape of the grid which is overlaid over the grid to determine the image mesh.
respect_small_datasets
When `PYAUTO_SMALL_DATASETS=1` is set, `shape` is capped per axis to the small-datasets
cap, matching the cap `Grid2D.uniform` and `Mask2D.circular` apply to the data. Pass
`False` to opt out for an image mesh whose resolution is load-bearing for the script.
"""

super().__init__()

shape = cap_mesh_shape_for_small_datasets(
shape, respect_small_datasets=respect_small_datasets
)

self.shape = (int(shape[0]), int(shape[1]))

def image_plane_mesh_grid_from(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ class RectangularBilinearAdaptDensity(RectangularRTUAdaptDensity):
def __init__(
self,
shape: Tuple[int, int] = (3, 3),
respect_small_datasets: bool = True,
):
"""
A rectangular mesh of pixels used to reconstruct a source on a regular
Expand Down Expand Up @@ -71,7 +72,9 @@ def __init__(
If either dimension is less than 3, as a minimum of 3×3 pixels
is required to define interior and boundary structure.
"""
super().__init__(shape=shape)
super().__init__(
shape=shape, respect_small_datasets=respect_small_datasets
)

@property
def interpolator_kwargs(self) -> dict:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ def __init__(
shape: Tuple[int, int] = (3, 3),
weight_power: float = 1.0,
weight_floor: float = 0.0,
respect_small_datasets: bool = True,
):
"""
A rectangular mesh of pixels used to reconstruct a source on a regular
Expand Down Expand Up @@ -69,6 +70,7 @@ def __init__(
shape=shape,
weight_power=weight_power,
weight_floor=weight_floor,
respect_small_datasets=respect_small_datasets,
)

@property
Expand Down
11 changes: 11 additions & 0 deletions autoarray/inversion/mesh/mesh/rectangular_rtu_adapt_density.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from autoarray.inversion.mesh.border_relocator import BorderRelocator

from autoarray.structures.grids import grid_2d_util
from autoarray.util.dataset_util import cap_mesh_shape_for_small_datasets

from autoarray import exc

Expand Down Expand Up @@ -71,6 +72,7 @@ def __init__(
shape: Tuple[int, int] = (3, 3),
bandwidth: Optional[float] = None,
n_knots: Optional[int] = None,
respect_small_datasets: bool = True,
):
"""
A rectangular mesh of pixels used to reconstruct a source on a regular
Expand Down Expand Up @@ -132,6 +134,11 @@ def __init__(
n_knots
Size of the fixed knot table used to invert the CDF. Defaults to
the kernel default.
respect_small_datasets
When ``PYAUTO_SMALL_DATASETS=1`` is set, `shape` is capped per axis
to the small-datasets cap, matching the cap `Grid2D.uniform` and
`Mask2D.circular` apply to the data. Pass ``False`` to opt out for a
mesh whose resolution is load-bearing for the script.

Raises
------
Expand All @@ -144,6 +151,10 @@ def __init__(
KERNEL_CDF_DEFAULT_KNOTS,
)

shape = cap_mesh_shape_for_small_datasets(
shape, respect_small_datasets=respect_small_datasets
)

if shape[0] <= 2 or shape[1] <= 2:
raise exc.MeshException(
"The rectangular pixelization must be at least dimensions 3x3"
Expand Down
8 changes: 7 additions & 1 deletion autoarray/inversion/mesh/mesh/rectangular_rtu_adapt_image.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ def __init__(
weight_floor: float = 0.0,
bandwidth: Optional[float] = None,
n_knots: Optional[int] = None,
respect_small_datasets: bool = True,
):
"""
A uniform rectangular mesh of pixels used to reconstruct a source on a
Expand Down Expand Up @@ -82,7 +83,12 @@ def __init__(
the kernel default.
"""

super().__init__(shape=shape, bandwidth=bandwidth, n_knots=n_knots)
super().__init__(
shape=shape,
bandwidth=bandwidth,
n_knots=n_knots,
respect_small_datasets=respect_small_datasets,
)

self.weight_power = weight_power
self.weight_floor = weight_floor
Expand Down
Loading
Loading