fix: NaN JAX gradient on MGE positive-only solves after #572 - #574
Merged
Merged
Conversation
Four 20x20 systems captured from autolens_workspace_test jax_grad/mge.py at PRNGKey perturbations 2, 10, 12, 14: the raw-mode gradient is NaN on each on 3de624b (eager 4/4, jit 2/4). Adds backward-pass convergence tests over these and the 8 SLaM #571 systems via raw_forward_backward_status. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…olve (#573) The raw forward solve (#572) stops at the data-scaled tolerance, leaving s*z ~1e-10..2.5e-9 >> nnls_target_kappa=1e-11; the relaxed-KKT solve on the Jacobi system then diverged to NaN at its 50-iteration cap on 4/16 jax_grad/mge.py points. The backward pass now runs <= 10 tight PDIP iterations on (Q_pc, q_pc) warm-started from the mapped iterate (kept only if converged), after which the relaxed solve converges in ~1 iteration. The primal / forward value is unchanged. solve_nnls gains an init=(x, s, z) warm start; the relaxed solve's converged flag is kept and exposed with the polish status through the new raw_forward_backward_status diagnostic. general.yaml comments updated (values unchanged). Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Review finding: on 3de624b the jitted gradient is NaN only on prng10/prng14 (prng2/prng12 pass jitted), and the backward-status test fails on import on main rather than on convergence, so it is not the red witness. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_019jDFQSNoi3ihaeM7ZJhfYL
Collaborator
Author
|
Merged by human |
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
Fixes NaN JAX gradients on mapper-less (MGE / linear light profile) positive-only inversions, introduced by #572 (the #571 raw-forward PDIP fix). Closes #573.
#572's
"raw"mode runs the forward PDIP on the raw system with a loose data-scaled tolerance, which leaves the complementaritys·zabout 1e-10 to 2.5e-9. The custom-vjp backward pass then calledsolve_relaxed_nnlson the Jacobi system atnnls_target_kappa=1e-11from that iterate. That is a push toward the boundary withz/saround 1e13-1e14; under jit it overshoots (s, z < 0), hits the 50-iteration cap and returns NaN, anddiff_nnlspasses the NaN to every gradient entry. On main, 4 of 16 perturbation points ofautolens_workspace_test/scripts/imaging/jax_grad/mge.pyare NaN eagerly, and the jitted value+grad is NaN at the script's own point, which is how it failed Heart Release Integrate (run 36108062907).Fix: in the backward pass only, polish the mapped iterate with at most
RAW_BACKWARD_POLISH_MAX_ITER=10warm-started PDIP iterations on(Q_pc, q_pc)at jaxnnls's tight tolerance before the relaxed solve. The polished point is kept only if it converged with finiteyands, z > 0; otherwise the unpolished iterate is used, as before. The primal is unchanged, so the forward value and the #571 forward-convergence fix are untouched: the jitted forward HLO is byte-identical to main.Why not raise the kappa instead: the alternative
kappa_eff = max(target_kappa, max(s·z))was measured and rejected. It left 6/48 SLaM systems NaN and was worse under jit (11/16), and on one system the relaxed solve reported "converged" withs < 0.jax_grad/mge.py, 16 keys finite (eager / jit)Corrective PR under Heart RED (human-authorized)
release validation FAILED (stage integrate). Re-read unchanged at ship time./prm; no release or rehearsal.autolens_workspace_testimaging/jax_grad/mge.pyfailure in that integrate run. The run's other two failures (autolens_workspaceweak/real_data/a2744.py,cluster/lenstool/modeling.py) were fixed by autolens_workspace#577 (merged).API Changes
Additive only; no signature break or default change.
solve_nnlsgains an optionalinit=(x, s, z)warm start (defaultNone= unchanged).raw_forward_backward_status(...)reports the backward pass's convergence (relaxed + polish flags and iterations), and a newRAW_BACKWARD_POLISH_MAX_ITERconstant.nnls_preconditioning_no_mapper: raw) gradients are now finite on the previously NaN points. Forward values are unchanged.See full details below.
Test Plan
files/mge_grad_nan_systems.npz(6.8 KB, 4 systems captured at the NaN points). The gradient test fails 6/8 on3de624b5and passes on the branch. Also added: backward-status tests over 8 SLaM + 4 captured systems.autolens_workspace_test/scripts/imaging/jax_grad/mge.pyunder the release profile (py3.12, jax 0.10.2, numpy 2.5.3): own point PASS; 16-key sweep 16/16 finite (eager and jit), FD 16/16; logL bit-identical to main.runtime/mge.pysingle-JIT forward: 31.22 vs 32.12 ms (noise; HLO identical)hazards/mge_nnls_capture.py: 48/48 converged on bothpyauto-heart smoke --root <task worktree>, branch PyAutoArray on PYTHONPATH): 159 passed across autofit / autogalaxy / autolens / autolens_workspace_test / euclid / howtolens, scripts and notebooks. 1 failure, pre-existing on main:autolens_workspace_test/scripts/interferometer/jax_likelihood/mge.pyvmap likelihood -45560751.36 vs pin -3152.65, identical on main3de624b5in the same env (not caused by this PR; flagged separately). autocti / autocti_test did not run (local GSL headers missing, environment only).Full API Changes (for automation & release notes)
Added
autoarray.util.jax_nnls.raw_forward_backward_status(Q_pc, q_pc, Q, q, D, target_kappa=1e-3, solver_tol=None, max_iter=50): returns(relaxed_converged, relaxed_iter, polish_converged, polish_iter)of the raw-mode backward passautoarray.util.jax_nnls.RAW_BACKWARD_POLISH_MAX_ITER = 10Changed Signature
autoarray.util.jax_nnls.solve_nnls(Q, q, solver_tol=None, max_iter=50, init=None): optional(x, s, z)warm startChanged Behaviour
solve_nnls_primal_raw_forward/ raw-mode positive-only JAX solves: the backward pass polishes the mapped forward iterate before the relaxed-KKT solve; gradients no longer NaN on the fix: NaN JAX gradient on MGE positive-only solves after #572 #573 points. Forward values unchanged.Migration
Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code