Skip to content

fix: NaN JAX gradient on MGE positive-only solves after #572 #573

Description

@Jammy2211

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

  • PyAutoArray (primary)
  • autolens_profiling — timing only, from a detached scratch worktree of origin/main (repo is claimed by other tasks; no commits)
  • autolens_workspace_test — unchanged; scripts/imaging/jax_grad/mge.py is the end-to-end witness

Branch Survey

Repository Current Branch Dirty?
./PyAutoArray main (3de624b) clean
./autolens_profiling main 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:

Decision harness (scratch script, not committed). Run both candidates and current main on:

  1. the 16 PRNGKey perturbation points of the jax_grad/mge.py model (reuse the diagnosis mge_diag.py);
  2. the 48 SLaM systems in test_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npz;
  3. 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);
  • relaxed-solve iterations and convergence;
  • value+grad wall time.

Pick rule:

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.
    • The existing fix(inversion): JAX PDIP positive-only solve fails to converge on SLaM MGE systems #571 forward-convergence tests stay green, unmodified.

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).
  • 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.
  • 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/*).

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_kappa comments
  • test_autoarray/inversion/inversion/test_nnls_mge_convergence.py + files/ fixtures

Original 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)

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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions