fix: finite, correct jax.grad at zero shear / multipole / ell_comps - #634
Merged
Merged
Conversation
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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Companion PR: PyAutoLabs/PyAutoLens#754
jax.gradreturned NaN at exactly zero forExternalShear(gamma_1,gamma_2),multipole_compsandell_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.convert.pytakessqrt(c0**2 + c1**2), whose derivative is 0/0 at the origin. A private_nudge_off_originhelper now moves the x-component by 1e-8 when both components are exactly zero, for traced JAX values only, viajax.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 PyAutoLenstest_static_lattice_jax.pyintact. Shear and multipole deflections are linear in their components, so the gradient at the origin is now exact and equals central finite differences.Isothermalclampedaxis_ratio <= 0.99999, so with the NaN gone itsell_compsgradient 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 ofarctan(s t)/sandarctanh(s t)/sins² = 1 - q², using a double-where so there are no NaN/inf gradients.Isothermal(ell_comps=(0, 0))now equalsIsothermalSphto fp64 round-off. Away from the limit (s² ≥ 1e-4) the deflections are bit-identical to main (NumPy, and jit with traced and constantell_comps).API Changes
None to signatures. Behaviour:
Isothermalno longer clamps its axis ratio, so a round or near-round (q > 0.99999)Isothermalgives exactly-circular deflections instead of those of q = 0.99999 at the clamp's orientation (a ~1e-5 relative change).jax.gradwith respect to shear / multipole / ellipticity components is finite and correct at exactly zero.See full details below.
Test Plan
test_convert.pytests (red on main: 8 failed): NumPy origin values pinned exactly;jax.gradfinite at (0,0) in fp64 + fp32; shear and multipole grad == central FDtest_isothermal.pytests (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_compsgrad == FD at (0,0) and (3e-6, -2e-6)test_autogalaxy: 1258 passedtest_autolens(with companion PyAutoLens PR): 762 passed, 1 xfailedFull API Changes (for automation & release notes)
Removed
Isothermal.axis_ratiooverride (themin(q, 0.99999)clamp); it now inheritsPowerLaw.axis_ratioChanged Behaviour
Isothermal.deflections_yx_2d_from— near-circular (1 - q² < 1e-4) uses a series form; exact at q = 1convert.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 gradientsAdded
convert._nudge_off_origin,convert._ORIGIN_NUDGE(private)Migration
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)Out of scope (follow-up)
MAX_ELLIP), chameleon and MGE/Gaussian; baresqrt(e0**2 + e1**2)innfw.py,dual_pseudo_isothermal_*,geometry_profiles.py:279Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code