Skip to content

fix(inversion): raw-forward PDIP for mapper-less positive-only JAX solves, surface convergence (#571) - #572

Merged
Jammy2211 merged 2 commits into
mainfrom
feature/mge-pdip-nnls-convergence
Sep 24, 2026
Merged

Jammy2211 merged 2 commits into
mainfrom
feature/mge-pdip-nnls-convergence

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Summary

Fixes #571. The JAX positive-only inversion solve (jaxnnls PDIP, max_iter=50) diverged on the SLaM source_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 (NumPy fnnls_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 inversion stats dict 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_from gains an optional preconditioning="jacobi"|"raw" argument (default "jacobi"). New Settings.nnls_preconditioning_no_mapper / config key nnls_preconditioning_no_mapper (default raw) selects the mode for mapper-less JAX inversions; set it to jacobi to 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 in autoarray.util.jax_nnls expose the convergence status. solve_nnls_primal keeps 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 controls
  • test_positive_only_dispatch.py — Mapper → jacobi, mapper-less → setting; 42 passed across both files
  • Capture re-run on the SLaM model with the fixed library: 0/48 unconverged, max |ΔlogL| 5.6e-7 vs NumPy
  • 14 mapper-less workspace_test JAX likelihood pins: max 2.5e-11 relative movement
  • GPU parity / vmap timing (A100) — follow-up profiling prompt

Heart RED development override

Heart is RED at ship time for reasons unrelated to this branch. Exact reasons from pyauto-heart readiness on 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.yaml

Branch 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 /prm with 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 no Mapper
  • config general.yaml: inversion.nnls_preconditioning_no_mapper: raw
  • autoarray.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" with solver="certified" raises
  • AbstractInversion.positive_only_preconditioning_used — "jacobi" if the inversion has a Mapper, else the setting above
  • autoarray.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-2
  • stats["converged"], stats["iterations"], stats["preconditioning"] recorded by the PDIP paths of reconstruction_positive_only_from (traced scalars, same out-dict channel as the certified path's certified/passes)

Changed Behaviour

  • JAX positive-only inversions containing no 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 a Mapper and the NumPy path are unchanged.

Migration

  • None required. To restore the old behaviour for mapper-less inversions: 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

Jammy2211 and others added 2 commits September 24, 2026 17:01
… 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
@Jammy2211 Jammy2211 added the pending-release PR queued for the next release build label Sep 24, 2026
@Jammy2211
Jammy2211 merged commit 3de624b into main Sep 24, 2026
3 checks passed
@Jammy2211
Jammy2211 deleted the feature/mge-pdip-nnls-convergence branch September 24, 2026 18:18
Jammy2211 added a commit that referenced this pull request Sep 25, 2026
fix: NaN JAX gradient on MGE positive-only solves after #572
@Jammy2211 Jammy2211 removed the pending-release PR queued for the next release build label Sep 26, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix(inversion): JAX PDIP positive-only solve fails to converge on SLaM MGE systems

1 participant