From a2175896cbd4ee91a734c6f9b3f76f3cd21ff464 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Sun, 27 Sep 2026 20:06:28 +0100 Subject: [PATCH] fix: finite, correct jax.grad at zero shear / multipole / ell_comps Nudge the x-component off the exact origin (traced JAX only, lax.select) in the convert.py polar conversions, whose sqrt has a 0/0 derivative there, and remove the Isothermal q<=0.99999 clamp in favour of a series form near q -> 1 so the ell_comps gradient at 0 matches finite differences. Closes #631 Co-Authored-By: Claude Opus 5.5 --- autogalaxy/convert.py | 40 +++++++ autogalaxy/interop/coolest/mass.py | 6 +- autogalaxy/profiles/mass/total/isothermal.py | 42 +++++-- .../profiles/mass/total/test_isothermal.py | 77 +++++++++++++ test_autogalaxy/test_convert.py | 106 ++++++++++++++++++ 5 files changed, 261 insertions(+), 10 deletions(-) diff --git a/autogalaxy/convert.py b/autogalaxy/convert.py index e53d514d6..36a4905e4 100644 --- a/autogalaxy/convert.py +++ b/autogalaxy/convert.py @@ -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]: """ @@ -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]), @@ -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)) @@ -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], diff --git a/autogalaxy/interop/coolest/mass.py b/autogalaxy/interop/coolest/mass.py index e08cf6a46..2e6786f09 100644 --- a/autogalaxy/interop/coolest/mass.py +++ b/autogalaxy/interop/coolest/mass.py @@ -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( diff --git a/autogalaxy/profiles/mass/total/isothermal.py b/autogalaxy/profiles/mass/total/isothermal.py index 3eb5fad9b..57d9f87e3 100644 --- a/autogalaxy/profiles/mass/total/isothermal.py +++ b/autogalaxy/profiles/mass/total/isothermal.py @@ -94,16 +94,18 @@ 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 @@ -111,7 +113,16 @@ def deflections_yx_2d_from(self, grid: aa.type.Grid2DLike, xp=np, **kwargs): """ 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 @@ -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 diff --git a/test_autogalaxy/profiles/mass/total/test_isothermal.py b/test_autogalaxy/profiles/mass/total/test_isothermal.py index f9ba63abd..cd2cb6a45 100644 --- a/test_autogalaxy/profiles/mass/total/test_isothermal.py +++ b/test_autogalaxy/profiles/mass/total/test_isothermal.py @@ -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) diff --git a/test_autogalaxy/test_convert.py b/test_autogalaxy/test_convert.py index f40c1017e..1d1d22ca4 100644 --- a/test_autogalaxy/test_convert.py +++ b/test_autogalaxy/test_convert.py @@ -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)