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 ): 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 000000000..7c71084fd Binary files /dev/null and b/test_autoarray/inversion/inversion/files/mge_grad_nan_systems.npz differ diff --git a/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py index 0d34d2632..1bfacd24f 100644 --- a/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py +++ b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py @@ -323,3 +323,95 @@ 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 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] + 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). + 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] + + 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