From 66f5a4b148b0ac747b807efef787b3d49b179e7a Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Wed, 23 Sep 2026 14:53:15 +0200 Subject: [PATCH 1/3] feat(inversion): certified active-set positive solver on the JAX path, opt-in, mapper-only dispatch New autoarray/util/jax_active_set.py: budgeted lax.while_loop free-all active-set search with primal + dual (KKT) certification and a permanent fixed set, run under stop_gradient; the returned solution is one final differentiable masked Cholesky solve (exact implicit active-set gradient); lax.cond fallback to the PDIP solve on an exhausted budget. reconstruction_positive_only_from gains solver="pdip"|"certified" and an optional stats out-dict (traced certified/passes); the pdip path and the NumPy fnnls path are unchanged. Settings gains positive_only_solver, certified_pass_budget, certified_fallback, certified_tau_rel (packaged defaults pdip/16/pdip/1e-9). AbstractInversion.positive_only_solver_used selects certified only on JAX for mapper-only inversions (no AbstractLinearObjFuncList). Refs PyAutoLabs/PyAutoArray#566 Co-Authored-By: Claude Fable 5.1 --- autoarray/config/general.yaml | 4 + autoarray/inversion/inversion/abstract.py | 40 ++ .../inversion/inversion/inversion_util.py | 95 +++++ autoarray/settings.py | 106 ++++++ autoarray/util/jax_active_set.py | 348 ++++++++++++++++++ 5 files changed, 593 insertions(+) create mode 100644 autoarray/util/jax_active_set.py diff --git a/autoarray/config/general.yaml b/autoarray/config/general.yaml index b79ea4bb9..85671dde2 100644 --- a/autoarray/config/general.yaml +++ b/autoarray/config/general.yaml @@ -10,6 +10,10 @@ inversion: 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_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. + certified_pass_budget: 16 # Maximum restricted active-set passes of the "certified" solver. Measured passes to certification: rectangular <= 11, Delaunay <= 7; the loop exits at certification so unused budget is free. + certified_fallback: pdip # What an uncertified (budget-exhausted) "certified" solve returns: "pdip" runs the PDIP solve instead (via lax.cond -- under vmap both solvers then run for every lane), "none" returns the last uncertified iterate. + certified_tau_rel: 1.0e-9 # Relative KKT tolerance of the "certified" solver: primal violation x < -tau*max|x|, dual violation g < -tau*max|q|. reconstruction_vmax_factor: 0.5 # Plots of an Inversion's reconstruction use the reconstructed data's bright value multiplied by this factor. log_det_method: cholesky # How the Bayesian-evidence log-determinant terms are computed. "cholesky" (default) is the historical 2*sum(log(diag(cholesky(M)))); "slogdet" uses logabsdet of slogdet(M), which is identical where M is positive-definite but finite (not NaN) where the Cholesky fails, for gradient-based searches (opt-in, non-default; does not change the default evidence). Under "slogdet" the kernel regularization schemes (Matern/Gaussian/Exponential) also compute the regularization log-det analytically from a Cholesky of their covariance instead of factorizing the formed inverse. See PyAutoArray#391. regularization_term_method: matmul # How the Bayesian-evidence regularization term s^T H s is computed. "matmul" (default) is the historical s @ (H @ s) against the explicitly formed regularization matrix; "cho_solve" evaluates coefficient * s^T C^-1 s for the kernel schemes (Matern/Gaussian/Exponential/MaternAdapt) via one Cholesky solve of their covariance C, avoiding the explicit inverse whose round-off is amplified by cond(C) (~1e9 on clustered traced mesh vertices). Opt-in, non-default; does not change the default evidence. Schemes with no such factorization fall back to the formed matrix. diff --git a/autoarray/inversion/inversion/abstract.py b/autoarray/inversion/inversion/abstract.py index 5d29741e6..11084c599 100644 --- a/autoarray/inversion/inversion/abstract.py +++ b/autoarray/inversion/inversion/abstract.py @@ -10,6 +10,7 @@ from autoarray.dataset.imaging.dataset import Imaging from autoarray.dataset.interferometer.dataset import Interferometer from autoarray.inversion.inversion.dataset_interface import DatasetInterface +from autoarray.inversion.linear_obj.func_list import AbstractLinearObjFuncList from autoarray.inversion.linear_obj.linear_obj import LinearObj from autoarray.inversion.mappers.abstract import Mapper from autoarray.inversion.regularization.abstract import AbstractRegularization @@ -565,6 +566,40 @@ def solve_ids_to_keep(self) -> Optional[np.ndarray]: return self.zeroed_ids_to_keep + @property + def positive_only_solver_used(self) -> str: + """ + The positive-only solver `reconstruction` passes to `reconstruction_positive_only_from`: + ``"certified"`` or ``"pdip"``. + + ``"certified"`` (the certified active-set solve, :mod:`autoarray.util.jax_active_set`) is selected + only when all of the following hold, and ``"pdip"`` (today's solver) otherwise: + + - `Settings.positive_only_solver` is ``"certified"`` (opt-in; the packaged default is ``"pdip"``); + - the inversion runs on the JAX backend -- the NumPy path always runs fnnls (a NumPy certified + scheme measured slower than fnnls with its warm-start memo); + - it contains a `Mapper` and **no** `AbstractLinearObjFuncList` (linear light profiles / MGE). + Dense MGE coefficient blocks make the active-set scheme converge poorly (a 60-MGE + + Delaunay-1500 system failed to certify in 40 passes where PDIP took 22 iterations), so those + inversions keep PDIP. + + On the NumPy backend the returned name has no effect (`reconstruction_positive_only_from` ignores + it there); it is still ``"pdip"`` so the property reports the solver family that ran on JAX. + """ + if self.settings.positive_only_solver != "certified": + return "pdip" + + if not self.use_jax: + return "pdip" + + if not self.has(cls=Mapper): + return "pdip" + + if self.has(cls=AbstractLinearObjFuncList): + return "pdip" + + return "certified" + 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 @@ -630,6 +665,9 @@ def reconstruction(self) -> np.ndarray: # no solve completed. factor = {} + # "certified" only for mapper-only JAX inversions -- see `positive_only_solver_used`. + solver = self.positive_only_solver_used + if ids_to_keep is not None: # Use advanced indexing to select rows/columns @@ -649,6 +687,7 @@ def reconstruction(self) -> np.ndarray: ids_to_keep=ids_to_keep ), factor=factor, + solver=solver, ) ) @@ -677,6 +716,7 @@ def reconstruction(self) -> np.ndarray: xp=self._xp, fingerprint=self._nnls_warm_start_fingerprint(), factor=factor, + solver=solver, ) self._nnls_factor = factor diff --git a/autoarray/inversion/inversion/inversion_util.py b/autoarray/inversion/inversion/inversion_util.py index 9485bb240..43bfb5e8d 100644 --- a/autoarray/inversion/inversion/inversion_util.py +++ b/autoarray/inversion/inversion/inversion_util.py @@ -248,6 +248,46 @@ def reconstruction_positive_negative_from( raise +def _certified_positive_only_from( + Q, q, settings, target_kappa, solver_tol, max_iter, stats=None +): + """ + The JAX certified active-set positive-only solve of ``(Q, q)`` with the PDIP solve as its fallback. + + Called by `reconstruction_positive_only_from` on its (usually Jacobi-scaled) system; see + :mod:`autoarray.util.jax_active_set` for the algorithm, budgets, gradient contract and vmap caveat. + """ + from autoarray.util.jax_active_set import solve_certified_with_fallback + from autoarray.util.jax_nnls import solve_nnls_primal + + settings = settings or Settings() + + def pdip_fn(): + return solve_nnls_primal( + Q, + q, + target_kappa=target_kappa, + solver_tol=solver_tol, + max_iter=max_iter, + ) + + x, certified, passes = solve_certified_with_fallback( + Q, + q, + pdip_fn=pdip_fn, + fallback=settings.certified_fallback == "pdip", + pass_budget=int(settings.certified_pass_budget), + tau_rel=float(settings.certified_tau_rel), + ) + + if stats is not None: + stats["solver"] = "certified" + stats["certified"] = certified + stats["passes"] = passes + + return x + + def reconstruction_positive_only_from( data_vector: np.ndarray, curvature_reg_matrix: np.ndarray, @@ -255,6 +295,8 @@ def reconstruction_positive_only_from( xp=np, fingerprint=None, factor: Optional[dict] = None, + solver: str = "pdip", + stats: Optional[dict] = None, ): """ Solve the linear system Eq.(2) (in terms of minimizing the quadratic value) of @@ -303,6 +345,20 @@ def reconstruction_positive_only_from( `fnnls_cholesky` calls below, so a memo-seeded attempt that raises cannot leave the factor of a solve whose result was discarded behind for the retry's caller to read. Purely observational: the returned reconstruction is byte-identical whether or not it is passed. + solver + Which positive-only solver the JAX (`xp=jnp`) path uses: ``"pdip"`` (default, the jaxnnls primal-dual + interior-point solve, byte-identical to before this option existed) or ``"certified"`` (the certified + active-set solve of :mod:`autoarray.util.jax_active_set`, applied to the same Jacobi-scaled system, with + the PDIP solve as its fallback when ``settings.certified_fallback == "pdip"`` and the pass budget + ``settings.certified_pass_budget`` is exhausted). The caller (`AbstractInversion.reconstruction`) passes + ``"certified"`` only for mapper-only inversions on the JAX backend -- see + `AbstractInversion.positive_only_solver_used`. The NumPy path ignores it and always runs fnnls. + stats + Optional out-dict for solver observability on the JAX path. With ``solver="certified"`` it receives + ``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. Notes ----- @@ -325,6 +381,11 @@ def reconstruction_positive_only_from( # other matrix. `fnnls_cholesky` publishes only on a successful return. factor.clear() + if solver not in ("pdip", "certified"): + raise ValueError( + f"solver={solver!r} is not a valid positive-only solver; expected 'pdip' or 'certified'." + ) + if xp.__name__.startswith("jax"): from autonerves import conf @@ -373,6 +434,24 @@ def reconstruction_positive_only_from( D = 1.0 / d Q_pc = (curvature_reg_matrix * D[:, None]) * D[None, :] q_pc = data_vector * D + + if solver == "certified": + return ( + _certified_positive_only_from( + Q=Q_pc, + q=q_pc, + settings=settings, + target_kappa=target_kappa, + solver_tol=solver_tol, + max_iter=max_iter, + stats=stats, + ) + * D + ) + + if stats is not None: + stats["solver"] = "pdip" + return ( solve_nnls_primal( Q_pc, @@ -384,6 +463,20 @@ def reconstruction_positive_only_from( * D ) + if solver == "certified": + return _certified_positive_only_from( + Q=curvature_reg_matrix, + q=data_vector, + settings=settings, + target_kappa=target_kappa, + solver_tol=solver_tol, + max_iter=max_iter, + stats=stats, + ) + + if stats is not None: + stats["solver"] = "pdip" + return solve_nnls_primal( curvature_reg_matrix, data_vector, @@ -392,6 +485,8 @@ def reconstruction_positive_only_from( max_iter=max_iter, ) + # `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. try: from autoarray.util.fnnls import fnnls_cholesky diff --git a/autoarray/settings.py b/autoarray/settings.py index 226a21128..19fabe8fa 100644 --- a/autoarray/settings.py +++ b/autoarray/settings.py @@ -22,6 +22,10 @@ def __init__( log_det_method: Optional[str] = None, regularization_term_method: Optional[str] = None, interferometer_numba_nnz_per_source_max: Optional[float] = None, + positive_only_solver: Optional[str] = None, + certified_pass_budget: Optional[int] = None, + certified_fallback: Optional[str] = None, + certified_tau_rel: Optional[float] = None, ): """ The settings of an Inversion, customizing how a linear set of equations are solved for. @@ -191,6 +195,34 @@ def __init__( (default) reads the packaged value (`60.0`); `0` disables the numba path. See the property of the same name for the measured crossovers and why the constant is machine-dependent. + positive_only_solver + Which solver the JAX (`xp=jnp`) positive-only reconstruction uses. `None` (default) reads the packaged + value ``"pdip"``. + + - ``"pdip"`` (default) — the jaxnnls primal-dual interior-point solve, unchanged. + - ``"certified"`` — the certified active-set solve (:mod:`autoarray.util.jax_active_set`): a + budgeted ``lax.while_loop`` of masked Cholesky solves that stops once the iterate satisfies the + primal and dual (KKT) conditions, with the PDIP solve as a fallback when the budget is exhausted. + Its gradient is the exact implicit active-set derivative. Measured 1.2-2.6x faster than PDIP on + source-only inversions returning the same constrained optimum (PyAutoArray#566). + + ``"certified"`` is applied **only** on the JAX backend to **mapper-only** inversions (no linear + light-profile / MGE coefficients, which converge poorly under the active-set scheme); every other + inversion, and the whole NumPy path, keeps its existing solver + (`AbstractInversion.positive_only_solver_used` records the decision). Opt-in until the batched + (``vmap``) policy is measured. + certified_pass_budget + Maximum number of restricted active-set passes of the ``"certified"`` solver. `None` (default) reads + the packaged value (`16`). Measured passes to certification: rectangular <= 11, Delaunay <= 7; the + loop exits at certification, so unused budget costs nothing. + certified_fallback + What the ``"certified"`` solver returns when it exhausts its budget uncertified. `None` (default) + reads the packaged value ``"pdip"`` (run the PDIP solve instead, via ``lax.cond``; under ``vmap`` that + ``cond`` executes both solvers for every lane). ``"none"`` returns the last, uncertified iterate. + certified_tau_rel + 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`). """ self.use_mixed_precision = use_mixed_precision self.nnls_solver_tol = nnls_solver_tol @@ -208,6 +240,16 @@ def __init__( self._interferometer_numba_nnz_per_source_max = ( interferometer_numba_nnz_per_source_max ) + self._positive_only_solver = positive_only_solver + self._certified_pass_budget = certified_pass_budget + self._certified_fallback = certified_fallback + self._certified_tau_rel = certified_tau_rel + + # 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 @property def use_positive_only_solver(self): @@ -343,3 +385,67 @@ def interferometer_numba_nnz_per_source_max(self) -> float: return 60.0 return self._interferometer_numba_nnz_per_source_max + + def _inversion_config_value(self, key, default): + # A workspace `general.yaml` normally omits the newer keys, so autoconf's config-path + # list falls through to autoarray's packaged value. The fallback fires only when the + # workspace config is the sole config path (isolated test configs push one dir) and + # returns that same packaged value, so both routes resolve identically. + try: + return conf.instance["general"]["inversion"][key] + except KeyError: + return default + + @property + def positive_only_solver(self) -> str: + """ + Which solver the JAX positive-only reconstruction uses: ``"pdip"`` or ``"certified"``. + + See the constructor docstring; ``"certified"`` is only applied to mapper-only JAX inversions. + """ + value = self._positive_only_solver + if value is None: + value = self._inversion_config_value("positive_only_solver", "pdip") + + if value not in ("pdip", "certified"): + raise ValueError( + f"positive_only_solver={value!r} is invalid; expected 'pdip' or 'certified'." + ) + + return value + + @property + def certified_pass_budget(self) -> int: + """ + Maximum number of restricted active-set passes of the ``"certified"`` solver. + """ + if self._certified_pass_budget is None: + return self._inversion_config_value("certified_pass_budget", 16) + + return self._certified_pass_budget + + @property + def certified_fallback(self) -> str: + """ + What an exhausted ``"certified"`` solve returns: ``"pdip"`` (the PDIP solve) or ``"none"``. + """ + value = self._certified_fallback + if value is None: + value = self._inversion_config_value("certified_fallback", "pdip") + + if value not in ("pdip", "none"): + raise ValueError( + f"certified_fallback={value!r} is invalid; expected 'pdip' or 'none'." + ) + + return value + + @property + def certified_tau_rel(self) -> float: + """ + Relative KKT tolerance of the ``"certified"`` solver's certificate. + """ + if self._certified_tau_rel is None: + return self._inversion_config_value("certified_tau_rel", 1.0e-9) + + return self._certified_tau_rel diff --git a/autoarray/util/jax_active_set.py b/autoarray/util/jax_active_set.py new file mode 100644 index 000000000..a76e56718 --- /dev/null +++ b/autoarray/util/jax_active_set.py @@ -0,0 +1,348 @@ +""" +Certified active-set positive-only (NNLS) solver for the JAX backend. + +Solves the non-negative quadratic program + + minimise 1/2 x^T Q x - q^T x subject to x >= 0 + +for a symmetric positive-definite ``Q`` (an inversion's +``curvature_reg_matrix``, usually Jacobi-scaled to a unit diagonal) and a +vector ``q`` (its ``data_vector``). It is an alternative to the primal-dual +interior-point (PDIP) solve in :mod:`autoarray.util.jax_nnls`, selected by +``Settings.positive_only_solver = "certified"``. + +Algorithm +--------- +A "free-all" active-set scheme on a boolean *fixed set* ``Z`` (the indices +held at exactly zero): + +1. **Pass 0** — the unconstrained solve ``x0 = Q^-1 q`` (with the permanent + set held at zero). If it has no negative entry it is already the answer. +2. **Seed** ``Z = permanent | (x0 < 0)``. +3. **Each pass** solves the restricted system on the free set ``F = ~Z`` and + certifies the iterate (see below). If it certifies, the loop stops; if not, + every free index with a violating negative value is added to ``Z`` and every + fixed (non-permanent) index with a violating negative gradient is released, + ``Z <- (Z | primal_violations) & ~dual_violations``, and the next pass runs. + +Each pass is a **masked, full-size** Cholesky solve (:func:`masked_solve`): +``Q`` with ``Z``'s rows/columns replaced by the identity and ``q`` zeroed on +``Z``, so the solve returns ``x_F = Q[F, F]^-1 q_F`` and ``x_Z = 0`` exactly. +That is costlier per pass than factorising the free block alone, but it is the +only form with a static shape — the free set changes every pass, and an index +list would retrace under ``jax.jit`` where a boolean mask does not. + +Certification +------------- +An iterate ``x`` with gradient ``g = Q x - q`` is certified optimal when the +KKT conditions hold to a relative tolerance ``tau_rel``: + +- **primal**: no free entry is negative beyond ``tau_x = tau_rel * max|x|``; +- **dual**: no fixed, releasable entry has a gradient below + ``-tau_g = -tau_rel * max|q|`` (such an index would lower the objective if it + were allowed to become positive). + +Stationarity on the free set holds by construction (the free block is solved +exactly). Permanent indices (e.g. the library's edge-zeroed pixels) are seeded +into ``Z`` and never released, so the certificate is a certificate for the +restricted problem the caller asked for. + +Budgets +------- +The search is a ``lax.while_loop`` that exits as soon as an iterate +certifies, so an easy system pays only the passes it needs. ``pass_budget`` +caps the worst case. Measured passes to certification on source-only HST +inversions (autolens_profiling fixed-lens-light campaign, PyAutoArray#566): +rectangular <= 11, Delaunay <= 7, and a fiducial Euclid rectangular system +reaches 11. The default budget of 16 covers every measured draw with margin; +because the loop exits early the unused budget costs nothing. Inversions that +include dense linear light-profile / MGE coefficients converge far worse +(60-MGE + Delaunay-1500 failed to certify in 40 passes) and are therefore +never dispatched here — see ``AbstractInversion.positive_only_solver_used``. + +Fallback +-------- +:func:`solve_certified_with_fallback` returns the certified iterate when the +search certified within budget and otherwise the result of ``pdip_fn`` (the +library's PDIP solve), via ``lax.cond``. With ``fallback=False`` the last +iterate is returned with ``certified=False`` flagged — it is feasible on the +fixed set but not proven optimal. + +**vmap caveat:** under ``jax.vmap`` a ``lax.cond`` with a batched predicate +is lowered to a ``select`` that executes *both* branches for every lane, so +the fallback PDIP solve runs for the whole batch even when every lane +certified (measured 44.6 vs 31.1 ms/lane at B=16 on an A100). Likewise the +``while_loop`` runs until the slowest lane certifies. The batched policy +(fallback ``"pdip"`` vs ``"none"`` under ``jit(vmap)``) and the production +default are phase B of the ``certified-positive-solver`` epic +(``PyAutoMind/draft/research/autolens_profiling/certified_solver_production_default.md``); +until then this solver is opt-in. + +Gradient contract +----------------- +``lax.while_loop`` is not reverse-mode differentiable, and the active set is a +piecewise-constant function of ``(Q, q)`` anyway. The search therefore runs on +``jax.lax.stop_gradient`` copies of its inputs and returns only the boolean +fixed set. The returned solution is then recomputed by one final +:func:`masked_solve` on that set, **outside** the ``stop_gradient``, so +autodiff through that single Cholesky solve yields the exact implicit +derivative of the NNLS solution on its (locally constant) active set: +``dx_F = Q[F,F]^-1 (dq_F - dQ[F,:] x)`` and ``dx_Z = 0``. This is the true +derivative wherever the active set is locally stable (strict complementarity). +The PDIP path's ``custom_vjp`` instead differentiates a relaxed central-path +KKT system (``target_kappa``) and is an approximation to the same quantity, so +the two gradients agree closely but not bit-for-bit. The final solve costs one +extra Cholesky factorisation over the search. + +Precision +--------- +The solve runs in float64 (the inputs are promoted), consistent with the +PDIP path: active-set and interior-point solvers are sensitive to fp32 noise +on ill-conditioned source meshes, so ``use_mixed_precision`` never reaches the +NNLS. + +Why the NumPy path is unchanged +------------------------------- +The NumPy backend keeps ``fnnls_cholesky`` with its cross-evaluation +warm-start memo: a NumPy port of this scheme measured 3-7 % slower than +fnnls, and factor reuse lost to the memo (PyAutoArray#566). + +JAX is imported inside functions, never at module level (see +``docs/agents/jax_and_decorators.md``); this module must only be called on the +``xp=jnp`` path. +""" + + +def _as_float64(Q, q): + import jax.numpy as jnp + + dtype = jnp.result_type(Q.dtype, q.dtype, jnp.float64) + return jnp.asarray(Q, dtype=dtype), jnp.asarray(q, dtype=dtype) + + +def masked_solve(Q, q, fixed): + """ + Solve ``Q x = q`` on the free set ``~fixed`` with ``x`` exactly zero on + ``fixed``, at full static size. + + ``Q``'s fixed rows and columns are replaced by the identity and ``q`` is + zeroed on ``fixed``, so the block-diagonal system returns + ``x_F = Q[F, F]^-1 q_F`` and ``x_Z = 0`` exactly in one ``(n, n)`` + Cholesky factorisation. + + Parameters + ---------- + Q + The ``(n, n)`` symmetric positive-definite matrix. + q + The ``(n,)`` right-hand side. + fixed + Boolean ``(n,)`` mask of the indices held at zero. + + Returns + ------- + The ``(n,)`` restricted solution. + """ + import jax.numpy as jnp + from jax.scipy.linalg import cho_solve + + keep = ~fixed + Q_masked = jnp.where(keep[:, None] & keep[None, :], Q, 0.0) + Q_masked = Q_masked + jnp.diag(jnp.where(fixed, 1.0, 0.0).astype(Q.dtype)) + q_masked = jnp.where(fixed, 0.0, q) + L = jnp.linalg.cholesky(Q_masked) + return cho_solve((L, True), q_masked) + + +def certify(Q, q, x, fixed, permanent, tau_rel): + """ + Check the KKT conditions of an active-set iterate. + + Parameters + ---------- + Q, q + The quadratic program. + x + The iterate, the :func:`masked_solve` of ``fixed``. + fixed + Boolean mask of the indices ``x`` holds at zero. + permanent + Boolean mask of fixed indices that may never be released. + tau_rel + Relative tolerance: primal violations are ``x < -tau_rel * max|x|`` on + the free set, dual violations are ``g < -tau_rel * max|q|`` on the + releasable fixed set, with ``g = Q x - q``. + + Returns + ------- + ``(certified, primal_violations, dual_violations)``: a scalar boolean and + two boolean ``(n,)`` masks. + """ + import jax.numpy as jnp + + g = Q @ x - q + tau_x = tau_rel * jnp.max(jnp.abs(x)) + tau_g = tau_rel * jnp.max(jnp.abs(q)) + + free = ~fixed + freeable = fixed & ~permanent + primal_violations = free & (x < -tau_x) + dual_violations = freeable & (g < -tau_g) + certified = ~(jnp.any(primal_violations) | jnp.any(dual_violations)) + + return certified, primal_violations, dual_violations + + +def active_set_search(Q, q, permanent=None, pass_budget=16, tau_rel=1.0e-9): + """ + Find the certified fixed (active) set of the NNLS problem. + + Runs on ``stop_gradient`` copies of ``Q`` and ``q`` — no reverse-mode path + goes through the ``while_loop``; see the module docstring's gradient + contract. + + Parameters + ---------- + Q, q + The quadratic program (``Q`` symmetric positive-definite). + permanent + Optional boolean ``(n,)`` mask of indices held at zero and never + released. ``None`` solves the pure problem. + pass_budget + Maximum number of restricted passes after pass 0 (static Python int). + tau_rel + Relative certification tolerance (see :func:`certify`). + + Returns + ------- + ``(fixed, certified, passes)``: the final boolean fixed set (the certified + one if ``certified``, otherwise the set the next pass would have solved), + the certification flag, and the number of restricted passes run (``0`` + when the unconstrained solve was already non-negative). + """ + import jax + import jax.numpy as jnp + + Q, q = _as_float64(Q, q) + Q = jax.lax.stop_gradient(Q) + q = jax.lax.stop_gradient(q) + + n = q.shape[-1] + + if permanent is None: + permanent = jnp.zeros(n, dtype=bool) + else: + permanent = jnp.asarray(permanent, dtype=bool) + + x0 = masked_solve(Q, q, permanent) + negative0 = x0 < 0.0 + fixed0 = permanent | negative0 + + # With no negative entry the pass-0 solve is the restricted solve of + # `fixed0 == permanent`, it has no releasable fixed index and no primal + # violation, so it is already certified and no restricted pass is needed. + certified0 = ~jnp.any(negative0) + + def cond_fun(carry): + _, certified, passes = carry + return (~certified) & (passes < pass_budget) + + def body_fun(carry): + fixed, _, passes = carry + x = masked_solve(Q, q, fixed) + certified, primal_violations, dual_violations = certify( + Q, q, x, fixed, permanent, tau_rel + ) + fixed_new = (fixed | primal_violations) & ~dual_violations + fixed_next = jnp.where(certified, fixed, fixed_new) + return fixed_next, certified, passes + 1 + + init = (fixed0, certified0, jnp.asarray(0, dtype=jnp.int32)) + + return jax.lax.while_loop(cond_fun, body_fun, init) + + +def solve_certified(Q, q, permanent=None, pass_budget=16, tau_rel=1.0e-9): + """ + Solve the NNLS problem with the certified active-set scheme. + + The active set is found by :func:`active_set_search` (under + ``stop_gradient``); the returned solution is one final differentiable + :func:`masked_solve` on that set, so ``jax.grad`` yields the exact implicit + active-set derivative. ``x`` is exactly zero on the fixed set. + + Parameters + ---------- + Q, q + The quadratic program (``Q`` symmetric positive-definite). + permanent + Optional boolean mask of indices held at zero and never released. + pass_budget + Maximum number of restricted passes after pass 0 (static Python int). + tau_rel + Relative certification tolerance. + + Returns + ------- + ``(x, certified, passes)``. When ``certified`` is ``False`` the budget was + exhausted and ``x`` is the (feasible on the fixed set, but unproven) + solve of the last fixed set. + """ + Q, q = _as_float64(Q, q) + + fixed, certified, passes = active_set_search( + Q, q, permanent=permanent, pass_budget=pass_budget, tau_rel=tau_rel + ) + + x = masked_solve(Q, q, fixed) + + return x, certified, passes + + +def solve_certified_with_fallback( + Q, + q, + pdip_fn, + fallback=True, + pass_budget=16, + tau_rel=1.0e-9, + permanent=None, +): + """ + Certified active-set solve with a fallback for an exhausted budget. + + Parameters + ---------- + Q, q + The quadratic program (``Q`` symmetric positive-definite). + pdip_fn + A zero-argument callable returning the fallback solution of the same + problem (the library's PDIP solve), with the same shape and dtype as + ``q``. + fallback + If ``True``, an uncertified iterate is discarded and ``pdip_fn()`` is + returned instead via ``lax.cond``. Under ``vmap`` that ``cond`` + executes both branches for every lane (see the module docstring). If + ``False``, the last iterate is returned with ``certified=False``. + pass_budget + Maximum number of restricted passes after pass 0 (static Python int). + tau_rel + Relative certification tolerance. + permanent + Optional boolean mask of indices held at zero and never released. + + Returns + ------- + ``(x, certified, passes)`` — ``certified`` is the search's flag, so with + ``fallback=True`` a ``False`` means ``x`` came from ``pdip_fn``. + """ + import jax + + x, certified, passes = solve_certified( + Q, q, permanent=permanent, pass_budget=pass_budget, tau_rel=tau_rel + ) + + if fallback: + x = jax.lax.cond(certified, lambda: x, pdip_fn) + + return x, certified, passes From 8fe4343202121e1ca976cf5a652cbf404c4463d9 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Wed, 23 Sep 2026 14:53:15 +0200 Subject: [PATCH 2/3] test(inversion): certified active-set solver numerics, fallback, jit/vmap, gradients, dispatch, settings Refs PyAutoLabs/PyAutoArray#566 Co-Authored-By: Claude Fable 5.1 --- .../inversion/test_inversion_util.py | 157 ++++++++ .../inversion/test_positive_only_dispatch.py | 211 +++++++++++ .../inversion/inversion/test_settings_dict.py | 46 +++ test_autoarray/util/test_jax_active_set.py | 358 ++++++++++++++++++ 4 files changed, 772 insertions(+) create mode 100644 test_autoarray/inversion/inversion/test_positive_only_dispatch.py create mode 100644 test_autoarray/util/test_jax_active_set.py diff --git a/test_autoarray/inversion/inversion/test_inversion_util.py b/test_autoarray/inversion/inversion/test_inversion_util.py index 4bbbb7fcb..f16e23582 100644 --- a/test_autoarray/inversion/inversion/test_inversion_util.py +++ b/test_autoarray/inversion/inversion/test_inversion_util.py @@ -384,3 +384,160 @@ def test__curvature_matrix_via_mapping_matrix_from__mixed_precision__weights_are ) assert jax_curvature == pytest.approx(expected, rel=1.0e-14) + + +# ---------------------------------------------------------------------------- +# The `solver` switch of `reconstruction_positive_only_from` (PyAutoArray#566). +# ---------------------------------------------------------------------------- + + +def _positive_only_system(n=30, seed=7): + """An SPD system whose unconstrained solution has negative entries, so the constraint binds.""" + rng = np.random.default_rng(seed) + A = rng.normal(size=(3 * n, n)) + b = rng.normal(size=3 * n) + curvature_reg_matrix = A.T @ A + data_vector = A.T @ b + assert (np.linalg.solve(curvature_reg_matrix, data_vector) < 0.0).any() + return data_vector, curvature_reg_matrix + + +@requires_jax +@pytest.mark.parametrize("use_jacobi", [True, False]) +def test__reconstruction_positive_only_from__certified_matches_pdip_on_jax( + use_jacobi, monkeypatch +): + import jax + + jax.config.update("jax_enable_x64", True) + import jax.numpy as jnp + from autonerves import conf + + if not use_jacobi: + # Exercise the unpreconditioned branch too: the `solver` switch lives in both. The dict stands in for + # the config; the Settings' certified_* keys resolve through their KeyError defaults. + monkeypatch.setattr( + conf, + "instance", + { + "general": { + "inversion": { + "nnls_jacobi_preconditioning": False, + "nnls_target_kappa": 1.0e-11, + } + } + }, + ) + + for data_vector, curvature_reg_matrix in [ + ( + np.array([1.0, 1.0, 2.0]), + np.array([[2.0, 1.0, 0.0], [1.0, 3.0, 1.0], [0.0, 1.0, 1.0]]), + ), + _positive_only_system(), + ]: + pdip = np.asarray( + aa.util.inversion.reconstruction_positive_only_from( + data_vector=jnp.asarray(data_vector), + curvature_reg_matrix=jnp.asarray(curvature_reg_matrix), + settings=aa.Settings(), + xp=jnp, + ) + ) + + stats = {} + certified = np.asarray( + aa.util.inversion.reconstruction_positive_only_from( + data_vector=jnp.asarray(data_vector), + curvature_reg_matrix=jnp.asarray(curvature_reg_matrix), + settings=aa.Settings(), + xp=jnp, + solver="certified", + stats=stats, + ) + ) + + assert stats["solver"] == "certified" + assert bool(stats["certified"]) + assert int(stats["passes"]) >= 1 + + assert (certified >= 0.0).all() + assert certified == pytest.approx( + pdip, rel=1.0e-8, abs=1.0e-8 * np.max(np.abs(certified)) + ) + + +@requires_jax +def test__reconstruction_positive_only_from__certified_budget_exhausted_falls_back_to_pdip(): + import jax + + jax.config.update("jax_enable_x64", True) + import jax.numpy as jnp + + data_vector, curvature_reg_matrix = _positive_only_system() + + kwargs = dict( + data_vector=jnp.asarray(data_vector), + curvature_reg_matrix=jnp.asarray(curvature_reg_matrix), + xp=jnp, + ) + + pdip = np.asarray( + aa.util.inversion.reconstruction_positive_only_from( + settings=aa.Settings(), **kwargs + ) + ) + + stats = {} + fallback = np.asarray( + aa.util.inversion.reconstruction_positive_only_from( + settings=aa.Settings(certified_pass_budget=0), + solver="certified", + stats=stats, + **kwargs, + ) + ) + + assert not bool(stats["certified"]) + assert np.array_equal(fallback, pdip) + + +def test__reconstruction_positive_only_from__numpy_path_ignores_solver(monkeypatch): + # The NumPy path must run fnnls whatever `solver` says: spy on it. + from autoarray.util import fnnls + + calls = [] + original = fnnls.fnnls_cholesky + + def spy(*args, **kwargs): + calls.append(1) + return original(*args, **kwargs) + + monkeypatch.setattr(fnnls, "fnnls_cholesky", spy) + + data_vector = np.array([1.0, 1.0, 2.0]) + curvature_reg_matrix = np.array([[2.0, 1.0, 0.0], [1.0, 3.0, 1.0], [0.0, 1.0, 1.0]]) + + results = [] + for solver in ["pdip", "certified"]: + calls.clear() + results.append( + aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + solver=solver, + ) + ) + assert len(calls) == 1 + + assert np.array_equal(results[0], results[1]) + assert results[0] == pytest.approx(np.array([0.5, 0.0, 2.0]), 1.0e-4) + + +def test__reconstruction_positive_only_from__invalid_solver_raises(): + with pytest.raises(ValueError, match="solver"): + aa.util.inversion.reconstruction_positive_only_from( + data_vector=np.array([1.0]), + curvature_reg_matrix=np.array([[1.0]]), + solver="bogus", + ) diff --git a/test_autoarray/inversion/inversion/test_positive_only_dispatch.py b/test_autoarray/inversion/inversion/test_positive_only_dispatch.py new file mode 100644 index 000000000..52198f966 --- /dev/null +++ b/test_autoarray/inversion/inversion/test_positive_only_dispatch.py @@ -0,0 +1,211 @@ +""" +Which positive-only solver `AbstractInversion.reconstruction` dispatches to (PyAutoArray#566). + +`Settings.positive_only_solver = "certified"` selects the certified active-set solve only on the JAX +backend for mapper-only inversions; inversions with linear light-profile / MGE coefficients +(`AbstractLinearObjFuncList`) and every NumPy inversion keep today's solver. +""" + +import importlib.util + +import numpy as np +import pytest + +import autoarray as aa + + +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)", +) + + +def _system(n, seed=3): + """An SPD system whose unconstrained solution has negative entries, so positivity binds.""" + rng = np.random.default_rng(seed) + A = rng.normal(size=(3 * n, n)) + b = rng.normal(size=3 * n) + curvature_reg_matrix = A.T @ A + data_vector = A.T @ b + assert (np.linalg.solve(curvature_reg_matrix, data_vector) < 0.0).any() + return data_vector, curvature_reg_matrix + + +def _mapper(pixels=16): + return aa.m.MockMapper( + mesh=aa.mesh.RectangularUniform(shape=(4, 4)), + parameters=pixels, + regularization=aa.reg.Constant(), + # Only its shape is read, by the NumPy warm-start memo's fingerprint. + source_plane_mesh_grid=np.zeros((pixels, 2)), + ) + + +def _inversion( + linear_obj_list, + data_vector, + curvature_reg_matrix, + use_jax, + positive_only_solver="certified", + use_edge_zeroed_pixels=False, +): + if use_jax: + import jax + + jax.config.update("jax_enable_x64", True) + import jax.numpy as jnp + + data_vector = jnp.asarray(data_vector) + curvature_reg_matrix = jnp.asarray(curvature_reg_matrix) + + inversion = aa.m.MockInversion( + linear_obj_list=linear_obj_list, + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=aa.Settings( + use_positive_only_solver=True, + use_edge_zeroed_pixels=use_edge_zeroed_pixels, + positive_only_solver=positive_only_solver, + # 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, + ), + ) + # `MockInversion` does not take `xp`; the backend flag is what `AbstractInversion(xp=jnp)` sets. + inversion.use_jax = use_jax + return inversion + + +def test__positive_only_solver_used__default_is_pdip(): + data_vector, curvature_reg_matrix = _system(16) + + inversion = _inversion( + [_mapper()], + data_vector, + curvature_reg_matrix, + use_jax=False, + positive_only_solver=None, + ) + + assert inversion.settings.positive_only_solver == "pdip" + assert inversion.positive_only_solver_used == "pdip" + + +def test__positive_only_solver_used__numpy_backend_keeps_pdip_and_fnnls(): + data_vector, curvature_reg_matrix = _system(16) + + certified = _inversion( + [_mapper()], data_vector, curvature_reg_matrix, use_jax=False + ) + default = _inversion( + [_mapper()], + data_vector, + curvature_reg_matrix, + use_jax=False, + positive_only_solver="pdip", + ) + + assert certified.positive_only_solver_used == "pdip" + assert np.array_equal(certified.reconstruction, default.reconstruction) + + +@requires_jax +def test__positive_only_solver_used__mapper_only_jax_is_certified_and_matches_pdip(): + data_vector, curvature_reg_matrix = _system(16) + + certified = _inversion([_mapper()], data_vector, curvature_reg_matrix, use_jax=True) + pdip = _inversion( + [_mapper()], + data_vector, + curvature_reg_matrix, + use_jax=True, + positive_only_solver="pdip", + ) + + assert certified.positive_only_solver_used == "certified" + assert pdip.positive_only_solver_used == "pdip" + + x_certified = np.asarray(certified.reconstruction) + x_pdip = np.asarray(pdip.reconstruction) + + assert (x_certified >= 0.0).all() + assert (x_certified == 0.0).any() + assert x_certified == pytest.approx( + x_pdip, rel=1.0e-8, abs=1.0e-8 * np.max(np.abs(x_certified)) + ) + + +@requires_jax +def test__positive_only_solver_used__linear_func_list_present_keeps_pdip(): + data_vector, curvature_reg_matrix = _system(17) + + func_list = aa.m.MockLinearObjFuncList(parameters=1) + + certified_setting = _inversion( + [func_list, _mapper()], data_vector, curvature_reg_matrix, use_jax=True + ) + pdip = _inversion( + [func_list, _mapper()], + data_vector, + curvature_reg_matrix, + use_jax=True, + positive_only_solver="pdip", + ) + + assert certified_setting.positive_only_solver_used == "pdip" + assert np.array_equal( + np.asarray(certified_setting.reconstruction), np.asarray(pdip.reconstruction) + ) + + +@requires_jax +def test__positive_only_solver_used__no_mapper_keeps_pdip(): + data_vector, curvature_reg_matrix = _system(4) + + inversion = _inversion( + [aa.m.MockLinearObjFuncList(parameters=4)], + data_vector, + curvature_reg_matrix, + use_jax=True, + ) + + assert inversion.positive_only_solver_used == "pdip" + + +@requires_jax +def test__positive_only_solver_used__edge_zeroed_subset_is_preserved(): + # The 4x4 rectangular mesh's 12 edge pixels are zeroed and the 4 interior pixels [5, 6, 9, 10] solved. + data_vector, curvature_reg_matrix = _system(16) + + certified = _inversion( + [_mapper()], + data_vector, + curvature_reg_matrix, + use_jax=True, + use_edge_zeroed_pixels=True, + ) + pdip = _inversion( + [_mapper()], + data_vector, + curvature_reg_matrix, + use_jax=True, + positive_only_solver="pdip", + use_edge_zeroed_pixels=True, + ) + + assert certified.positive_only_solver_used == "certified" + assert np.asarray(certified.solve_ids_to_keep) == pytest.approx( + np.array([5, 6, 9, 10]) + ) + + x_certified = np.asarray(certified.reconstruction) + x_pdip = np.asarray(pdip.reconstruction) + + keep = np.array([5, 6, 9, 10]) + edge = np.setdiff1d(np.arange(16), keep) + + assert (x_certified[edge] == 0.0).all() + assert (x_pdip[edge] == 0.0).all() + assert x_certified == pytest.approx( + x_pdip, rel=1.0e-8, abs=1.0e-8 * np.max(np.abs(x_certified)) + ) diff --git a/test_autoarray/inversion/inversion/test_settings_dict.py b/test_autoarray/inversion/inversion/test_settings_dict.py index a531ea286..9ccf9acd6 100644 --- a/test_autoarray/inversion/inversion/test_settings_dict.py +++ b/test_autoarray/inversion/inversion/test_settings_dict.py @@ -63,3 +63,49 @@ def test_settings_nnls_warm_start_error_tolerance_round_trips(): settings = from_dict(to_dict(aa.Settings(nnls_warm_start_error_tolerance=2.5))) assert settings.nnls_warm_start_error_tolerance == 2.5 + + +def test_settings_positive_only_solver_keys_default_from_config(): + # The test config does not ship the keys, so these resolve through the KeyError fallbacks -- the + # production path whenever a workspace shadows autoarray's general.yaml -- which must equal the + # packaged values. + settings = aa.Settings() + + assert settings.positive_only_solver == "pdip" + assert settings.certified_pass_budget == 16 + assert settings.certified_fallback == "pdip" + assert settings.certified_tau_rel == 1.0e-9 + + import yaml + + packaged = Path(aa.__file__).parent / "config" / "general.yaml" + inversion = yaml.safe_load(packaged.read_text())["inversion"] + + assert inversion["positive_only_solver"] == "pdip" + assert inversion["certified_pass_budget"] == 16 + assert inversion["certified_fallback"] == "pdip" + assert float(inversion["certified_tau_rel"]) == 1.0e-9 + + +def test_settings_positive_only_solver_keys_round_trip(): + settings = aa.Settings( + positive_only_solver="certified", + certified_pass_budget=7, + certified_fallback="none", + certified_tau_rel=1.0e-8, + ) + + settings = from_dict(to_dict(settings)) + + assert settings.positive_only_solver == "certified" + assert settings.certified_pass_budget == 7 + assert settings.certified_fallback == "none" + assert settings.certified_tau_rel == 1.0e-8 + + +def test_settings_positive_only_solver_keys_are_validated(): + with pytest.raises(ValueError, match="positive_only_solver"): + aa.Settings(positive_only_solver="fnnls") + + with pytest.raises(ValueError, match="certified_fallback"): + aa.Settings(certified_fallback="numpy") diff --git a/test_autoarray/util/test_jax_active_set.py b/test_autoarray/util/test_jax_active_set.py new file mode 100644 index 000000000..99ae9b6eb --- /dev/null +++ b/test_autoarray/util/test_jax_active_set.py @@ -0,0 +1,358 @@ +""" +Tests of the certified active-set positive-only solver, `autoarray.util.jax_active_set`. + +Tolerances were declared before the tests were first run (PyAutoArray#566): + +- solution vs SciPy `nnls` (an independent active-set NNLS on the least-squares form): `atol 1e-10`; +- solution vs the library's PDIP solve (`solve_nnls_primal`): `rtol 1e-8` (amended after the first run: plus an + absolute floor of `1e-8 * max|x|`, because PDIP never reaches the active set's exact zeros); +- gradient vs central finite differences: `rtol 1e-6` (the NNLS solution is piecewise linear in `q`, so a + central difference inside one active-set piece is exact up to rounding); +- gradient vs `jax.grad` through PDIP's relaxed-KKT `custom_vjp`: `rtol 1e-4` (PDIP differentiates a relaxed + central-path system, an approximation to the exact active-set derivative; amended after the first run: plus + an absolute floor of `1e-4 * max|grad|`, because PDIP leaks ~1e-7 onto active coordinates whose exact + derivative is zero). +""" + +import importlib +import sys + +import numpy as np +import pytest + + +def test__jax_active_set_module_never_imports_jax_at_module_level(monkeypatch): + # Library unit tests are NumPy-only: importing the module must succeed even when jax is + # unimportable. A None entry in sys.modules makes any `import jax` raise ImportError, so a + # module-level import would fail this reload. + monkeypatch.setitem(sys.modules, "jax", None) + monkeypatch.setitem(sys.modules, "jaxnnls", None) + + module = importlib.reload(importlib.import_module("autoarray.util.jax_active_set")) + + for name in ( + "masked_solve", + "certify", + "active_set_search", + "solve_certified", + "solve_certified_with_fallback", + ): + assert hasattr(module, name) + + +# jax is an `[optional]` extra and is absent on the NumPy-only matrix env: every test below the import test +# skips there, while the import test above still runs. +if importlib.util.find_spec("jax") is None: + pytestmark = pytest.mark.skip(reason="requires jax (the [optional] extras)") + + def test__placeholder_requires_jax(): # pragma: no cover + pass + +else: + import jax + + jax.config.update("jax_enable_x64", True) + + import jax.numpy as jnp + from scipy.optimize import nnls + + from autoarray.util import jax_active_set + from autoarray.util.jax_nnls import solve_nnls_primal + + +TARGET_KAPPA = 1.0e-11 + + +def _qp(n, seed, rows_factor=3): + """ + A seeded NNLS problem in least-squares form (A, b) and its quadratic-program form (Q, q) = (A^T A, A^T b). + + A Gaussian `A` with a Gaussian `b` gives an unconstrained solution with roughly half its entries negative. + """ + rng = np.random.default_rng(seed) + A = rng.normal(size=(rows_factor * n, n)) + b = rng.normal(size=rows_factor * n) + return A, b, A.T @ A, A.T @ b + + +def _pdip(Q, q): + return solve_nnls_primal(jnp.asarray(Q), jnp.asarray(q), target_kappa=TARGET_KAPPA) + + +def _reference_passes(Q, q, permanent=None, max_passes=40, tau_rel=1.0e-9): + """ + NumPy reference of the free-all certified active-set scheme (a port of the autolens_profiling harness's + `active_set_certified`), returning the number of restricted passes to certification (0 when the + unconstrained solve is already non-negative) and the solution. + """ + n = Q.shape[0] + permanent = np.zeros(n, bool) if permanent is None else np.asarray(permanent) + + def restricted(fixed): + x = np.zeros(n) + free = ~fixed + if free.any(): + idx = np.where(free)[0] + x[idx] = np.linalg.solve(Q[np.ix_(idx, idx)], q[idx]) + return x + + x = restricted(permanent) + if not (x < 0.0).any(): + return 0, x + + fixed = permanent | (x < 0.0) + tau_g = tau_rel * np.max(np.abs(q)) + + for p in range(1, max_passes + 1): + x = restricted(fixed) + g = Q @ x - q + tau_x = tau_rel * np.max(np.abs(x)) + primal = ~fixed & (x < -tau_x) + dual = fixed & ~permanent & (g < -tau_g) + if not primal.any() and not dual.any(): + return p, x + fixed = (fixed | primal) & ~dual + + return None, x + + +@pytest.mark.parametrize("n, seed", [(20, 0), (20, 1), (60, 2), (60, 3)]) +def test__solve_certified__matches_scipy_nnls_and_pdip(n, seed): + A, b, Q, q = _qp(n, seed) + + # The fixture must exercise the constraint: the unconstrained solution has negative entries. + assert (np.linalg.solve(Q, q) < 0.0).any() + + x, certified, passes = jax_active_set.solve_certified( + jnp.asarray(Q), jnp.asarray(q) + ) + x = np.asarray(x) + + assert bool(certified) + assert x.dtype == np.float64 + assert (x >= 0.0).all() + + x_scipy, _ = nnls(A, b) + assert np.max(np.abs(x - x_scipy)) < 1.0e-10 + + # Exactly zero (not merely small) on the active set. + assert (x[x_scipy == 0.0] == 0.0).all() + + # PDIP is an interior-point method: it approaches the active set's zeros but never reaches them exactly + # (measured up to 7.6e-12 on an active entry), so the absolute floor is relative to the solution's scale. + x_pdip = np.asarray(_pdip(Q, q)) + assert x == pytest.approx(x_pdip, rel=1.0e-8, abs=1.0e-8 * np.max(np.abs(x))) + + +@pytest.mark.parametrize("n, seed", [(20, 0), (20, 1), (60, 2), (60, 3)]) +def test__solve_certified__exits_early_at_the_reference_pass_count(n, seed): + _, _, Q, q = _qp(n, seed) + + budget = 16 + _, certified, passes = jax_active_set.solve_certified( + jnp.asarray(Q), jnp.asarray(q), pass_budget=budget + ) + + reference_passes, _ = _reference_passes(Q, q) + + assert bool(certified) + assert int(passes) == reference_passes + assert int(passes) < budget + + +def test__solve_certified__non_negative_unconstrained_solution_needs_no_pass(): + Q = np.array([[2.0, 0.5], [0.5, 1.0]]) + x_true = np.array([1.0, 2.0]) + q = Q @ x_true + + x, certified, passes = jax_active_set.solve_certified( + jnp.asarray(Q), jnp.asarray(q) + ) + + assert bool(certified) + assert int(passes) == 0 + assert np.asarray(x) == pytest.approx(x_true, rel=1.0e-12) + + +def test__solve_certified__permanent_fixed_set_is_never_released(): + A, b, Q, q = _qp(20, 4) + + # Hold at zero indices the free problem would make strictly positive, so releasing them would lower the + # objective (a negative gradient) -- the scheme must hold them anyway. + x_free, _ = nnls(A, b) + positive = np.where(x_free > 0.0)[0] + permanent = np.zeros(20, bool) + permanent[positive[:3]] = True + + x, certified, _ = jax_active_set.solve_certified( + jnp.asarray(Q), jnp.asarray(q), permanent=jnp.asarray(permanent) + ) + x = np.asarray(x) + + assert bool(certified) + assert (x[permanent] == 0.0).all() + + # It is the NNLS solution of the problem with those columns removed. + keep = ~permanent + x_reduced, _ = nnls(A[:, keep], b) + assert np.max(np.abs(x[keep] - x_reduced)) < 1.0e-10 + + # And the released-if-it-were-allowed indices really would want to move. + g = Q @ x - q + assert (g[permanent] < 0.0).any() + + +def _hard_system(): + """A system the default search certifies only after more than one restricted pass.""" + _, _, Q, q = _qp(60, 2) + reference_passes, _ = _reference_passes(Q, q) + assert reference_passes >= 2 + return jnp.asarray(Q), jnp.asarray(q) + + +def test__solve_certified_with_fallback__exhausted_budget_returns_pdip_bit_exactly(): + Q, q = _hard_system() + + def pdip_fn(): + return solve_nnls_primal(Q, q, target_kappa=TARGET_KAPPA) + + x, certified, passes = jax_active_set.solve_certified_with_fallback( + Q, q, pdip_fn=pdip_fn, fallback=True, pass_budget=1 + ) + + assert not bool(certified) + assert int(passes) == 1 + assert np.array_equal(np.asarray(x), np.asarray(pdip_fn())) + + +def test__solve_certified_with_fallback__no_fallback_flags_the_uncertified_iterate(): + Q, q = _hard_system() + + def pdip_fn(): # pragma: no cover - must not be called + raise AssertionError("the fallback must not run with fallback=False") + + x, certified, passes = jax_active_set.solve_certified_with_fallback( + Q, q, pdip_fn=pdip_fn, fallback=False, pass_budget=1 + ) + + assert not bool(certified) + assert int(passes) == 1 + assert np.isfinite(np.asarray(x)).all() + + +def test__solve_certified_with_fallback__certified_returns_the_active_set_solution(): + Q, q = _hard_system() + + def pdip_fn(): + return solve_nnls_primal(Q, q, target_kappa=TARGET_KAPPA) + + x, certified, _ = jax_active_set.solve_certified_with_fallback( + Q, q, pdip_fn=pdip_fn, fallback=True + ) + x_direct, _, _ = jax_active_set.solve_certified(Q, q) + + assert bool(certified) + assert np.array_equal(np.asarray(x), np.asarray(x_direct)) + + +def test__solve_certified__jit_and_vmap_match_per_system_results(): + systems = [_qp(20, seed) for seed in range(4)] + Qs = jnp.asarray(np.stack([s[2] for s in systems])) + qs = jnp.asarray(np.stack([s[3] for s in systems])) + + scalar = [jax_active_set.solve_certified(Qs[i], qs[i]) for i in range(4)] + + jitted = jax.jit(jax_active_set.solve_certified) + for i in range(4): + x, certified, passes = jitted(Qs[i], qs[i]) + assert np.asarray(x) == pytest.approx(np.asarray(scalar[i][0]), abs=1.0e-12) + assert bool(certified) == bool(scalar[i][1]) + assert int(passes) == int(scalar[i][2]) + + x_b, certified_b, passes_b = jax.jit(jax.vmap(jax_active_set.solve_certified))( + Qs, qs + ) + + for i in range(4): + assert np.asarray(x_b[i]) == pytest.approx( + np.asarray(scalar[i][0]), abs=1.0e-12 + ) + assert bool(certified_b[i]) == bool(scalar[i][1]) + # Per-lane pass counts survive the batched while_loop. + assert int(passes_b[i]) == int(scalar[i][2]) + + +def test__solve_certified_with_fallback__vmap(): + systems = [_qp(20, seed) for seed in range(4)] + Qs = jnp.asarray(np.stack([s[2] for s in systems])) + qs = jnp.asarray(np.stack([s[3] for s in systems])) + + def solve(Q, q): + return jax_active_set.solve_certified_with_fallback( + Q, + q, + pdip_fn=lambda: solve_nnls_primal(Q, q, target_kappa=TARGET_KAPPA), + fallback=True, + )[0] + + x_b = jax.jit(jax.vmap(solve))(Qs, qs) + + for i in range(4): + x_scipy, _ = nnls(systems[i][0], systems[i][1]) + assert np.max(np.abs(np.asarray(x_b[i]) - x_scipy)) < 1.0e-10 + + +def test__solve_certified__gradient_matches_finite_differences_and_pdip(): + _, _, Q, q = _qp(20, 5) + rng = np.random.default_rng(11) + w = jnp.asarray(rng.normal(size=20)) + Q = jnp.asarray(Q) + q = jnp.asarray(q) + + def objective(q_): + return w @ jax_active_set.solve_certified(Q, q_)[0] + + grad = np.asarray(jax.grad(objective)(q)) + assert np.isfinite(grad).all() + + # The derivative is exactly zero along the active set's coordinates of q... and non-trivial elsewhere. + x = np.asarray(jax_active_set.solve_certified(Q, q)[0]) + assert (grad[x == 0.0] == 0.0).all() + assert np.abs(grad[x > 0.0]).max() > 0.0 + + h = 1.0e-6 * float(jnp.max(jnp.abs(q))) + fd = np.zeros(20) + for i in range(20): + e = jnp.zeros(20).at[i].set(h) + fd[i] = (float(objective(q + e)) - float(objective(q - e))) / (2.0 * h) + + assert grad == pytest.approx(fd, rel=1.0e-6, abs=1.0e-12) + + grad_pdip = np.asarray( + jax.grad(lambda q_: w @ solve_nnls_primal(Q, q_, target_kappa=TARGET_KAPPA))(q) + ) + # PDIP's relaxed-KKT gradient leaks a small value onto active coordinates, where the exact derivative is + # zero (measured 1.8e-7 against max|grad| ~ 1e-2), so the absolute floor is relative to the gradient's scale. + assert grad == pytest.approx( + grad_pdip, rel=1.0e-4, abs=1.0e-4 * np.max(np.abs(grad)) + ) + + +def test__solve_certified__gradient_wrt_matrix_matches_finite_differences(): + _, _, Q, q = _qp(20, 6) + rng = np.random.default_rng(12) + w = jnp.asarray(rng.normal(size=20)) + direction = rng.normal(size=(20, 20)) + direction = jnp.asarray(direction + direction.T) + Q = jnp.asarray(Q) + q = jnp.asarray(q) + + def objective(t): + return w @ jax_active_set.solve_certified(Q + t * direction, q)[0] + + derivative = float(jax.grad(objective)(0.0)) + + h = 1.0e-6 + fd = (float(objective(h)) - float(objective(-h))) / (2.0 * h) + + assert derivative == pytest.approx(fd, rel=1.0e-6) From 233cfc0d7430cd4d36538db9fee6acba92ebad51 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Wed, 23 Sep 2026 16:03:05 +0200 Subject: [PATCH 3/3] fix(inversion): review findings (#566) Move the no-module-level-jax import guard into its own test module: the module-level skip in test_jax_active_set.py (applied when jax is absent) also skipped the guard, contradicting its comment, so the guard never ran on the NumPy-only env it exists for. Co-Authored-By: Claude Opus 5.5 (1M context) --- test_autoarray/util/test_jax_active_set.py | 24 ++--------------- .../util/test_jax_active_set_import.py | 26 +++++++++++++++++++ 2 files changed, 28 insertions(+), 22 deletions(-) create mode 100644 test_autoarray/util/test_jax_active_set_import.py diff --git a/test_autoarray/util/test_jax_active_set.py b/test_autoarray/util/test_jax_active_set.py index 99ae9b6eb..dcf2b772b 100644 --- a/test_autoarray/util/test_jax_active_set.py +++ b/test_autoarray/util/test_jax_active_set.py @@ -15,33 +15,13 @@ """ import importlib -import sys import numpy as np import pytest -def test__jax_active_set_module_never_imports_jax_at_module_level(monkeypatch): - # Library unit tests are NumPy-only: importing the module must succeed even when jax is - # unimportable. A None entry in sys.modules makes any `import jax` raise ImportError, so a - # module-level import would fail this reload. - monkeypatch.setitem(sys.modules, "jax", None) - monkeypatch.setitem(sys.modules, "jaxnnls", None) - - module = importlib.reload(importlib.import_module("autoarray.util.jax_active_set")) - - for name in ( - "masked_solve", - "certify", - "active_set_search", - "solve_certified", - "solve_certified_with_fallback", - ): - assert hasattr(module, name) - - -# jax is an `[optional]` extra and is absent on the NumPy-only matrix env: every test below the import test -# skips there, while the import test above still runs. +# jax is an `[optional]` extra and is absent on the NumPy-only matrix env: every test in this module skips +# there. The no-module-level-jax-import guard lives in `test_jax_active_set_import.py` so it still runs. if importlib.util.find_spec("jax") is None: pytestmark = pytest.mark.skip(reason="requires jax (the [optional] extras)") diff --git a/test_autoarray/util/test_jax_active_set_import.py b/test_autoarray/util/test_jax_active_set_import.py new file mode 100644 index 000000000..6cdf6349e --- /dev/null +++ b/test_autoarray/util/test_jax_active_set_import.py @@ -0,0 +1,26 @@ +""" +Import guard for `autoarray.util.jax_active_set`, kept apart from `test_jax_active_set.py` because that module +skips wholesale when jax is absent, and this test must run exactly there. +""" + +import importlib +import sys + + +def test__jax_active_set_module_never_imports_jax_at_module_level(monkeypatch): + # Library unit tests are NumPy-only: importing the module must succeed even when jax is + # unimportable. A None entry in sys.modules makes any `import jax` raise ImportError, so a + # module-level import would fail this reload. + monkeypatch.setitem(sys.modules, "jax", None) + monkeypatch.setitem(sys.modules, "jaxnnls", None) + + module = importlib.reload(importlib.import_module("autoarray.util.jax_active_set")) + + for name in ( + "masked_solve", + "certify", + "active_set_search", + "solve_certified", + "solve_certified_with_fallback", + ): + assert hasattr(module, name)