diff --git a/autoarray/config/general.yaml b/autoarray/config/general.yaml index 4d75ee632..6ac374c1b 100644 --- a/autoarray/config/general.yaml +++ b/autoarray/config/general.yaml @@ -8,6 +8,8 @@ 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_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. 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 ba5b8beb1..43b915100 100644 --- a/autoarray/inversion/inversion/abstract.py +++ b/autoarray/inversion/inversion/abstract.py @@ -1,4 +1,5 @@ import copy +import hashlib import warnings import numpy as np @@ -534,6 +535,37 @@ def solve_ids_to_keep(self) -> Optional[np.ndarray]: return self.zeroed_ids_to_keep + 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 + cross-evaluation passive-set memo (`autoarray.inversion.inversion.nnls_memo`) only + reuses a seed across evaluations whose reconstruction entries mean the same thing. + + Deliberately built from shapes, NOT from the curvature matrix or data vector: the + memo exists to carry a passive set between *neighbouring* parameter points, whose + matrices differ. What must not change is the mapping from index to source pixel -- + the mesh sizes and the parameter-vector length. Under edge zeroing the passive set + lives in the subset index space, so `ids_to_keep` is part of the identity too. + + Returns `None` on the JAX path, which does not run fnnls at all. + """ + if self._xp.__name__.startswith("jax"): + return None + + parts = [str(tuple(np.shape(self.data_vector)))] + + for mapper in self.cls_list_from(cls=Mapper): + parts.append(str(tuple(np.shape(mapper.source_plane_mesh_grid)))) + + if ids_to_keep is not None: + parts.append( + hashlib.sha256( + np.ascontiguousarray(ids_to_keep, dtype=np.int64).tobytes() + ).hexdigest() + ) + + return "|".join(parts) + @cached_property def reconstruction(self) -> np.ndarray: """ @@ -576,6 +608,9 @@ def reconstruction(self) -> np.ndarray: curvature_reg_matrix=curvature_reg_matrix, settings=self.settings, xp=self._xp, + fingerprint=self._nnls_warm_start_fingerprint( + ids_to_keep=ids_to_keep + ), ) ) @@ -599,6 +634,7 @@ def reconstruction(self) -> np.ndarray: curvature_reg_matrix=self.curvature_reg_matrix, settings=self.settings, xp=self._xp, + fingerprint=self._nnls_warm_start_fingerprint(), ) return inversion_util.reconstruction_positive_negative_from( diff --git a/autoarray/inversion/inversion/inversion_util.py b/autoarray/inversion/inversion/inversion_util.py index 84d22719e..50caf1533 100644 --- a/autoarray/inversion/inversion/inversion_util.py +++ b/autoarray/inversion/inversion/inversion_util.py @@ -258,6 +258,7 @@ def reconstruction_positive_only_from( curvature_reg_matrix: np.ndarray, settings: Settings = None, xp=np, + fingerprint=None, ): """ Solve the linear system Eq.(2) (in terms of minimizing the quadratic value) of @@ -296,6 +297,19 @@ def reconstruction_positive_only_from( settings Controls the settings of the inversion, for this function where the solution is checked to not be all the same values.\ + fingerprint + Identifies the index space this solve's passive set lives in, enabling the cross-evaluation warm-start + memo (`Settings.nnls_warm_start_memo`) on the NumPy path. `None` disables the memo for this call. + + Notes + ----- + On the NumPy path this function writes two keys into the `stats` dict it passes to `fnnls_cholesky` + that `fnnls_cholesky` itself knows nothing about: ``seed_source`` (``"memo"`` if the solve that + produced the returned reconstruction started from a memo seed, ``"dense"`` if it started from the + sign of the unconstrained dense solve) and ``warm_start_fallback`` (`True` if a memo seed breached + `Settings.nnls_warm_start_error_tolerance` and its entry was dropped, so the next solve for that key + restarts dense). They are set after `fnnls_cholesky` returns, so a diagnostic wrapping the solver + must read the dict it handed in *after* the evaluation, not at the point the solver returns. Returns ------- @@ -371,13 +385,99 @@ def reconstruction_positive_only_from( try: from autoarray.util.fnnls import fnnls_cholesky + from autoarray.inversion.inversion import nnls_memo - return fnnls_cholesky( - curvature_reg_matrix, - (data_vector).T, - P_initial=np.linalg.solve(curvature_reg_matrix, data_vector) > 0, + use_memo = ( + settings is not None + and settings.nnls_warm_start_memo + and nnls_memo.memo_enabled() + and fingerprint is not None + ) + + stats = {} + + key = ( + nnls_memo.memo_key(n=data_vector.shape[0], fingerprint=fingerprint) + if use_memo + else None ) + n = data_vector.shape[0] + + entry = nnls_memo.passive_set_get(key=key, n=n) if use_memo else None + + stats["seed_source"] = "dense" + stats["warm_start_fallback"] = False + + if entry is not None: + try: + reconstruction = fnnls_cholesky( + curvature_reg_matrix, + (data_vector).T, + P_initial=entry.passive_set, + stats=stats, + ) + stats["seed_source"] = "memo" + except (RuntimeError, np.linalg.LinAlgError, ValueError): + # A seed from a previous evaluation is a guess about a + # different matrix, so it can factorise badly where the + # dense-sign start would not. That must cost one retry, never a + # resample: fall back to exactly the un-memoized computation + # before the InversionException path below is reached. + reconstruction = None + else: + reconstruction = None + + if reconstruction is None: + reconstruction = fnnls_cholesky( + curvature_reg_matrix, + (data_vector).T, + P_initial=np.linalg.solve(curvature_reg_matrix, data_vector) > 0, + stats=stats, + ) + stats["seed_source"] = "dense" + + if use_memo: + error_fraction = stats["warm_start_errors"] / max(n, 1) + + if stats["seed_source"] == "memo": + # A memo seed is judged against the dense-sign start it + # replaced, not against an absolute error count: the absolute + # fraction does not separate seeds that save iterations from + # seeds that cost them, the ratio to the dense-sign reference + # does. A seed that breaches the tolerance is discarded, so the + # next solve for this key restarts dense and refreshes the + # reference -- no stale seed can be dragged through a run in a + # regime the reference was never measured in. + tolerance = settings.nnls_warm_start_error_tolerance + + guard_active = ( + tolerance is not None and np.isfinite(tolerance) and tolerance > 0.0 + ) + + if ( + guard_active + and error_fraction > tolerance * entry.dense_error_fraction + ): + nnls_memo.memo_drop(key=key) + stats["warm_start_fallback"] = True + else: + # The reference describes the dense-sign start, so it is + # carried forward unchanged; only a dense solve refreshes it. + nnls_memo.passive_set_put( + key=key, + passive_set=stats["passive_set"], + dense_error_fraction=entry.dense_error_fraction, + ) + else: + nnls_memo.passive_set_put( + key=key, + passive_set=stats["passive_set"], + dense_error_fraction=error_fraction, + ) + + return reconstruction + except (RuntimeError, np.linalg.LinAlgError, ValueError) as e: if is_test_mode(): # See reconstruction_positive_negative_from: benign dummy in test mode, diff --git a/autoarray/inversion/inversion/nnls_memo.py b/autoarray/inversion/inversion/nnls_memo.py new file mode 100644 index 000000000..ef4f919c1 --- /dev/null +++ b/autoarray/inversion/inversion/nnls_memo.py @@ -0,0 +1,123 @@ +import os + +import numpy as np +from typing import Dict, NamedTuple, Optional + +# Cross-evaluation memo for the positive-only (fnnls) solve's passive set. +# +# The Bro & De Jong active set iteration is warm-started from a guess at which +# reconstruction entries are non-zero. The production guess -- the sign of the +# unconstrained dense solve -- gets ~150 of ~1560 entries wrong on a Delaunay +# euclid fit, and each wrong entry costs an active-set iteration in a solve +# that is ~70% of the whole numba CPU likelihood evaluation. Successive +# sampler evaluations sit close together in parameter space and share nearly +# all of their passive set, so seeding from the previous evaluation's FINAL +# passive set is a much better guess. +# +# Entries are keyed by the index space the passive set lives in (mesh/data +# shapes, plus the edge-zeroed subset when one is in use) -- never by the +# matrix values, since the whole point is to hit across nearby parameter +# points. A wrong seed cannot corrupt the answer: the NNLS optimum is unique +# and the solver adds and removes entries until the KKT conditions hold, so a +# stale entry costs iterations, not correctness. Forked pool workers inherit a +# copy at fork and diverge from there, which is fine for the same reason. +# +# An entry therefore carries TWO things: the passive set to seed from, and +# `dense_error_fraction` -- the error fraction of the most recent solve for +# that key that started from the DENSE-SIGN guess. That number is the +# reference the fallback guard in `reconstruction_positive_only_from` measures +# a seed against: the PyAutoArray#498 robustness matrix showed the absolute +# error fraction of a seed does NOT separate seeds that save iterations from +# seeds that cost them (helpful and unhelpful cells overlap at 0.048-0.138), +# but the ratio of the seed's fraction to the dense-sign start's does (helpful +# cells never exceed 0.89, the worst seed reaches 1.42). The reference is +# per-key and self-calibrating, so a solve regime far outside anything the +# matrix probed cannot drag a stale seed through a whole run: once a seed is +# that much worse than the dense-sign start, the entry is dropped +# (`memo_drop`) and the next solve for that key restarts dense, refreshing the +# reference. +# +# Disable with AUTOARRAY_NNLS_WARM_START=0. + + +class MemoEntry(NamedTuple): + """ + One memoized solve: the passive set to seed the next solve for this key + from, and the error fraction of the most recent dense-sign-started solve + for the same key (the reference the fallback guard compares a seed to). + """ + + passive_set: np.ndarray + dense_error_fraction: float + + +_nnls_passive_set_memo: Dict[str, MemoEntry] = {} + +_NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES = 8 + + +def memo_enabled() -> bool: + """ + Whether the passive-set memo is active in this process. + """ + return os.environ.get("AUTOARRAY_NNLS_WARM_START", "1") != "0" + + +def memo_key(n: int, fingerprint) -> str: + """ + The memo key for a solve of size `n` whose index space is described by + `fingerprint` (see `AbstractInversion._nnls_warm_start_fingerprint`). + """ + return f"{n}:{fingerprint}" + + +def passive_set_get(key: str, n: int) -> Optional[MemoEntry]: + """ + The memoized entry for `key` -- its passive set and dense-sign reference + error fraction -- or None on a miss. + + An entry whose indices do not all fit a size-`n` solve is a miss, not a + hit: `n` is already part of the key, so this only fires if a caller + fingerprints two different index spaces identically, and a miss is always + a safe outcome. + """ + entry = _nnls_passive_set_memo.get(key) + + if entry is None: + return None + + if entry.passive_set.size and entry.passive_set.max() >= n: + return None + + return entry + + +def passive_set_put( + key: str, passive_set: np.ndarray, dense_error_fraction: float +) -> None: + """ + Store a solve's final passive set alongside the dense-sign reference error + fraction to carry forward, evicting the oldest entry once the memo is full + (FIFO; the memo tracks one inversion's recent history, not a working set + worth ranking). + """ + stored = np.asarray(passive_set, dtype=int).copy() + stored.setflags(write=False) + + if ( + key not in _nnls_passive_set_memo + and len(_nnls_passive_set_memo) >= _NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES + ): + _nnls_passive_set_memo.pop(next(iter(_nnls_passive_set_memo))) + + _nnls_passive_set_memo[key] = MemoEntry( + passive_set=stored, dense_error_fraction=float(dense_error_fraction) + ) + + +def memo_drop(key: str) -> None: + """ + Forget `key`, so the next solve for it restarts from the dense-sign guess + and refreshes the reference error fraction. A no-op if the key is absent. + """ + _nnls_passive_set_memo.pop(key, None) diff --git a/autoarray/settings.py b/autoarray/settings.py index 5ad59387e..0c6f88ee4 100644 --- a/autoarray/settings.py +++ b/autoarray/settings.py @@ -17,6 +17,8 @@ def __init__( no_regularization_add_to_curvature_diag_value: float = None, nnls_solver_tol: Optional[float] = None, nnls_max_iter: Optional[int] = None, + nnls_warm_start_memo: Optional[bool] = None, + nnls_warm_start_error_tolerance: Optional[float] = None, log_det_method: Optional[str] = None, regularization_term_method: Optional[str] = None, ): @@ -99,6 +101,31 @@ def __init__( jaxnnls's own hard-coded cap of 50 (production HST pixelization+MGE systems converge in ~19-21 iterations). Under `vmap` the solve runs until the slowest lane in the batch converges, so this also caps the worst-case batched cost. Only the JAX (`xp=jnp`) path honors this. + nnls_warm_start_memo + Whether the NumPy / numba positive-only (NNLS) solve warm-starts its active-set iteration from the + passive set of the *previous* likelihood evaluation, held in a small process-local memo + (:mod:`autoarray.inversion.inversion.nnls_memo`). Successive sampler evaluations sit close together in + parameter space and share most of their passive set, so the seeded start removes active-set iterations + from the solve, which is ~70% of a numba CPU likelihood evaluation on production Delaunay meshes. The + NNLS optimum is unique, so the reconstruction is unchanged. `None` (default) reads the packaged value + (`true`); setting ``AUTOARRAY_NNLS_WARM_START=0`` is the process-wide kill-switch that disables the + memo whatever this setting says. Only the NumPy (`xp=np`) fnnls path honors this; the JAX path + ignores it. + nnls_warm_start_error_tolerance + The relative guard on a memo seed's quality. Each memo entry carries a reference: the error + fraction (warm-start errors / solve size) of the most recent solve for that key which started + from the dense-sign guess. A memo-seeded solve whose own error fraction exceeds + ``tolerance * reference`` is judged worse than simply starting dense, so its entry is discarded + and the next solve for that key restarts from the dense-sign start, refreshing the reference. + The guard is relative because the absolute error fraction does not separate helpful from + unhelpful seeds (the two populations overlap at 0.048-0.138 in the PyAutoArray#498 32-cell + robustness matrix) whereas the ratio to the dense-sign start does (helpful cells top out at + 0.89, the worst seed measured reaches 1.42). `None` (default) reads the packaged value + (``1.5``), chosen above that worst observed ratio so the guard is protective against regimes + far outside the matrix rather than flapping inside it. The guard is **disabled** by any value + that is not finite and positive -- ``float("inf")`` is the idiomatic choice, and ``0`` or a + negative value disables it too rather than meaning "drop everything". Only the NumPy + (`xp=np`) fnnls path honors this; the JAX path ignores it. log_det_method Which computation is used for the two Bayesian-evidence log-determinant terms (``log_det_curvature_reg_matrix_term`` and ``log_det_regularization_matrix_term``). `None` @@ -151,6 +178,8 @@ def __init__( self.use_mixed_precision = use_mixed_precision self.nnls_solver_tol = nnls_solver_tol self.nnls_max_iter = nnls_max_iter + self._nnls_warm_start_memo = nnls_warm_start_memo + self._nnls_warm_start_error_tolerance = nnls_warm_start_error_tolerance self._use_positive_only_solver = use_positive_only_solver self._use_edge_zeroed_pixels = use_edge_zeroed_pixels self._use_border_relocator = use_border_relocator @@ -201,6 +230,44 @@ def no_regularization_add_to_curvature_diag_value(self): return self._no_regularization_add_to_curvature_diag_value + @property + def nnls_warm_start_memo(self) -> bool: + """ + Whether the NumPy fnnls solve is warm-started from the previous evaluation's passive set. + """ + if self._nnls_warm_start_memo is None: + try: + return conf.instance["general"]["inversion"]["nnls_warm_start_memo"] + except KeyError: + # Workspaces ship their own general.yaml that shadows autoarray's + # and will not carry this key, so in practice this fallback *is* + # the production default: on, matching the packaged value. + return True + + return self._nnls_warm_start_memo + + @property + def nnls_warm_start_error_tolerance(self) -> float: + """ + How much worse than the dense-sign start a memo seed may be before its entry is dropped. + + A seeded solve whose error fraction exceeds this multiple of the entry's dense-sign reference + fraction is discarded, so the next solve for that key restarts dense. Any value that is not + finite and positive disables the guard. + """ + if self._nnls_warm_start_error_tolerance is None: + try: + return conf.instance["general"]["inversion"][ + "nnls_warm_start_error_tolerance" + ] + except KeyError: + # Workspaces ship their own general.yaml that shadows autoarray's + # and will not carry this key, so in practice this fallback *is* + # the production default, matching the packaged value. + return 1.5 + + return self._nnls_warm_start_error_tolerance + @property def log_det_method(self): if self._log_det_method is None: diff --git a/autoarray/util/fnnls.py b/autoarray/util/fnnls.py index 01ea17047..6d726aac6 100644 --- a/autoarray/util/fnnls.py +++ b/autoarray/util/fnnls.py @@ -1,5 +1,7 @@ import numpy as np +from typing import Optional + from autoarray.util.cholesky_funcs import ( _cho_solve_buffer, cholinsertlast_inplace, @@ -26,9 +28,22 @@ def fnnls_cholesky( ZTZ, ZTx, P_initial=np.zeros(0, dtype=int), + stats: Optional[dict] = None, ): """ Similar to fnnls, but use solving the lstsq problem by updating Cholesky factorisation. + + Parameters + ---------- + P_initial + The warm-start passive set, either as a length-n boolean mask (what the + production dense-sign start hands over) or as an integer index array. + stats + If a dict is passed it is filled on return with the solve's diagnostics: + `outer_iterations`, `inner_iterations`, `passive_set` (the final passive + indices, in the order they were added), `n_passive` and + `warm_start_errors` (how many entries the warm start got wrong). Purely + observational -- the returned solution is unaffected. """ from scipy import linalg as slg @@ -46,14 +61,6 @@ def fnnls_cholesky( ZTx = np.asarray(ZTx) P_initial = np.asarray(P_initial) - lstsq = lambda A, x: slg.solve( - A, - x, - assume_a="pos", - overwrite_a=True, - overwrite_b=True, - ) - n = np.shape(ZTZ)[0] epsilon = 2.2204e-16 tolerance = epsilon * n @@ -62,8 +69,23 @@ def fnnls_cholesky( loop_count = 0 loop_count2 = 0 + # `P_initial` arrives either as a boolean mask (the production dense-sign + # start, and the memo's re-seeded passive set once expanded) or as an + # integer index array (the tests, and the historical call signature). + # Normalise both to the pair the algorithm actually uses -- the mask `P` + # and the insertion-ordered index array `P_inorder` -- once, here, rather + # than re-deriving one from the other at each use. P = np.zeros(n, dtype=bool) - P[P_initial] = True + + if P_initial.dtype == bool: + P[:] = P_initial + P_inorder = np.where(P_initial)[0].astype(int) + else: + P[P_initial] = True + P_inorder = P_initial.astype(int) + + P_initial_mask = P.copy() + d = np.zeros(n) w = ZTx - (ZTZ) @ d s_chol = np.zeros(n) @@ -80,13 +102,56 @@ def fnnls_cholesky( U_buffer = np.zeros((n, n)) k_active = 0 - if P_initial.shape[0] != 0: - P_number = np.arange(len(P), dtype="int") - P_inorder = P_number[P_initial] - s_chol[P] = lstsq((ZTZ)[P][:, P], (ZTx)[P]) - d = s_chol.clip(min=0) - else: - P_inorder = np.array([], dtype="int") + if P_inorder.size != 0: + # Factorise the warm-start passive set ONCE and keep the factor: the + # outer loop below then only ever extends it by one column + # (`cholinsertlast_inplace`). Previously the warm start did a dense + # `slg.solve` here and the first outer iteration threw the result away + # to rebuild the whole factor from scratch -- an O(k^3) factorisation + # of ~1000 columns on every likelihood evaluation. + U = slg.cholesky(ZTZ[P_inorder][:, P_inorder]) + k_active = U.shape[0] + U_buffer[:k_active, :k_active] = U + + s_chol[P_inorder] = _cho_solve_buffer(U_buffer, k_active, ZTx[P_inorder]) + + # A warm start whose passive set contains entries with a non-positive + # unconstrained solution must be repaired BEFORE the outer loop: the + # old code merely clipped, and if `P` happened to be all-True the outer + # `while (not np.all(P))` never ran and the clipped -- wrong -- vector + # was returned. The dense-sign start could not reach that state; a + # memo-seeded start can. + # + # The repair is deliberately NOT `fix_constraint_cholesky`. That step + # interpolates from the previous feasible iterate, and a warm start has + # none: with `d` the clipped `s_chol`, every violator gives + # `d[q] - s_chol[q] == 0`, so its `alpha` is 0/0 or x/0 and nan/inf + # propagates into the whole solution. The alpha -> 0 limit of that step + # is exactly "drop every violator and re-solve", so do that directly. + # It terminates (`P` strictly shrinks) and cannot increase the + # objective (the surviving subspace still contains the zero vector), so + # the outer loop receives a feasible iterate exactly as it expects. + while P_inorder.size and np.min(s_chol[P_inorder]) <= tolerance: + id_delete = np.where(s_chol[P_inorder] <= tolerance)[0] + + k_active = choldeleteindexes_inplace(U_buffer, k_active, id_delete) + + P[P_inorder[id_delete]] = False + P_inorder = np.delete(P_inorder, id_delete) + + s_chol[~P] = 0.0 + + if P_inorder.size: + s_chol[P_inorder] = _cho_solve_buffer( + U_buffer, k_active, ZTx[P_inorder] + ) + + loop_count2 += 1 + if loop_count2 > 10000: + raise RuntimeError + + d = s_chol.copy() + w = ZTx - (ZTZ) @ d # P_inorder is similar as P. They are both used to select solutions in the passive set. # P_inorder saves the `indexes` of those passive solutions. @@ -103,8 +168,9 @@ def fnnls_cholesky( idmax = np.argmax(w * ~P) P_inorder = np.append(P_inorder, int(idmax)) - if loop_count == 0: - # We need to initialize the Cholesky factorisation, U, for the first loop. + if k_active == 0: + # Cold start (or a passive set emptied by the constraint fixer): + # there is no factor to extend, so build the 1 x 1 one. U = slg.cholesky(ZTZ[P_inorder][:, P_inorder]) k_active = U.shape[0] U_buffer[:k_active, :k_active] = U @@ -164,6 +230,13 @@ def fnnls_cholesky( f"normal-equations matrix is singular to working precision." ) + if stats is not None: + stats["outer_iterations"] = loop_count + stats["inner_iterations"] = loop_count2 + stats["passive_set"] = P_inorder.copy() + stats["n_passive"] = int(P_inorder.size) + stats["warm_start_errors"] = int(np.count_nonzero(P_initial_mask != P)) + return d diff --git a/test_autoarray/inversion/inversion/test_nnls_memo.py b/test_autoarray/inversion/inversion/test_nnls_memo.py new file mode 100644 index 000000000..f5f8a35fd --- /dev/null +++ b/test_autoarray/inversion/inversion/test_nnls_memo.py @@ -0,0 +1,205 @@ +""" +The cross-evaluation memo for the positive-only (fnnls) solve's passive set +(`nnls_memo.py`, wired in through `reconstruction_positive_only_from` and +`AbstractInversion.reconstruction`). + +The memo only ever changes how many active-set iterations the solve takes: the +NNLS optimum is unique, so a memoized reconstruction must equal an un-memoized +one to round-off, including in the edge-zeroed subset branch where the passive +set lives in the subset index space. +""" + +import numpy as np +import pytest + +import autoarray as aa + +from autoarray.inversion.inversion import nnls_memo +from autoarray.inversion.inversion.nnls_memo import ( + _NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES, + _nnls_passive_set_memo, + memo_drop, + memo_key, + passive_set_get, + passive_set_put, +) + + +@pytest.fixture(autouse=True) +def _clean_memo(): + _nnls_passive_set_memo.clear() + yield + _nnls_passive_set_memo.clear() + + +def _normal_equations(seed, n=8, n_data=20): + """A system whose unconstrained solution has negative components, so the + passive set is a strict subset and a seed can be wrong about it.""" + rng = np.random.default_rng(seed) + Z = rng.normal(size=(n_data, n)) + x = Z @ rng.normal(size=n) + rng.normal(size=n_data) + return Z.T @ Z, Z.T @ x + + +class SubsetInversion(aa.m.MockInversion): + """Pins the edge-zeroed subset branch of `reconstruction` without building a + real mesh: `solve_ids_to_keep` is the only thing that branch consults.""" + + def __init__(self, ids_to_keep, **kwargs): + super().__init__(**kwargs) + self._ids_to_keep = ids_to_keep + + @property + def solve_ids_to_keep(self): + return self._ids_to_keep + + +def _inversion_from( + curvature_reg_matrix, data_vector, memo, ids_to_keep=None, tolerance=None +): + n = data_vector.shape[0] + + kwargs = dict( + linear_obj_list=[ + aa.m.MockMapper(source_plane_mesh_grid=np.zeros((n, 2)), parameters=n) + ], + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=aa.Settings( + use_positive_only_solver=True, + use_edge_zeroed_pixels=False, + nnls_warm_start_memo=memo, + nnls_warm_start_error_tolerance=tolerance, + ), + ) + + if ids_to_keep is None: + return aa.m.MockInversion(**kwargs) + + return SubsetInversion(ids_to_keep=ids_to_keep, **kwargs) + + +@pytest.mark.parametrize("seed", [0, 1, 2]) +def test__memoized_reconstruction__matches_unmemoized(seed): + curvature_reg_matrix, data_vector = _normal_equations(seed) + + expected = _inversion_from( + curvature_reg_matrix, data_vector, memo=False + ).reconstruction + + # The first memoized solve populates the memo; the second consumes it. + for _ in range(2): + reconstruction = _inversion_from( + curvature_reg_matrix, data_vector, memo=True + ).reconstruction + + assert reconstruction == pytest.approx(expected, rel=1e-10, abs=1e-12) + + assert len(_nnls_passive_set_memo) == 1 + + +@pytest.mark.parametrize("seed", [0, 1, 2]) +def test__memoized_reconstruction__subset_branch__matches_unmemoized(seed): + curvature_reg_matrix, data_vector = _normal_equations(seed) + + ids_to_keep = np.array([0, 2, 3, 5, 6, 7]) + + expected = _inversion_from( + curvature_reg_matrix, data_vector, memo=False, ids_to_keep=ids_to_keep + ).reconstruction + + for _ in range(2): + reconstruction = _inversion_from( + curvature_reg_matrix, data_vector, memo=True, ids_to_keep=ids_to_keep + ).reconstruction + + assert reconstruction == pytest.approx(expected, rel=1e-10, abs=1e-12) + + # The subset solve is of size len(ids_to_keep), and its passive set indexes + # the subset -- not the full parameter vector. + (key,) = _nnls_passive_set_memo + assert key.startswith(f"{len(ids_to_keep)}:") + assert np.all(_nnls_passive_set_memo[key].passive_set < len(ids_to_keep)) + + +def test__fingerprint__changes_with_ids_to_keep(): + # The subset passive set indexes the subset, so a different `ids_to_keep` + # is a different index space and must not reuse the seed. + curvature_reg_matrix, data_vector = _normal_equations(0) + + fingerprint_a = _inversion_from( + curvature_reg_matrix, data_vector, memo=True + )._nnls_warm_start_fingerprint(ids_to_keep=np.array([0, 2, 3])) + + fingerprint_b = _inversion_from( + curvature_reg_matrix, data_vector, memo=True + )._nnls_warm_start_fingerprint(ids_to_keep=np.array([0, 2, 4])) + + fingerprint_full = _inversion_from( + curvature_reg_matrix, data_vector, memo=True + )._nnls_warm_start_fingerprint() + + assert fingerprint_a != fingerprint_b + assert fingerprint_a != fingerprint_full + + +def test__passive_set_put__evicts_the_oldest_entry_when_full(): + for i in range(_NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES + 2): + passive_set_put( + key=f"key_{i}", passive_set=np.array([i]), dense_error_fraction=0.1 + ) + + assert len(_nnls_passive_set_memo) == _NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES + assert "key_0" not in _nnls_passive_set_memo + assert "key_1" not in _nnls_passive_set_memo + assert f"key_{_NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES + 1}" in _nnls_passive_set_memo + + +def test__passive_set_put__stores_a_read_only_copy_and_the_reference_fraction(): + passive_set = np.array([0, 3, 4]) + + passive_set_put(key="key", passive_set=passive_set, dense_error_fraction=0.25) + + passive_set[0] = 99 + + entry = passive_set_get(key="key", n=5) + + assert np.array_equal(entry.passive_set, np.array([0, 3, 4])) + assert entry.dense_error_fraction == 0.25 + with pytest.raises(ValueError): + entry.passive_set[0] = 1 + + +def test__passive_set_get__miss_on_out_of_range_indices_and_unknown_key(): + passive_set_put( + key="key", passive_set=np.array([0, 3, 4]), dense_error_fraction=0.0 + ) + + assert passive_set_get(key="key", n=5) is not None + assert passive_set_get(key="key", n=4) is None + assert passive_set_get(key="other", n=5) is None + + +def test__memo_drop__forgets_the_key_and_is_a_no_op_when_absent(): + passive_set_put(key="key", passive_set=np.array([0, 2]), dense_error_fraction=0.1) + + memo_drop(key="key") + + assert passive_set_get(key="key", n=5) is None + + memo_drop(key="key") + + +def test__memo_enabled__reads_the_environment(monkeypatch): + monkeypatch.delenv("AUTOARRAY_NNLS_WARM_START", raising=False) + assert nnls_memo.memo_enabled() is True + + monkeypatch.setenv("AUTOARRAY_NNLS_WARM_START", "0") + assert nnls_memo.memo_enabled() is False + + monkeypatch.setenv("AUTOARRAY_NNLS_WARM_START", "1") + assert nnls_memo.memo_enabled() is True + + +def test__memo_key__separates_solve_sizes(): + assert memo_key(n=3, fingerprint="mesh") != memo_key(n=4, fingerprint="mesh") diff --git a/test_autoarray/inversion/inversion/test_settings_dict.py b/test_autoarray/inversion/inversion/test_settings_dict.py index 453b324d3..a531ea286 100644 --- a/test_autoarray/inversion/inversion/test_settings_dict.py +++ b/test_autoarray/inversion/inversion/test_settings_dict.py @@ -4,7 +4,7 @@ from pathlib import Path import autoarray as aa -from autonerves.dictable import from_dict, output_to_json, from_json +from autonerves.dictable import from_dict, to_dict, output_to_json, from_json @pytest.fixture(name="settings_dict") @@ -32,3 +32,34 @@ def test_file(): assert isinstance(from_json(filename), aa.Settings) finally: os.remove(filename) + + +def test_settings_nnls_warm_start_memo_round_trips(): + # The field is serialised through the property (not the private attribute), + # so a `None` default resolves to the packaged config value on the way out + # and must come back as the same explicit boolean. + assert aa.Settings().nnls_warm_start_memo is True + assert aa.Settings(nnls_warm_start_memo=True).nnls_warm_start_memo is True + assert aa.Settings(nnls_warm_start_memo=False).nnls_warm_start_memo is False + + settings = from_dict(to_dict(aa.Settings(nnls_warm_start_memo=True))) + + assert settings.nnls_warm_start_memo is True + + +def test_settings_nnls_warm_start_error_tolerance_round_trips(): + # The test config does not ship the key, so the default resolves through + # the KeyError fallback -- which is also the production path whenever a + # workspace shadows autoarray's general.yaml. + assert aa.Settings().nnls_warm_start_error_tolerance == 1.5 + assert ( + aa.Settings(nnls_warm_start_error_tolerance=2.5).nnls_warm_start_error_tolerance + == 2.5 + ) + assert aa.Settings( + nnls_warm_start_error_tolerance=float("inf") + ).nnls_warm_start_error_tolerance == float("inf") + + settings = from_dict(to_dict(aa.Settings(nnls_warm_start_error_tolerance=2.5))) + + assert settings.nnls_warm_start_error_tolerance == 2.5 diff --git a/test_autoarray/util/test_cholesky_inplace.py b/test_autoarray/util/test_cholesky_inplace.py index e8c962207..fb73d42ee 100644 --- a/test_autoarray/util/test_cholesky_inplace.py +++ b/test_autoarray/util/test_cholesky_inplace.py @@ -193,3 +193,129 @@ def test__fnnls_cholesky__accepts_jax_arrays(seed): assert np.all(np.asarray(d_jax) >= 0.0) assert np.asarray(d_jax) == pytest.approx(d_np, rel=1e-6, abs=1e-8) + + +def _mixed_sign_normal_equations(seed, n=30, n_data=50): + """A system whose unconstrained solution has many negative components, so + the non-negativity constraints bind and the passive set is a strict + subset.""" + rng = np.random.default_rng(seed) + Z = rng.normal(size=(n_data, n)) + x = Z @ rng.normal(size=n) + rng.normal(size=n_data) + return Z.T @ Z, Z.T @ x + + +@pytest.mark.parametrize("seed", [0, 1, 2]) +def test__fnnls_cholesky__stats_are_self_consistent(seed): + ZTZ, ZTx = _mixed_sign_normal_equations(seed) + + stats = {} + d = fnnls_cholesky(ZTZ, ZTx, stats=stats) + + assert set(stats) == { + "outer_iterations", + "inner_iterations", + "passive_set", + "n_passive", + "warm_start_errors", + } + assert stats["n_passive"] == len(stats["passive_set"]) + assert np.array_equal(np.sort(stats["passive_set"]), np.where(d > 0)[0]) + assert stats["outer_iterations"] > 0 + # A cold start's "warm start" is the empty passive set, so every finally + # passive entry counts as an error. + assert stats["warm_start_errors"] == stats["n_passive"] + + +@pytest.mark.parametrize("seed", [0, 1, 2]) +def test__fnnls_cholesky__warm_start_from_the_true_support__reports_no_errors(seed): + ZTZ, ZTx = _mixed_sign_normal_equations(seed) + + d_cold = fnnls_cholesky(ZTZ, ZTx) + + stats = {} + d_warm = fnnls_cholesky(ZTZ, ZTx, P_initial=d_cold > 0, stats=stats) + + assert stats["warm_start_errors"] == 0 + # Seeded at the optimum the solver has nothing to do: no index has a + # positive gradient, so the active-set loop never runs. + assert stats["outer_iterations"] == 0 + assert stats["inner_iterations"] == 0 + assert d_warm == pytest.approx(d_cold, rel=1e-10, abs=1e-12) + + +@pytest.mark.parametrize("seed", [0, 1, 2]) +def test__fnnls_cholesky__factorisation_seeded_warm_start__matches_cold_start(seed): + # The warm start now factorises its passive set once and hands that factor + # straight to the active-set loop (instead of a dense solve thrown away and + # rebuilt). The solution must be untouched. + ZTZ, ZTx = _mixed_sign_normal_equations(seed) + + d_cold = fnnls_cholesky(ZTZ, ZTx) + + P_initial = slg.solve(ZTZ.copy(), ZTx.copy(), assume_a="pos") > 0 + d_warm = fnnls_cholesky(ZTZ, ZTx, P_initial=P_initial) + + assert d_warm == pytest.approx(d_cold, rel=1e-10, abs=1e-12) + + +@pytest.mark.parametrize("seed", [0, 1, 2]) +def test__fnnls_cholesky__badly_wrong_warm_start__still_converges(seed): + # A memo-seeded passive set is a guess about a *different* matrix, so it can + # be arbitrarily wrong. The NNLS optimum is unique, so every seed must land + # on the same solution -- only the iteration count may differ. + ZTZ, ZTx = _mixed_sign_normal_equations(seed) + n = ZTZ.shape[0] + + d_cold = fnnls_cholesky(ZTZ, ZTx) + + rng = np.random.default_rng(100 + seed) + flipped = (d_cold > 0).copy() + flip = rng.random(n) < 0.3 + flipped[flip] = ~flipped[flip] + + for P_initial in [flipped, np.ones(n, dtype=bool)]: + stats = {} + d_warm = fnnls_cholesky(ZTZ, ZTx, P_initial=P_initial, stats=stats) + + assert d_warm == pytest.approx(d_cold, rel=1e-8, abs=1e-10) + assert stats["warm_start_errors"] == np.count_nonzero( + P_initial != (d_warm > 0) + ) + + +def test__fnnls_cholesky__all_true_warm_start_with_negative_components(): + # The trap the pre-loop constraint fix exists to close: with `P` all True + # the outer `while (not np.all(P))` never runs, so before the fix a warm + # start whose unconstrained solution has negative components returned the + # merely-clipped vector -- a wrong answer, silently. + ZTZ = np.array([[2.0, 1.0, 0.0], [1.0, 3.0, 1.0], [0.0, 1.0, 1.0]]) + ZTx = np.array([1.0, 1.0, 2.0]) + + # Unconstrained solution is [1, -1, 3]: entry 1 must leave the passive set. + assert np.linalg.solve(ZTZ, ZTx)[1] < 0.0 + + d = fnnls_cholesky(ZTZ, ZTx, P_initial=np.ones(3, dtype=bool)) + + assert d == pytest.approx(np.array([0.5, 0.0, 2.0]), rel=1e-10, abs=1e-12) + + +@pytest.mark.parametrize("seed", [0, 1, 2]) +def test__fnnls_cholesky__mask_and_index_warm_starts_are_equivalent(seed): + ZTZ, ZTx = _mixed_sign_normal_equations(seed) + + mask = slg.solve(ZTZ.copy(), ZTx.copy(), assume_a="pos") > 0 + + stats_mask = {} + d_mask = fnnls_cholesky(ZTZ, ZTx, P_initial=mask, stats=stats_mask) + + stats_index = {} + d_index = fnnls_cholesky( + ZTZ, ZTx, P_initial=np.where(mask)[0], stats=stats_index + ) + + assert d_mask == pytest.approx(d_index, rel=1e-12, abs=1e-14) + assert stats_mask["warm_start_errors"] == stats_index["warm_start_errors"] + assert np.array_equal( + np.sort(stats_mask["passive_set"]), np.sort(stats_index["passive_set"]) + ) diff --git a/test_autoarray/util/test_jax_nnls.py b/test_autoarray/util/test_jax_nnls.py index 796bd5030..8ee16ca74 100644 --- a/test_autoarray/util/test_jax_nnls.py +++ b/test_autoarray/util/test_jax_nnls.py @@ -43,9 +43,7 @@ def test__reconstruction_positive_only_from__numpy_path_ignores_knobs(): # knob-carrying settings. 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]] - ) + curvature_reg_matrix = np.array([[2.0, 1.0, 0.0], [1.0, 3.0, 1.0], [0.0, 1.0, 1.0]]) for settings in [None, aa.Settings(nnls_solver_tol=1e-6, nnls_max_iter=30)]: reconstruction = aa.util.inversion.reconstruction_positive_only_from( @@ -58,3 +56,313 @@ def test__reconstruction_positive_only_from__numpy_path_ignores_knobs(): # Unconstrained solution is [1, -1, 3]; the NNLS solution zeroes the # negative component and re-solves the free ones. assert reconstruction == pytest.approx(np.array([0.5, 0.0, 2.0]), 1.0e-4) + + +@pytest.fixture(autouse=True) +def _clear_nnls_memo(): + from autoarray.inversion.inversion.nnls_memo import _nnls_passive_set_memo + + _nnls_passive_set_memo.clear() + yield + _nnls_passive_set_memo.clear() + + +def _small_positive_only_system(): + # Unconstrained solution is [1, -1, 3]; the NNLS solution is [0.5, 0, 2]. + 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]]) + return data_vector, curvature_reg_matrix + + +def test__reconstruction_positive_only_from__warm_start_memo_records_the_passive_set(): + from autoarray.inversion.inversion.nnls_memo import ( + _nnls_passive_set_memo, + memo_key, + ) + + data_vector, curvature_reg_matrix = _small_positive_only_system() + + reconstruction = aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=aa.Settings(nnls_warm_start_memo=True), + fingerprint="mesh", + ) + + assert reconstruction == pytest.approx(np.array([0.5, 0.0, 2.0]), 1.0e-4) + + key = memo_key(n=3, fingerprint="mesh") + + entry = _nnls_passive_set_memo[key] + + assert np.array_equal(entry.passive_set, np.array([0, 2])) + # The dense-sign start of this system is exactly right, so the reference + # error fraction it hands the guard is zero. + assert entry.dense_error_fraction == 0.0 + + +def test__reconstruction_positive_only_from__warm_start_memo_seeds_the_next_solve( + monkeypatch, +): + from autoarray.inversion.inversion import nnls_memo + + data_vector, curvature_reg_matrix = _small_positive_only_system() + settings = aa.Settings(nnls_warm_start_memo=True) + + first = aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=settings, + fingerprint="mesh", + ) + + # Spy on the memo lookup so the second solve is shown to actually consume + # the seed, not merely to leave the memo populated. + seeds = [] + passive_set_get = nnls_memo.passive_set_get + + def _spy(**kwargs): + seed = passive_set_get(**kwargs) + seeds.append(seed) + return seed + + monkeypatch.setattr(nnls_memo, "passive_set_get", _spy) + + second = aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=settings, + fingerprint="mesh", + ) + + assert len(seeds) == 1 + assert np.array_equal(seeds[0].passive_set, np.array([0, 2])) + assert second == pytest.approx(first, rel=1e-10, abs=1e-12) + + +def test__reconstruction_positive_only_from__warm_start_memo_on_by_default(): + from autoarray.inversion.inversion.nnls_memo import ( + _nnls_passive_set_memo, + memo_key, + ) + + data_vector, curvature_reg_matrix = _small_positive_only_system() + + # The memo ships on, so default settings plus a fingerprint memoize. + reconstruction = aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=aa.Settings(), + fingerprint="mesh", + ) + + assert reconstruction == pytest.approx(np.array([0.5, 0.0, 2.0]), 1.0e-4) + assert list(_nnls_passive_set_memo) == [memo_key(n=3, fingerprint="mesh")] + + +def test__reconstruction_positive_only_from__warm_start_memo_opt_outs(): + from autoarray.inversion.inversion.nnls_memo import _nnls_passive_set_memo + + data_vector, curvature_reg_matrix = _small_positive_only_system() + + # No settings object at all, and an explicit opt-out, both leave the memo + # untouched even though the default is on. + for settings in [None, aa.Settings(nnls_warm_start_memo=False)]: + aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=settings, + fingerprint="mesh", + ) + + # A caller that supplies no fingerprint cannot be memoized either, since + # there is nothing identifying the index space the passive set lives in. + aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=aa.Settings(), + ) + + assert _nnls_passive_set_memo == {} + + +def test__reconstruction_positive_only_from__warm_start_memo_disabled_by_env( + monkeypatch, +): + from autoarray.inversion.inversion.nnls_memo import _nnls_passive_set_memo + + monkeypatch.setenv("AUTOARRAY_NNLS_WARM_START", "0") + + data_vector, curvature_reg_matrix = _small_positive_only_system() + + reconstruction = aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=aa.Settings(nnls_warm_start_memo=True), + fingerprint="mesh", + ) + + assert reconstruction == pytest.approx(np.array([0.5, 0.0, 2.0]), 1.0e-4) + assert _nnls_passive_set_memo == {} + + +# =================================================================== +# Relative fallback guard on a memo seed (Settings.nnls_warm_start_error_tolerance) +# =================================================================== + + +def _solve_capturing_stats(monkeypatch, settings, fingerprint="mesh"): + """ + One `reconstruction_positive_only_from` on the small positive-only system, + returning (reconstruction, stats). + + `seed_source` / `warm_start_fallback` are written into the stats dict AFTER + `fnnls_cholesky` returns, so the dict must be held by reference and read + once the call has finished -- reading it inside the wrapper would see the + solver's keys only. + """ + import autoarray.util.fnnls as fnnls_mod + + original = fnnls_mod.fnnls_cholesky + captured = [] + + def _wrapped(ZTZ, ZTx, P_initial=np.zeros(0, dtype=int), stats=None): + captured.append(stats) + return original(ZTZ, ZTx, P_initial, stats=stats) + + monkeypatch.setattr(fnnls_mod, "fnnls_cholesky", _wrapped) + + data_vector, curvature_reg_matrix = _small_positive_only_system() + + reconstruction = aa.util.inversion.reconstruction_positive_only_from( + data_vector=data_vector, + curvature_reg_matrix=curvature_reg_matrix, + settings=settings, + fingerprint=fingerprint, + ) + + monkeypatch.setattr(fnnls_mod, "fnnls_cholesky", original) + + return reconstruction, captured[-1] + + +def test__warm_start_guard__seed_worse_than_tolerance_is_dropped_and_next_solve_is_dense( + monkeypatch, +): + from autoarray.inversion.inversion.nnls_memo import ( + _nnls_passive_set_memo, + memo_key, + passive_set_put, + ) + + key = memo_key(n=3, fingerprint="mesh") + + # A deliberately wrong seed against a deliberately small reference: the + # true passive set is [0, 2], so seeding [1] gets all three entries wrong + # (fraction 1.0) against a dense-sign reference of 0.1 -- 1.0 > 1.5 * 0.1. + passive_set_put(key=key, passive_set=np.array([1]), dense_error_fraction=0.1) + + settings = aa.Settings(nnls_warm_start_memo=True) + assert settings.nnls_warm_start_error_tolerance == 1.5 + + reconstruction, stats = _solve_capturing_stats(monkeypatch, settings) + + assert reconstruction == pytest.approx(np.array([0.5, 0.0, 2.0]), 1.0e-4) + assert stats["seed_source"] == "memo" + assert stats["warm_start_fallback"] is True + assert _nnls_passive_set_memo == {} + + # With the entry dropped, the next solve for the key restarts from the + # dense-sign start and refreshes the reference. + _, stats = _solve_capturing_stats(monkeypatch, settings) + + assert stats["seed_source"] == "dense" + assert stats["warm_start_fallback"] is False + + entry = _nnls_passive_set_memo[key] + + assert np.array_equal(entry.passive_set, np.array([0, 2])) + assert entry.dense_error_fraction == 0.0 + + +def test__warm_start_guard__seed_within_tolerance_keeps_the_entry_and_the_reference( + monkeypatch, +): + from autoarray.inversion.inversion.nnls_memo import ( + _nnls_passive_set_memo, + memo_key, + passive_set_put, + ) + + key = memo_key(n=3, fingerprint="mesh") + + # Seeding every entry passive gets exactly one of three wrong (fraction + # 1/3), which is inside 1.5 * 0.5. + passive_set_put(key=key, passive_set=np.array([0, 1, 2]), dense_error_fraction=0.5) + + reconstruction, stats = _solve_capturing_stats( + monkeypatch, aa.Settings(nnls_warm_start_memo=True) + ) + + assert reconstruction == pytest.approx(np.array([0.5, 0.0, 2.0]), 1.0e-4) + assert stats["seed_source"] == "memo" + assert stats["warm_start_fallback"] is False + + entry = _nnls_passive_set_memo[key] + + assert np.array_equal(entry.passive_set, np.array([0, 2])) + # Only a dense-sign solve refreshes the reference, so it is carried through + # the seeded solve unchanged. + assert entry.dense_error_fraction == 0.5 + + +@pytest.mark.parametrize("tolerance", [float("inf"), 0.0, -1.0]) +def test__warm_start_guard__disabled_tolerance_never_drops(monkeypatch, tolerance): + from autoarray.inversion.inversion.nnls_memo import ( + _nnls_passive_set_memo, + memo_key, + passive_set_put, + ) + + key = memo_key(n=3, fingerprint="mesh") + + passive_set_put(key=key, passive_set=np.array([1]), dense_error_fraction=0.1) + + _, stats = _solve_capturing_stats( + monkeypatch, + aa.Settings( + nnls_warm_start_memo=True, nnls_warm_start_error_tolerance=tolerance + ), + ) + + assert stats["seed_source"] == "memo" + assert stats["warm_start_fallback"] is False + assert _nnls_passive_set_memo[key].dense_error_fraction == 0.1 + + +def test__warm_start_guard__a_perfect_dense_reference_does_not_breach_on_a_perfect_seed( + monkeypatch, +): + from autoarray.inversion.inversion.nnls_memo import ( + _nnls_passive_set_memo, + memo_key, + ) + + settings = aa.Settings(nnls_warm_start_memo=True) + + # The dense-sign start of this system is exact, so the reference is 0.0 and + # the guard degenerates to `frac > 0`. A seed that is also exact must not + # breach it -- a perfect dense start is cheap to keep, not a reason to drop. + _, stats = _solve_capturing_stats(monkeypatch, settings) + + assert stats["seed_source"] == "dense" + + key = memo_key(n=3, fingerprint="mesh") + + assert _nnls_passive_set_memo[key].dense_error_fraction == 0.0 + + _, stats = _solve_capturing_stats(monkeypatch, settings) + + assert stats["seed_source"] == "memo" + assert stats["warm_start_fallback"] is False + assert key in _nnls_passive_set_memo