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
1 change: 1 addition & 0 deletions autoarray/config/general.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ inversion:
use_border_relocator: false # If True, by default a pixelization's border is used to relocate all pixels outside its border to the border.
nnls_jacobi_preconditioning: true # If True (default), the curvature matrix passed to jaxnnls.solve_nnls_primal is Jacobi-preconditioned (D Q D y = D q, x = D y). Fixes NaN backward-pass gradients on ill-conditioned Q and roughly halves forward solve time. Set False to restore the raw unpreconditioned solve.
nnls_target_kappa: 1.0e-11 # Central-path relaxation parameter passed to jaxnnls.solve_nnls_primal. Larger values smooth the relaxed-KKT backward pass and prevent NaN gradients on ill-conditioned Q; smaller values tighten the primal solve. Verified finite gradients across all MGE/rectangular/delaunay pipelines (imaging + interferometer) with scale invariance over 5 orders of magnitude in noise. jaxnnls's own default (1e-3) is too aggressive for the backward pass.
nnls_preconditioning_no_mapper: raw # How the JAX positive-only PDIP solve scales inversions with NO mapper (linear light profiles / MGE only). "raw" (default) runs the forward solve on the un-preconditioned system with a data-scaled KKT tolerance (1e-2 * n * eps * max(1, max|data_vector|)) and keeps the Jacobi-space relaxed-KKT gradient; "jacobi" uses the Jacobi-preconditioned solve. Jacobi scaling of signal-free MGE columns (diagonal = the no-regularization floor) made the PDIP dual diverge on 14/48 SLaM source_lp[1] points (PyAutoArray#571). Inversions with a mapper always use jacobi; the NumPy path is unaffected.
nnls_warm_start_memo: true # If True (default), the NumPy/numba positive-only (fnnls) solve warm-starts its active set from the previous likelihood evaluation's passive set, cutting active-set iterations on successive sampler evaluations. The NNLS optimum is unique so the reconstruction is unchanged. On by default as of PyAutoArray#498, measured on the euclid+hst Delaunay-1250 fiducial (9.9x / 4.0x fewer active-set iterations on successive evaluations, reconstruction unchanged). Set false, or AUTOARRAY_NNLS_WARM_START=0, to disable. JAX path unaffected.
nnls_warm_start_error_tolerance: 1.5 # Relative quality guard on a warm-start memo seed. Each memo entry remembers the error fraction of the most recent dense-sign-started solve for that key; a seeded solve whose own error fraction exceeds this multiple of that reference is dropped, so the next solve restarts from the dense-sign start and refreshes the reference. Default 1.5 sits above the worst seed/dense error-fraction ratio seen in the PyAutoArray#498 32-cell robustness matrix (1.42), so it is protective against unmeasured regimes rather than flapping. Any non-finite or non-positive value (e.g. .inf) disables the guard. NumPy/numba fnnls path only.
positive_only_solver: pdip # Which solver the JAX (xp=jnp) positive-only reconstruction uses. "pdip" (default) is the jaxnnls interior-point solve; "certified" is the certified active-set solve (budgeted masked-Cholesky passes that stop once the KKT conditions certify, exact implicit gradient, PDIP fallback), measured 1.2-2.6x faster on source-only inversions (PyAutoArray#566). Applied only to mapper-only JAX inversions (MGE / linear light profiles keep PDIP); the NumPy path always uses fnnls. Opt-in until the batched (vmap) policy is measured.
Expand Down
27 changes: 27 additions & 0 deletions autoarray/inversion/inversion/abstract.py
Original file line number Diff line number Diff line change
Expand Up @@ -600,6 +600,31 @@ def positive_only_solver_used(self) -> str:

return "certified"

@property
def positive_only_preconditioning_used(self) -> str:
"""
How `reconstruction` asks `reconstruction_positive_only_from` to scale the JAX PDIP solve:
``"jacobi"`` or ``"raw"`` (PyAutoArray#571).

- Inversions containing a `Mapper` always use ``"jacobi"`` (today's Jacobi-preconditioned solve).
- Inversions with **no** `Mapper` (linear light profiles / MGE only) use
`Settings.nnls_preconditioning_no_mapper` (packaged default ``"raw"``): Jacobi scaling turns their
signal-free Gaussian columns, whose diagonal is only the no-regularization floor, into degenerate
coordinates on which the PDIP dual diverges (14/48 near-truth SLaM `source_lp[1]` points hit the
iteration cap with wrong log-likelihoods), while the raw solve converges on all of them.
- Whenever the certified solver is used the answer is ``"jacobi"`` (it only runs on mapper-only
inversions anyway).

The NumPy path runs fnnls and ignores the value.
"""
if self.has(cls=Mapper):
return "jacobi"

if self.positive_only_solver_used != "pdip":
return "jacobi"

return self.settings.nnls_preconditioning_no_mapper

def _nnls_warm_start_fingerprint(self, ids_to_keep=None) -> Optional[str]:
"""
Identify the index space this inversion's positive-only solve works in, so the
Expand Down Expand Up @@ -688,6 +713,7 @@ def reconstruction(self) -> np.ndarray:
),
factor=factor,
solver=solver,
preconditioning=self.positive_only_preconditioning_used,
)
)

Expand Down Expand Up @@ -717,6 +743,7 @@ def reconstruction(self) -> np.ndarray:
fingerprint=self._nnls_warm_start_fingerprint(),
factor=factor,
solver=solver,
preconditioning=self.positive_only_preconditioning_used,
)

self._nnls_factor = factor
Expand Down
93 changes: 72 additions & 21 deletions autoarray/inversion/inversion/inversion_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -297,6 +297,7 @@ def reconstruction_positive_only_from(
factor: Optional[dict] = None,
solver: str = "pdip",
stats: Optional[dict] = None,
preconditioning: str = "jacobi",
):
"""
Solve the linear system Eq.(2) (in terms of minimizing the quadratic value) of
Expand Down Expand Up @@ -358,7 +359,22 @@ def reconstruction_positive_only_from(
``certified`` (whether the active-set search certified within budget; ``False`` means the fallback or an
uncertified iterate was returned) and ``passes`` (restricted passes run) as *traced* JAX scalars, so it is
safe under ``jax.jit`` / ``vmap`` -- read them as outputs of the traced function or through
``jax.debug.callback``. Every call also records ``solver``. It never changes the returned reconstruction.
``jax.debug.callback``. With ``solver="pdip"`` it receives ``converged`` (``1`` if the PDIP KKT residual
met its tolerance within the iteration cap, ``0`` if the cap was hit and the returned reconstruction is
the unconverged iterate) and ``iterations`` (PDIP iterations run), also as traced JAX scalars, plus
``preconditioning``. Every call also records ``solver``. It never changes the returned reconstruction.
preconditioning
How the JAX PDIP solve (``solver="pdip"``) scales the system (PyAutoArray#571). ``"jacobi"`` (default,
byte-identical to before this option existed): the Jacobi-preconditioned solve ``(D Q D) y = D q``
governed by the ``nnls_jacobi_preconditioning`` config key. ``"raw"``: the forward PDIP solve runs on
the un-preconditioned ``(Q, q)`` with the data-scaled tolerance
:func:`autoarray.util.jax_nnls.data_scaled_solver_tol` (or ``settings.nnls_solver_tol`` if set), and the
gradient is the Jacobi-space relaxed-KKT pass as in ``"jacobi"`` -- see
:func:`autoarray.util.jax_nnls.solve_nnls_primal_raw_forward`. Jacobi scaling makes the signal-free
columns of linear-object-only (MGE) inversions, whose diagonal is only the
``no_regularization_add_to_curvature_diag_value`` floor, degenerate coordinates on which the PDIP dual
diverges; the caller (`AbstractInversion.reconstruction`) therefore passes ``"raw"`` for inversions with
no `Mapper` -- see `AbstractInversion.positive_only_preconditioning_used`. Ignored on the NumPy path.

Notes
-----
Expand Down Expand Up @@ -386,11 +402,25 @@ def reconstruction_positive_only_from(
f"solver={solver!r} is not a valid positive-only solver; expected 'pdip' or 'certified'."
)

if preconditioning not in ("jacobi", "raw"):
raise ValueError(
f"preconditioning={preconditioning!r} is invalid; expected 'jacobi' or 'raw'."
)

if preconditioning == "raw" and solver != "pdip":
raise ValueError(
"preconditioning='raw' applies only to solver='pdip'; the certified solver always runs on the "
"Jacobi-scaled system."
)

if xp.__name__.startswith("jax"):

from autonerves import conf

from autoarray.util.jax_nnls import solve_nnls_primal
from autoarray.util.jax_nnls import (
solve_nnls_primal_raw_forward,
solve_nnls_primal_with_status,
)

try:
use_jacobi = conf.instance["general"]["inversion"][
Expand All @@ -403,9 +433,7 @@ def reconstruction_positive_only_from(
use_jacobi = True

try:
target_kappa = conf.instance["general"]["inversion"][
"nnls_target_kappa"
]
target_kappa = conf.instance["general"]["inversion"]["nnls_target_kappa"]
except KeyError:
# Workspaces ship their own general.yaml that shadows autoarray's;
# fall back to the same value autoarray's general.yaml declares.
Expand All @@ -424,6 +452,34 @@ def reconstruction_positive_only_from(
if max_iter is None:
max_iter = 50

def _record_pdip(converged, iterations):
if stats is not None:
stats["solver"] = "pdip"
stats["preconditioning"] = preconditioning
stats["converged"] = converged
stats["iterations"] = iterations

if preconditioning == "raw":
# Same Jacobi quantities as below: the backward pass runs on the scaled system, the forward
# solve on the raw one (see `solve_nnls_primal_raw_forward`).
d = xp.sqrt(xp.diag(curvature_reg_matrix))
D = 1.0 / d
Q_pc = (curvature_reg_matrix * D[:, None]) * D[None, :]
q_pc = data_vector * D

y, converged, iterations = solve_nnls_primal_raw_forward(
Q_pc,
q_pc,
curvature_reg_matrix,
data_vector,
D,
target_kappa=target_kappa,
solver_tol=solver_tol,
max_iter=max_iter,
)
_record_pdip(converged, iterations)
return y * D

if use_jacobi:
# Ill-conditioned Q makes jaxnnls's relaxed-KKT backward pass
# produce NaN gradients. Rescale Q so its diagonal is unit:
Expand All @@ -449,19 +505,15 @@ def reconstruction_positive_only_from(
* D
)

if stats is not None:
stats["solver"] = "pdip"

return (
solve_nnls_primal(
Q_pc,
q_pc,
target_kappa=target_kappa,
solver_tol=solver_tol,
max_iter=max_iter,
)
* D
x, converged, iterations = solve_nnls_primal_with_status(
Q_pc,
q_pc,
target_kappa=target_kappa,
solver_tol=solver_tol,
max_iter=max_iter,
)
_record_pdip(converged, iterations)
return x * D

if solver == "certified":
return _certified_positive_only_from(
Expand All @@ -474,16 +526,15 @@ def reconstruction_positive_only_from(
stats=stats,
)

if stats is not None:
stats["solver"] = "pdip"

return solve_nnls_primal(
x, converged, iterations = solve_nnls_primal_with_status(
curvature_reg_matrix,
data_vector,
target_kappa=target_kappa,
solver_tol=solver_tol,
max_iter=max_iter,
)
_record_pdip(converged, iterations)
return x

# `solver` is deliberately ignored on the NumPy path: a NumPy port of the certified active-set scheme
# measured 3-7 % slower than fnnls and lost to its warm-start memo (PyAutoArray#566), so fnnls stays.
Expand Down
41 changes: 41 additions & 0 deletions autoarray/settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ def __init__(
certified_pass_budget: Optional[int] = None,
certified_fallback: Optional[str] = None,
certified_tau_rel: Optional[float] = None,
nnls_preconditioning_no_mapper: Optional[str] = None,
):
"""
The settings of an Inversion, customizing how a linear set of equations are solved for.
Expand Down Expand Up @@ -223,6 +224,23 @@ def __init__(
Relative KKT tolerance of the ``"certified"`` solver's certificate: primal violations are
``x < -tau_rel * max|x|``, dual violations ``g < -tau_rel * max|q|``. `None` (default) reads the
packaged value (`1.0e-9`).
nnls_preconditioning_no_mapper
How the JAX positive-only PDIP solve scales an inversion **with no `Mapper`** (linear light profiles /
MGE only). `None` (default) reads the packaged value ``"raw"``.

- ``"raw"`` (default) -- the forward PDIP solve runs on the un-preconditioned system with a
data-scaled KKT tolerance (``1e-2 * n * eps_pdip * max(1, max|data_vector|)``, or
`nnls_solver_tol` if set); the gradient is the same Jacobi-space relaxed-KKT pass as ``"jacobi"``.
On the SLaM `source_lp[1]` MGE model (2 x 20 lens + 20 source Gaussians) Jacobi scaling made 14/48
near-truth points hit the 50-iteration cap with wrong log-likelihoods (signal-free Gaussian columns,
whose diagonal is only `no_regularization_add_to_curvature_diag_value`, become degenerate
coordinates on which the PDIP dual diverges); the raw solve converges on all of them in 16-19
iterations (PyAutoArray#571).
- ``"jacobi"`` -- the Jacobi-preconditioned solve, as for mapper inversions.

Inversions containing a `Mapper` always use ``"jacobi"``
(`AbstractInversion.positive_only_preconditioning_used` records the decision). The NumPy path always
runs fnnls and ignores this.
"""
self.use_mixed_precision = use_mixed_precision
self.nnls_solver_tol = nnls_solver_tol
Expand All @@ -244,12 +262,15 @@ def __init__(
self._certified_pass_budget = certified_pass_budget
self._certified_fallback = certified_fallback
self._certified_tau_rel = certified_tau_rel
self._nnls_preconditioning_no_mapper = nnls_preconditioning_no_mapper

# Validate explicit values eagerly, so a typo fails at construction rather than deep inside a fit.
if positive_only_solver is not None:
self.positive_only_solver
if certified_fallback is not None:
self.certified_fallback
if nnls_preconditioning_no_mapper is not None:
self.nnls_preconditioning_no_mapper

@property
def use_positive_only_solver(self):
Expand Down Expand Up @@ -449,3 +470,23 @@ def certified_tau_rel(self) -> float:
return self._inversion_config_value("certified_tau_rel", 1.0e-9)

return self._certified_tau_rel

@property
def nnls_preconditioning_no_mapper(self) -> str:
"""
How the JAX PDIP solve scales an inversion with no `Mapper`: ``"raw"`` or ``"jacobi"``.

See the constructor docstring; inversions with a `Mapper` always use ``"jacobi"``.
"""
value = self._nnls_preconditioning_no_mapper
if value is None:
value = self._inversion_config_value(
"nnls_preconditioning_no_mapper", "raw"
)

if value not in ("raw", "jacobi"):
raise ValueError(
f"nnls_preconditioning_no_mapper={value!r} is invalid; expected 'raw' or 'jacobi'."
)

return value
Loading
Loading