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
4 changes: 4 additions & 0 deletions autoarray/config/general.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
40 changes: 40 additions & 0 deletions autoarray/inversion/inversion/abstract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -649,6 +687,7 @@ def reconstruction(self) -> np.ndarray:
ids_to_keep=ids_to_keep
),
factor=factor,
solver=solver,
)
)

Expand Down Expand Up @@ -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
Expand Down
95 changes: 95 additions & 0 deletions autoarray/inversion/inversion/inversion_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -248,13 +248,55 @@ 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,
settings: Settings = None,
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
Expand Down Expand Up @@ -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
-----
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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
Expand Down
106 changes: 106 additions & 0 deletions autoarray/settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand All @@ -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):
Expand Down Expand Up @@ -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
Loading
Loading