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: 2 additions & 2 deletions autoarray/config/general.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,8 @@ inversion:
no_regularization_add_to_curvature_diag_value : 1.0e-3 # The default value added to the curvature matrix's diagonal when regularization is not applied to a linear object, which prevents inversion's failing due to the matrix being singular.
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_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_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
124 changes: 101 additions & 23 deletions autoarray/util/jax_nnls.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,11 @@
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.
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.

JAX is imported inside functions, never at module level (see
``docs/agents/jax_and_decorators.md``); this module must only be imported
Expand All @@ -39,7 +43,7 @@
from functools import lru_cache


def solve_nnls(Q, q, solver_tol=None, max_iter=50):
def solve_nnls(Q, q, solver_tol=None, max_iter=50, init=None):
"""
Solve the non-negative least squares problem with the jaxnnls PDIP
algorithm, with configurable convergence tolerance and iteration cap.
Expand All @@ -59,6 +63,10 @@ def solve_nnls(Q, q, solver_tol=None, max_iter=50):
``min(n * eps * 5e3, 1e-2)``.
max_iter
Maximum number of PDIP iterations (jaxnnls hard-codes 50).
init
Optional ``(x, s, z)`` warm start (strictly positive ``s`` and ``z``)
replacing jaxnnls's ``initialize``. ``None`` (default) is the upstream
cold start.

Returns
-------
Expand All @@ -69,7 +77,7 @@ def solve_nnls(Q, q, solver_tol=None, max_iter=50):
import jax.numpy as jnp
from jaxnnls.pdip import EPSILON, initialize, pdip_pc_step

x, s, z = initialize(Q, q)
x, s, z = initialize(Q, q) if init is None else init

if solver_tol is None:
solver_tol = jax.lax.min(Q.shape[0] * EPSILON, 1e-2)
Expand Down Expand Up @@ -179,10 +187,57 @@ 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


def _raw_forward_backward_point(
Q_pc, q_pc, Q, q, D, target_kappa, 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)``.
"""
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.
yp, sp, zp, polish_converged, polish_iter = solve_nnls(
Q_pc, q_pc, max_iter=RAW_BACKWARD_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)))

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


@lru_cache(maxsize=None)
def _solve_nnls_raw_forward_with(target_kappa, solver_tol, max_iter):
"""
Build (and cache) the ``"raw"``-mode solver (PyAutoArray#571).
Build (and cache) the ``"raw"``-mode solver (PyAutoArray#571, #573).

The returned function takes the Jacobi-scaled system ``(Q_pc, q_pc)``
(``Q_pc = D Q D``, ``q_pc = D q``) together with the raw system ``(Q, q)``
Expand All @@ -196,38 +251,43 @@ 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.
- **Backward:** exactly today's Jacobi-mode pass, i.e. the relaxed-KKT implicit
derivative on ``Q_pc``, started from the mapped iterate. 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 those, and of ``D``, on the raw inputs.
- **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
``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,
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
those, and of ``D``, on the raw inputs.

The backward-pass convergence is observable through
:func:`raw_forward_backward_status`.
"""
import jax
import jax.numpy as jnp
from jaxnnls.diff_qp import diff_nnls
from jaxnnls.pdip_relaxed import solve_relaxed_nnls

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

def primal(Q_pc, q_pc, Q, q, D):
y, _, _, converged, pdip_iter = raw_solve(Q, q, D)
return y, converged, pdip_iter
return x / D, converged, pdip_iter

def forward(Q_pc, q_pc, Q, q, D):
y, sy, zy, converged, pdip_iter = raw_solve(Q, q, D)
yr, sr, zr, _, _ = solve_relaxed_nnls(
Q_pc, q_pc, y, sy, zy, target_kappa=target_kappa
out, (yr, sr, zr), status = _raw_forward_backward_point(
Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter
)
return (y, converged, pdip_iter), (Q_pc, yr, sr, zr, Q, q, D)
return out, (Q_pc, yr, sr, zr, status[0], Q, q, D)

def backward(res, output_grad):
Q_pc, yr, sr, zr, Q, q, D = res
Q_pc, yr, sr, zr, _, Q, q, D = res
dQ_pc, dq_pc = diff_nnls(Q_pc, yr, sr, zr, output_grad[0])
return dQ_pc, dq_pc, jnp.zeros_like(Q), jnp.zeros_like(q), jnp.zeros_like(D)

Expand All @@ -236,6 +296,24 @@ def backward(res, output_grad):
return primal


def raw_forward_backward_status(
Q_pc, q_pc, Q, q, D, target_kappa=1e-3, solver_tol=None, max_iter=50
):
"""
Diagnostic (not differentiable): the convergence of the ``"raw"`` mode's
backward-pass preparation for one system, as the integer tuple
``(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).
Arguments are those of :func:`solve_nnls_primal_raw_forward`.
"""
return _raw_forward_backward_point(
Q_pc, q_pc, Q, q, D, target_kappa, solver_tol, max_iter
)[2]


def solve_nnls_primal_raw_forward(
Q_pc, q_pc, Q, q, D, target_kappa=1e-3, solver_tol=None, max_iter=50
):
Expand Down
8 changes: 8 additions & 0 deletions test_autoarray/inversion/inversion/files/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,3 +7,11 @@
`autolens_profiling/scripts/imaging/hazards/mge_nnls_capture.py` (autolens_profiling d6926af), run
2026-09-24 on CPU fp64 with PyAutoArray 7fa8d2714f, PyAutoGalaxy 70a61e26cd, PyAutoLens 86054bbc19,
PyAutoFit a736840127, jax 0.10.2.
- `mge_grad_nan_systems.npz` — 4 positive-only systems `Q_<k>` (20x20) / `q_<k>` captured from the
autolens_workspace_test `scripts/imaging/jax_grad/mge.py` model (MGE source, NFWSph + ExternalShear) at
`physical_values_from_prior_medians + jax.random.uniform(PRNGKey(p), minval=0.01, maxval=0.05)` for
p = 2, 10, 12, 14 (PyAutoArray#573): on PyAutoArray 3de624b5 the `"raw"`-mode gradient is NaN on each (the
relaxed-KKT backward solve diverges from the loose raw-forward iterate); `meta` holds the per-system JSON.
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.
Binary file not shown.
92 changes: 92 additions & 0 deletions test_autoarray/inversion/inversion/test_nnls_mge_convergence.py
Original file line number Diff line number Diff line change
Expand Up @@ -323,3 +323,95 @@ def test__control__well_conditioned_pdip_unchanged(jnp, n, seed):
np.testing.assert_allclose(
x_raw, x_scipy, rtol=0, atol=1e-8 * np.abs(x_scipy).max()
)


# ---------------------------------------------------------------------------------------------------------------
# PyAutoArray#573: NaN gradients of the "raw" mode.
#
# `files/mge_grad_nan_systems.npz` holds 4 (20 x 20) systems captured from the autolens_workspace_test
# `jax_grad/mge.py` model (MGE source, NFWSph + shear) at the PRNGKey perturbations 2, 10, 12 and 14 (see
# `files/README.md`). The raw forward solve stops at the data-scaled tolerance with s * z ~ 1e-10 .. 2.5e-9, far
# above `nnls_target_kappa = 1e-11`; the relaxed-KKT solve on the Jacobi system then has to push toward the
# boundary from z / s ~ 1e13 and hits its 50-iteration cap with NaN (or "converges" with s < 0), so the gradient
# is NaN. The fix polishes the mapped iterate with a few tight PDIP iterations on the Jacobi system first.
# ---------------------------------------------------------------------------------------------------------------

GRAD_NAN_FIXTURE = Path(__file__).parent / "files" / "mge_grad_nan_systems.npz"


def _load_grad_nan_systems():
with np.load(GRAD_NAN_FIXTURE) as data:
meta = json.loads(str(data["meta"]))
systems = [
(np.asarray(data[f"Q_{s['key']}"]), np.asarray(data[f"q_{s['key']}"]))
for s in meta["systems"]
]
return meta, systems


GRAD_NAN_META, GRAD_NAN_SYSTEMS = _load_grad_nan_systems()
GRAD_NAN_IDS = [f"prng{s['prng_key']}" for s in GRAD_NAN_META["systems"]]


def test__grad_nan_fixture_is_the_captured_jax_grad_mge_set():
assert GRAD_NAN_FIXTURE.stat().st_size < 50_000
assert [s["prng_key"] for s in GRAD_NAN_META["systems"]] == [2, 10, 12, 14]
for Q, q in GRAD_NAN_SYSTEMS:
assert Q.shape == (20, 20) and q.shape == (20,)


@requires_jax
@pytest.mark.parametrize("jit", [False, True], ids=["eager", "jit"])
@pytest.mark.parametrize("index", range(len(GRAD_NAN_SYSTEMS)), ids=GRAD_NAN_IDS)
def test__raw_mode_gradient_is_finite_on_the_captured_grad_nan_systems(jnp, index, jit):
"""Red on PyAutoArray 3de624b5 (#572): the gradient is NaN on all four systems eagerly, and under jit on
prng10 / prng14 (the jitted NaN is rounding-sensitive; prng2 / prng12 happen to pass jitted on main).
"""
import jax

Q, q = GRAD_NAN_SYSTEMS[index]
w = jnp.linspace(0.5, 1.5, q.shape[0])

grad = jax.grad(lambda Q_, q_: w @ _raw(jnp, Q_, q_), argnums=(0, 1))
if jit:
grad = jax.jit(grad)
gQ, gq = grad(jnp.asarray(Q), jnp.asarray(q))

assert np.all(np.isfinite(np.asarray(gQ))) and np.all(np.isfinite(np.asarray(gq)))
assert np.any(np.asarray(gq) != 0.0)


def _backward_status(jnp, Q, q):
from autoarray.util.jax_nnls import raw_forward_backward_status

Qj, qj = jnp.asarray(Q), jnp.asarray(q)
Q_pc, q_pc, D = (jnp.asarray(a) for a in _jacobi(Q, q))
return [
int(v)
for v in raw_forward_backward_status(
Q_pc, q_pc, Qj, qj, D, target_kappa=1.0e-11, max_iter=PRODUCTION_MAX_ITER
)
]


@requires_jax
@pytest.mark.parametrize(
"system",
[("slam", k) for k in KEYS]
+ [("grad_nan", i) for i in range(len(GRAD_NAN_SYSTEMS))],
ids=[f"slam-{i}" for i in IDS] + [f"grad_nan-{i}" for i in GRAD_NAN_IDS],
)
def test__raw_mode_backward_pass_converges(jnp, system):
"""The backward pass reports convergence: the tight polish of the mapped iterate converges (measured <= 6
iterations) and the relaxed-KKT solve then converges well inside its 50-iteration cap (measured 1).
Not a red-on-main witness (``raw_forward_backward_status`` is new with #573); the gradient test is.
"""
kind, index = system
Q, q = (SYSTEMS if kind == "slam" else GRAD_NAN_SYSTEMS)[index]

relaxed_converged, relaxed_iter, polish_converged, polish_iter = _backward_status(
jnp, Q, q
)

assert polish_converged == 1, polish_iter
assert relaxed_converged == 1 and relaxed_iter < PRODUCTION_MAX_ITER, relaxed_iter
Loading