Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions autoarray/config/general.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
36 changes: 36 additions & 0 deletions autoarray/inversion/inversion/abstract.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import copy
import hashlib
import warnings

import numpy as np
Expand Down Expand Up @@ -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:
"""
Expand Down Expand Up @@ -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
),
)
)

Expand All @@ -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(
Expand Down
108 changes: 104 additions & 4 deletions autoarray/inversion/inversion/inversion_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
-------
Expand Down Expand Up @@ -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,
Expand Down
123 changes: 123 additions & 0 deletions autoarray/inversion/inversion/nnls_memo.py
Original file line number Diff line number Diff line change
@@ -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)
Loading
Loading