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
40 changes: 40 additions & 0 deletions autogalaxy/convert.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,37 @@
# Stated together here rather than left as unrelated literals in separate files.
ELL_COMPS_MAGNITUDE_CLAMP = 0.999

# The amount the x-component of a pair of components is nudged by when both are
# exactly zero, for traced JAX values only (see `_nudge_off_origin`).
_ORIGIN_NUDGE = 1e-8


def _nudge_off_origin(c0, c1, xp=np):
"""
Returns the x-component `c1` of a pair of components (`ell_comps[1]`, `gamma_1`, `multipole_comps[1]`), nudged
by `_ORIGIN_NUDGE` when both components are exactly zero, for traced JAX values only.

Every polar conversion below computes `sqrt(c0**2 + c1**2)`, whose derivative at the origin is 0/0, so
`jax.grad` returns NaN there even though the profiles themselves are smooth through it (shear and multipole
deflections are linear in their components, so the gradient at the nudged point is exact). The nudged value is
`c1 + nudge`, so the derivative through `c1` is unchanged, and positive, so the angle `arctan2(0, +nudge) = 0`
matches the existing guarded value at the origin.

Only traced values carry a gradient, so NumPy and concrete JAX values are returned unchanged: staging the nudge
onto a constant component alters how XLA compiles the downstream arithmetic and moves jitted deflections by an
ulp, which flips bit-exact ties in the point solver (`test_static_lattice_jax.py` in PyAutoLens pins one). For
the same reason traced values go through `jax.lax.select` rather than `jnp.where` or an additive offset, both of
which were measured to shift that pinned tie; away from the origin the traced value is selected unchanged.
"""
if xp.__name__.startswith("jax"):
import jax

if isinstance(c0, jax.core.Tracer) or isinstance(c1, jax.core.Tracer):
return jax.lax.select(
xp.logical_and(c0 == 0, c1 == 0), c1 + _ORIGIN_NUDGE, c1
)
return c1


def ell_comps_from(axis_ratio: float, angle: float, xp=np) -> Tuple[float, float]:
"""
Expand Down Expand Up @@ -75,6 +106,8 @@ def axis_ratio_and_angle_from(
ell_comps
The elliptical components of the light or mass profile which are converted to an angle.
"""
ell_comps = (ell_comps[0], _nudge_off_origin(ell_comps[0], ell_comps[1], xp=xp))

angle = 0.5 * xp.arctan2(
ell_comps[0],
xp.where(xp.logical_and(ell_comps[0] == 0, ell_comps[1] == 0), 1.0, ell_comps[1]),
Expand Down Expand Up @@ -215,6 +248,8 @@ def shear_magnitude_and_angle_from(
gamma_2
The gamma 2 component of the shear.
"""
gamma_1 = _nudge_off_origin(gamma_2, gamma_1, xp=xp)

angle = (
0.5
* xp.arctan2(gamma_2, xp.where(xp.logical_and(gamma_1 == 0, gamma_2 == 0), 1.0, gamma_1))
Expand Down Expand Up @@ -320,6 +355,11 @@ def multipole_k_m_and_phi_m_from(
-------
The normalization and angle parameters of the multipole.
"""
multipole_comps = (
multipole_comps[0],
_nudge_off_origin(multipole_comps[0], multipole_comps[1], xp=xp),
)

phi_m = (
xp.arctan2(
multipole_comps[0],
Expand Down
6 changes: 3 additions & 3 deletions autogalaxy/interop/coolest/mass.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,9 +118,9 @@ def _isothermal_from(parameters: Dict) -> Isothermal:
einstein_radius = float(
einstein_radius_ag_from(theta_E=parameters["theta_E"], axis_ratio=q, slope=2.0)
)
# An exactly-round COOLEST profile maps to the spherical class — the
# elliptical Isothermal clips its axis ratio to 0.99999 for the stability
# of its analytic deflections, so it is not numerically exact at q = 1.
# An exactly-round COOLEST profile maps to the spherical class, the
# dedicated q = 1 profile (the elliptical Isothermal agrees with it there to
# fp64 round-off, but the spherical class states the intent).
if q == 1.0:
return IsothermalSph(centre=centre, einstein_radius=einstein_radius)
return Isothermal(
Expand Down
42 changes: 35 additions & 7 deletions autogalaxy/profiles/mass/total/isothermal.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,24 +94,35 @@ def __init__(
slope=2.0,
)

def axis_ratio(self, xp=np):
axis_ratio = super().axis_ratio(xp=xp)
return xp.minimum(axis_ratio, 0.99999)

@aa.decorators.to_vector_yx
@aa.decorators.transform(rotate_back=True)
def deflections_yx_2d_from(self, grid: aa.type.Grid2DLike, xp=np, **kwargs):
"""
r"""
Calculate the deflection angles on a grid of (y,x) arc-second coordinates.

With :math:`s^2 = 1 - q^2`, :math:`t_x = x / \Psi` and :math:`t_y = y / \Psi` the deflections are
:math:`2 b q \arctan(s t_x) / s` and :math:`2 b q \,{\rm arctanh}(s t_y) / s`. Both ratios tend to
:math:`t` as :math:`q \to 1`, but the closed form is 0/0 there, so near the circular limit they are
evaluated as their Taylor series in :math:`s^2`. This makes the deflections, and their gradient with respect
to the ellipticity components, correct all the way to :math:`q = 1`.

Parameters
----------
grid
The grid of (y,x) arc-second coordinates the deflection angles are computed on.
"""

axis_ratio = self.axis_ratio(xp)
sqrt_one_minus_q2 = xp.sqrt(1 - axis_ratio**2)
one_minus_q2 = 1 - axis_ratio**2

# Below this s^2 the series is used. It is truncated after the s^6 t^7 term, so for |t_y| <= 1 and
# |t_x| <= 1 / q ~ 1 (as Psi >= |y| and Psi >= q |x|) its relative error is below s^8 / 9 ~ 1e-17 at the
# threshold, under fp64 round-off. Above it the closed form is well conditioned: s >= 1e-2.
small = one_minus_q2 < 1.0e-4

# Double-where: the closed form is evaluated at a safe s^2 wherever the series is used, so its unused branch
# never produces an inf / NaN value or gradient at q = 1.
sqrt_one_minus_q2 = xp.sqrt(xp.where(small, 1.0, one_minus_q2))

factor = (
2.0 * self.einstein_radius_rescaled(xp) * axis_ratio / sqrt_one_minus_q2
Expand All @@ -125,7 +136,24 @@ def deflections_yx_2d_from(self, grid: aa.type.Grid2DLike, xp=np, **kwargs):
deflection_x = xp.arctan(
xp.divide(xp.multiply(sqrt_one_minus_q2, grid.array[:, 1]), psi)
)
return xp.multiply(factor, xp.vstack((deflection_y, deflection_x)).T)
deflections = xp.multiply(factor, xp.vstack((deflection_y, deflection_x)).T)

t_y = grid.array[:, 0] / psi
t_x = grid.array[:, 1] / psi

# arctanh(s t) / s = t + s^2 t^3 / 3 + s^4 t^5 / 5 + s^6 t^7 / 7 + ...
# arctan(s t) / s = t - s^2 t^3 / 3 + s^4 t^5 / 5 - s^6 t^7 / 7 + ...
u_y = one_minus_q2 * t_y**2
u_x = -one_minus_q2 * t_x**2
series_y = t_y * (1.0 + u_y * (1.0 / 3.0 + u_y * (1.0 / 5.0 + u_y / 7.0)))
series_x = t_x * (1.0 + u_x * (1.0 / 3.0 + u_x * (1.0 / 5.0 + u_x / 7.0)))

deflections_series = xp.multiply(
2.0 * self.einstein_radius_rescaled(xp) * axis_ratio,
xp.vstack((series_y, series_x)).T,
)

return xp.where(small, deflections_series, deflections)

@aa.decorators.to_vector_yx
@aa.decorators.transform
Expand Down
77 changes: 77 additions & 0 deletions test_autogalaxy/profiles/mass/total/test_isothermal.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,3 +203,80 @@ def test__shear_yx_2d_from__matches_via_hessian():
np.testing.assert_allclose(
np.asarray(shear_analytic), np.asarray(shear_via_hessian), rtol=1e-3, atol=1e-6
)


def _deflections_old_closed_form(mass, grid):
"""
The pre-series closed-form SIE deflections (Kormann et al. 1994), evaluated with the profile's own axis ratio
for an unrotated profile centred on the origin. Well conditioned for every q < 1 used below.
"""
q = mass.axis_ratio()
s = np.sqrt(1.0 - q**2)
factor = 2.0 * mass.einstein_radius_rescaled() * q / s
y, x = grid.array[:, 0], grid.array[:, 1]
psi = np.sqrt(q**2 * x**2 + y**2 + 1e-16)
return np.vstack((factor * np.arctanh(s * y / psi), factor * np.arctan(s * x / psi))).T


@pytest.mark.parametrize("axis_ratio", [0.5, 0.9, 0.9999, 1.0 - 1.0e-7])
def test__deflections_yx_2d_from__matches_closed_form_up_to_circular_limit(axis_ratio):
mass = ag.mp.Isothermal(
centre=(0.0, 0.0),
ell_comps=ag.convert.ell_comps_from(axis_ratio=axis_ratio, angle=0.0),
einstein_radius=1.3,
)

assert mass.axis_ratio() == pytest.approx(axis_ratio, rel=1.0e-12)

deflections = mass.deflections_yx_2d_from(grid=grid)

assert deflections.array == pytest.approx(
_deflections_old_closed_form(mass, grid), rel=1.0e-12
)


def test__deflections_yx_2d_from__circular_limit_equals_isothermal_sph():
ell = ag.mp.Isothermal(centre=(0.1, -0.2), ell_comps=(0.0, 0.0), einstein_radius=1.3)
sph = ag.mp.IsothermalSph(centre=(0.1, -0.2), einstein_radius=1.3)

assert ell.deflections_yx_2d_from(grid=grid).array == pytest.approx(
sph.deflections_yx_2d_from(grid=grid).array, rel=1.0e-12
)


@pytest.mark.parametrize("ell_comps", [(0.0, 0.0), (3.0e-6, -2.0e-6)])
def test__deflections_yx_2d_from__jax_grad_ell_comps_near_circular_matches_finite_difference(
ell_comps,
):
"""
At and near ell_comps = (0, 0) the gradient of the deflections with respect to the ellipticity components must
be the true one, not zero (an axis-ratio clamp) or NaN (the polar conversion at the origin).
"""
jax = pytest.importorskip("jax")
import jax.numpy as jnp

with jax.enable_x64(True):
small_grid = ag.Grid2D.uniform(shape_native=(3, 3), pixel_scales=0.3)
weights = jnp.asarray(np.random.default_rng(1).normal(size=(9, 2)))

def f(e):
mass = ag.mp.Isothermal(
centre=(0.01, 0.02), ell_comps=(e[0], e[1]), einstein_radius=1.0
)
return jnp.sum(
mass.deflections_yx_2d_from(grid=small_grid, xp=jnp).array * weights
)

e = jnp.array(ell_comps)
grad = np.asarray(jax.grad(f)(e))

h = 1.0e-4
fd = np.array(
[
(f(e + jnp.array([h, 0.0])) - f(e - jnp.array([h, 0.0]))) / (2 * h),
(f(e + jnp.array([0.0, h])) - f(e - jnp.array([0.0, h]))) / (2 * h),
]
)

assert np.all(np.isfinite(grad))
assert grad == pytest.approx(fd, abs=1.0e-6)
106 changes: 106 additions & 0 deletions test_autogalaxy/test_convert.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,3 +170,109 @@ def test__multipole_comps_from():
multipole_comps = ag.convert.multipole_comps_from(k_m=0.14142135, phi_m=112.5, m=2)

assert multipole_comps == pytest.approx((-0.1, -0.1), abs=1e-3)


def test__polar_conversions__numpy_origin_values_unchanged():
axis_ratio, angle = ag.convert.axis_ratio_and_angle_from(ell_comps=(0.0, 0.0))

assert axis_ratio == 1.0
assert angle == 0.0

magnitude, angle = ag.convert.shear_magnitude_and_angle_from(
gamma_1=0.0, gamma_2=0.0
)

assert magnitude == 0.0
assert angle == 0.0

k_m, phi_m = ag.convert.multipole_k_m_and_phi_m_from(
multipole_comps=(0.0, 0.0), m=4
)

assert k_m == 0.0
assert phi_m == 0.0


def _deflection_objectives():
"""
Scalar objectives (weighted sums of deflections on a tiny grid) of the profiles whose polar conversions take a
square root of their components, as functions of those two components.
"""
import jax.numpy as jnp
import numpy as np

grid = ag.Grid2D.uniform(shape_native=(3, 3), pixel_scales=0.3)
weights = jnp.asarray(np.random.default_rng(1).normal(size=(9, 2)))

def shear(c):
mass = ag.mp.ExternalShear(gamma_1=c[0], gamma_2=c[1])
return jnp.sum(mass.deflections_yx_2d_from(grid=grid, xp=jnp).array * weights)

def multipole(c):
mass = ag.mp.PowerLawMultipole(
centre=(0.01, 0.02),
einstein_radius=1.0,
slope=2.0,
m=4,
multipole_comps=(c[0], c[1]),
)
return jnp.sum(mass.deflections_yx_2d_from(grid=grid, xp=jnp).array * weights)

def isothermal(c):
mass = ag.mp.Isothermal(
centre=(0.01, 0.02), ell_comps=(c[0], c[1]), einstein_radius=1.0
)
return jnp.sum(mass.deflections_yx_2d_from(grid=grid, xp=jnp).array * weights)

return {"shear": shear, "multipole": multipole, "isothermal": isothermal}


@pytest.mark.parametrize("name", ["shear", "multipole"])
def test__polar_conversions__jax_grad_finite_at_origin_fp64(name):
jax = pytest.importorskip("jax")
import jax.numpy as jnp
import numpy as np

with jax.enable_x64(True):
f = _deflection_objectives()[name]
grad = jax.grad(f)(jnp.zeros(2))

assert np.all(np.isfinite(np.asarray(grad)))


@pytest.mark.parametrize("name", ["shear", "multipole", "isothermal"])
def test__polar_conversions__jax_grad_finite_at_origin_fp32(name):
jax = pytest.importorskip("jax")
import jax.numpy as jnp
import numpy as np

with jax.enable_x64(False):
f = _deflection_objectives()[name]
grad = jax.grad(f)(jnp.zeros(2, dtype=jnp.float32))

assert np.all(np.isfinite(np.asarray(grad)))


@pytest.mark.parametrize("name", ["shear", "multipole"])
def test__polar_conversions__jax_grad_at_origin_matches_finite_difference(name):
"""
Shear and multipole deflections are linear in their components, so the gradient at the origin is well defined
and equal to the central finite difference there.
"""
jax = pytest.importorskip("jax")
import jax.numpy as jnp
import numpy as np

with jax.enable_x64(True):
f = _deflection_objectives()[name]
grad = np.asarray(jax.grad(f)(jnp.zeros(2)))

h = 1.0e-5
fd = np.array(
[
(f(jnp.array([h, 0.0])) - f(jnp.array([-h, 0.0]))) / (2 * h),
(f(jnp.array([0.0, h])) - f(jnp.array([0.0, -h]))) / (2 * h),
]
)

assert grad == pytest.approx(fd, abs=1.0e-6)
Loading