diff --git a/autoarray/config/general.yaml b/autoarray/config/general.yaml index 85671dde2..e6ba12033 100644 --- a/autoarray/config/general.yaml +++ b/autoarray/config/general.yaml @@ -8,6 +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. + 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_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/abstract.py b/autoarray/inversion/inversion/abstract.py index 11084c599..79c27677e 100644 --- a/autoarray/inversion/inversion/abstract.py +++ b/autoarray/inversion/inversion/abstract.py @@ -600,6 +600,31 @@ def positive_only_solver_used(self) -> str: return "certified" + @property + def positive_only_preconditioning_used(self) -> str: + """ + How `reconstruction` asks `reconstruction_positive_only_from` to scale the JAX PDIP solve: + ``"jacobi"`` or ``"raw"`` (PyAutoArray#571). + + - Inversions containing a `Mapper` always use ``"jacobi"`` (today's Jacobi-preconditioned solve). + - Inversions with **no** `Mapper` (linear light profiles / MGE only) use + `Settings.nnls_preconditioning_no_mapper` (packaged default ``"raw"``): Jacobi scaling turns their + signal-free Gaussian columns, whose diagonal is only the no-regularization floor, into degenerate + coordinates on which the PDIP dual diverges (14/48 near-truth SLaM `source_lp[1]` points hit the + iteration cap with wrong log-likelihoods), while the raw solve converges on all of them. + - Whenever the certified solver is used the answer is ``"jacobi"`` (it only runs on mapper-only + inversions anyway). + + The NumPy path runs fnnls and ignores the value. + """ + if self.has(cls=Mapper): + return "jacobi" + + if self.positive_only_solver_used != "pdip": + return "jacobi" + + return self.settings.nnls_preconditioning_no_mapper + def _nnls_warm_start_fingerprint(self, ids_to_keep=None) -> Optional[str]: """ Identify the index space this inversion's positive-only solve works in, so the @@ -688,6 +713,7 @@ def reconstruction(self) -> np.ndarray: ), factor=factor, solver=solver, + preconditioning=self.positive_only_preconditioning_used, ) ) @@ -717,6 +743,7 @@ def reconstruction(self) -> np.ndarray: fingerprint=self._nnls_warm_start_fingerprint(), factor=factor, solver=solver, + preconditioning=self.positive_only_preconditioning_used, ) self._nnls_factor = factor diff --git a/autoarray/inversion/inversion/inversion_util.py b/autoarray/inversion/inversion/inversion_util.py index 43bfb5e8d..d94668707 100644 --- a/autoarray/inversion/inversion/inversion_util.py +++ b/autoarray/inversion/inversion/inversion_util.py @@ -297,6 +297,7 @@ def reconstruction_positive_only_from( factor: Optional[dict] = None, solver: str = "pdip", stats: Optional[dict] = None, + preconditioning: str = "jacobi", ): """ Solve the linear system Eq.(2) (in terms of minimizing the quadratic value) of @@ -358,7 +359,22 @@ def reconstruction_positive_only_from( ``certified`` (whether the active-set search certified within budget; ``False`` means the fallback or an uncertified iterate was returned) and ``passes`` (restricted passes run) as *traced* JAX scalars, so it is safe under ``jax.jit`` / ``vmap`` -- read them as outputs of the traced function or through - ``jax.debug.callback``. Every call also records ``solver``. It never changes the returned reconstruction. + ``jax.debug.callback``. With ``solver="pdip"`` it receives ``converged`` (``1`` if the PDIP KKT residual + met its tolerance within the iteration cap, ``0`` if the cap was hit and the returned reconstruction is + the unconverged iterate) and ``iterations`` (PDIP iterations run), also as traced JAX scalars, plus + ``preconditioning``. Every call also records ``solver``. It never changes the returned reconstruction. + preconditioning + How the JAX PDIP solve (``solver="pdip"``) scales the system (PyAutoArray#571). ``"jacobi"`` (default, + 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 + 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 + ``no_regularization_add_to_curvature_diag_value`` floor, degenerate coordinates on which the PDIP dual + diverges; the caller (`AbstractInversion.reconstruction`) therefore passes ``"raw"`` for inversions with + no `Mapper` -- see `AbstractInversion.positive_only_preconditioning_used`. Ignored on the NumPy path. Notes ----- @@ -386,11 +402,25 @@ def reconstruction_positive_only_from( f"solver={solver!r} is not a valid positive-only solver; expected 'pdip' or 'certified'." ) + if preconditioning not in ("jacobi", "raw"): + raise ValueError( + f"preconditioning={preconditioning!r} is invalid; expected 'jacobi' or 'raw'." + ) + + if preconditioning == "raw" and solver != "pdip": + raise ValueError( + "preconditioning='raw' applies only to solver='pdip'; the certified solver always runs on the " + "Jacobi-scaled system." + ) + if xp.__name__.startswith("jax"): from autonerves import conf - from autoarray.util.jax_nnls import solve_nnls_primal + from autoarray.util.jax_nnls import ( + solve_nnls_primal_raw_forward, + solve_nnls_primal_with_status, + ) try: use_jacobi = conf.instance["general"]["inversion"][ @@ -403,9 +433,7 @@ def reconstruction_positive_only_from( use_jacobi = True try: - target_kappa = conf.instance["general"]["inversion"][ - "nnls_target_kappa" - ] + target_kappa = conf.instance["general"]["inversion"]["nnls_target_kappa"] except KeyError: # Workspaces ship their own general.yaml that shadows autoarray's; # fall back to the same value autoarray's general.yaml declares. @@ -424,6 +452,34 @@ def reconstruction_positive_only_from( if max_iter is None: max_iter = 50 + def _record_pdip(converged, iterations): + if stats is not None: + stats["solver"] = "pdip" + stats["preconditioning"] = preconditioning + stats["converged"] = converged + stats["iterations"] = iterations + + if preconditioning == "raw": + # Same Jacobi quantities as below: the backward pass runs on the scaled system, the forward + # solve on the raw one (see `solve_nnls_primal_raw_forward`). + d = xp.sqrt(xp.diag(curvature_reg_matrix)) + D = 1.0 / d + Q_pc = (curvature_reg_matrix * D[:, None]) * D[None, :] + q_pc = data_vector * D + + y, converged, iterations = solve_nnls_primal_raw_forward( + Q_pc, + q_pc, + curvature_reg_matrix, + data_vector, + D, + target_kappa=target_kappa, + solver_tol=solver_tol, + max_iter=max_iter, + ) + _record_pdip(converged, iterations) + return y * D + if use_jacobi: # Ill-conditioned Q makes jaxnnls's relaxed-KKT backward pass # produce NaN gradients. Rescale Q so its diagonal is unit: @@ -449,19 +505,15 @@ def reconstruction_positive_only_from( * D ) - if stats is not None: - stats["solver"] = "pdip" - - return ( - solve_nnls_primal( - Q_pc, - q_pc, - target_kappa=target_kappa, - solver_tol=solver_tol, - max_iter=max_iter, - ) - * D + x, converged, iterations = solve_nnls_primal_with_status( + Q_pc, + q_pc, + target_kappa=target_kappa, + solver_tol=solver_tol, + max_iter=max_iter, ) + _record_pdip(converged, iterations) + return x * D if solver == "certified": return _certified_positive_only_from( @@ -474,16 +526,15 @@ def reconstruction_positive_only_from( stats=stats, ) - if stats is not None: - stats["solver"] = "pdip" - - return solve_nnls_primal( + x, converged, iterations = solve_nnls_primal_with_status( curvature_reg_matrix, data_vector, target_kappa=target_kappa, solver_tol=solver_tol, max_iter=max_iter, ) + _record_pdip(converged, iterations) + return x # `solver` is deliberately ignored on the NumPy path: a NumPy port of the certified active-set scheme # measured 3-7 % slower than fnnls and lost to its warm-start memo (PyAutoArray#566), so fnnls stays. diff --git a/autoarray/settings.py b/autoarray/settings.py index 19fabe8fa..ef18c180c 100644 --- a/autoarray/settings.py +++ b/autoarray/settings.py @@ -26,6 +26,7 @@ def __init__( certified_pass_budget: Optional[int] = None, certified_fallback: Optional[str] = None, certified_tau_rel: Optional[float] = None, + nnls_preconditioning_no_mapper: Optional[str] = None, ): """ The settings of an Inversion, customizing how a linear set of equations are solved for. @@ -223,6 +224,23 @@ def __init__( Relative KKT tolerance of the ``"certified"`` solver's certificate: primal violations are ``x < -tau_rel * max|x|``, dual violations ``g < -tau_rel * max|q|``. `None` (default) reads the packaged value (`1.0e-9`). + nnls_preconditioning_no_mapper + How the JAX positive-only PDIP solve scales an inversion **with no `Mapper`** (linear light profiles / + MGE only). `None` (default) reads the packaged value ``"raw"``. + + - ``"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"``. + 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 + coordinates on which the PDIP dual diverges); the raw solve converges on all of them in 16-19 + iterations (PyAutoArray#571). + - ``"jacobi"`` -- the Jacobi-preconditioned solve, as for mapper inversions. + + Inversions containing a `Mapper` always use ``"jacobi"`` + (`AbstractInversion.positive_only_preconditioning_used` records the decision). The NumPy path always + runs fnnls and ignores this. """ self.use_mixed_precision = use_mixed_precision self.nnls_solver_tol = nnls_solver_tol @@ -244,12 +262,15 @@ def __init__( self._certified_pass_budget = certified_pass_budget self._certified_fallback = certified_fallback self._certified_tau_rel = certified_tau_rel + self._nnls_preconditioning_no_mapper = nnls_preconditioning_no_mapper # Validate explicit values eagerly, so a typo fails at construction rather than deep inside a fit. if positive_only_solver is not None: self.positive_only_solver if certified_fallback is not None: self.certified_fallback + if nnls_preconditioning_no_mapper is not None: + self.nnls_preconditioning_no_mapper @property def use_positive_only_solver(self): @@ -449,3 +470,23 @@ def certified_tau_rel(self) -> float: return self._inversion_config_value("certified_tau_rel", 1.0e-9) return self._certified_tau_rel + + @property + def nnls_preconditioning_no_mapper(self) -> str: + """ + How the JAX PDIP solve scales an inversion with no `Mapper`: ``"raw"`` or ``"jacobi"``. + + See the constructor docstring; inversions with a `Mapper` always use ``"jacobi"``. + """ + value = self._nnls_preconditioning_no_mapper + if value is None: + value = self._inversion_config_value( + "nnls_preconditioning_no_mapper", "raw" + ) + + if value not in ("raw", "jacobi"): + raise ValueError( + f"nnls_preconditioning_no_mapper={value!r} is invalid; expected 'raw' or 'jacobi'." + ) + + return value diff --git a/autoarray/util/jax_nnls.py b/autoarray/util/jax_nnls.py index e7abd179f..3748f1bce 100644 --- a/autoarray/util/jax_nnls.py +++ b/autoarray/util/jax_nnls.py @@ -20,6 +20,15 @@ runs until the slowest lane converges, so ``max_iter`` also caps the worst-case batched cost. +Convergence is observable (PyAutoArray#571): :func:`solve_nnls_primal_with_status` +returns the PDIP ``converged`` flag and iteration count next to ``x`` (the +plain :func:`solve_nnls_primal` keeps its drop-in signature). The ``"raw"`` +positive-only mode (:func:`solve_nnls_primal_raw_forward`) runs the forward +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. + JAX is imported inside functions, never at module level (see ``docs/agents/jax_and_decorators.md``); this module must only be imported on the ``xp=jnp`` path. The solver knobs are static closure parameters — @@ -76,38 +85,85 @@ def converged_check(inputs): return x, s, z, converged, pdip_iter +# The data-scaled tolerance of the "raw" (un-preconditioned) mode is this fraction of jaxnnls's +# own ``n * EPSILON`` rule, multiplied by ``max(1, max|q|)``. Measured on the SLaM MGE fixture +# (PyAutoArray#571): factor 1 stops ~1e-11 (relative objective) short of fnnls, 1e-2 reaches +# <= 4e-13 for 1-2 extra iterations (16-19 in total), 1e-3 / 1e-4 buy one more digit per iteration. +DATA_SCALED_TOL_FACTOR = 1.0e-2 + + +def data_scaled_solver_tol(q): + """ + The convergence tolerance of the ``"raw"`` positive-only mode: + ``DATA_SCALED_TOL_FACTOR * n * EPSILON * max(1, max|q|)``. + + jaxnnls's own rule ``n * EPSILON`` is absolute, so on an unscaled system + (curvature entries ~1e7) it is unreachable in floating point; scaling it by + the data vector makes it a relative KKT tolerance. ``q`` may be traced. + """ + import jax.numpy as jnp + from jaxnnls.pdip import EPSILON + + return ( + DATA_SCALED_TOL_FACTOR + * q.shape[0] + * EPSILON + * jnp.maximum(1.0, jnp.max(jnp.abs(q))) + ) + + @lru_cache(maxsize=None) def _solve_nnls_primal_with(target_kappa, solver_tol, max_iter): """ Build (and cache) the differentiable primal solver for one static - setting of the knobs. The returned function takes only (Q, q), so the - custom-vjp backward pass returns exactly (dQ, dq). + setting of the knobs. The returned function takes only (Q, q) and returns + ``(x, converged, pdip_iter)``; the custom-vjp backward pass returns + exactly (dQ, dq) from the cotangent of ``x`` (the integer status outputs + carry no cotangent). """ import jax from jaxnnls.diff_qp import diff_nnls from jaxnnls.pdip_relaxed import solve_relaxed_nnls def primal(Q, q): - return solve_nnls(Q, q, solver_tol=solver_tol, max_iter=max_iter)[0] + x, _, _, converged, pdip_iter = solve_nnls( + Q, q, solver_tol=solver_tol, max_iter=max_iter + ) + return x, converged, pdip_iter def forward(Q, q): - x, s, z, _, _ = solve_nnls(Q, q, solver_tol=solver_tol, max_iter=max_iter) + x, s, z, converged, pdip_iter = solve_nnls( + Q, q, solver_tol=solver_tol, max_iter=max_iter + ) # Relax the solution with vanilla Newton steps on the relaxed KKT # conditions; only the backward pass consumes the relaxed variables. - xr, sr, zr, _, _ = solve_relaxed_nnls( - Q, q, x, s, z, target_kappa=target_kappa - ) - return x, (Q, xr, sr, zr) + xr, sr, zr, _, _ = solve_relaxed_nnls(Q, q, x, s, z, target_kappa=target_kappa) + return (x, converged, pdip_iter), (Q, xr, sr, zr) - def backward(res, input_grad): + def backward(res, output_grad): Q, xr, sr, zr = res - return diff_nnls(Q, xr, sr, zr, input_grad) + return diff_nnls(Q, xr, sr, zr, output_grad[0]) primal = jax.custom_vjp(primal) primal.defvjp(forward, backward) return primal +def solve_nnls_primal_with_status( + Q, q, target_kappa=1e-3, solver_tol=None, max_iter=50 +): + """ + As :func:`solve_nnls_primal`, but also returns the PDIP convergence flag + and iteration count: ``(x, converged, pdip_iter)``. + + ``x`` (value and gradient) is identical to :func:`solve_nnls_primal`; + ``converged`` (``1`` if the KKT residual met the tolerance within + ``max_iter``) and ``pdip_iter`` are non-differentiable integer outputs, so + they are safe to return from ``jax.jit`` / ``vmap``-ed code. + """ + return _solve_nnls_primal_with(target_kappa, solver_tol, max_iter)(Q, q) + + def solve_nnls_primal(Q, q, target_kappa=1e-3, solver_tol=None, max_iter=50): """ Solve the non-negative least squares problem, differentiable via the @@ -115,6 +171,80 @@ def solve_nnls_primal(Q, q, target_kappa=1e-3, solver_tol=None, max_iter=50): Drop-in replacement for ``jaxnnls.solve_nnls_primal`` with two extra knobs; at their defaults (``solver_tol=None``, ``max_iter=50``) the - forward solve and gradients are identical to upstream. + forward solve and gradients are identical to upstream. Use + :func:`solve_nnls_primal_with_status` to also get the convergence flag. """ - return _solve_nnls_primal_with(target_kappa, solver_tol, max_iter)(Q, q) + return solve_nnls_primal_with_status( + Q, q, target_kappa=target_kappa, solver_tol=solver_tol, max_iter=max_iter + )[0] + + +@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). + + 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)`` + and ``D``, and returns ``(y, converged, pdip_iter)`` with ``y`` the solution + of the scaled system, so ``x = D * y``. + + - **Forward:** the PDIP solve runs on the *raw* ``(Q, q)`` with a + data-scaled tolerance (:func:`data_scaled_solver_tol`, unless + ``solver_tol`` is given), and its iterate is mapped to the scaled + coordinates (``y = x / D``, slack ``s / D``, dual ``z * D``). On + 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. + """ + 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): + 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 + ) + 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 + + 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 + ) + return (y, converged, pdip_iter), (Q_pc, yr, sr, zr, Q, q, D) + + def backward(res, output_grad): + 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) + + primal = jax.custom_vjp(primal) + primal.defvjp(forward, backward) + return primal + + +def solve_nnls_primal_raw_forward( + Q_pc, q_pc, Q, q, D, target_kappa=1e-3, solver_tol=None, max_iter=50 +): + """ + The ``"raw"`` positive-only mode: forward PDIP on the raw system with a + data-scaled tolerance, backward pass on the Jacobi-scaled system. Returns + ``(y, converged, pdip_iter)`` with ``x = D * y``; see + :func:`_solve_nnls_raw_forward_with`. + """ + return _solve_nnls_raw_forward_with(target_kappa, solver_tol, max_iter)( + Q_pc, q_pc, Q, q, D + ) diff --git a/test_autoarray/inversion/inversion/files/README.md b/test_autoarray/inversion/inversion/files/README.md new file mode 100644 index 000000000..887f535f8 --- /dev/null +++ b/test_autoarray/inversion/inversion/files/README.md @@ -0,0 +1,9 @@ +# Inversion test fixtures + +- `mge_slam_nnls_systems.npz` — 8 positive-only systems `Q_` (curvature_reg_matrix, 60x60) / `q_` + (data_vector) captured from the SLaM `source_lp[1]` MGE model on the autolens_profiling HST dataset + (PyAutoArray#571): keys 0-4 never converge under JAX PDIP in 200 iterations, 5-6 hit the 50-iteration cap + but converge by 200, 7 is healthy (19 iterations); `meta` holds the per-system JSON. Generator: + `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. diff --git a/test_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npz b/test_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npz new file mode 100644 index 000000000..a97768290 Binary files /dev/null and b/test_autoarray/inversion/inversion/files/mge_slam_nnls_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 new file mode 100644 index 000000000..0d34d2632 --- /dev/null +++ b/test_autoarray/inversion/inversion/test_nnls_mge_convergence.py @@ -0,0 +1,325 @@ +""" +Regression tests for the JAX positive-only (PDIP NNLS) solve on real SLaM MGE systems (PyAutoArray#571). + +The fixture `files/mge_slam_nnls_systems.npz` holds 8 `(curvature_reg_matrix, data_vector)` systems captured +from the SLaM `source_lp[1]` MGE model (2 x 20 lens Gaussians with `sigma_min = pixel_scale / 10` plus 20 +source Gaussians, 60 linear columns) exactly as the JAX likelihood hands them to +`reconstruction_positive_only_from` (see `files/README.md` for the generator and versions): + +- keys 0-4: the Jacobi-preconditioned PDIP solve never converges, even with a 200-iteration cap; +- keys 5-6: it hits the production 50-iteration cap but converges by 200; +- key 7: a healthy system (19 iterations). + +Mechanism (issue comment "Step 3 diagnosis"): Jacobi scaling turns the signal-free source-Gaussian columns, +whose diagonal is only the no-regularization floor, into degenerate coordinates on which the PDIP dual +diverges. The fix is the ``preconditioning="raw"`` mode, which the inversion dispatches for mapper-less +inversions: the forward solve runs on the raw system with a data-scaled tolerance. With it, every fixture +system must converge within the production cap and reach the NumPy `fnnls_cholesky` objective +`0.5 x^T Q x - q^T x`, single, end-to-end and under `vmap`, with finite gradients. The Jacobi mode must now +*report* its non-convergence (`stats["converged"] == 0`), and the control test pins today's Jacobi answer +bit-identically on well-conditioned random systems. + +Tolerances were declared before the first run: objective within `1e-8 * |obj_fnnls| + 1e-8` (the objectives +are negative, ~ -9.7e5, so the bound is additive rather than the multiplicative `obj * (1 + 1e-8)` of the +issue plan, which would demand PDIP beat fnnls); control solution vs SciPy `nnls` within `1e-8 * max|x|`. +""" + +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 +from autoarray.util.fnnls import fnnls_cholesky + + +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_slam_nnls_systems.npz" +PRODUCTION_MAX_ITER = 50 +OBJECTIVE_RTOL = 1.0e-8 +OBJECTIVE_ATOL = 1.0e-8 + + +def _load_systems(): + with np.load(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 + + +META, SYSTEMS = _load_systems() +KEYS = [s["key"] for s in META["systems"]] +IDS = [f"{s['key']}-{s['category']}" for s in META["systems"]] + + +def _objective(Q, q, x): + return float(0.5 * x @ Q @ x - q @ x) + + +def _fnnls(Q, q): + """The NumPy oracle, started exactly as `reconstruction_positive_only_from` starts it.""" + return np.asarray( + fnnls_cholesky(Q, q, P_initial=np.linalg.solve(Q, q) > 0), dtype=float + ) + + +def _jacobi(Q, q): + """The Jacobi scaling `reconstruction_positive_only_from` applies on the JAX path.""" + D = 1.0 / np.sqrt(np.diag(Q)) + return Q * D[:, None] * D[None, :], q * D, D + + +def _assert_objective_reaches_fnnls(Q, q, x): + obj_fnnls = _objective(Q, q, _fnnls(Q, q)) + assert np.all(np.isfinite(x)), "PDIP returned a non-finite solution" + obj = _objective(Q, q, x) + assert obj <= obj_fnnls + OBJECTIVE_RTOL * abs(obj_fnnls) + OBJECTIVE_ATOL, ( + obj, + obj_fnnls, + ) + + +@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 + # make every fixture "converge" at a 6e-4 KKT tolerance and hide the bug. + assert EPSILON < 1.0e-10, EPSILON + + return jnp + + +def test__fixture_is_the_captured_slam_mge_set(): + assert FIXTURE.stat().st_size < 1_000_000 + assert [s["category"] for s in META["systems"]] == 5 * [ + "never_converges_cap200" + ] + 2 * ["cap50_hit_converges_by_200"] + ["healthy"] + for Q, q in SYSTEMS: + assert Q.shape == (60, 60) and q.shape == (60,) + np.testing.assert_allclose(Q, Q.T, rtol=0, atol=1e-8 * np.abs(Q).max()) + + +def _raw(jnp, Q, q, stats=None, settings=None): + return inversion_util.reconstruction_positive_only_from( + data_vector=jnp.asarray(q), + curvature_reg_matrix=jnp.asarray(Q), + settings=settings or aa.Settings(), + xp=jnp, + stats=stats, + preconditioning="raw", + ) + + +@requires_jax +@pytest.mark.parametrize("key", KEYS, ids=IDS) +def test__raw_pdip_converges_within_production_cap(jnp, key): + """The "raw" solve (un-preconditioned, data-scaled tolerance) converges on every fixture system and reaches + the fnnls objective to 1e-12 relative (measured <= 4e-13).""" + from autoarray.util.jax_nnls import data_scaled_solver_tol, solve_nnls + + Q, q = SYSTEMS[key] + Qj, qj = jnp.asarray(Q), jnp.asarray(q) + + x, _, _, converged, pdip_iter = solve_nnls( + Qj, qj, solver_tol=data_scaled_solver_tol(qj), max_iter=PRODUCTION_MAX_ITER + ) + + assert int(converged) == 1, f"PDIP did not converge ({int(pdip_iter)} iterations)" + assert int(pdip_iter) < PRODUCTION_MAX_ITER + x = np.asarray(x) + assert np.all(np.isfinite(x)) + obj_fnnls = _objective(Q, q, _fnnls(Q, q)) + assert abs(_objective(Q, q, x) - obj_fnnls) <= 1.0e-12 * abs(obj_fnnls) + + +@requires_jax +@pytest.mark.parametrize("key", KEYS, ids=IDS) +def test__reconstruction_positive_only_from__jax_matches_numpy_objective(jnp, key): + Q, q = SYSTEMS[key] + settings = aa.Settings() + stats = {} + + x_jax = np.asarray(_raw(jnp, Q, q, stats=stats, settings=settings)) + x_np = inversion_util.reconstruction_positive_only_from( + data_vector=q, curvature_reg_matrix=Q, settings=settings, xp=np + ) + + assert np.all(np.isfinite(x_jax)), "JAX reconstruction is non-finite" + obj_jax = _objective(Q, q, x_jax) + obj_np = _objective(Q, q, np.asarray(x_np)) + assert abs(obj_jax - obj_np) <= OBJECTIVE_RTOL * abs(obj_np), (obj_jax, obj_np) + + assert stats["solver"] == "pdip" + assert stats["preconditioning"] == "raw" + assert int(stats["converged"]) == 1 + assert int(stats["iterations"]) < PRODUCTION_MAX_ITER + + +@requires_jax +def test__jacobi_mode_reports_non_convergence_on_the_witness_systems(jnp): + """F1: the Jacobi mode still fails on fixture keys 0-6 (the mechanism is unchanged), but the failure is no + longer silent -- `stats["converged"]` is 0 and the iteration count is the cap. Key 7 converges. + """ + for key, (Q, q) in enumerate(SYSTEMS): + stats = {} + inversion_util.reconstruction_positive_only_from( + data_vector=jnp.asarray(q), + curvature_reg_matrix=jnp.asarray(Q), + settings=aa.Settings(), + xp=jnp, + stats=stats, + ) + assert stats["preconditioning"] == "jacobi" + expected = 1 if META["systems"][key]["category"] == "healthy" else 0 + assert int(stats["converged"]) == expected, key + if expected == 0: + assert int(stats["iterations"]) == PRODUCTION_MAX_ITER + + +@requires_jax +def test__raw_pdip_converges_under_vmap(jnp): + import jax + + def solve(Q, q): + stats = {} + x = inversion_util.reconstruction_positive_only_from( + data_vector=q, + curvature_reg_matrix=Q, + settings=aa.Settings(), + xp=jnp, + stats=stats, + preconditioning="raw", + ) + return x, stats["converged"], stats["iterations"] + + Qs = jnp.asarray(np.stack([Q for Q, _ in SYSTEMS])) + qs = jnp.asarray(np.stack([q for _, q in SYSTEMS])) + + x, converged, iterations = jax.jit(jax.vmap(solve))(Qs, qs) + + converged = np.asarray(converged) + iterations = np.asarray(iterations) + failed = [ + k for k in KEYS if converged[k] != 1 or iterations[k] >= PRODUCTION_MAX_ITER + ] + assert not failed, ( + f"failed lanes {failed}; converged {converged.tolist()}; " + f"iterations {iterations.tolist()}" + ) + for k, (Q, q) in enumerate(SYSTEMS): + _assert_objective_reaches_fnnls(Q, q, np.asarray(x[k])) + + +@requires_jax +@pytest.mark.parametrize("key", KEYS, ids=IDS) +def test__raw_mode_gradient_is_finite_and_non_zero(jnp, key): + """The raw mode keeps the Jacobi-space relaxed-KKT backward pass: a relaxed-KKT pass on the raw Q gives NaN + gradients on keys 2-4, and the Jacobi-mode gradient itself is NaN on key 1 (its forward solve diverged). + """ + import jax + + Q, q = SYSTEMS[key] + w = jnp.linspace(0.5, 1.5, q.shape[0]) + + gQ, gq = jax.grad(lambda Q_, q_: w @ _raw(jnp, Q_, q_), argnums=(0, 1))( + 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) + + +@requires_jax +def test__raw_mode_gradient_matches_jacobi_mode_on_the_healthy_system(jnp): + """Where the Jacobi forward solve converges (key 7) the two modes differentiate the same relaxed system from + nearly the same iterate: measured max relative difference 2.8e-3.""" + import jax + + key = [s["key"] for s in META["systems"] if s["category"] == "healthy"][0] + Q, q = SYSTEMS[key] + w = jnp.linspace(0.5, 1.5, q.shape[0]) + + def jacobi(Q_, q_): + return w @ inversion_util.reconstruction_positive_only_from( + data_vector=q_, curvature_reg_matrix=Q_, settings=aa.Settings(), xp=jnp + ) + + g_raw = np.asarray(jax.grad(lambda q_: w @ _raw(jnp, Q, q_))(jnp.asarray(q))) + g_jac = np.asarray(jax.grad(lambda q_: jacobi(jnp.asarray(Q), q_))(jnp.asarray(q))) + + assert np.abs(g_raw - g_jac).max() <= 1.0e-2 * np.abs(g_jac).max() + + +def _random_system(n, seed): + """A seeded, well-conditioned NNLS problem (as `test_jax_active_set._qp`): roughly half the unconstrained + solution is negative, so positivity binds.""" + rng = np.random.default_rng(seed) + A = rng.normal(size=(3 * n, n)) + b = rng.normal(size=3 * n) + return A, b, A.T @ A, A.T @ b + + +@requires_jax +@pytest.mark.parametrize("n, seed", [(12, 0), (30, 1), (60, 2)]) +def test__control__well_conditioned_pdip_unchanged(jnp, n, seed): + """ + Must pass before and after any fix, bit-identically: the library JAX path is compared against the + upstream jaxnnls solve of the same Jacobi-scaled system, recomputed in-test (today's answer). + """ + import jaxnnls + from jaxnnls.pdip import solve_nnls as upstream_solve_nnls + from scipy.optimize import nnls + + A, b, Q, q = _random_system(n, seed) + Q_pc, q_pc, D = _jacobi(Q, q) + + _, _, _, converged, pdip_iter = upstream_solve_nnls( + jnp.asarray(Q_pc), jnp.asarray(q_pc) + ) + assert int(converged) == 1 and int(pdip_iter) < PRODUCTION_MAX_ITER + + expected = ( + np.asarray( + jaxnnls.solve_nnls_primal( + jnp.asarray(Q_pc), jnp.asarray(q_pc), target_kappa=1.0e-11 + ) + ) + * D + ) + + x = np.asarray( + inversion_util.reconstruction_positive_only_from( + data_vector=jnp.asarray(q), + curvature_reg_matrix=jnp.asarray(Q), + settings=aa.Settings(), + xp=jnp, + ) + ) + + np.testing.assert_array_equal(x, expected) + + x_scipy, _ = nnls(A, b) + np.testing.assert_allclose(x, x_scipy, rtol=0, atol=1e-8 * np.abs(x_scipy).max()) + + stats = {} + x_raw = np.asarray(_raw(jnp, Q, q, stats=stats)) + assert int(stats["converged"]) == 1 + np.testing.assert_allclose( + x_raw, x_scipy, rtol=0, atol=1e-8 * np.abs(x_scipy).max() + ) diff --git a/test_autoarray/inversion/inversion/test_positive_only_dispatch.py b/test_autoarray/inversion/inversion/test_positive_only_dispatch.py index 52198f966..810083e27 100644 --- a/test_autoarray/inversion/inversion/test_positive_only_dispatch.py +++ b/test_autoarray/inversion/inversion/test_positive_only_dispatch.py @@ -48,6 +48,7 @@ def _inversion( use_jax, positive_only_solver="certified", use_edge_zeroed_pixels=False, + nnls_preconditioning_no_mapper=None, ): if use_jax: import jax @@ -66,6 +67,7 @@ def _inversion( use_positive_only_solver=True, use_edge_zeroed_pixels=use_edge_zeroed_pixels, positive_only_solver=positive_only_solver, + nnls_preconditioning_no_mapper=nnls_preconditioning_no_mapper, # Off so two NumPy solves of the same system are bit-comparable: with the memo on, the second # solve starts from the first's passive set and can differ in the last ulp. nnls_warm_start_memo=False, @@ -209,3 +211,106 @@ def test__positive_only_solver_used__edge_zeroed_subset_is_preserved(): assert x_certified == pytest.approx( x_pdip, rel=1.0e-8, abs=1.0e-8 * np.max(np.abs(x_certified)) ) + + +def test__positive_only_preconditioning_used__mapper_inversions_keep_jacobi(): + data_vector, curvature_reg_matrix = _system(17) + func_list = aa.m.MockLinearObjFuncList(parameters=1) + + for linear_obj_list in ([_mapper()], [func_list, _mapper()]): + for solver in ("pdip", "certified"): + inversion = _inversion( + linear_obj_list, + data_vector[: 16 + len(linear_obj_list) - 1], + curvature_reg_matrix[ + : 16 + len(linear_obj_list) - 1, : 16 + len(linear_obj_list) - 1 + ], + use_jax=False, + positive_only_solver=solver, + ) + assert inversion.positive_only_preconditioning_used == "jacobi" + + +def test__positive_only_preconditioning_used__no_mapper_uses_settings_default_raw(): + data_vector, curvature_reg_matrix = _system(4) + + default = _inversion( + [aa.m.MockLinearObjFuncList(parameters=4)], + data_vector, + curvature_reg_matrix, + use_jax=False, + positive_only_solver="pdip", + ) + forced = _inversion( + [aa.m.MockLinearObjFuncList(parameters=4)], + data_vector, + curvature_reg_matrix, + use_jax=False, + positive_only_solver="pdip", + nnls_preconditioning_no_mapper="jacobi", + ) + + assert default.settings.nnls_preconditioning_no_mapper == "raw" + assert default.positive_only_preconditioning_used == "raw" + assert forced.positive_only_preconditioning_used == "jacobi" + + +def test__settings__nnls_preconditioning_no_mapper_rejects_unknown_value(): + with pytest.raises(ValueError): + aa.Settings(nnls_preconditioning_no_mapper="diag") + + +@requires_jax +def test__positive_only_preconditioning_used__no_mapper_jax_reconstruction_is_the_raw_solve(): + import jax.numpy as jnp + + from autoarray.inversion.inversion import inversion_util + + data_vector, curvature_reg_matrix = _system(4) + + inversion = _inversion( + [aa.m.MockLinearObjFuncList(parameters=4)], + data_vector, + curvature_reg_matrix, + use_jax=True, + positive_only_solver="pdip", + ) + + expected = inversion_util.reconstruction_positive_only_from( + data_vector=jnp.asarray(data_vector), + curvature_reg_matrix=jnp.asarray(curvature_reg_matrix), + settings=inversion.settings, + xp=jnp, + preconditioning="raw", + ) + + assert np.array_equal(np.asarray(inversion.reconstruction), np.asarray(expected)) + assert np.asarray(expected) == pytest.approx( + np.asarray( + inversion_util.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=inversion.settings, + xp=np, + ) + ), + abs=1.0e-8, + ) + + +@requires_jax +def test__reconstruction_positive_only_from__raw_rejects_certified(): + import jax.numpy as jnp + + from autoarray.inversion.inversion import inversion_util + + data_vector, curvature_reg_matrix = _system(4) + + with pytest.raises(ValueError): + inversion_util.reconstruction_positive_only_from( + data_vector=jnp.asarray(data_vector), + curvature_reg_matrix=jnp.asarray(curvature_reg_matrix), + xp=jnp, + solver="certified", + preconditioning="raw", + )