fix(inversion): raw-forward PDIP for mapper-less positive-only JAX solves, surface convergence (#571) - #572
Merged
Merged
Conversation
… systems (#571) Adds 8 real (curvature_reg_matrix, data_vector) systems captured from the SLaM source_lp[1] MGE model (autolens_profiling scripts/imaging/hazards/mge_nnls_capture.py, 158 KB .npz) and test_nnls_mge_convergence.py: the Jacobi-preconditioned PDIP solve must converge inside the 50-iteration cap and reach the fnnls objective, single, end-to-end through reconstruction_positive_only_from, and under vmap. Red on main by design: 15 failed / 6 passed (7 cap-hit systems x 2 parametrized tests + the vmap test). A control on 3 seeded well-conditioned systems pins today's PDIP answer bit-identically and passes. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
…ce convergence (#571) On linear-object-only (MGE) inversions, Jacobi scaling turns signal-free Gaussian columns (whose diagonal is only the no-regularization floor) into degenerate coordinates on which the PDIP dual diverges: 14/48 near-truth SLaM source_lp[1] points hit the 50-iteration cap and returned wrong log-likelihoods, and 5 never converged. - reconstruction_positive_only_from gains preconditioning="jacobi"|"raw" (default "jacobi", byte-identical: x, jitted x and gradients checked on 11 systems). "raw" runs the forward PDIP on the un-preconditioned system with a data-scaled tolerance (1e-2 * n * EPSILON * max(1, max|q|)). It keeps today's Jacobi-space relaxed-KKT backward pass, because a raw-Q backward pass gives NaN gradients on 3 of the 8 fixture systems. - AbstractInversion.positive_only_preconditioning_used: mapper inversions keep "jacobi"; mapper-less ones use Settings.nnls_preconditioning_no_mapper (config key nnls_preconditioning_no_mapper, default "raw"). - F1: solve_nnls_primal_with_status returns (x, converged, pdip_iter), with no cotangents for the integer flags. solve_nnls_primal keeps its drop-in signature for downstream scripts. The PDIP paths record stats["converged"] and stats["iterations"] as traced scalars (the same channel as the certified path's certified/passes). - The red regression tests are now green: all 8 fixture systems converge in 16-19 iterations, single, end-to-end and under vmap, with the fnnls objective reached to <=1e-12 relative and finite, non-zero gradients. The Jacobi mode now reports converged=0 on the witness systems. Dispatch tests added. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01WAqZd1NiPsVYPSAP8tX4YL
This was referenced Sep 24, 2026
This was referenced Sep 25, 2026
Jammy2211
added a commit
that referenced
this pull request
Sep 25, 2026
fix: NaN JAX gradient on MGE positive-only solves after #572
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 #571. The JAX positive-only inversion solve (jaxnnls PDIP,
max_iter=50) diverged on the SLaMsource_lp[1]MGE model (2 lens bases x 20 Gaussians + 20 source Gaussians): on 14/48 near-truth parameter vectors it hit the cap and silently returned the unconverged iterate, giving log-likelihoods from -2e5 down to -1e138 (NumPyfnnls_cholesky: 15k-22k). Mechanism (issue comment "Step 3 diagnosis"): signal-free source Gaussian columns whose diagonal is only the 1e-3 regularisation floor become degenerate unit coordinates after Jacobi scaling; the slack collapses on the first step and the dual blows up. It is neither conditioning nor collinearity.Fix: inversions without a Mapper (MGE / linear-light-profile-only stages) now run the PDIP forward solve on the raw curvature matrix with a data-scaled tolerance
1e-2 · n · EPSILON · max(1, ‖q‖∞), and keep the existing Jacobi-space relaxed-KKT backward pass (a pure raw backward gives NaN gradients on 3 fixture systems). Inversions with a Mapper keep the Jacobi path, bit-identical to main (x, jitted x and gradients on 11 systems). The convergence flag and iteration count are surfaced through the inversionstatsdict so a non-converged solve is no longer silent. The NumPy path is untouched.Evidence: red regression fixture (8 real SLaM systems, 158 KB) fails 15/21 on main and passes after the fix; full suite 1681 passed; re-running the 48-vector capture (autolens_profiling
scripts/imaging/hazards/mge_nnls_capture.py) gives 0/48 unconverged with max |logL_jax − logL_numpy| = 5.6e-7; 14 mapper-less autolens_workspace_test JAX pin scripts move by ≤ 2.5e-11 relative (4 fail identically before and after: pre-existing vmap mismatch). CPU timing neutral. GPU timing/parity not measured.API Changes
Additive and default-preserving for every existing caller.
reconstruction_positive_only_fromgains an optionalpreconditioning="jacobi"|"raw"argument (default"jacobi"). NewSettings.nnls_preconditioning_no_mapper/ config keynnls_preconditioning_no_mapper(defaultraw) selects the mode for mapper-less JAX inversions; set it tojacobito recover the previous behaviour. Behaviour change: mapper-less JAX positive-only inversions now converge where they previously returned garbage; converged results move by ≤ 1e-12 relative in the NNLS objective, and their gradients change at the ~1e-3 level (previous relaxed-KKT approximation, now started from the raw solution). New helpers inautoarray.util.jax_nnlsexpose the convergence status.solve_nnls_primalkeeps its x-only signature.See full details below.
Test Plan
python -m pytest test_autoarray/— 1681 passed (worktree, JAX fp64)test_autoarray/inversion/inversion/test_nnls_mge_convergence.py— red on main (15 failed / 6 passed), green after the fix; includes vmap+jit, gradient-finite and bit-identity controlstest_positive_only_dispatch.py— Mapper → jacobi, mapper-less → setting; 42 passed across both filesHeart RED development override
Heart is RED at ship time for reasons unrelated to this branch. Exact reasons from
pyauto-heart readinesson 2026-09-24:release validation FAILED (stage integrate)workspace validation not passing (4 failed, cloud#35579888156: autolens notebooks/cluster/modeling.ipynb, autolens notebooks/weak/a2744.ipynb, autolens scripts/cluster/modeling.py, +1 more)manifest drift: hub organism blurb (organs present) — 7 mismatch(es) vs PyAutoMind/repos.yamlBranch gates passed: full suite 1681 passed; regression fixture green; capture 0/48; pins ≤ 2.5e-11. Human authorisation, given live in the session in reply to the override request naming #571: "I authorise you to continue". Permits commit/push/PR-open only; merge requires a separate
/prmwith green checks. This branch does not repair Heart.Full API Changes (for automation & release notes)
Added
autoarray.settings.Settings.nnls_preconditioning_no_mapper—"raw"(default, via config) or"jacobi"; preconditioning used by JAX positive-only solves in inversions with noMappergeneral.yaml: inversion.nnls_preconditioning_no_mapper: rawautoarray.inversion.inversion.inversion_util.reconstruction_positive_only_from(..., preconditioning="jacobi")— new optional kwarg;"raw"runs the forward PDIP on the unscaled system with the data-scaled tolerance and the Jacobi-space backward pass;"raw"withsolver="certified"raisesAbstractInversion.positive_only_preconditioning_used—"jacobi"if the inversion has aMapper, else the setting aboveautoarray.util.jax_nnls.solve_nnls_primal_with_status(Q, q, ...)→(x, converged, pdip_iter)(integer outputs carry no cotangent)autoarray.util.jax_nnls.solve_nnls_primal_raw_forward(...),data_scaled_solver_tol(Q, q),DATA_SCALED_TOL_FACTOR = 1e-2stats["converged"],stats["iterations"],stats["preconditioning"]recorded by the PDIP paths ofreconstruction_positive_only_from(traced scalars, same out-dict channel as the certified path'scertified/passes)Changed Behaviour
Mapper(MGE /lp_linear-only) use the raw-forward PDIP by default. Previously-converged systems move ≤ 1e-12 relative in objective; previously-unconverged systems now return the correct solution; gradients on these inversions change at the ~1e-3 level. Inversions with aMapperand the NumPy path are unchanged.Migration
Settings(nnls_preconditioning_no_mapper="jacobi")or set the config key.Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code
https://claude.ai/code/session_01WAqZd1NiPsVYPSAP8tX4YL