You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
PyAutoArray #572 (the #571 raw-forward PDIP fix) made the JAX gradient of mapper-less (MGE) positive-only inversions NaN on a fraction of parameter points: the raw forward solve stops at a loose data-scaled tolerance, and the backward relaxed-KKT solve (nnls_target_kappa=1e-11) then has to push toward the boundary from a z/s ~1e13-14 iterate and overshoots under jit. This fails Heart Release Integrate (2026-09-25, autolens_workspace_test/scripts/imaging/jax_grad/mge.py) and is the last blocker for the release. The forward likelihood is unaffected.
Plan
Build a scratch decision harness: current main vs (a) polish with a few tight PDIP iterations on the Jacobi system before the relaxed solve, vs (b) effective kappa = max(target_kappa, max(s·z)).
untracked dataset/abell_1201 only; claimed by certified-solver-phase-c1-lane-rate, interferometer-mge-breakdown
./autolens_workspace_test
main
clean
Suggested branch:feature/mge-nnls-grad-nan
Implementation Steps
Approach: pick (a) or (b) from evidence, in one bounded experiment
Both are candidates for the step inside forward() between raw_solve and solve_relaxed_nnls:
(a) Polish: run a few solve_nnls iterations on (Q_pc, q_pc) at the tight default tolerance (n·eps), warm-started from the mapped iterate. Then call the relaxed solve. solve_nnls has no warm-start today, so this needs a small init argument threaded into its while_loop state.
(b) Effective kappa:kappa_eff = max(target_kappa, max(sy*zy)), so the relaxed solve only moves away from the boundary. No extra solve. The cost is slightly more smoothing in the gradient on affected points.
Decision harness (scratch script, not committed). Run both candidates and current main on:
the 16 PRNGKey perturbation points of the jax_grad/mge.py model (reuse the diagnosis mge_diag.py);
the 48 SLaM systems in test_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npz;
the captured failing relaxed input relax_in_first.npz.
For each candidate, record:
gradients finite (all points);
agreement with finite differences (the check mge.py already does, as rel. error);
forward y / log-likelihood unchanged vs main (must be bit-identical or within 1e-12: the forward path is untouched);
Take (a) only if (b) degrades FD agreement and (a) converges on all 48 systems.
If neither passes, stop and report to you. Don't fall back to the revert on my own call.
Implementation (chosen candidate)
autoarray/util/jax_nnls.py:
Implement the chosen step in _solve_nnls_raw_forward_with.forward.
Surface the relaxed solve's converged flag: stop discarding it (yr, sr, zr, _, _). Keep it in the residuals, and expose a small non-custom-vjp diagnostic helper, raw_forward_backward_status(Q_pc, q_pc, Q, q, D, …), that returns the relaxed-solve (converged, iters) for tests and profiling.
Update the docstrings for the new step.
The custom_vjp primal/forward output signature is unchanged, so no caller changes are needed.
autoarray/config/general.yaml: update the nnls_preconditioning_no_mapper / nnls_target_kappa comment to describe the backward-pass behaviour. Values stay the same.
Tests in test_autoarray/inversion/inversion/test_nnls_mge_convergence.py, reusing the existing npz loader:
The gradient of sum(y) (or a fixed random cotangent) through solve_nnls_primal_raw_forward is finite on all 48 SLaM systems, and relaxed-status converged.
A small captured-failure fixture (the k=2 system, saved into files/ next to the existing npz and noted in its README) that is NaN on main and finite after the fix. Run it red on unfixed source first so it's a real witness.
A 16-key perturbation sweep on a small synthetic MGE-like system, if one can be built cheaply and fails on main. If not, the 48 SLaM systems plus the captured case cover the sweep requirement.
Runtime check against autolens_profiling (required before shipping)
Use a scratch detached worktree of lens/autolens_profiling at origin/main. Run the same local CPU fp64 HST settings as the pinned results, main vs branch, alternating (A-B-A-B) to control for drift:
scripts/imaging/likelihood_runtime/mge.py: forward likelihood time. Expected unchanged, since the primal is untouched.
Value+gradient time for the MGE likelihood (jax.value_and_grad of the same fit, via the gradient path the runtime script exposes, or a thin scratch wrapper around it).
The deliverable is a table of main vs branch with medians and spread. Gate: any slowdown beyond noise (over 3% on value+grad, or any on forward) is reported to you before shipping. A Monitor watches the progress file throughout. No results committed to autolens_profiling, since it's claimed.
Verification
python -m pytest test_autoarray/inversion/inversion/test_nnls_mge_convergence.py test_autoarray/util/test_jax_nnls.py test_autoarray/inversion/inversion/test_positive_only_dispatch.py -q, then the full PyAutoArray suite.
End-to-end: autolens_workspace_test/scripts/imaging/jax_grad/mge.py under the release env, at its own key and swept over 16 keys, on the branch: all finite.
/smoke_test on the MGE jax scripts downstream (autolens_workspace_test imaging/jax_likelihood/mge*.py, jax_grad/*).
do it properly so go ahead with these, noting that in terms of run times we need to monitor any slowly against autolens_profiling - Proper fix, either of:
run a few tight solver iterations before the gradient solve; or
raise the gradient solve's target to at least the gap the forward solve leaves, so it never has to push toward the boundary.
Context
Heart Release Integrate 2026-09-25 (PyAutoHeart run 36108062907) failed on autolens_workspace_test/scripts/imaging/jax_grad/mge.py: "Gradient contains non-finite values".
Diagnosed to PyAutoArray #572 (merge 3de624b): autoarray/util/jax_nnls.py _solve_nnls_raw_forward_with — raw forward stops at data_scaled_solver_tol
(~5.5e-8) leaving s·z ~1e-10..2.5e-9; backward solve_relaxed_nnls(Q_pc, ..., target_kappa=1e-11)
must push toward the boundary at z/s ~1e13-14; jit while_loop overshoots → NaN at the
50-iteration cap. 4/16 perturbation keys NaN on main, 0/12 pre-#572. Forward likelihood unaffected.
This is the last failure blocking the release (autolens_workspace#577 fixed the other two).
Keep the #571 forward-convergence fix. Choose (a) polish vs (b) effective kappa on evidence;
surface the relaxed solve's converged flag; regression test sweeping perturbation points;
monitor runtime vs autolens_profiling before shipping.
Overview
PyAutoArray #572 (the #571 raw-forward PDIP fix) made the JAX gradient of mapper-less (MGE) positive-only inversions NaN on a fraction of parameter points: the raw forward solve stops at a loose data-scaled tolerance, and the backward relaxed-KKT solve (
nnls_target_kappa=1e-11) then has to push toward the boundary from a z/s ~1e13-14 iterate and overshoots under jit. This fails Heart Release Integrate (2026-09-25,autolens_workspace_test/scripts/imaging/jax_grad/mge.py) and is the last blocker for the release. The forward likelihood is unaffected.Plan
Detailed implementation plan
Affected Repositories
scripts/imaging/jax_grad/mge.pyis the end-to-end witnessBranch Survey
Suggested branch:
feature/mge-nnls-grad-nanImplementation Steps
Approach: pick (a) or (b) from evidence, in one bounded experiment
Both are candidates for the step inside
forward()betweenraw_solveandsolve_relaxed_nnls:solve_nnlsiterations on(Q_pc, q_pc)at the tight default tolerance (n·eps), warm-started from the mapped iterate. Then call the relaxed solve.solve_nnlshas no warm-start today, so this needs a smallinitargument threaded into itswhile_loopstate.kappa_eff = max(target_kappa, max(sy*zy)), so the relaxed solve only moves away from the boundary. No extra solve. The cost is slightly more smoothing in the gradient on affected points.Decision harness (scratch script, not committed). Run both candidates and current main on:
jax_grad/mge.pymodel (reuse the diagnosismge_diag.py);test_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npz;relax_in_first.npz.For each candidate, record:
mge.pyalready does, as rel. error);y/ log-likelihood unchanged vs main (must be bit-identical or within 1e-12: the forward path is untouched);Pick rule:
Implementation (chosen candidate)
autoarray/util/jax_nnls.py:_solve_nnls_raw_forward_with.forward.convergedflag: stop discarding it (yr, sr, zr, _, _). Keep it in the residuals, and expose a small non-custom-vjp diagnostic helper,raw_forward_backward_status(Q_pc, q_pc, Q, q, D, …), that returns the relaxed-solve(converged, iters)for tests and profiling.custom_vjpprimal/forward output signature is unchanged, so no caller changes are needed.autoarray/config/general.yaml: update thennls_preconditioning_no_mapper/nnls_target_kappacomment to describe the backward-pass behaviour. Values stay the same.test_autoarray/inversion/inversion/test_nnls_mge_convergence.py, reusing the existing npz loader:sum(y)(or a fixed random cotangent) throughsolve_nnls_primal_raw_forwardis finite on all 48 SLaM systems, and relaxed-status converged.files/next to the existing npz and noted in its README) that is NaN on main and finite after the fix. Run it red on unfixed source first so it's a real witness.Runtime check against autolens_profiling (required before shipping)
Use a scratch detached worktree of
lens/autolens_profilingatorigin/main. Run the same local CPU fp64 HST settings as the pinned results, main vs branch, alternating (A-B-A-B) to control for drift:scripts/imaging/likelihood_runtime/mge.py: forward likelihood time. Expected unchanged, since the primal is untouched.jax.value_and_gradof the same fit, via the gradient path the runtime script exposes, or a thin scratch wrapper around it).scripts/imaging/hazards/mge_nnls_capture.py: confirms fix(inversion): JAX PDIP positive-only solve fails to converge on SLaM MGE systems #571 stays fixed (48/48 converge).The deliverable is a table of main vs branch with medians and spread. Gate: any slowdown beyond noise (over 3% on value+grad, or any on forward) is reported to you before shipping. A Monitor watches the progress file throughout. No results committed to autolens_profiling, since it's claimed.
Verification
python -m pytest test_autoarray/inversion/inversion/test_nnls_mge_convergence.py test_autoarray/util/test_jax_nnls.py test_autoarray/inversion/inversion/test_positive_only_dispatch.py -q, then the full PyAutoArray suite.autolens_workspace_test/scripts/imaging/jax_grad/mge.pyunder the release env, at its own key and swept over 16 keys, on the branch: all finite./smoke_teston the MGE jax scripts downstream (autolens_workspace_testimaging/jax_likelihood/mge*.py,jax_grad/*).Key Files
autoarray/util/jax_nnls.py—_solve_nnls_raw_forward_with(forward step + relaxed status)autoarray/config/general.yaml—nnls_preconditioning_no_mapper/nnls_target_kappacommentstest_autoarray/inversion/inversion/test_nnls_mge_convergence.py+files/fixturesOriginal Prompt
Click to expand starting prompt
fix: NaN JAX gradient on MGE (mapper-less) positive-only solves after #572
Type: bug
Target: @PyAutoArray
Autonomy: human-required
Original request (verbatim, 2026-09-25)
Context
Heart Release Integrate 2026-09-25 (PyAutoHeart run 36108062907) failed on
autolens_workspace_test/scripts/imaging/jax_grad/mge.py: "Gradient contains non-finite values".Diagnosed to PyAutoArray #572 (merge 3de624b):
autoarray/util/jax_nnls.py_solve_nnls_raw_forward_with— raw forward stops atdata_scaled_solver_tol(~5.5e-8) leaving s·z ~1e-10..2.5e-9; backward
solve_relaxed_nnls(Q_pc, ..., target_kappa=1e-11)must push toward the boundary at z/s ~1e13-14; jit while_loop overshoots → NaN at the
50-iteration cap. 4/16 perturbation keys NaN on main, 0/12 pre-#572. Forward likelihood unaffected.
This is the last failure blocking the release (autolens_workspace#577 fixed the other two).
Keep the #571 forward-convergence fix. Choose (a) polish vs (b) effective kappa on evidence;
surface the relaxed solve's converged flag; regression test sweeping perturbation points;
monitor runtime vs autolens_profiling before shipping.