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

  1. Jammy2211 commented on Sep 25, 2026

    @Jammy2211
    CollaboratorAuthor

    Corrective-PR authorization (Heart RED, human-authorized)

    Heart RED reason (verbatim, pyauto-heart readiness, 2026-09-25T15:59:20Z; re-read unchanged at 18:05:02Z):

    release validation FAILED (stage integrate)

    Authorization: the human (@Jammy2211), live in the Claude Code CLI session on 2026-09-25, replied to the agent-surfaced reason and the request "authorize the corrective PR for PyAutoArray#573 against Heart RED 'release validation FAILED (stage integrate)'" with, verbatim:

    I authorize on heart red

    Scope, per PyAutoBrain/AUTONOMY.md "Corrective-PR exception for Heart RED": commit, push and one pending-release PR for this issue only. No merge (a separate /prm), no release, no rehearsal.

    Causal mapping: RED reason → Release Integrate run 36108062907 failed autolens_workspace_test/scripts/imaging/jax_grad/mge.py ("Gradient contains non-finite values") → this issue (the #572 raw-forward mode leaves s·z far above nnls_target_kappa, so the relaxed-KKT backward solve diverges to NaN) → the plan above → diff: autoarray/util/jax_nnls.py polish before the relaxed solve, plus regression tests. The other two failures in that run (autolens_workspace weak/real_data/a2744.py, cluster/lenstool/modeling.py) were fixed by autolens_workspace#577 (merged b490ba43).

    Tests: red-first regression fixture (6/8 cases fail on 3de624b5, pass on the branch); targeted files 77 passed; full PyAutoArray suite 1702 passed; jax_grad/mge.py end-to-end under the release env: 16/16 perturbation keys finite (eager and jit), FD 16/16 (max rel 1.59e-6), logL bit-identical to main; autolens_profiling runtime: no slowdown beyond noise (value+grad −2.7%, forward HLO byte-identical); #571 capture 48/48 converged on both.

    Validation plan: downstream workspace smoke (all six) on the branch before PR; after merge, Heart stays RED until a fresh Release Integrate run passes on new wheels (nightly or re-dispatch — a human call). Release stays blocked until then.

  2. Jammy2211 commented on Sep 25, 2026

    @Jammy2211
    CollaboratorAuthor

    Library PR Created

    • PR: fix: NaN JAX gradient on MGE positive-only solves after #572 #574 (pending-release, corrective under Heart RED release validation FAILED (stage integrate))
    • Fix: backward-pass polish (RAW_BACKWARD_POLISH_MAX_ITER=10) before the relaxed-KKT solve; kappa_eff alternative measured and rejected (6/48 SLaM NaN). Forward unchanged.
    • Evidence: full suite 1702 PASS; red-first fixture; jax_grad/mge.py 16/16 finite eager+jit; runtime vs autolens_profiling no slowdown; independent review CLEAN; smoke 159 PASS + 1 pre-existing-on-main failure (autolens_workspace_test interferometer/jax_likelihood/mge.py — identical on 3de624b, not this PR).
    • Workspace impact: none (API additive, no workspace callers) → option (iii).
    • Next: human /prm to merge; then a fresh Release Integrate run to clear Heart RED.
  3. Jammy2211 commented on Sep 25, 2026

    @Jammy2211
    CollaboratorAuthor

    Shipped

    • Merged: fix: NaN JAX gradient on MGE positive-only solves after #572 #574 → 5f8a8dee (human /prm, corrective under Heart RED release validation FAILED (stage integrate); all 3 CI legs green).
    • Issue left open deliberately (corrective-PR policy): close once a fresh Release Integrate run on wheels containing 5f8a8dee shows autolens_workspace_test/scripts/imaging/jax_grad/mge.py PASS.
    • Unreleased: pending-release stays on the PR until a release publishes.
  4. Jammy2211 commented on Sep 25, 2026

    @Jammy2211
    CollaboratorAuthor

    Verified in release validation — closing

    Release Integrate run 36179179279 on TestPyPI 2026.9.25.1.dev78401 (PyAutoArray 5f8a8dee, rehearsal 36177451534): success, 709 passed / 0 failed / 0 timeouts. autolens_test, imaging (containing jax_grad/mge.py) passed, as did autolens, weak and autolens, cluster (autolens_workspace#577). Ingested: Heart RED release validation FAILED (stage integrate) is cleared (now YELLOW). pending-release stays on #574 until a release publishes.

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