You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
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.
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
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.
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.
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.
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.
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.
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.
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.
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).
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).
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.
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).
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.
Numerics: converged points must stay within fp tolerance of current results; only the failing points may change (towards fnnls).
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.
Overview
The JAX positive-only inversion solve (jaxnnls primal-dual interior point,
max_iter=50) fails to converge on the real SLaMsource_lp[1]MGE model (2 lens bases x 20 Gaussians withsigma_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 NumPyfnnls_cholesky); with cap 200, 9 converge (72-195 it) and 5 never do. Under Nautilusvmap(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_profilingimaging/mgecell (20+20 Gaussians, mass fixed) hides it (1/32).Plan
fnnls_choleskyobjective on every fixture system, single and under vmap.Detailed implementation plan
Affected Repositories
Branch Survey
Suggested branch:
feature/mge-pdip-nnls-convergenceWorktree:
~/Code/PyAutoLabs-wt/mge-pdip-nnls-convergence/Code facts (main 7fa8d27)
reconstruction_positive_only_from—autoarray/inversion/inversion/inversion_util.py:291-486; JAX branch :390; confignnls_jacobi_preconditioning(True),nnls_target_kappa(1e-11);settings.nnls_solver_tol/nnls_max_iter(None → 50); JacobiD = 1/sqrt(diag Q)with no tiny-diagonal guard (:429-438); returns onlyx,statsgetssolver="pdip"only.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(jaxnnlspdip.py:32,206) → 6.7e-11 for n=60; returns(x, s, z, converged, pdip_iter)but both callers discardconverged(:91, :94); custom_vjp :79-108.fnnls_cholesky(autoarray/util/fnnls.py:27) viainversion_util.py:516-570.autoarray/util/jax_active_set.py(solve_certified:265,solve_certified_with_fallback:302); dispatch refuses any inversion containingAbstractLinearObjFuncList(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.autoarray/settings.py(:18nnls_solver_tol, :19nnls_max_iter, :25-28 positive_only_solver/certified_*); yamlautoarray/config/general.yaml:5-16.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 guardrequires_jax = pytest.mark.skipif(find_spec("jax") is None); x64 enabled inline. No test compares JAX PDIP with fnnls or checks the cap.Implementation Steps
scripts/imaging/hazards/mge_nnls_capture.py: build the SLaMsource_lp[1]model as inautolens_workspace/scripts/imaging/features/pixelization/slam.py:66-116on the profiling HST dataset; sample a seeded ±0.005 near-truth box; wrapreconstruction_positive_only_fromto record the(curvature_reg_matrix, data_vector)passed (post-regularisation, pre-Jacobi), PDIPconverged/pdip_iterat cap 50 and 200, fnnls objective and logL. Writeresults/hazards/component/mge/nnls_capture_slam_hst_<version>.json+ a compressed.npzof 8 systems (5 never-converging, 2 cap-hit-converge-by-200, 1 healthy; ~29 KB raw each). Copy the.npztotest_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npzwith a README line naming the generator and versions.test_autoarray/inversion/inversion/test_nnls_mge_convergence.py(JAX-guarded, x64): per systemx_np = fnnls_cholesky(Q, q);x, s, z, converged, it = solve_nnls(Q_pc, q_pc)mirroring the Jacobi path; assertconverged == 1,it < 50,obj(x*D) <= obj(x_np)*(1+1e-8)+1e-8withobj = 0.5 xᵀQx − qᵀx; end-to-endreconstruction_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.cond(Q),cond(Q_pc), min diag, column collinearity, fnnls active-set size; unrollpdip_pc_step200 steps logging KKT inf-norm andmax|x|→ plateau above tolerance vs blow-up; toggles: Jacobi off, relativesolver_tol,max_iter200;solve_certified(Q, q)directly on each system. Post the mechanism note here; checkpoint with the human before step 4.converged/pdip_iterthrough_solve_nnls_primal_with;reconstruction_positive_only_fromwritesstats["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 (extendabstract.py:584-596only fornot 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.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.ship_library(Heart RED today:integrate:fail+ 3 unrelated autolens cluster/weak script failures → human development override at ship time); filedraft/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, discardedconvergedautoarray/inversion/inversion/inversion_util.py—reconstruction_positive_only_from, Jacobi, statsautoarray/inversion/inversion/abstract.py— solver dispatch/refusalautoarray/util/jax_active_set.py— certified solver (candidate fallback)autoarray/settings.py,autoarray/config/general.yaml— knobstest_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:
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_workspacescripts/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-likesimple__no_lens_light-style dataset used by autolens_profiling.max_iter=50cap ofsolve_nnls_primal(jaxnnls primal-dual interior point,lax.while_loop,autoarray/util/jax_nnls.py:33-78, called fromreconstruction_positive_only_frominautoarray/inversion/inversion/inversion_util.py:391-466after Jacobi preconditioning) and return log-likelihoods from -2e5 down to -1e138. NumPy fnnls on the same curvature/data vector gives sane values (15k-22k).imaging/mgecell (20+20 Gaussians, mass FIXED) hides this: 1/32 bad points. Any regression test must use the 2-basis SLaM model.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).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 containsAbstractLinearObjFuncList(autoarray/inversion/inversion/abstract.py:578-598), so the MGE path cannot opt into it today.Ask
convergedis true).nnls_target_kappavs 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.imaging/mgeruntime 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