diff --git a/autoarray/config/general.yaml b/autoarray/config/general.yaml index 5df05d173..916e85045 100644 --- a/autoarray/config/general.yaml +++ b/autoarray/config/general.yaml @@ -8,7 +8,7 @@ inversion: 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. 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_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|)), polishes that iterate with <= 10 tight warm-started PDIP iterations on the Jacobi system and returns the polished iterate (the loose tolerance also judges complementarity, so unpolished it left 11.5 % of the fnnls flux on inactive columns of the euclid vis_lp system; polished <= 3.3e-4 over 81 corpus systems, ~5 extra iterations, PyAutoArray#594), and keeps the Jacobi-space relaxed-KKT gradient started from the same polished point (without the polish 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/inversion/inversion/inversion_util.py b/autoarray/inversion/inversion/inversion_util.py index d94668707..76454a552 100644 --- a/autoarray/inversion/inversion/inversion_util.py +++ b/autoarray/inversion/inversion/inversion_util.py @@ -368,7 +368,9 @@ def reconstruction_positive_only_from( byte-identical to before this option existed): the Jacobi-preconditioned solve ``(D Q D) y = D q`` governed by the ``nnls_jacobi_preconditioning`` config key. ``"raw"``: the forward PDIP solve runs on the un-preconditioned ``(Q, q)`` with the data-scaled tolerance - :func:`autoarray.util.jax_nnls.data_scaled_solver_tol` (or ``settings.nnls_solver_tol`` if set), and the + :func:`autoarray.util.jax_nnls.data_scaled_solver_tol` (or ``settings.nnls_solver_tol`` if set), its iterate + is polished by a few tight PDIP iterations on the Jacobi-scaled system and the polished iterate is returned + (PyAutoArray#594; ``stats["converged"]`` / ``stats["iterations"]`` describe the raw forward solve), and the gradient is the Jacobi-space relaxed-KKT pass as in ``"jacobi"`` -- see :func:`autoarray.util.jax_nnls.solve_nnls_primal_raw_forward`. Jacobi scaling makes the signal-free columns of linear-object-only (MGE) inversions, whose diagonal is only the diff --git a/autoarray/settings.py b/autoarray/settings.py index ef18c180c..54bf6360d 100644 --- a/autoarray/settings.py +++ b/autoarray/settings.py @@ -230,7 +230,11 @@ def __init__( - ``"raw"`` (default) -- the forward PDIP solve runs on the un-preconditioned system with a data-scaled KKT tolerance (``1e-2 * n * eps_pdip * max(1, max|data_vector|)``, or - `nnls_solver_tol` if set); the gradient is the same Jacobi-space relaxed-KKT pass as ``"jacobi"``. + `nnls_solver_tol` if set), polished by at most 10 tight warm-started PDIP iterations on the + Jacobi-scaled system; the polished iterate is the returned reconstruction (unpolished, the loose + tolerance left 11.5 % of the fnnls flux on inactive columns of the euclid vis_lp system, + PyAutoArray#594) and the start of the gradient, the same Jacobi-space relaxed-KKT pass as + ``"jacobi"``. On the SLaM `source_lp[1]` MGE model (2 x 20 lens + 20 source Gaussians) Jacobi scaling made 14/48 near-truth points hit the 50-iteration cap with wrong log-likelihoods (signal-free Gaussian columns, whose diagonal is only `no_regularization_add_to_curvature_diag_value`, become degenerate diff --git a/autoarray/util/jax_nnls.py b/autoarray/util/jax_nnls.py index df21160a0..3d374434d 100644 --- a/autoarray/util/jax_nnls.py +++ b/autoarray/util/jax_nnls.py @@ -27,11 +27,18 @@ 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. 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. +no-regularization floor) makes the PDIP dual diverge. The raw iterate is then +*polished* with a few tight, warm-started PDIP iterations on the Jacobi system, and +the polished iterate is both the forward value and the start of the backward +pass's relaxed-KKT solve. The data-scaled stop judges complementarity ``s * z`` +against the same loose threshold, so a column fnnls holds at zero can keep +``x ~ tol / z`` while "converged": on the euclid vis_lp system the unpolished +iterate left 11.5 % of the fnnls flux on reference-inactive columns (a +5.8 % +latent source flux), invisible to the objective and the log likelihood; polished, +that is <= 3.3e-4 across 81 corpus systems for ~5 extra iterations +(PyAutoArray#594). The polish first went into the backward pass only, where the +relaxed-KKT solve otherwise diverges to NaN from the loose forward tolerance +(PyAutoArray#573); :func:`raw_forward_backward_status` reports its convergence. JAX is imported inside functions, never at module level (see ``docs/agents/jax_and_decorators.md``); this module must only be imported @@ -187,51 +194,75 @@ 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 +# The ``"raw"`` mode polishes the mapped raw-forward iterate with at most this many PDIP iterations on the +# Jacobi-scaled system at jaxnnls's own tight tolerance; the polished iterate is the forward value (PyAutoArray#594) +# and the backward pass's starting point (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_POLISH_MAX_ITER = 10 +# The name under which the polish cap was introduced (#573, backward pass only); kept for external importers. +RAW_BACKWARD_POLISH_MAX_ITER = RAW_POLISH_MAX_ITER -def _raw_forward_backward_point( - Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter -): +def _raw_forward_polished(Q_pc, q_pc, Q, q, D, 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)``. + The forward value of the ``"raw"`` mode: the raw PDIP solve, mapped to the + Jacobi coordinates and polished there (shared by the custom-vjp primal and + forward passes, so every call path returns the same ``y`` bit-for-bit). + + Returns ``(y_out, converged, pdip_iter, (yp, sp, zp), ok, polish_iter)``: + ``y_out`` the forward value (the polished iterate if ``ok``, else the mapped + raw iterate), ``converged`` / ``pdip_iter`` the raw forward solve's flag and + iteration count, ``(yp, sp, zp)`` the point the backward pass starts from + (selected by the same ``ok``), ``ok`` whether the polish was accepted and + ``polish_iter`` its iteration count. """ 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. + # Polish (PyAutoArray#573, #594): the data-scaled tolerance also judges complementarity, so it leaves + # s * z ~ 1e-10 .. 1e-9 -- a column fnnls holds at zero can keep x ~ tol / z (11.5 % of the reference flux on + # inactive columns of the euclid vis_lp system), and the relaxed solve of the backward pass would have to push + # toward the boundary from z / s ~ 1e13, where 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. 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) + Q_pc, q_pc, max_iter=RAW_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))) + return yp, converged, pdip_iter, (yp, sp, zp), ok, polish_iter + + +def _raw_forward_backward_point( + Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter +): + """ + The forward value 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)`` with ``y`` the polished forward value + (:func:`_raw_forward_polished`) and ``converged`` / ``pdip_iter`` those of the + raw forward solve, the relaxed point ``(yr, sr, zr)`` the backward pass + differentiates at, and the status ``(relaxed_converged, relaxed_iter, + polish_converged, polish_iter)``. + """ + from jaxnnls.pdip_relaxed import solve_relaxed_nnls + + y_out, converged, pdip_iter, (yp, sp, zp), ok, polish_iter = _raw_forward_polished( + Q_pc, q_pc, Q, q, D, solver_tol, max_iter + ) 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 + return (y_out, converged, pdip_iter), (yr, sr, zr), status @lru_cache(maxsize=None) @@ -251,16 +282,25 @@ 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. + The mapped iterate is then *polished*: at most + :data:`RAW_POLISH_MAX_ITER` PDIP iterations on ``(Q_pc, q_pc)`` at + jaxnnls's tight tolerance, warm-started from it (kept only if the polish + converges with a finite, strictly interior point). The polished iterate is + the returned ``y`` (PyAutoArray#594): the data-scaled stop judges + complementarity ``s * z`` against the loose tolerance, so an unpolished + column fnnls holds at zero can keep ``x ~ tol / z`` -- 11.5 % of the fnnls + flux on inactive columns of the euclid vis_lp system, invisible to the + objective and log likelihood; polished, <= 3.3e-4 over 81 corpus systems + for ~5 extra iterations. The custom-vjp primal and forward rule compute it + through the same :func:`_raw_forward_polished`, so plain, jitted and + differentiated calls return the same value. ``converged`` and + ``pdip_iter`` are those of the raw forward solve. - **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 + Jacobi mode), started from the polished point. Without the polish the + loose forward tolerance leaves ``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, + one iteration. 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 @@ -274,11 +314,10 @@ def _solve_nnls_raw_forward_with(target_kappa, solver_tol, max_iter): from jaxnnls.diff_qp import diff_nnls def primal(Q_pc, q_pc, Q, q, D): - tol = data_scaled_solver_tol(q) if solver_tol is None else solver_tol - x, _, _, converged, pdip_iter = solve_nnls( - Q, q, solver_tol=tol, max_iter=max_iter + y_out, converged, pdip_iter, _, _, _ = _raw_forward_polished( + Q_pc, q_pc, Q, q, D, solver_tol, max_iter ) - return x / D, converged, pdip_iter + return y_out, converged, pdip_iter def forward(Q_pc, q_pc, Q, q, D): out, (yr, sr, zr), status = _raw_forward_backward_point( @@ -305,8 +344,9 @@ def raw_forward_backward_status( ``(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). + taken at; ``polish_*`` the tight warm-started PDIP polish before it, + whose iterate is also the forward value (``polish_converged == 0`` means the + mapped raw iterate was used unpolished, for both). Arguments are those of :func:`solve_nnls_primal_raw_forward`. """ return _raw_forward_backward_point( @@ -319,7 +359,8 @@ def solve_nnls_primal_raw_forward( ): """ The ``"raw"`` positive-only mode: forward PDIP on the raw system with a - data-scaled tolerance, backward pass on the Jacobi-scaled system. Returns + data-scaled tolerance, polished on the Jacobi-scaled system, backward pass + on the Jacobi-scaled system. Returns ``(y, converged, pdip_iter)`` with ``x = D * y``; see :func:`_solve_nnls_raw_forward_with`. """ diff --git a/test_autoarray/inversion/inversion/files/README.md b/test_autoarray/inversion/inversion/files/README.md index 77a5bbc0f..dbf386900 100644 --- a/test_autoarray/inversion/inversion/files/README.md +++ b/test_autoarray/inversion/inversion/files/README.md @@ -15,3 +15,13 @@ 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. +- `mge_solver_reference_systems.npz` — 9 positive-only MGE systems (60x60 fp64) with their fnnls reference + solutions, for the raw-forward PDIP amplitude regression (PyAutoArray#594): keys `Q_` / `q_` / + `x_ref_` for names `k0`..`k7` (the 8 #571 SLaM `source_lp[1]` systems, identical to + `mge_slam_nnls_systems.npz`) and `euclid_vis_lp` (the euclid pipeline vis_lp system behind + `test_latent_euclid_variables_traces_under_jax_jit`, captured with PyAutoArray d4298445). `meta` is a JSON + string with per-system `group`, `source_column_index_list`, `max_abs_q`, `cond_Q`, `category` and reference + active-column counts, plus provenance. Copied verbatim (no recomputation) from the autolens_profiling solver + corpus `results/lens/solver/corpus/{slam_fixture_571,euclid_vis_lp}.npz` + `manifest.json` at + autolens_profiling 3ad68af (phase-1 record `complete/2026/09/linear-solver-accuracy-study.md`) by a one-off + script outside the repo that checks symmetry of `Q`, `x_ref >= 0` and `max|q|` against the manifest. diff --git a/test_autoarray/inversion/inversion/files/mge_solver_reference_systems.npz b/test_autoarray/inversion/inversion/files/mge_solver_reference_systems.npz new file mode 100644 index 000000000..0bda72980 Binary files /dev/null and b/test_autoarray/inversion/inversion/files/mge_solver_reference_systems.npz differ diff --git a/test_autoarray/inversion/inversion/test_nnls_raw_forward_amplitude.py b/test_autoarray/inversion/inversion/test_nnls_raw_forward_amplitude.py new file mode 100644 index 000000000..bcfb0fcae --- /dev/null +++ b/test_autoarray/inversion/inversion/test_nnls_raw_forward_amplitude.py @@ -0,0 +1,310 @@ +""" +Amplitude regression for the forward value of the ``"raw"`` positive-only PDIP mode (PyAutoArray#594). + +The mapper-less JAX positive-only solve (``preconditioning="raw"``, PyAutoArray#571/#572) stops on an absolute +infinity-norm KKT residual with the data-scaled tolerance ``1e-2 * n * EPSILON * max(1, max|q|)``. The +complementarity ``s * z`` is judged against that same threshold, so a column fnnls holds at zero can keep +``x ~ tol / z`` while the solve reports "converged": the objective, the log likelihood and the KKT residual are +all blind to it, but the amplitudes (and latent fluxes derived from them) are not. The phase-1 accuracy study +(autolens_profiling #355, record ``complete/2026/09/linear-solver-accuracy-study.md``) measured it on the euclid +vis_lp system as 11.5 % of the reference flux left on reference-inactive columns (``total_source_flux`` ++5.76 %, the red euclid-pipeline latent jit-vs-eager test), and >1e-3 source-flux errors on 4/8 #571 systems. + +The fixture ``files/mge_solver_reference_systems.npz`` holds the 8 #571 SLaM systems (``k0``..``k7``) and the +euclid vis_lp system with their fnnls reference solutions ``x_ref`` (see ``files/README.md``). Each system is +solved through the library entry point :func:`autoarray.util.jax_nnls.solve_nnls_primal_raw_forward` (built as +``inversion_util.reconstruction_positive_only_from`` builds it) and through +``reconstruction_positive_only_from(..., preconditioning="raw")`` itself, and must satisfy: + +- ``converged == 1``, ``iterations < 50``, finite; +- ``flux_inactive_rel = sum(x[x_ref <= 1e-6 max x_ref]) / sum(x_ref) <= 1e-3``; +- ``|flux_rel_all| = |sum x - sum x_ref| / |sum x_ref| <= 1e-3``; +- ``|flux_rel_source|`` (the same over ``source_column_index_list``) ``<= 1e-3`` on ``k0``..``k7`` only. The + euclid system is excluded from this one metric: its source columns carry only ~0.4 % of the reference flux + (a single active column), so the relative source flux is noise-dominated (+5e-2 even for the polished + solve); its end-to-end latent is covered by the euclid-pipeline test (polished: +7.5e-5); +- ``jax.jit`` of the reconstruction equals the eager value exactly, the custom_vjp primal equals its + differentiated forward value exactly (eager and jitted), and ``jax.grad`` of ``sum(y)`` with + respect to ``q`` is finite and non-zero (a mis-wired custom_vjp returns zeros silently). + +Thresholds were fixed from the phase-1 rows before this test was written. Raw forward (unpolished): euclid +inactive 0.115, all 0.115 (others inactive <= 5.4e-5); source k1 1.64e-3, k2 1.22e-3, k3 4.04e-3, k5 6.44e-3 +(k0 5.2e-4, k4 6.7e-4, k6 9.2e-4, k7 9.5e-4). Polished forward: inactive <= 3.3e-4 (euclid), source <= 5e-5 +on k0..k7, all <= 3.3e-4. The amplitude-max criterion of the task prompt is dropped: its worst value (0.219) +is shared by every accurate candidate, i.e. it measures fnnls noise on near-flat directions. + +Red on the unfixed base (PyAutoArray bd03e09e, whose solver equals d4298445): euclid ``inactive`` and ``all``, +and ``source`` on k1, k2, k3, k5 -- on both paths. The jit == eager, primal == forward and gradient cases are green there too (the base +returns the unpolished ``y`` from both the primal and the fwd rule). +""" + +import importlib.util +import json +from pathlib import Path + +import numpy as np +import pytest + +import autoarray as aa +from autoarray.inversion.inversion import inversion_util + + +requires_jax = pytest.mark.skipif( + importlib.util.find_spec("jax") is None, + reason="requires jax (installed via the [optional] extras; absent on the NumPy-only matrix env)", +) + +FIXTURE = Path(__file__).parent / "files" / "mge_solver_reference_systems.npz" +PRODUCTION_MAX_ITER = 50 +TARGET_KAPPA = ( + 1.0e-11 # autoarray general.yaml ``nnls_target_kappa``, as the dispatch reads it +) +INACTIVE_REL = 1.0e-6 +FLUX_TOL = 1.0e-3 + + +def _load_systems(): + with np.load(FIXTURE) as data: + meta = json.loads(str(data["meta"])) + systems = { + s["name"]: tuple( + np.asarray(data[f"{p}_{s['name']}"]) for p in ("Q", "q", "x_ref") + ) + for s in meta["systems"] + } + return meta, systems + + +META, SYSTEMS = _load_systems() +SYSTEM_META = {s["name"]: s for s in META["systems"]} +NAMES = list(SYSTEM_META) +SLAM_NAMES = [n for n in NAMES if SYSTEM_META[n]["group"] == "slam_fixture_571"] +PATHS = ["library_entry", "reconstruction_positive_only_from"] + + +@pytest.fixture(scope="module") +def jnp(): + import jax + + jax.config.update("jax_enable_x64", True) + + import jax.numpy as jnp + from jaxnnls.pdip import EPSILON + + # jaxnnls fixes its tolerance scale at import time from the default dtype; a float32-era import would + # loosen every tolerance and hide the bias. + assert EPSILON < 1.0e-10, EPSILON + + return jnp + + +def _entry(jnp, Q, q): + """`solve_nnls_primal_raw_forward`, built exactly as the raw branch of `reconstruction_positive_only_from`.""" + from autoarray.util.jax_nnls import solve_nnls_primal_raw_forward + + d = jnp.sqrt(jnp.diag(Q)) + D = 1.0 / d + Q_pc = (Q * D[:, None]) * D[None, :] + q_pc = q * D + y, converged, iterations = solve_nnls_primal_raw_forward( + Q_pc, + q_pc, + Q, + q, + D, + target_kappa=TARGET_KAPPA, + solver_tol=None, + max_iter=PRODUCTION_MAX_ITER, + ) + return y, D, converged, iterations + + +def _entry_solver(jnp, Q, q): + """The ``y`` output of `solve_nnls_primal_raw_forward` as a function of its five inputs, and those inputs.""" + from autoarray.util.jax_nnls import solve_nnls_primal_raw_forward + + D = 1.0 / jnp.sqrt(jnp.diag(Q)) + args = ((Q * D[:, None]) * D[None, :], q * D, Q, q, D) + + def solve(*a): + return solve_nnls_primal_raw_forward( + *a, target_kappa=TARGET_KAPPA, solver_tol=None, max_iter=PRODUCTION_MAX_ITER + )[0] + + return solve, args + + +def _dispatch(jnp, Q, q, stats=None): + return inversion_util.reconstruction_positive_only_from( + data_vector=q, + curvature_reg_matrix=Q, + settings=aa.Settings(), + xp=jnp, + stats=stats, + preconditioning="raw", + solver="pdip", + ) + + +_SOLVED = {} + + +def _solve(jnp, path, name): + """(x, converged, iterations) for one system through one path, cached per module.""" + if (path, name) not in _SOLVED: + Q, q, _ = SYSTEMS[name] + Qj, qj = jnp.asarray(Q), jnp.asarray(q) + if path == "library_entry": + y, D, converged, iterations = _entry(jnp, Qj, qj) + x = y * D + else: + stats = {} + x = _dispatch(jnp, Qj, qj, stats=stats) + assert stats["solver"] == "pdip" and stats["preconditioning"] == "raw" + converged, iterations = stats["converged"], stats["iterations"] + _SOLVED[(path, name)] = (np.asarray(x), int(converged), int(iterations)) + return _SOLVED[(path, name)] + + +def _flux_inactive_rel(x, x_ref): + inactive = x_ref <= INACTIVE_REL * x_ref.max() + return float(np.sum(x[inactive]) / np.sum(x_ref)) + + +def _flux_rel(x, x_ref, index=slice(None)): + return float((np.sum(x[index]) - np.sum(x_ref[index])) / abs(np.sum(x_ref[index]))) + + +def test__fixture_is_the_reference_system_set(): + assert FIXTURE.stat().st_size < 400_000 + assert NAMES == [f"k{k}" for k in range(8)] + ["euclid_vis_lp"] + assert META["provenance"]["corpus_commit"] == "3ad68af" + for name, (Q, q, x_ref) in SYSTEMS.items(): + assert Q.shape == (60, 60) and q.shape == (60,) and x_ref.shape == (60,) + np.testing.assert_allclose(Q, Q.T, rtol=0, atol=1e-8 * np.abs(Q).max()) + assert np.all(x_ref >= 0.0) + assert np.isclose(np.abs(q).max(), SYSTEM_META[name]["max_abs_q"], rtol=1e-12) + + +@requires_jax +@pytest.mark.parametrize("name", NAMES) +@pytest.mark.parametrize("path", PATHS) +def test__raw_forward_converges(jnp, path, name): + x, converged, iterations = _solve(jnp, path, name) + + assert converged == 1, f"not converged ({iterations} iterations)" + assert iterations < PRODUCTION_MAX_ITER + assert np.all(np.isfinite(x)) + + +@requires_jax +@pytest.mark.parametrize("name", NAMES) +@pytest.mark.parametrize("path", PATHS) +def test__raw_forward_inactive_column_flux(jnp, path, name): + x, _, _ = _solve(jnp, path, name) + value = _flux_inactive_rel(x, SYSTEMS[name][2]) + + assert value <= FLUX_TOL, f"flux_inactive_rel = {value:.3e}" + + +@requires_jax +@pytest.mark.parametrize("name", NAMES) +@pytest.mark.parametrize("path", PATHS) +def test__raw_forward_total_flux(jnp, path, name): + x, _, _ = _solve(jnp, path, name) + value = _flux_rel(x, SYSTEMS[name][2]) + + assert abs(value) <= FLUX_TOL, f"flux_rel_all = {value:.3e}" + + +@requires_jax +@pytest.mark.parametrize("name", SLAM_NAMES) +@pytest.mark.parametrize("path", PATHS) +def test__raw_forward_source_flux(jnp, path, name): + """k0..k7 only: see the module docstring for why the euclid system is excluded from this metric.""" + x, _, _ = _solve(jnp, path, name) + value = _flux_rel( + x, SYSTEMS[name][2], index=SYSTEM_META[name]["source_column_index_list"] + ) + + assert abs(value) <= FLUX_TOL, f"flux_rel_source = {value:.3e}" + + +@requires_jax +@pytest.mark.parametrize("name", NAMES) +def test__raw_forward_jit_matches_eager(jnp, name): + """``jax.jit`` of the reconstruction returns the eager value: end-to-end to a few ULP, and exactly for + the solver alone. + + End-to-end, XLA's CPU code generation may reassociate the Jacobi scaling and the final ``x / D`` around + the solver, which moves entries by a few ULP and differs between runners: the GitHub Actions Python 3.12 + leg (jax 0.11.2, same as the green 3.13 leg) reproduced eager to 5.4e-11 relative / 9e-12 absolute on + every fixture system (PyAutoArray#595), so the end-to-end check is a tight ``allclose`` rather than + bit-exact. The solver-alone check below stays bit-exact. + """ + import jax + + Q, q, _ = SYSTEMS[name] + Qj, qj = jnp.asarray(Q), jnp.asarray(q) + + def f(Q_, q_): + return _dispatch(jnp, Q_, q_) + + eager = np.asarray(f(Qj, qj)) + np.testing.assert_allclose( + np.asarray(jax.jit(f)(Qj, qj)), + eager, + rtol=1e-9, + atol=1e-12 * np.abs(eager).max(), + ) + + # The solver alone, on inputs built outside the traced function: tracing the Jacobi scaling together with + # the final ``x / D`` lets XLA reassociate them and moves ``y`` by 1 ULP (seen on the unfixed base too), + # which is not what this test is about. + solve, args = _entry_solver(jnp, Qj, qj) + np.testing.assert_array_equal( + np.asarray(jax.jit(solve)(*args)), np.asarray(solve(*args)) + ) + + +@requires_jax +@pytest.mark.parametrize("jit", [False, True], ids=["eager", "jit"]) +@pytest.mark.parametrize("name", NAMES) +def test__raw_forward_primal_matches_differentiated_forward(jnp, name, jit): + """The custom_vjp primal (plain calls) and its fwd rule (any differentiated call) return the same ``y`` + bit-for-bit, so a value does not change when a gradient is taken through it.""" + import jax + + Q, q, _ = SYSTEMS[name] + solve, args = _entry_solver(jnp, jnp.asarray(Q), jnp.asarray(q)) + + def primal_out(*a): + return solve(*a) + + def fwd_out(*a): + return jax.vjp(solve, *a)[0] + + if jit: + primal_out, fwd_out = jax.jit(primal_out), jax.jit(fwd_out) + + np.testing.assert_array_equal( + np.asarray(fwd_out(*args)), np.asarray(primal_out(*args)) + ) + + +@requires_jax +@pytest.mark.parametrize("name", NAMES) +def test__raw_forward_gradient_is_finite_and_non_zero(jnp, name): + import jax + + Q, q, _ = SYSTEMS[name] + Qj = jnp.asarray(Q) + + def f(q_): + y, _, _, _ = _entry(jnp, Qj, q_) + return jnp.sum(y) + + for grad in (jax.grad(f), jax.jit(jax.grad(f))): + gq = np.asarray(grad(jnp.asarray(q))) + assert np.all(np.isfinite(gq)) + assert np.any(gq != 0.0)