Skip to content

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

Description

@Jammy2211

Overview

The JAX positive-only inversion solve (jaxnnls primal-dual interior point, max_iter=50) fails to converge on the real SLaM source_lp[1] MGE model (2 lens bases x 20 Gaussians with sigma_min = pixel_scale/10 + 20 source Gaussians = 60 linear columns) and returns the unconverged iterate silently. At 14/48 parameter vectors within ±0.005 of truth the log-likelihood is wrong (-2e5 down to -1e138 vs 15k-22k from NumPy fnnls_cholesky); with cap 200, 9 converge (72-195 it) and 5 never do. Under Nautilus vmap (n_batch 50) one bad lane also drives every batch to the full cap (NNLS ≈ 28% of a GPU batch). Found in the 2026-09-24 MGE likelihood speed audit; the autolens_profiling imaging/mge cell (20+20 Gaussians, mass fixed) hides it (1/32).

Plan

  • Capture the failing (Q, q) linear systems from the SLaM model into a small fixture via a capture script in autolens_profiling.
  • Add a PyAutoArray regression test that is red on main: JAX PDIP must report converged and match the fnnls_cholesky objective on every fixture system, single and under vmap.
  • Diagnose the mechanism on the fixture: unreachable absolute KKT tolerance (6.7e-11 at cond ~1e11) vs genuine divergence in the KKT Cholesky; effect of Jacobi preconditioning; whether the certified active-set solver certifies these pure-MGE systems.
  • Checkpoint with the human on mechanism + fix choice.
  • Fix so the solve never silently returns an unconverged iterate; converged systems stay within fp tolerance of today; NumPy path untouched.
  • Ship library PR (library-first, Heart RED needs the development override); file the follow-on SLaM timing prompt and re-check downstream MGE JAX likelihood pins.
Detailed implementation plan

Affected Repositories

  • PyAutoArray (primary)
  • autolens_profiling (capture script + hazard JSON only; parallel-claim, disjoint files)

Branch Survey

Repository Current Branch Dirty?
array/PyAutoArray main (7fa8d27) clean, unclaimed
lens/autolens_profiling main clean (untracked dataset/abell_1201 pre-existing); claimed by certified-solver-phase-c1-lane-rate, hst-gpu-residue-p4, point-source-cpu-p3 → parallel-claim with disjoint file set

Suggested branch: feature/mge-pdip-nnls-convergence
Worktree: ~/Code/PyAutoLabs-wt/mge-pdip-nnls-convergence/

Code facts (main 7fa8d27)

  • Entry reconstruction_positive_only_from — autoarray/inversion/inversion/inversion_util.py:291-486; JAX branch :390; config nnls_jacobi_preconditioning (True), nnls_target_kappa (1e-11); settings.nnls_solver_tol / nnls_max_iter (None → 50); Jacobi D = 1/sqrt(diag Q) with no tiny-diagonal guard (:429-438); returns only x, stats gets solver="pdip" only.
  • Loop solve_nnls — autoarray/util/jax_nnls.py:33-76: lax.while_loop(pdip_iter < max_iter and converged == 0, pdip_pc_step); stop rule is an absolute KKT inf-norm < min(n*EPSILON, 1e-2), EPSILON = eps*5e3 (jaxnnls pdip.py:32,206) → 6.7e-11 for n=60; returns (x, s, z, converged, pdip_iter) but both callers discard converged (:91, :94); custom_vjp :79-108.
  • NumPy oracle: fnnls_cholesky (autoarray/util/fnnls.py:27) via inversion_util.py:516-570.
  • Certified active-set solver autoarray/util/jax_active_set.py (solve_certified :265, solve_certified_with_fallback :302); dispatch refuses any inversion containing AbstractLinearObjFuncList (inversion/inversion/abstract.py:584-596) because a joint 60-MGE + Delaunay-1500 system failed to certify in 40 passes; pure-MGE systems never tested.
  • Settings autoarray/settings.py (:18 nnls_solver_tol, :19 nnls_max_iter, :25-28 positive_only_solver/certified_*); yaml autoarray/config/general.yaml:5-16.
  • Tests: test_autoarray/util/test_jax_nnls.py, test_jax_active_set.py, test_autoarray/inversion/inversion/test_inversion_util.py (:263+, :390+), test_positive_only_dispatch.py; JAX guard requires_jax = pytest.mark.skipif(find_spec("jax") is None); x64 enabled inline. No test compares JAX PDIP with fnnls or checks the cap.
  • Related: PyAutoArray#377 (pixelized knife-edge flips, ΔLL ~1e-3) — different symptom, cross-reference only.

Implementation Steps

  1. Fixture capture (autolens_profiling) — new scripts/imaging/hazards/mge_nnls_capture.py: build the SLaM source_lp[1] model as in autolens_workspace/scripts/imaging/features/pixelization/slam.py:66-116 on the profiling HST dataset; sample a seeded ±0.005 near-truth box; wrap reconstruction_positive_only_from to record the (curvature_reg_matrix, data_vector) passed (post-regularisation, pre-Jacobi), PDIP converged/pdip_iter at cap 50 and 200, fnnls objective and logL. Write results/hazards/component/mge/nnls_capture_slam_hst_<version>.json + a compressed .npz of 8 systems (5 never-converging, 2 cap-hit-converge-by-200, 1 healthy; ~29 KB raw each). Copy the .npz to test_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npz with a README line naming the generator and versions.
  2. Red regression test (PyAutoArray) — test_autoarray/inversion/inversion/test_nnls_mge_convergence.py (JAX-guarded, x64): per system x_np = fnnls_cholesky(Q, q); x, s, z, converged, it = solve_nnls(Q_pc, q_pc) mirroring the Jacobi path; assert converged == 1, it < 50, obj(x*D) <= obj(x_np)*(1+1e-8)+1e-8 with obj = 0.5 xᵀQx − qᵀx; end-to-end reconstruction_positive_only_from(..., xp=jnp, stats=stats) matches the NumPy call within 1e-8 relative in objective and (post-fix) stats["converged"]; vmap variant over the 8 systems. Run red on main first; record counts here.
  3. Diagnosis — per system cond(Q), cond(Q_pc), min diag, column collinearity, fnnls active-set size; unroll pdip_pc_step 200 steps logging KKT inf-norm and max|x| → plateau above tolerance vs blow-up; toggles: Jacobi off, relative solver_tol, max_iter 200; solve_certified(Q, q) directly on each system. Post the mechanism note here; checkpoint with the human before step 4.
  4. Fix — F1 (always): keep converged/pdip_iter through _solve_nnls_primal_with; reconstruction_positive_only_from writes stats["converged"], stats["iterations"]; non-converged solves never pass through silently. F2 (if plateau): problem-scaled tolerance (max(n·EPSILON, rel·‖q_pc‖∞) or relative KKT residual), verified bit-identical on converging systems. F3 (if blow-up): lax.cond(converged, x, solve_certified(...)[0]) fallback for linear-object-only inversions, gated on step 3 showing certification on pure-MGE systems (extend abstract.py:584-596 only for not has(Mapper); the joint refusal stays); else return the best-KKT iterate tracked in the loop. Never lower the cap, add a ridge to Q, or alter the NumPy path. Docstrings/yaml (try/except KeyError) for any new knob.
  5. Verification — python -m pytest test_autoarray/ green; existing PDIP-vs-scipy and certified tests unchanged; random well-conditioned systems bit-identical before/after (asserted in the new test); autolens_workspace_test JAX MGE likelihood pins against the worktree library (pins may only move at previously failing vectors); re-run step 1 capture with the fixed library: 0/48 unconverged, logL within 1 nat of fnnls.
  6. Ship + follow-ons — ship_library (Heart RED today: integrate:fail + 3 unrelated autolens cluster/weak script failures → human development override at ship time); file draft/research/autolens_profiling/mge_nnls_fix_slam_timing.md (SLaM 60-column single + vmap16, CPU/RTX/A100, NNLS share of a batch).

Key Files

  • autoarray/util/jax_nnls.py — PDIP loop, custom_vjp, discarded converged
  • autoarray/inversion/inversion/inversion_util.py — reconstruction_positive_only_from, Jacobi, stats
  • autoarray/inversion/inversion/abstract.py — solver dispatch/refusal
  • autoarray/util/jax_active_set.py — certified solver (candidate fallback)
  • autoarray/settings.py, autoarray/config/general.yaml — knobs
  • test_autoarray/inversion/inversion/test_nnls_mge_convergence.py (new), test_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npz (new)
  • autolens_profiling/scripts/imaging/hazards/mge_nnls_capture.py (new)

Original Prompt

Click to expand starting prompt

JAX positive-only (PDIP NNLS) solve does not converge on the SLaM source_lp[1] MGE model and returns wrong log-likelihoods

Type: bug
Priority: high
Repos:

  • PyAutoArray
  • autolens_profiling
    Witness: A PyAutoArray regression test builds the 2-basis SLaM MGE model (2 x 20 lens Gaussians with sigma_min = pixel_scale/10 plus 20 source Gaussians) and, at every one of a fixed set of >= 48 near-truth parameter vectors, the JAX positive-only reconstruction reports converged=True and its log-likelihood agrees with NumPy fnnls within 1 nat on CPU fp64; the test fails red on current main (14/48 unconverged).

Target: PyAutoArray

Witness (2026-09-24 audit, autolens 2aaa1c1a8 / autoarray 681938a, jax 0.10.2, CPU fp64)

The SLaM source_lp[1] imaging model (autolens_workspace scripts/imaging/features/pixelization/slam.py:66-116: lens light = 2 bases x 20 Gaussians, sigma_min = pixel_scale/10; source = 20 Gaussians; free Isothermal + ExternalShear; 60 linear columns, 17 free parameters) was evaluated at 48 parameter vectors within +/-0.005 of the simulated truth on the HST-like simple__no_lens_light-style dataset used by autolens_profiling.

  • 14/48 points hit the max_iter=50 cap of solve_nnls_primal (jaxnnls primal-dual interior point, lax.while_loop, autoarray/util/jax_nnls.py:33-78, called from reconstruction_positive_only_from in autoarray/inversion/inversion/inversion_util.py:391-466 after Jacobi preconditioning) and return log-likelihoods from -2e5 down to -1e138. NumPy fnnls on the same curvature/data vector gives sane values (15k-22k).
  • Raising the cap to 200: 9 of the 14 converge (72-195 iterations), 5 never converge and return NaN.
  • Converged results are not always right either: one GPU point was 1 nat off fnnls.
  • Results differ between CPU, GPU, single and vmap evaluation and under 1e-16 perturbations of the images. The Jacobi-preconditioned curvature matrix has condition number ~1.5e11. Suspect: near-collinear narrow Gaussians across the two lens bases.
  • The autolens_profiling imaging/mge cell (20+20 Gaussians, mass FIXED) hides this: 1/32 bad points. Any regression test must use the 2-basis SLaM model.
  • Under Nautilus use_jax_vmap=True (default, autofit/non_linear/search/nest/nautilus/search.py:198, n_batch=50) one failing lane makes every batch run all 50 iterations: NNLS is ~28% of a GPU vmap batch and ~44% of a single A100 evaluation, so convergence is also the largest remaining GPU speed lever (~15-17% of a batch).
  • Lowering the cap is not safe: 15 iterations already moves the objective by 2e-8 at converging points.
  • The certified active-set solver (autoarray/util/jax_active_set.py, PyAutoArray#566/feat(inversion): certified active-set positive solver on the JAX path, opt-in, mapper-only dispatch (#566) #567) is refused whenever the inversion contains AbstractLinearObjFuncList (autoarray/inversion/inversion/abstract.py:578-598), so the MGE path cannot opt into it today.

Ask

  1. Reproduce with a regression test on the 2-basis SLaM MGE model (compare the JAX positive-only reconstruction / log-likelihood against fnnls at a set of near-truth vectors; assert agreement within a nats tolerance and that converged is true).
  2. Isolate the root cause (conditioning of the collinear 2-basis block vs PDIP tolerance/nnls_target_kappa vs cap) and fix convergence: candidates are better preconditioning, a tolerance/adaptive-cap policy, a fallback when the loop hits the cap (never return the unconverged iterate as a likelihood), or extending the certified active-set solver dispatch to linear-object-function inversions.
  3. Numerics: converged points must stay within fp tolerance of current results; only the failing points may change (towards fnnls).
  4. Report the CPU and GPU timing impact on the SLaM model single and vmap16 (autolens_profiling imaging/mge runtime cell is NOT representative; measure the 60-column model).

Audit scripts and JSON artefacts (nnls_fail.py, nnls_fail_slam_cpu.json, lanes.py) exist in the 2026-09-24 session scratch dir; regenerate rather than rely on them.

🤖 Generated with Claude Code

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