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
97 changes: 81 additions & 16 deletions autoarray/structures/grids/uniform_2d.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,9 @@
from __future__ import annotations
import contextlib
import os
import numpy as np
from pathlib import Path
from typing import List, Optional, Tuple, Union
from typing import Any, List, Optional, Tuple, Union

from autonerves import conf
from autonerves.fitsable import ndarray_via_fits_from
Expand All @@ -18,6 +19,59 @@

from autoarray import exc
from autoarray import type as ty
from autoarray import validate


def _shift_and_rotate_is_constant(offset: Any, angle: Any) -> bool:
"""
Whether a shift-and-rotate of a grid by ``offset`` / ``angle`` is a constant --
that is, whether both are concrete numbers rather than traced model parameters.

``grid_offset`` and ``grid_rotation_angle`` come from a ``DatasetModel``, so they
are Python floats in the overwhelmingly common case (no dataset model, or a fixed
one) and JAX tracers only when a fit leaves them free. The positive test is
:func:`autoarray.validate.is_concrete_scalar`, the same gate the constructor
guards use; a tracer is not an ``int`` / ``float`` / ``np.number`` and so fails it.
"""
if not validate.is_concrete_scalar(angle):
return False

if not isinstance(offset, (tuple, list, np.ndarray)) or len(offset) != 2:
return False

return all(validate.is_concrete_scalar(value) for value in offset)


def _compile_time_eval_context(offset: Any, angle: Any, xp):
"""
The context a constant shift-and-rotate is evaluated in.

Under ``jax.jit`` **every** ``jax.numpy`` call is staged into the jaxpr, even one
whose operands are all concrete: ``jnp.asarray(numpy_grid) - jnp.array((0.0, 0.0))``
inside a trace returns a ``DynamicJaxprTracer``, not an array. The shifted and
rotated grid of a fit with no free ``grid_offset`` is therefore a compile-time
constant that JAX nevertheless recomputes on every likelihood evaluation -- and
this stack disables XLA's constant folding (``--xla_disable_hlo_passes=constant_folding``,
set by ``autonerves/jax_wrapper.py`` for compile-time reasons), so XLA does not
fold it away either.

``jax.ensure_compile_time_eval()`` restores eager evaluation for the operations
inside it, so the grid comes out as a concrete ``jax.Array``. The arithmetic is
identical -- it is the same operations on the same values, executed now rather
than staged -- and every downstream consumer that only reads the coordinates is
handed a constant it can act on, which is what lets
``autogalaxy.profiles.mass.abstract.deflections_memo`` fold a fixed-geometry
deflection field out of the trace entirely.

Numpy is unaffected (numpy has no trace to stage into), and a traced ``offset`` or
``angle`` falls back to the ordinary staged path, where it belongs.
"""
if xp is np or not _shift_and_rotate_is_constant(offset, angle):
return contextlib.nullcontext()

import jax

return jax.ensure_compile_time_eval()


class Grid2D(Structure):
Expand Down Expand Up @@ -772,28 +826,39 @@ def subtracted_and_rotated_from(
y'' = y' cos(theta) + x' sin(theta)
x'' = x' cos(theta) - y' sin(theta)

__Trace-time constant__

When ``offset`` and ``angle`` are concrete numbers -- which they are unless a
fit leaves ``grid_offset`` / ``grid_rotation_angle`` free -- the whole
calculation is a constant, and on the JAX backend it is evaluated eagerly
inside ``jax.ensure_compile_time_eval()`` rather than staged into the jaxpr
(see :func:`_compile_time_eval_context`). The returned grid is then a concrete
``jax.Array``, which is what allows the fixed-geometry deflection memo
downstream to fold its field out of the trace. The arithmetic is unchanged.

Parameters
----------
offset
The (y, x) offset subtracted from every grid coordinate before rotation.
angle
The rotation angle in degrees. Positive values rotate counter-clockwise.
"""
offset_array = xp.array(offset)
angle_rad = xp.deg2rad(angle)
cos_a = xp.cos(angle_rad)
sin_a = xp.sin(angle_rad)

def _shift_and_rotate(grid_array):
shifted = grid_array - offset_array
sy = shifted[:, 0]
sx = shifted[:, 1]
ry = sx * sin_a + sy * cos_a
rx = sx * cos_a - sy * sin_a
return xp.stack((ry, rx), axis=-1)

grid_rotated = _shift_and_rotate(self.array)
over_sampled_rotated = _shift_and_rotate(self.over_sampled.array)
with _compile_time_eval_context(offset=offset, angle=angle, xp=xp):
offset_array = xp.array(offset)
angle_rad = xp.deg2rad(angle)
cos_a = xp.cos(angle_rad)
sin_a = xp.sin(angle_rad)

def _shift_and_rotate(grid_array):
shifted = grid_array - offset_array
sy = shifted[:, 0]
sx = shifted[:, 1]
ry = sx * sin_a + sy * cos_a
rx = sx * cos_a - sy * sin_a
return xp.stack((ry, rx), axis=-1)

grid_rotated = _shift_and_rotate(self.array)
over_sampled_rotated = _shift_and_rotate(self.over_sampled.array)

mask = Mask2D(
mask=self.mask,
Expand Down
52 changes: 52 additions & 0 deletions test_autoarray/structures/grids/test_uniform_2d.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import contextlib
from pathlib import Path
import numpy as np
import pytest
Expand Down Expand Up @@ -882,6 +883,57 @@ def test__subtracted_and_rotated_from__shift_first_then_rotate():
assert rotated.array == pytest.approx(expected, 1.0e-4)


def test__shift_and_rotate_is_constant__concrete_offset_and_angle_only():
from autoarray.structures.grids import uniform_2d

assert uniform_2d._shift_and_rotate_is_constant(offset=(0.5, -0.5), angle=0.0)
assert uniform_2d._shift_and_rotate_is_constant(
offset=np.array([0.5, -0.5]), angle=np.float64(30.0)
)

# A value that is not a concrete number -- a JAX tracer reaches this the same way
# a string does, by failing `validate.is_concrete_scalar` -- is not a constant.
assert not uniform_2d._shift_and_rotate_is_constant(offset=(0.5, "x"), angle=0.0)
assert not uniform_2d._shift_and_rotate_is_constant(offset=(0.5, -0.5), angle=None)
assert not uniform_2d._shift_and_rotate_is_constant(offset=(0.5,), angle=0.0)
assert not uniform_2d._shift_and_rotate_is_constant(offset=0.5, angle=0.0)


def test__subtracted_and_rotated_from__numpy_is_unchanged_by_the_constant_gate():
"""
The compile-time-eval context is a JAX-only concern: on numpy the method is the
plain shift-and-rotate it has always been, for a concrete *and* a non-concrete
offset, and it never enters a context that would import jax.
"""
from autoarray.structures.grids import uniform_2d

grid = aa.Grid2D.uniform(shape_native=(3, 3), pixel_scales=1.0, over_sample_size=2)

rotated = grid.subtracted_and_rotated_from(offset=(1.0, 2.0), angle=90.0)

shifted = grid.array - np.array([1.0, 2.0])
expected = np.stack((shifted[:, 1], -shifted[:, 0]), axis=-1)

assert rotated.array == pytest.approx(expected, 1.0e-8)
assert isinstance(rotated.array, np.ndarray)

over_shifted = grid.over_sampled.array - np.array([1.0, 2.0])
over_expected = np.stack((over_shifted[:, 1], -over_shifted[:, 0]), axis=-1)

assert rotated.over_sampled.array == pytest.approx(over_expected, 1.0e-8)

# The numpy backend takes the null context whatever the offset is.
context = uniform_2d._compile_time_eval_context(
offset=(1.0, 2.0), angle=90.0, xp=np
)
assert isinstance(context, contextlib.nullcontext)

context = uniform_2d._compile_time_eval_context(
offset=(1.0, None), angle=90.0, xp=np
)
assert isinstance(context, contextlib.nullcontext)


def test__over_sampled__sub_size_1_is_the_slim_grid():
mask = aa.Mask2D(
mask=[
Expand Down
Loading