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: 1 addition & 1 deletion autoarray/config/general.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +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. The relaxed solve must start from an iterate whose complementarity s*z is not far above this value; the "raw" no-mapper mode polishes its forward iterate to ensure that (PyAutoArray#573).
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, whose relaxed solve starts from the forward iterate polished by <= 10 tight warm-started PDIP iterations on the Jacobi system (without the polish the loose forward tolerance leaves s*z far above nnls_target_kappa and the relaxed solve diverged to NaN gradients on 4/16 jax_grad/mge.py points, PyAutoArray#573); "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_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|)), polishes that iterate with <= 10 tight warm-started PDIP iterations on the Jacobi system and returns the polished iterate (the loose tolerance also judges complementarity, so unpolished it left 11.5 % of the fnnls flux on inactive columns of the euclid vis_lp system; polished <= 3.3e-4 over 81 corpus systems, ~5 extra iterations, PyAutoArray#594), and keeps the Jacobi-space relaxed-KKT gradient started from the same polished point (without the polish the relaxed solve diverged to NaN gradients on 4/16 jax_grad/mge.py points, PyAutoArray#573); "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
4 changes: 3 additions & 1 deletion autoarray/inversion/inversion/inversion_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -368,7 +368,9 @@ def reconstruction_positive_only_from(
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
:func:`autoarray.util.jax_nnls.data_scaled_solver_tol` (or ``settings.nnls_solver_tol`` if set), its iterate
is polished by a few tight PDIP iterations on the Jacobi-scaled system and the polished iterate is returned
(PyAutoArray#594; ``stats["converged"]`` / ``stats["iterations"]`` describe the raw forward solve), 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
Expand Down
6 changes: 5 additions & 1 deletion autoarray/settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -230,7 +230,11 @@ def __init__(

- ``"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"``.
`nnls_solver_tol` if set), polished by at most 10 tight warm-started PDIP iterations on the
Jacobi-scaled system; the polished iterate is the returned reconstruction (unpolished, the loose
tolerance left 11.5 % of the fnnls flux on inactive columns of the euclid vis_lp system,
PyAutoArray#594) and the start of the gradient, 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
Expand Down
125 changes: 83 additions & 42 deletions autoarray/util/jax_nnls.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,11 +27,18 @@
solve on the un-preconditioned system with a data-scaled tolerance
(:func:`data_scaled_solver_tol`) and keeps the Jacobi-space backward pass. It
exists because Jacobi scaling of signal-free MGE columns (diagonal = the
no-regularization floor) makes the PDIP dual diverge. Its backward pass polishes
the mapped forward iterate with a few tight, warm-started PDIP iterations on the
Jacobi system before the relaxed-KKT solve, which otherwise diverges to NaN from
the loose forward tolerance (PyAutoArray#573); :func:`raw_forward_backward_status`
reports that pass's convergence.
no-regularization floor) makes the PDIP dual diverge. The raw iterate is then
*polished* with a few tight, warm-started PDIP iterations on the Jacobi system, and
the polished iterate is both the forward value and the start of the backward
pass's relaxed-KKT solve. The data-scaled stop judges complementarity ``s * z``
against the same loose threshold, so a column fnnls holds at zero can keep
``x ~ tol / z`` while "converged": on the euclid vis_lp system the unpolished
iterate left 11.5 % of the fnnls flux on reference-inactive columns (a +5.8 %
latent source flux), invisible to the objective and the log likelihood; polished,
that is <= 3.3e-4 across 81 corpus systems for ~5 extra iterations
(PyAutoArray#594). The polish first went into the backward pass only, where the
relaxed-KKT solve otherwise diverges to NaN from the loose forward tolerance
(PyAutoArray#573); :func:`raw_forward_backward_status` reports its convergence.

JAX is imported inside functions, never at module level (see
``docs/agents/jax_and_decorators.md``); this module must only be imported
Expand Down Expand Up @@ -187,51 +194,75 @@ def solve_nnls_primal(Q, q, target_kappa=1e-3, solver_tol=None, max_iter=50):
)[0]


# The backward pass of the ``"raw"`` mode first polishes the mapped raw-forward iterate with at most this many
# PDIP iterations on the Jacobi-scaled system at jaxnnls's own tight tolerance (PyAutoArray#573). Measured on the
# SLaM MGE fixture, the 48 SLaM ``source_lp[1]`` systems and the jax_grad/mge.py points: 4-6 iterations.
RAW_BACKWARD_POLISH_MAX_ITER = 10
# The ``"raw"`` mode polishes the mapped raw-forward iterate with at most this many PDIP iterations on the
# Jacobi-scaled system at jaxnnls's own tight tolerance; the polished iterate is the forward value (PyAutoArray#594)
# and the backward pass's starting point (PyAutoArray#573). Measured on the SLaM MGE fixture, the 48 SLaM
# ``source_lp[1]`` systems and the jax_grad/mge.py points: 4-6 iterations.
RAW_POLISH_MAX_ITER = 10
# The name under which the polish cap was introduced (#573, backward pass only); kept for external importers.
RAW_BACKWARD_POLISH_MAX_ITER = RAW_POLISH_MAX_ITER


def _raw_forward_backward_point(
Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter
):
def _raw_forward_polished(Q_pc, q_pc, Q, q, D, solver_tol, max_iter):
"""
The forward solve and the relaxed-KKT point of the ``"raw"`` mode (shared by
the custom-vjp forward pass and :func:`raw_forward_backward_status`).

Returns ``(y, converged, pdip_iter)`` of the raw forward solve (mapped to the
Jacobi coordinates), the relaxed point ``(yr, sr, zr)`` the backward pass
differentiates at, and the status ``(relaxed_converged, relaxed_iter,
polish_converged, polish_iter)``.
The forward value of the ``"raw"`` mode: the raw PDIP solve, mapped to the
Jacobi coordinates and polished there (shared by the custom-vjp primal and
forward passes, so every call path returns the same ``y`` bit-for-bit).

Returns ``(y_out, converged, pdip_iter, (yp, sp, zp), ok, polish_iter)``:
``y_out`` the forward value (the polished iterate if ``ok``, else the mapped
raw iterate), ``converged`` / ``pdip_iter`` the raw forward solve's flag and
iteration count, ``(yp, sp, zp)`` the point the backward pass starts from
(selected by the same ``ok``), ``ok`` whether the polish was accepted and
``polish_iter`` its iteration count.
"""
import jax.numpy as jnp
from jaxnnls.pdip_relaxed import solve_relaxed_nnls

tol = data_scaled_solver_tol(q) if solver_tol is None else solver_tol
x, s, z, converged, pdip_iter = solve_nnls(Q, q, solver_tol=tol, max_iter=max_iter)
y, sy, zy = x / D, s / D, z * D

# Polish (PyAutoArray#573): the data-scaled tolerance leaves s * z ~ 1e-10 .. 1e-9, far above
# ``target_kappa``, so the relaxed solve below would have to push toward the boundary from z / s ~ 1e13
# and its fixed 50-iteration while_loop overshoots to NaN. A few tight PDIP iterations on the scaled
# system, warm-started from the mapped iterate, bring s * z down to the jaxnnls tolerance first. If the
# polish does not converge (the scaled dual is what diverges on #571's systems from a cold start), the
# mapped iterate is kept, i.e. the pre-polish behaviour.
# Polish (PyAutoArray#573, #594): the data-scaled tolerance also judges complementarity, so it leaves
# s * z ~ 1e-10 .. 1e-9 -- a column fnnls holds at zero can keep x ~ tol / z (11.5 % of the reference flux on
# inactive columns of the euclid vis_lp system), and the relaxed solve of the backward pass would have to push
# toward the boundary from z / s ~ 1e13, where its fixed 50-iteration while_loop overshoots to NaN. A few
# tight PDIP iterations on the scaled system, warm-started from the mapped iterate, bring s * z down to the
# jaxnnls tolerance. If the polish does not converge (the scaled dual is what diverges on #571's systems from
# a cold start), the mapped iterate is kept, i.e. the pre-polish behaviour.
yp, sp, zp, polish_converged, polish_iter = solve_nnls(
Q_pc, q_pc, max_iter=RAW_BACKWARD_POLISH_MAX_ITER, init=(y, sy, zy)
Q_pc, q_pc, max_iter=RAW_POLISH_MAX_ITER, init=(y, sy, zy)
)
ok = jnp.logical_and(
polish_converged == 1,
jnp.all(jnp.isfinite(yp)) & jnp.all(sp > 0) & jnp.all(zp > 0),
)
yp, sp, zp = (jnp.where(ok, a, b) for a, b in ((yp, y), (sp, sy), (zp, zy)))
return yp, converged, pdip_iter, (yp, sp, zp), ok, polish_iter


def _raw_forward_backward_point(
Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter
):
"""
The forward value and the relaxed-KKT point of the ``"raw"`` mode (shared by
the custom-vjp forward pass and :func:`raw_forward_backward_status`).

Returns ``(y, converged, pdip_iter)`` with ``y`` the polished forward value
(:func:`_raw_forward_polished`) and ``converged`` / ``pdip_iter`` those of the
raw forward solve, the relaxed point ``(yr, sr, zr)`` the backward pass
differentiates at, and the status ``(relaxed_converged, relaxed_iter,
polish_converged, polish_iter)``.
"""
from jaxnnls.pdip_relaxed import solve_relaxed_nnls

y_out, converged, pdip_iter, (yp, sp, zp), ok, polish_iter = _raw_forward_polished(
Q_pc, q_pc, Q, q, D, solver_tol, max_iter
)
yr, sr, zr, relaxed_converged, relaxed_iter = solve_relaxed_nnls(
Q_pc, q_pc, yp, sp, zp, target_kappa=target_kappa
)
status = (relaxed_converged, relaxed_iter, ok.astype(int), polish_iter)
return (y, converged, pdip_iter), (yr, sr, zr), status
return (y_out, converged, pdip_iter), (yr, sr, zr), status


@lru_cache(maxsize=None)
Expand All @@ -251,16 +282,25 @@ def _solve_nnls_raw_forward_with(target_kappa, solver_tol, max_iter):
linear-object-only (MGE) systems, Jacobi scaling turns signal-free columns
whose diagonal is only the no-regularization floor into degenerate
coordinates that make the PDIP dual diverge; the raw solve does not.
The mapped iterate is then *polished*: at most
:data:`RAW_POLISH_MAX_ITER` PDIP iterations on ``(Q_pc, q_pc)`` at
jaxnnls's tight tolerance, warm-started from it (kept only if the polish
converges with a finite, strictly interior point). The polished iterate is
the returned ``y`` (PyAutoArray#594): the data-scaled stop judges
complementarity ``s * z`` against the loose tolerance, so an unpolished
column fnnls holds at zero can keep ``x ~ tol / z`` -- 11.5 % of the fnnls
flux on inactive columns of the euclid vis_lp system, invisible to the
objective and log likelihood; polished, <= 3.3e-4 over 81 corpus systems
for ~5 extra iterations. The custom-vjp primal and forward rule compute it
through the same :func:`_raw_forward_polished`, so plain, jitted and
differentiated calls return the same value. ``converged`` and
``pdip_iter`` are those of the raw forward solve.
- **Backward:** the relaxed-KKT implicit derivative on ``Q_pc`` (as the
Jacobi mode), started from the mapped iterate after a *polish*: at most
:data:`RAW_BACKWARD_POLISH_MAX_ITER` PDIP iterations on ``(Q_pc, q_pc)``
at jaxnnls's tight tolerance, warm-started from the mapped iterate
(kept only if it converges). Without it the loose forward tolerance
leaves complementarity ``s * z`` orders of magnitude above
Jacobi mode), started from the polished point. Without the polish the
loose forward tolerance leaves ``s * z`` orders of magnitude above
``target_kappa`` and the relaxed solve diverges to NaN on a fraction of
points (PyAutoArray#573); with it the relaxed solve converges in about
one iteration. The primal ``y`` is the unpolished forward solution, so
the forward value is unchanged. The relaxed-KKT pass on the raw,
one iteration. The relaxed-KKT pass on the raw,
ill-conditioned ``Q`` produces NaN gradients, which is why Jacobi scaling
was introduced. ``(Q, q, D)`` get zero cotangents: ``y`` depends only on
``(Q_pc, q_pc)``, and the caller's autodiff carries the dependence of
Expand All @@ -274,11 +314,10 @@ def _solve_nnls_raw_forward_with(target_kappa, solver_tol, max_iter):
from jaxnnls.diff_qp import diff_nnls

def primal(Q_pc, q_pc, Q, q, D):
tol = data_scaled_solver_tol(q) if solver_tol is None else solver_tol
x, _, _, converged, pdip_iter = solve_nnls(
Q, q, solver_tol=tol, max_iter=max_iter
y_out, converged, pdip_iter, _, _, _ = _raw_forward_polished(
Q_pc, q_pc, Q, q, D, solver_tol, max_iter
)
return x / D, converged, pdip_iter
return y_out, converged, pdip_iter

def forward(Q_pc, q_pc, Q, q, D):
out, (yr, sr, zr), status = _raw_forward_backward_point(
Expand All @@ -305,8 +344,9 @@ def raw_forward_backward_status(
``(relaxed_converged, relaxed_iter, polish_converged, polish_iter)``.

``relaxed_*`` describe the relaxed-KKT solve whose point the gradient is
taken at; ``polish_*`` the tight warm-started PDIP polish before it
(``polish_converged == 0`` means the mapped iterate was used unpolished).
taken at; ``polish_*`` the tight warm-started PDIP polish before it,
whose iterate is also the forward value (``polish_converged == 0`` means the
mapped raw iterate was used unpolished, for both).
Arguments are those of :func:`solve_nnls_primal_raw_forward`.
"""
return _raw_forward_backward_point(
Expand All @@ -319,7 +359,8 @@ def solve_nnls_primal_raw_forward(
):
"""
The ``"raw"`` positive-only mode: forward PDIP on the raw system with a
data-scaled tolerance, backward pass on the Jacobi-scaled system. Returns
data-scaled tolerance, polished on the Jacobi-scaled system, backward pass
on the Jacobi-scaled system. Returns
``(y, converged, pdip_iter)`` with ``x = D * y``; see
:func:`_solve_nnls_raw_forward_with`.
"""
Expand Down
10 changes: 10 additions & 0 deletions test_autoarray/inversion/inversion/files/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -15,3 +15,13 @@
Captured 2026-09-25 on CPU fp64 via a `jax.debug.callback` on `reconstruction_positive_only_from`, with
autolens_workspace_test 5ec64413d2, PyAutoArray 3de624b5b9, PyAutoGalaxy 70a61e26cd, PyAutoLens 86054bbc19,
PyAutoFit dd9fbe0aab, jax 0.10.2, numpy 2.5.3.
- `mge_solver_reference_systems.npz` — 9 positive-only MGE systems (60x60 fp64) with their fnnls reference
solutions, for the raw-forward PDIP amplitude regression (PyAutoArray#594): keys `Q_<name>` / `q_<name>` /
`x_ref_<name>` for names `k0`..`k7` (the 8 #571 SLaM `source_lp[1]` systems, identical to
`mge_slam_nnls_systems.npz`) and `euclid_vis_lp` (the euclid pipeline vis_lp system behind
`test_latent_euclid_variables_traces_under_jax_jit`, captured with PyAutoArray d4298445). `meta` is a JSON
string with per-system `group`, `source_column_index_list`, `max_abs_q`, `cond_Q`, `category` and reference
active-column counts, plus provenance. Copied verbatim (no recomputation) from the autolens_profiling solver
corpus `results/lens/solver/corpus/{slam_fixture_571,euclid_vis_lp}.npz` + `manifest.json` at
autolens_profiling 3ad68af (phase-1 record `complete/2026/09/linear-solver-accuracy-study.md`) by a one-off
script outside the repo that checks symmetry of `Q`, `x_ref >= 0` and `max|q|` against the manifest.
Binary file not shown.
Loading
Loading