From 0cef674e1810c47e8afac76b34e34aa7779f4fc1 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Fri, 25 Sep 2026 17:14:48 +0100 Subject: [PATCH 1/4] test: red regression fixture for NaN raw-mode MGE gradients (#573) Four 20x20 systems captured from autolens_workspace_test jax_grad/mge.py at PRNGKey perturbations 2, 10, 12, 14: the raw-mode gradient is NaN on each on 3de624b5 (eager 4/4, jit 2/4). Adds backward-pass convergence tests over these and the 8 SLaM #571 systems via raw_forward_backward_status. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../inversion/inversion/files/README.md | 8 ++ .../inversion/files/mge_grad_nan_systems.npz | Bin 0 -> 6846 bytes .../inversion/test_nnls_mge_convergence.py | 88 ++++++++++++++++++ 3 files changed, 96 insertions(+) create mode 100644 test_autoarray/inversion/inversion/files/mge_grad_nan_systems.npz diff --git a/test_autoarray/inversion/inversion/files/README.md b/test_autoarray/inversion/inversion/files/README.md index 887f535f8..77a5bbc0f 100644 --- a/test_autoarray/inversion/inversion/files/README.md +++ b/test_autoarray/inversion/inversion/files/README.md @@ -7,3 +7,11 @@ `autolens_profiling/scripts/imaging/hazards/mge_nnls_capture.py` (autolens_profiling d6926af), run 2026-09-24 on CPU fp64 with PyAutoArray 7fa8d2714f, PyAutoGalaxy 70a61e26cd, PyAutoLens 86054bbc19, PyAutoFit a736840127, jax 0.10.2. +- `mge_grad_nan_systems.npz` — 4 positive-only systems `Q_` (20x20) / `q_` captured from the + autolens_workspace_test `scripts/imaging/jax_grad/mge.py` model (MGE source, NFWSph + ExternalShear) at + `physical_values_from_prior_medians + jax.random.uniform(PRNGKey(p), minval=0.01, maxval=0.05)` for + p = 2, 10, 12, 14 (PyAutoArray#573): on PyAutoArray 3de624b5 the `"raw"`-mode gradient is NaN on each (the + relaxed-KKT backward solve diverges from the loose raw-forward iterate); `meta` holds the per-system JSON. + Captured 2026-09-25 on CPU fp64 via a `jax.debug.callback` on `reconstruction_positive_only_from`, with + autolens_workspace_test 5ec64413d2, PyAutoArray 3de624b5b9, PyAutoGalaxy 70a61e26cd, PyAutoLens 86054bbc19, + PyAutoFit dd9fbe0aab, jax 0.10.2, numpy 2.5.3. diff --git a/test_autoarray/inversion/inversion/files/mge_grad_nan_systems.npz b/test_autoarray/inversion/inversion/files/mge_grad_nan_systems.npz new file mode 100644 index 0000000000000000000000000000000000000000..7c71084fd06db745f8097ad8a1bc5c5133ea0bad GIT binary patch literal 6846 zcmb`MbySpF7sio}3rI_gl%SGQ!hnQ;bT^EEfJiq(jf4y(2#9njrF56Tfb@{kFqFj5 z9YYSEc(3=W_x}04v(7qe*33T7yXRf&H+w%@O%Vh09vT`N?zhE+c1a2fx%}(GK_f>~ z2l8;)JA$sET}MMB{PvA+7bfPP`(~V>lBNe;NB7OlOhV5)lp<)FSRBWyopG2 zTm4-lln*3D4Q?-JWMFpE30I=m!SFF3)4xs5C~;4fcDJFEUgEn(hOf!wnRu;D2`}8X zX)xTm7s$1r#gCRk7l#ZVT_pwrL5w-0*RieIzT=1l2T=Y96x>)~M z+t3tU+JJEmc2W}+VZdBkO2iI_v1{eWzIa|8Jljb|uFKvMovAg2!zEtG8Q&>4<|F=e zeTwz_gHI8!JnhFi*#MfZ1yiTx*=b!czIz3`mr44(b zE?qqWGIAf&nX0G5JD;yl51b@Hc`geTft9&(cxMDwu{8(`Dd;qO8$?7;NrVZoz3os; zWiF{)~@ha(lJ$9XR_CHkKKR$goud?$wP)26Pv6s$LcLsXXU?Z|)aA(XolUlcg`} zO|>2c<<5rTFutRzX#&}KS?uQwP?UXiIB3bXkb8<$O>FJ01i~t)Cm)w@bC*rKk(u)w zGIslgRZtv|pLVi`!Nj*3T4E(4M-?RGqrDdi#HF|$sj;nzpI2=#t0-3!cjuOcG*%=i0*|Rf#w^G@^qdKz$K)gg z;e7Hv6+8+DZJIK<;0+&x7CNShWZQLW4^4zLv)=3H33bS*RPqW+#7T0NC ztM^tUFf?{#rhPaF8UoOUZD_clQLp8Deh?*7VMq~Y)Y}47vI$1%{M2v5J$aPEg7zBR?WAlz86*F>doS zWgn1RWHd`*^l3ur#sj(K@97}6qPpe?Uqj@_Kx~3B2>=V zAmWu<7Np#u&8A)nRYV<1DN>|~Xb}x+=qNIS{+yG}6aTIa9~qieSnfGJ!JCsE@p<~s zGv>ObKi_)1)Ar5oYzvtS0<~Y@ao?!n;m`1J`Uf62ul)&)sy{ch!qNrR*f%Wj`tGUs7ujG8;Waqte~|9a<$e`ex}eT6xiD8V>Q}fE z`gxFfkEpvgfdPKG<2p`5QY!sSC#Pkp<;<@gyzcxg&9w2gCO9x~gT68Nd=E~R}%FpSl19^XEO^fynH-{pJRXBRr=$k9eSWhpzDx~y(VVx!YwRUaz4!N zM>;Z9kNIx3a!@((_q<~DK62}dCO&IA+e}e~if!+sP$$vPv)7b^6@$Ca>U~3H3;mJ! z5UM&LjLgaNZ0HXh!XykQ>tNqd@Fnq377f|*C2j2e)7BFQRL)HjE6o9H682bWmnj5{ z%{vZ{b8rbP9-cxV`eR+ToitB`ZNfCDu=jWJOKs*|s13P!m&RNPBVemmc)AqQ5pO=T z+poo{Qp+jQEX*VrwN3R2gu!h|5kqEo42_|IB2DS}R)+5r2-;R%ET$%rCA?1t^57z7 z3W#M2u?-8m6PLgY-nByKWnM`nb?s*P`Ye8=jCk4t`Q^=$+i9y9TNiBd2Hj69G0*U( z1I24Uc^ihPc>|5IZ+#8|K~w6yM1UMVCgP6%{MBxqa*U1lvh3eKvV9<#9=_D|kQ(CR3_;4xAJ`B{j!Qi*Y`d-m}Up>A!z%jqCgaZ}AuF%lioB5tu zX3;756!YxPJYdVCVEt9EjB@bsnA(tXgx?~K^2#Y*ozV)_QAmB0Tic{rSkm%_i4<#m zUVk_EDS)Ys|HAqGYi7>%TjnqN!@yJAA?#)5afR7XdQ3xWUOr`n2myX^D^A(_TFrJ4j`=I`lUm zDWfcF>_OPxsgXW9Fd@w~Wa|teYP428Ah#2rV)lyiGmZL|XnnlR=R5<|SgGQpdE#yD zXvp2}QwK~D`At=Z;>YK?oi>JWWRhI{2ktUs&33SfG(|&W+lHOK+?S<5cE_j{!0AiC zEpA5=BU6duE=icZ0?vh&=Hi79ayaG968XCoy-f{r89Q5r@Q+Y=^{TrJLlq(sf{f=_ zs&foqyblPxWlZMk@{{SR3V|xxQwH0(GoKRoT=dL0a6B|9CYy~bv1rzi3zY@@IUdjU zoAY<5#QeQww!d(1kbwNx-7Td3dSa&EfE}wG`ck69B*b(K-oxH$Se+EU6TMq;GUx3{ zfeQHPF5QUwb&@`b=US#r9KFkQ)6(+nTuUMldbal9&1})vv(7tJcl@85mCg0OJQAEB?QH zyqHB<%j#n(VJYTm8AL5vD>`x&?jN~V`$X@YVJA=P!ucX;{dkRxkC1oT63sD#=qsSQ zxTP$04=~Ryxq0JMJ%oKb&X*FWh#(*inhhzcF=2AM^h+eL%?u>t zCo+-yjgD>ig3^9YRvpOqI~~LSt|Tuthur8ogs_^OaBfSzz-)PiT|-EuRMmXs*WLUu z=*4Z9+szTYlb|9y9(y%MN>fZKrt)ZXY9Q%+vidRBCY`yF{MXmqY1lVlTIc8?p^i&8 zPrn}G_`n#~`j8Ik{i;#kjcLUl`SX=R3#2knW0Kx<$j5(}QRnm2lGVe8G-vVgEyC_* zA3w=6O5{2@gR1#FWYxbYhLJe}y^f5F_PFKNvBPObVj(}dZ zYXLcr}_(kQ%j7x*q>vU<9|Sx2vhi3(j&hs0XK~a3M~n1BQl#swqN7k0Hn8^q$MB~_3Q9^6*658n-~5`ad3GQFkgRW08?(zB`2 z$FAZmy&~xX7LEqByfm)#CD;=4_1QQ3f(M*=3Ad79mBZQz*=wd_;+{iSiQTw3HnUP3 zcW*>loXO#%M3R&noHJq^)RZy3Pw(%IENJW0@_IFPQ=pL1T(LewItwTB>Mvg&;WQAsSMxq#a)$(oO|({BWPP|n^-fJWi|dI4<>{z|g;4RH^nvY!OMT%f z&sX1|8R2vpsMfv2cKsHLp(?EPELaO%k^V$1t8uCB_K5P^0bhTv+S`qK8c1g&_w7tmPx z35vs;2|SF)1Y+0Da2p_g-=B#BJf0bzAn2cfIP(Qo7+zl+h^`}R*qQ98;5mcS^?>x4 zp!{5|WBg#&A;At31V(WHw27=FYQTGis)u9`y{Mm?jgs00uv>VxYspC2-a9)bY$w6| z8C}=9W;c|aoi|=k)|!|UfiWAX+YFq;8<$%F91tLILD$B5>JmAUBqO*}OxmFY1df$( zpOvG}eD%6(O?vvnq=-OK%toBw)VNFo6Sq<)Llm9bxJjlZzM4r8Yi|$Wnu^y%IYOF{ zbJ{ypKmS76iV`E+Bj>m{`wGnCnsddH1W~3}9K*Nchr` zrd_q_Xnw4la7xMTH-NS_Nwd+*+lGbLAxEPlg7JAS^jkO&!db*tTzI^jz?w!Lvfxk1wVMz(eY z5*<10*FRmJ6uJ~`{#gF}S?)#Lxv5{1o+5W=(#9?Th!NPEa>Aiv!9#3Hj~lW8&8HX? zOLVMUHS|EQJyQnQA^%Ze{&BkB?u&}|y*xTe`nM`exB{$~&WlHWyD2f_BbizeqSFG9B70*lCn1KrtdUNv8bLK?&$E=9oyoHFUiP z#v&#*BT)6x#X`VNl(UXY zmiL3h-ue6GaFJv|Mt6#{m#Gi`wlb9SYBL-AXEGJUAHk7Qicattb&=ZHWEDGAuCvyWVq)-;G zx5X_`mgu3jdTOEziS1%RpP{cAjPVU1yE6lvUTqPbI94ILql5I=o^#ZgIn#BuVIxc^m5s zj>$vL*MoWTg?XF`g2{?Z@%BcereCh!FjD@(B|@JKNi7qszp zq8K@Rm8B~TQsn9CR?=X2_mi2`8|89D2`_B&V_(XB2eyxl=?6&*!$?J13*^e^v59zy ze6!-1;+jQ>;QrtN-yxqA@Q|oM_e8KPrZIcew1ry2*g@f4g|03X5`Hr8bQTGKeBr5xIBhzfdROxWWbK<2|!^R7% z-tOSg3%QEz0k*jh$NaHEktFGt#s@gE@w}sYai&PCpdrCib!u(MYc8r45(o1_Zg@ur?AcGrb_eweE_=6fEn}WCO9`f?#3VH*G0Yd+>Uq(BJ_Lh$XF2L zJp$&M&C}#tKOijS2)Ej#NATex5or4r%HR!93o;-Uy3IRN={~!&rFPEWX|+<#Qg2=q z?-;fX!cjuJGIs(K*G}Gd*@X|KF{^y0e;|9=SFi*yf~ZWlPBbcpl^u_kSW(Rekgf0r z`{b`2lX&vCzXd*@AK|w-Cd!CP>L^Rhy*Q4HoGoZw?;JWYKD^g<&Lbok3}fD|@7d=N zCu@8xEK$r_Pg2n5{xo@;u5pQ0wP2*de*ILWV2wtvK3u0czc92q4=idkzlpvrIytMR zcn$p?#;>!C@5733gA1I$%CxKP#dk0AugS*mL02B;l}UNUr(c5r-)f$}9L?W@uAIQ{ zR^Pvr?qZ>Ro1uL71%D5^as#gD_irU*yep;u47%#7zqif*Qc}RbQu^Lw|E_dZ#$MH_ z|5hR*zEb+npsOnIdztqyCCqQLpucLu-<7U_|KB;OkY6c%FC>2tx&r4F7_T6CHHG?d ey|^Yyf%Z2>)fBO?uRg*1_L}+ja9}sC-u(~jT6Pxz literal 0 HcmV?d00001 diff --git a/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py index 0d34d2632..f2a0da030 100644 --- a/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py +++ b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py @@ -323,3 +323,91 @@ def test__control__well_conditioned_pdip_unchanged(jnp, n, seed): np.testing.assert_allclose( x_raw, x_scipy, rtol=0, atol=1e-8 * np.abs(x_scipy).max() ) + + +# --------------------------------------------------------------------------------------------------------------- +# PyAutoArray#573: NaN gradients of the "raw" mode. +# +# `files/mge_grad_nan_systems.npz` holds 4 (20 x 20) systems captured from the autolens_workspace_test +# `jax_grad/mge.py` model (MGE source, NFWSph + shear) at the PRNGKey perturbations 2, 10, 12 and 14 (see +# `files/README.md`). The raw forward solve stops at the data-scaled tolerance with s * z ~ 1e-10 .. 2.5e-9, far +# above `nnls_target_kappa = 1e-11`; the relaxed-KKT solve on the Jacobi system then has to push toward the +# boundary from z / s ~ 1e13 and hits its 50-iteration cap with NaN (or "converges" with s < 0), so the gradient +# is NaN. The fix polishes the mapped iterate with a few tight PDIP iterations on the Jacobi system first. +# --------------------------------------------------------------------------------------------------------------- + +GRAD_NAN_FIXTURE = Path(__file__).parent / "files" / "mge_grad_nan_systems.npz" + + +def _load_grad_nan_systems(): + with np.load(GRAD_NAN_FIXTURE) as data: + meta = json.loads(str(data["meta"])) + systems = [ + (np.asarray(data[f"Q_{s['key']}"]), np.asarray(data[f"q_{s['key']}"])) + for s in meta["systems"] + ] + return meta, systems + + +GRAD_NAN_META, GRAD_NAN_SYSTEMS = _load_grad_nan_systems() +GRAD_NAN_IDS = [f"prng{s['prng_key']}" for s in GRAD_NAN_META["systems"]] + + +def test__grad_nan_fixture_is_the_captured_jax_grad_mge_set(): + assert GRAD_NAN_FIXTURE.stat().st_size < 50_000 + assert [s["prng_key"] for s in GRAD_NAN_META["systems"]] == [2, 10, 12, 14] + for Q, q in GRAD_NAN_SYSTEMS: + assert Q.shape == (20, 20) and q.shape == (20,) + + +@requires_jax +@pytest.mark.parametrize("jit", [False, True], ids=["eager", "jit"]) +@pytest.mark.parametrize("index", range(len(GRAD_NAN_SYSTEMS)), ids=GRAD_NAN_IDS) +def test__raw_mode_gradient_is_finite_on_the_captured_grad_nan_systems(jnp, index, jit): + """Red on PyAutoArray 3de624b5 (#572): the gradient is NaN on all four systems, eager and jitted.""" + import jax + + Q, q = GRAD_NAN_SYSTEMS[index] + w = jnp.linspace(0.5, 1.5, q.shape[0]) + + grad = jax.grad(lambda Q_, q_: w @ _raw(jnp, Q_, q_), argnums=(0, 1)) + if jit: + grad = jax.jit(grad) + gQ, gq = grad(jnp.asarray(Q), jnp.asarray(q)) + + assert np.all(np.isfinite(np.asarray(gQ))) and np.all(np.isfinite(np.asarray(gq))) + assert np.any(np.asarray(gq) != 0.0) + + +def _backward_status(jnp, Q, q): + from autoarray.util.jax_nnls import raw_forward_backward_status + + Qj, qj = jnp.asarray(Q), jnp.asarray(q) + Q_pc, q_pc, D = (jnp.asarray(a) for a in _jacobi(Q, q)) + return [ + int(v) + for v in raw_forward_backward_status( + Q_pc, q_pc, Qj, qj, D, target_kappa=1.0e-11, max_iter=PRODUCTION_MAX_ITER + ) + ] + + +@requires_jax +@pytest.mark.parametrize( + "system", + [("slam", k) for k in KEYS] + + [("grad_nan", i) for i in range(len(GRAD_NAN_SYSTEMS))], + ids=[f"slam-{i}" for i in IDS] + [f"grad_nan-{i}" for i in GRAD_NAN_IDS], +) +def test__raw_mode_backward_pass_converges(jnp, system): + """The backward pass reports convergence: the tight polish of the mapped iterate converges (measured <= 6 + iterations) and the relaxed-KKT solve then converges well inside its 50-iteration cap (measured 1).""" + kind, index = system + Q, q = (SYSTEMS if kind == "slam" else GRAD_NAN_SYSTEMS)[index] + + relaxed_converged, relaxed_iter, polish_converged, polish_iter = _backward_status( + jnp, Q, q + ) + + assert polish_converged == 1, polish_iter + assert relaxed_converged == 1 and relaxed_iter < PRODUCTION_MAX_ITER, relaxed_iter From 36ac9b4f2121109185fed26a51b3928688416844 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Fri, 25 Sep 2026 17:14:48 +0100 Subject: [PATCH 2/4] fix: polish the raw-forward iterate before the relaxed-KKT backward solve (#573) The raw forward solve (#572) stops at the data-scaled tolerance, leaving s*z ~1e-10..2.5e-9 >> nnls_target_kappa=1e-11; the relaxed-KKT solve on the Jacobi system then diverged to NaN at its 50-iteration cap on 4/16 jax_grad/mge.py points. The backward pass now runs <= 10 tight PDIP iterations on (Q_pc, q_pc) warm-started from the mapped iterate (kept only if converged), after which the relaxed solve converges in ~1 iteration. The primal / forward value is unchanged. solve_nnls gains an init=(x, s, z) warm start; the relaxed solve's converged flag is kept and exposed with the polish status through the new raw_forward_backward_status diagnostic. general.yaml comments updated (values unchanged). Co-Authored-By: Claude Opus 5.5 (1M context) --- autoarray/config/general.yaml | 4 +- autoarray/util/jax_nnls.py | 124 +++++++++++++++++++++++++++------- 2 files changed, 103 insertions(+), 25 deletions(-) diff --git a/autoarray/config/general.yaml b/autoarray/config/general.yaml index e6ba12033..5df05d173 100644 --- a/autoarray/config/general.yaml +++ b/autoarray/config/general.yaml @@ -7,8 +7,8 @@ inversion: no_regularization_add_to_curvature_diag_value : 1.0e-3 # The default value added to the curvature matrix's diagonal when regularization is not applied to a linear object, which prevents inversion's failing due to the matrix being singular. use_border_relocator: false # If True, by default a pixelization's border is used to relocate all pixels outside its border to the border. nnls_jacobi_preconditioning: true # If True (default), the curvature matrix passed to jaxnnls.solve_nnls_primal is Jacobi-preconditioned (D Q D y = D q, x = D y). Fixes NaN backward-pass gradients on ill-conditioned Q and roughly halves forward solve time. Set False to restore the raw unpreconditioned solve. - nnls_target_kappa: 1.0e-11 # Central-path relaxation parameter passed to jaxnnls.solve_nnls_primal. Larger values smooth the relaxed-KKT backward pass and prevent NaN gradients on ill-conditioned Q; smaller values tighten the primal solve. Verified finite gradients across all MGE/rectangular/delaunay pipelines (imaging + interferometer) with scale invariance over 5 orders of magnitude in noise. jaxnnls's own default (1e-3) is too aggressive for the backward pass. - nnls_preconditioning_no_mapper: raw # How the JAX positive-only PDIP solve scales inversions with NO mapper (linear light profiles / MGE only). "raw" (default) runs the forward solve on the un-preconditioned system with a data-scaled KKT tolerance (1e-2 * n * eps * max(1, max|data_vector|)) and keeps the Jacobi-space relaxed-KKT gradient; "jacobi" uses the Jacobi-preconditioned solve. Jacobi scaling of signal-free MGE columns (diagonal = the no-regularization floor) made the PDIP dual diverge on 14/48 SLaM source_lp[1] points (PyAutoArray#571). Inversions with a mapper always use jacobi; the NumPy path is unaffected. + nnls_target_kappa: 1.0e-11 # Central-path relaxation parameter passed to jaxnnls.solve_nnls_primal. Larger values smooth the relaxed-KKT backward pass and prevent NaN gradients on ill-conditioned Q; smaller values tighten the primal solve. Verified finite gradients across all MGE/rectangular/delaunay pipelines (imaging + interferometer) with scale invariance over 5 orders of magnitude in noise. jaxnnls's own default (1e-3) is too aggressive for the backward pass. The relaxed solve must start from an iterate whose complementarity s*z is not far above this value; the "raw" no-mapper mode polishes its forward iterate to ensure that (PyAutoArray#573). + nnls_preconditioning_no_mapper: raw # How the JAX positive-only PDIP solve scales inversions with NO mapper (linear light profiles / MGE only). "raw" (default) runs the forward solve on the un-preconditioned system with a data-scaled KKT tolerance (1e-2 * n * eps * max(1, max|data_vector|)) and keeps the Jacobi-space relaxed-KKT gradient, whose relaxed solve starts from the forward iterate polished by <= 10 tight warm-started PDIP iterations on the Jacobi system (without the polish the loose forward tolerance leaves s*z far above nnls_target_kappa and the relaxed solve diverged to NaN gradients on 4/16 jax_grad/mge.py points, PyAutoArray#573); "jacobi" uses the Jacobi-preconditioned solve. Jacobi scaling of signal-free MGE columns (diagonal = the no-regularization floor) made the PDIP dual diverge on 14/48 SLaM source_lp[1] points (PyAutoArray#571). Inversions with a mapper always use jacobi; the NumPy path is unaffected. nnls_warm_start_memo: true # If True (default), the NumPy/numba positive-only (fnnls) solve warm-starts its active set from the previous likelihood evaluation's passive set, cutting active-set iterations on successive sampler evaluations. The NNLS optimum is unique so the reconstruction is unchanged. On by default as of PyAutoArray#498, measured on the euclid+hst Delaunay-1250 fiducial (9.9x / 4.0x fewer active-set iterations on successive evaluations, reconstruction unchanged). Set false, or AUTOARRAY_NNLS_WARM_START=0, to disable. JAX path unaffected. nnls_warm_start_error_tolerance: 1.5 # Relative quality guard on a warm-start memo seed. Each memo entry remembers the error fraction of the most recent dense-sign-started solve for that key; a seeded solve whose own error fraction exceeds this multiple of that reference is dropped, so the next solve restarts from the dense-sign start and refreshes the reference. Default 1.5 sits above the worst seed/dense error-fraction ratio seen in the PyAutoArray#498 32-cell robustness matrix (1.42), so it is protective against unmeasured regimes rather than flapping. Any non-finite or non-positive value (e.g. .inf) disables the guard. NumPy/numba fnnls path only. positive_only_solver: pdip # Which solver the JAX (xp=jnp) positive-only reconstruction uses. "pdip" (default) is the jaxnnls interior-point solve; "certified" is the certified active-set solve (budgeted masked-Cholesky passes that stop once the KKT conditions certify, exact implicit gradient, PDIP fallback), measured 1.2-2.6x faster on source-only inversions (PyAutoArray#566). Applied only to mapper-only JAX inversions (MGE / linear light profiles keep PDIP); the NumPy path always uses fnnls. Opt-in until the batched (vmap) policy is measured. diff --git a/autoarray/util/jax_nnls.py b/autoarray/util/jax_nnls.py index 3748f1bce..df21160a0 100644 --- a/autoarray/util/jax_nnls.py +++ b/autoarray/util/jax_nnls.py @@ -27,7 +27,11 @@ solve on the un-preconditioned system with a data-scaled tolerance (:func:`data_scaled_solver_tol`) and keeps the Jacobi-space backward pass. It exists because Jacobi scaling of signal-free MGE columns (diagonal = the -no-regularization floor) makes the PDIP dual diverge. +no-regularization floor) makes the PDIP dual diverge. Its backward pass polishes +the mapped forward iterate with a few tight, warm-started PDIP iterations on the +Jacobi system before the relaxed-KKT solve, which otherwise diverges to NaN from +the loose forward tolerance (PyAutoArray#573); :func:`raw_forward_backward_status` +reports that pass's convergence. JAX is imported inside functions, never at module level (see ``docs/agents/jax_and_decorators.md``); this module must only be imported @@ -39,7 +43,7 @@ from functools import lru_cache -def solve_nnls(Q, q, solver_tol=None, max_iter=50): +def solve_nnls(Q, q, solver_tol=None, max_iter=50, init=None): """ Solve the non-negative least squares problem with the jaxnnls PDIP algorithm, with configurable convergence tolerance and iteration cap. @@ -59,6 +63,10 @@ def solve_nnls(Q, q, solver_tol=None, max_iter=50): ``min(n * eps * 5e3, 1e-2)``. max_iter Maximum number of PDIP iterations (jaxnnls hard-codes 50). + init + Optional ``(x, s, z)`` warm start (strictly positive ``s`` and ``z``) + replacing jaxnnls's ``initialize``. ``None`` (default) is the upstream + cold start. Returns ------- @@ -69,7 +77,7 @@ def solve_nnls(Q, q, solver_tol=None, max_iter=50): import jax.numpy as jnp from jaxnnls.pdip import EPSILON, initialize, pdip_pc_step - x, s, z = initialize(Q, q) + x, s, z = initialize(Q, q) if init is None else init if solver_tol is None: solver_tol = jax.lax.min(Q.shape[0] * EPSILON, 1e-2) @@ -179,10 +187,57 @@ def solve_nnls_primal(Q, q, target_kappa=1e-3, solver_tol=None, max_iter=50): )[0] +# The backward pass of the ``"raw"`` mode first polishes the mapped raw-forward iterate with at most this many +# PDIP iterations on the Jacobi-scaled system at jaxnnls's own tight tolerance (PyAutoArray#573). Measured on the +# SLaM MGE fixture, the 48 SLaM ``source_lp[1]`` systems and the jax_grad/mge.py points: 4-6 iterations. +RAW_BACKWARD_POLISH_MAX_ITER = 10 + + +def _raw_forward_backward_point( + Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter +): + """ + The forward solve and the relaxed-KKT point of the ``"raw"`` mode (shared by + the custom-vjp forward pass and :func:`raw_forward_backward_status`). + + Returns ``(y, converged, pdip_iter)`` of the raw forward solve (mapped to the + Jacobi coordinates), the relaxed point ``(yr, sr, zr)`` the backward pass + differentiates at, and the status ``(relaxed_converged, relaxed_iter, + polish_converged, polish_iter)``. + """ + import jax.numpy as jnp + from jaxnnls.pdip_relaxed import solve_relaxed_nnls + + tol = data_scaled_solver_tol(q) if solver_tol is None else solver_tol + x, s, z, converged, pdip_iter = solve_nnls(Q, q, solver_tol=tol, max_iter=max_iter) + y, sy, zy = x / D, s / D, z * D + + # Polish (PyAutoArray#573): the data-scaled tolerance leaves s * z ~ 1e-10 .. 1e-9, far above + # ``target_kappa``, so the relaxed solve below would have to push toward the boundary from z / s ~ 1e13 + # and its fixed 50-iteration while_loop overshoots to NaN. A few tight PDIP iterations on the scaled + # system, warm-started from the mapped iterate, bring s * z down to the jaxnnls tolerance first. If the + # polish does not converge (the scaled dual is what diverges on #571's systems from a cold start), the + # mapped iterate is kept, i.e. the pre-polish behaviour. + yp, sp, zp, polish_converged, polish_iter = solve_nnls( + Q_pc, q_pc, max_iter=RAW_BACKWARD_POLISH_MAX_ITER, init=(y, sy, zy) + ) + ok = jnp.logical_and( + polish_converged == 1, + jnp.all(jnp.isfinite(yp)) & jnp.all(sp > 0) & jnp.all(zp > 0), + ) + yp, sp, zp = (jnp.where(ok, a, b) for a, b in ((yp, y), (sp, sy), (zp, zy))) + + yr, sr, zr, relaxed_converged, relaxed_iter = solve_relaxed_nnls( + Q_pc, q_pc, yp, sp, zp, target_kappa=target_kappa + ) + status = (relaxed_converged, relaxed_iter, ok.astype(int), polish_iter) + return (y, converged, pdip_iter), (yr, sr, zr), status + + @lru_cache(maxsize=None) def _solve_nnls_raw_forward_with(target_kappa, solver_tol, max_iter): """ - Build (and cache) the ``"raw"``-mode solver (PyAutoArray#571). + Build (and cache) the ``"raw"``-mode solver (PyAutoArray#571, #573). The returned function takes the Jacobi-scaled system ``(Q_pc, q_pc)`` (``Q_pc = D Q D``, ``q_pc = D q``) together with the raw system ``(Q, q)`` @@ -196,38 +251,43 @@ def _solve_nnls_raw_forward_with(target_kappa, solver_tol, max_iter): linear-object-only (MGE) systems, Jacobi scaling turns signal-free columns whose diagonal is only the no-regularization floor into degenerate coordinates that make the PDIP dual diverge; the raw solve does not. - - **Backward:** exactly today's Jacobi-mode pass, i.e. the relaxed-KKT implicit - derivative on ``Q_pc``, started from the mapped iterate. The relaxed-KKT - pass on the raw, ill-conditioned ``Q`` produces NaN gradients, which is why - Jacobi scaling was introduced. ``(Q, q, D)`` get zero cotangents: ``y`` - depends only on ``(Q_pc, q_pc)``, and the caller's autodiff carries the - dependence of those, and of ``D``, on the raw inputs. + - **Backward:** the relaxed-KKT implicit derivative on ``Q_pc`` (as the + Jacobi mode), started from the mapped iterate after a *polish*: at most + :data:`RAW_BACKWARD_POLISH_MAX_ITER` PDIP iterations on ``(Q_pc, q_pc)`` + at jaxnnls's tight tolerance, warm-started from the mapped iterate + (kept only if it converges). Without it the loose forward tolerance + leaves complementarity ``s * z`` orders of magnitude above + ``target_kappa`` and the relaxed solve diverges to NaN on a fraction of + points (PyAutoArray#573); with it the relaxed solve converges in about + one iteration. The primal ``y`` is the unpolished forward solution, so + the forward value is unchanged. The relaxed-KKT pass on the raw, + ill-conditioned ``Q`` produces NaN gradients, which is why Jacobi scaling + was introduced. ``(Q, q, D)`` get zero cotangents: ``y`` depends only on + ``(Q_pc, q_pc)``, and the caller's autodiff carries the dependence of + those, and of ``D``, on the raw inputs. + + The backward-pass convergence is observable through + :func:`raw_forward_backward_status`. """ import jax import jax.numpy as jnp from jaxnnls.diff_qp import diff_nnls - from jaxnnls.pdip_relaxed import solve_relaxed_nnls - def raw_solve(Q, q, D): + def primal(Q_pc, q_pc, Q, q, D): tol = data_scaled_solver_tol(q) if solver_tol is None else solver_tol - x, s, z, converged, pdip_iter = solve_nnls( + x, _, _, converged, pdip_iter = solve_nnls( Q, q, solver_tol=tol, max_iter=max_iter ) - return x / D, s / D, z * D, converged, pdip_iter - - def primal(Q_pc, q_pc, Q, q, D): - y, _, _, converged, pdip_iter = raw_solve(Q, q, D) - return y, converged, pdip_iter + return x / D, converged, pdip_iter def forward(Q_pc, q_pc, Q, q, D): - y, sy, zy, converged, pdip_iter = raw_solve(Q, q, D) - yr, sr, zr, _, _ = solve_relaxed_nnls( - Q_pc, q_pc, y, sy, zy, target_kappa=target_kappa + out, (yr, sr, zr), status = _raw_forward_backward_point( + Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter ) - return (y, converged, pdip_iter), (Q_pc, yr, sr, zr, Q, q, D) + return out, (Q_pc, yr, sr, zr, status[0], Q, q, D) def backward(res, output_grad): - Q_pc, yr, sr, zr, Q, q, D = res + Q_pc, yr, sr, zr, _, Q, q, D = res dQ_pc, dq_pc = diff_nnls(Q_pc, yr, sr, zr, output_grad[0]) return dQ_pc, dq_pc, jnp.zeros_like(Q), jnp.zeros_like(q), jnp.zeros_like(D) @@ -236,6 +296,24 @@ def backward(res, output_grad): return primal +def raw_forward_backward_status( + Q_pc, q_pc, Q, q, D, target_kappa=1e-3, solver_tol=None, max_iter=50 +): + """ + Diagnostic (not differentiable): the convergence of the ``"raw"`` mode's + backward-pass preparation for one system, as the integer tuple + ``(relaxed_converged, relaxed_iter, polish_converged, polish_iter)``. + + ``relaxed_*`` describe the relaxed-KKT solve whose point the gradient is + taken at; ``polish_*`` the tight warm-started PDIP polish before it + (``polish_converged == 0`` means the mapped iterate was used unpolished). + Arguments are those of :func:`solve_nnls_primal_raw_forward`. + """ + return _raw_forward_backward_point( + Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter + )[2] + + def solve_nnls_primal_raw_forward( Q_pc, q_pc, Q, q, D, target_kappa=1e-3, solver_tol=None, max_iter=50 ): From ffdc0ae414d2ef74f0cdce3c65b63ada94cde9b9 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Fri, 25 Sep 2026 17:29:31 +0100 Subject: [PATCH 3/4] style: black the #573 regression tests Co-Authored-By: Claude Opus 5.5 (1M context) --- .../inversion/inversion/test_nnls_mge_convergence.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py index f2a0da030..bb384f2a2 100644 --- a/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py +++ b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py @@ -401,7 +401,8 @@ def _backward_status(jnp, Q, q): ) def test__raw_mode_backward_pass_converges(jnp, system): """The backward pass reports convergence: the tight polish of the mapped iterate converges (measured <= 6 - iterations) and the relaxed-KKT solve then converges well inside its 50-iteration cap (measured 1).""" + iterations) and the relaxed-KKT solve then converges well inside its 50-iteration cap (measured 1). + """ kind, index = system Q, q = (SYSTEMS if kind == "slam" else GRAD_NAN_SYSTEMS)[index] From 44762fad5f22f407af160c8afba6b02ad25ef5a4 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Fri, 25 Sep 2026 19:12:39 +0100 Subject: [PATCH 4/4] test: state the #573 red cases precisely in the regression docstrings Review finding: on 3de624b5 the jitted gradient is NaN only on prng10/prng14 (prng2/prng12 pass jitted), and the backward-status test fails on import on main rather than on convergence, so it is not the red witness. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_019jDFQSNoi3ihaeM7ZJhfYL --- .../inversion/inversion/test_nnls_mge_convergence.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py index bb384f2a2..1bfacd24f 100644 --- a/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py +++ b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py @@ -364,7 +364,9 @@ def test__grad_nan_fixture_is_the_captured_jax_grad_mge_set(): @pytest.mark.parametrize("jit", [False, True], ids=["eager", "jit"]) @pytest.mark.parametrize("index", range(len(GRAD_NAN_SYSTEMS)), ids=GRAD_NAN_IDS) def test__raw_mode_gradient_is_finite_on_the_captured_grad_nan_systems(jnp, index, jit): - """Red on PyAutoArray 3de624b5 (#572): the gradient is NaN on all four systems, eager and jitted.""" + """Red on PyAutoArray 3de624b5 (#572): the gradient is NaN on all four systems eagerly, and under jit on + prng10 / prng14 (the jitted NaN is rounding-sensitive; prng2 / prng12 happen to pass jitted on main). + """ import jax Q, q = GRAD_NAN_SYSTEMS[index] @@ -402,6 +404,7 @@ def _backward_status(jnp, Q, q): def test__raw_mode_backward_pass_converges(jnp, system): """The backward pass reports convergence: the tight polish of the mapped iterate converges (measured <= 6 iterations) and the relaxed-KKT solve then converges well inside its 50-iteration cap (measured 1). + Not a red-on-main witness (``raw_forward_backward_status`` is new with #573); the gradient test is. """ kind, index = system Q, q = (SYSTEMS if kind == "slam" else GRAD_NAN_SYSTEMS)[index]