Skip to content

fix: finite, correct jax.grad at zero shear / multipole / ell_comps - #634

Merged
Jammy2211 merged 1 commit into
mainfrom
feature/jax-grad-nan-zero-components
Sep 27, 2026
Merged

Jammy2211 merged 1 commit into
mainfrom
feature/jax-grad-nan-zero-components

Conversation

@Jammy2211

@Jammy2211 Jammy2211 commented Sep 27, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Companion PR: PyAutoLabs/PyAutoLens#754

jax.grad returned NaN at exactly zero for ExternalShear (gamma_1, gamma_2), multipole_comps and ell_comps, so a gradient search starting at prior medians of 0 got NaN gradients (the inference benchmark centred those priors on 1e-3 to dodge it). Closes #631.

  • NaN cause: every polar conversion in convert.py takes sqrt(c0**2 + c1**2), whose derivative is 0/0 at the origin. A private _nudge_off_origin helper now moves the x-component by 1e-8 when both components are exactly zero, for traced JAX values only, via jax.lax.select. NumPy and concrete JAX values are untouched (bit-identical), and so are jitted graphs with constant components, which keeps the point-solver tie pinned in PyAutoLens test_static_lattice_jax.py intact. Shear and multipole deflections are linear in their components, so the gradient at the origin is now exact and equals central finite differences.
  • Isothermal: the elliptical Isothermal clamped axis_ratio <= 0.99999, so with the NaN gone its ell_comps gradient at 0 was finite but wrong (≈ -267 vs finite-difference -0.53). The clamp is removed. Near q → 1 the deflections are evaluated as the Taylor series of arctan(s t)/s and arctanh(s t)/s in s² = 1 - q², using a double-where so there are no NaN/inf gradients. Isothermal(ell_comps=(0, 0)) now equals IsothermalSph to fp64 round-off. Away from the limit (s² ≥ 1e-4) the deflections are bit-identical to main (NumPy, and jit with traced and constant ell_comps).

API Changes

None to signatures. Behaviour: Isothermal no longer clamps its axis ratio, so a round or near-round (q > 0.99999) Isothermal gives exactly-circular deflections instead of those of q = 0.99999 at the clamp's orientation (a ~1e-5 relative change). jax.grad with respect to shear / multipole / ellipticity components is finite and correct at exactly zero.
See full details below.

Test Plan

  • New test_convert.py tests (red on main: 8 failed): NumPy origin values pinned exactly; jax.grad finite at (0,0) in fp64 + fp32; shear and multipole grad == central FD
  • New test_isothermal.py tests (red before the Isothermal fix: 4 failed): closed-form match at q ∈ {0.5, 0.9, 0.9999, 1-1e-7}; circular limit == IsothermalSph; ell_comps grad == FD at (0,0) and (3e-6, -2e-6)
  • test_autogalaxy: 1258 passed
  • test_autolens (with companion PyAutoLens PR): 762 passed, 1 xfailed
Full API Changes (for automation & release notes)

Removed

  • Isothermal.axis_ratio override (the min(q, 0.99999) clamp); it now inherits PowerLaw.axis_ratio

Changed Behaviour

  • Isothermal.deflections_yx_2d_from — near-circular (1 - q² < 1e-4) uses a series form; exact at q = 1
  • convert.axis_ratio_and_angle_from, convert.shear_magnitude_and_angle_from, convert.multipole_k_m_and_phi_m_from (and everything routed through them) — traced JAX inputs at exactly (0, 0) are nudged by 1e-8 on the x-component, giving finite and correct gradients

Added

  • convert._nudge_off_origin, convert._ORIGIN_NUDGE (private)

Migration

  • Tests pinning numbers from a round Isothermal (ell_comps = (0, 0)) will move ~1e-5 relative, or more in hypersensitive pixelization and point-solver fixtures; repin to the circular values (companion PyAutoLens PR does this for 6 tests)
  • The benchmark's 1e-3 prior centring for shear / multipole / ell_comps can be dropped after release

Out of scope (follow-up)

  • The same clamp class in dPIE (MAX_ELLIP), chameleon and MGE/Gaussian; bare sqrt(e0**2 + e1**2) in nfw.py, dual_pseudo_isothermal_*, geometry_profiles.py:279

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

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 <noreply@anthropic.com>
@Jammy2211 Jammy2211 added the pending-release PR queued for the next release build label Sep 27, 2026
@Jammy2211
Jammy2211 merged commit c7fc595 into main Sep 27, 2026
4 checks passed
@Jammy2211
Jammy2211 deleted the feature/jax-grad-nan-zero-components branch September 27, 2026 19:26
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

pending-release PR queued for the next release build

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix: jax.grad NaN at exactly-zero shear / multipole / ell_comps

1 participant