Skip to content

feat(inversion): certified active-set positive solver on the JAX path, opt-in, mapper-only dispatch (phase A) #566

Description

@Jammy2211

Overview

Promote the certified active-set positive-only solver from the autolens_profiling harness into PyAutoArray — phase A of the certified-positive-solver epic: the library implementation, opt-in, with structure-aware dispatch. The fixed-lens-light and HST-GPU-residue campaigns measured this solver returning the library's own constrained optimum (pins ≤ 5.6e-10 on every leg) at 1.2–2.6× the speed of the JAX PDIP on source-only inversions (A100 whole call 31.7 vs 50.8 ms), but every gain is still a harness monkeypatch. Phase A ships it behind Settings/config (positive_only_solver: pdip remains the default), selected only on the JAX backend for mapper-only inversions (MGE-inclusive systems never certified and keep PDIP; the NumPy path is untouched because the NumPy certified scheme measured slower than fnnls). The active set is searched under stop_gradient and the final masked Cholesky solve is autodiffed, giving the exact implicit active-set gradient. Phase B — the production-composition jit(vmap) benchmark and the default flip — is filed as PyAutoMind/draft/research/autolens_profiling/certified_solver_production_default.md and is not this issue.

Plan

  1. Add a JAX certified active-set positive solver to PyAutoArray as a new util module: budgeted lax.while_loop that stops at certification (no wasted factorizations), primal + dual (KKT) certification, permanent fixed set, and a lax.cond fallback to the existing PDIP when the budget is exhausted.
  2. Make it differentiable the right way: the active set is found under stop_gradient, then one final masked Cholesky solve on that set is differentiated by autodiff — the exact implicit active-set derivative — and tested against finite differences and against PDIP's custom_vjp gradient.
  3. Wire it behind Settings/config as opt-in (positive_only_solver: pdip stays the default): certified is selected only on the JAX backend for mapper-only inversions (no linear light-profile/MGE coefficients); MGE-inclusive systems and the whole NumPy path keep today's solvers unchanged. Solver observability (certified flag, passes) is exposed without changing return types.
  4. Tests: numerics vs a SciPy NNLS reference and vs PDIP, forced fallback bit-exact, edge-zero subsetting preserved, jit and vmap execution, gradient parity, dispatch by composition and backend, settings/config round trip, NumPy path untouched.
  5. Validate downstream on the JAX parity scripts with the setting flipped on through a workspace config override (no workspace commits), and record scalar-jit timings on the RTX for the PR body. The production default and the batched (vmap) policy are phase B, filed as a new draft prompt referencing this issue.
Detailed implementation plan

Affected Repositories

  • PyAutoArray (primary, only)

Work Classification

Library (source + tests + config). Additive: new optional kwargs and settings; no API break.

Branch Survey

Repository Current Branch Dirty?
./PyAutoArray main (22e6d60) clean

Suggested branch: feature/certified-positive-solver
Worktree root: ~/Code/PyAutoLabs-wt/certified-positive-solver/ (conflict guard: no active claim on PyAutoArray; the unregistered delaunay-area-magnification-audit worktree is legacy, not a claim)

Health at planning time

Heart RED (release validation FAILED (stage integrate)) plus YELLOW workspace-validation and manifest-drift reasons — unrelated to this change. Shipping a library PR will require the development-only human Heart-RED override; merge stays /prm behind the library freeze gate. PyAutoArray #553–#556 are merged but unreleased (pending-release chain).

Classification, branch, worktree

  • Work type: library (PyAutoArray only). No API break: new optional kwargs and settings only.
  • Branch feature/certified-positive-solver; worktree ~/Code/PyAutoLabs-wt/certified-positive-solver/PyAutoArray (via /start_library). Conflict guard: no conflict (no active.md claim on PyAutoArray; the other PyAutoArray worktree delaunay-area-magnification-audit is unregistered — legacy, note only).
  • Mind: rescope the existing prompt with a dated "Scoped 2026-09-23" block (phase A here, phase B filed separately), issue in PyAutoArray, row certified-positive-solver, status: library-dev; new draft draft/research/autolens_profiling/certified_solver_production_default.md (phase B).

Step 1 — autoarray/util/jax_active_set.py (new; never imports jax at module level, like jax_nnls.py)

  • masked_solve(Q, q, fixed): port of _masked_solve (identity on fixed rows/cols, zeroed RHS, cholesky + cho_solve), static shapes.
  • certify(Q, q, x, fixed, permanent, tau_rel) → (certified, primal_violations, dual_violations) per the harness tests (g = Qx − q; tau_x = tau_rel·max|x|, tau_g = tau_rel·max|q|; permanent never freeable).
  • active_set_search(Q, q, permanent, pass_budget, tau_rel): pass 0 unconstrained, seed Z = permanent | (x0 < 0), then lax.while_loop(cond=(~certified) & (passes < budget), body=one masked solve + certify + free_all update); returns (Z_final, certified, passes). Wrapped in stop_gradient on its inputs so no reverse-mode path goes through the loop (while_loop is not reverse-differentiable).
  • solve_certified(Q, q, permanent=None, pass_budget=16, tau_rel=1e-9) → (x, certified, passes): runs the search, then x = masked_solve(Q, q, Z_final) outside the stop_gradient so autodiff of the final Cholesky solve yields the implicit active-set gradient; x is exactly zero on Z. Everything computed in float64 (the NNLS stays fp64 under use_mixed_precision, as the Settings docstring already states).
  • solve_certified_with_fallback(Q, q, pdip_fn, fallback: bool, ...): lax.cond(certified, lambda: x, pdip_fn) when fallback, else returns the last iterate with certified=False flagged. Document the vmap both-branches cost in the docstring and point at phase B.

Step 2 — inversion_util.reconstruction_positive_only_from (JAX branch, :328-393)

  • New kwarg solver: str = "pdip" ("pdip" | "certified"), read by the caller from settings. In the JAX branch, after the existing Jacobi scaling, branch: pdip → today's solve_nnls_primal (byte-identical behaviour); certified → jax_active_set.solve_certified_with_fallback(Q_pc, q_pc, pdip_fn=lambda: solve_nnls_primal(Q_pc, q_pc, ...), fallback=settings.certified_fallback == "pdip", pass_budget=settings.certified_pass_budget, tau_rel=settings.certified_tau_rel), then unscale result * D as today. permanent=None (the caller already subsets to solve_ids_to_keep; keep the kwarg for callers that do not).
  • Observability: a module-level last_certified_stats is not jit-safe; instead return path unchanged and expose certified/passes via an optional stats: dict out-parameter populated with traced arrays (the caller may jax.debug.callback them), mirroring the existing NumPy stats convention. No change to the return type.
  • NumPy branch: untouched (solver ignored with a one-line comment citing the measured no-win).

Step 3 — AbstractInversion.reconstruction (abstract.py:599-691) and Settings

  • Resolve solver = self._positive_only_solver once: "certified" iff settings.positive_only_solver == "certified" and self.use_jax and self.has(cls=Mapper) and not self.has(cls=AbstractLinearObjFuncList); else "pdip". Pass solver= in both the subset and full branches. Import AbstractLinearObjFuncList (the module imports only LinearObj/Mapper today). Record the decision on the instance (self.positive_only_solver_used) for tests/observability.
  • Settings (autoarray/settings.py): constructor kwargs positive_only_solver=None, certified_pass_budget=None, certified_fallback=None, certified_tau_rel=None; properties with the try/except KeyError fallback pattern; defaults pdip, 16, pdip, 1.0e-9. Budget 16 is the measured every-draw-certifies bound (rect worst 11, Euclid rect 11) with margin — cheap because the loop exits at certification. Add the keys under inversion: in autoarray/config/general.yaml with one-line comments (what each does, "opt-in until phase B", the vmap caveat). Extend the settings dict round-trip.
  • use_positive_only_solver=False path unchanged; solve_ids_to_keep semantics unchanged (edge zeroing preserved).

Step 4 — tests (test_autoarray/)

  • test_autoarray/util/test_jax_active_set.py (pytest.importorskip("jax"), jax_enable_x64): (a) on seeded random SPD QPs (n = 20, 60; some with negative unconstrained solutions) solve_certified matches scipy.optimize.nnls to atol 1e-10 and PDIP (solve_nnls_primal) to rtol 1e-8; (b) passes ≤ budget and the loop stops early (assert passes equals the harness's NumPy reference active_set_certified count when ported into the test as a small helper, or simply < budget on easy systems); (c) permanent fixed indices stay exactly zero and are never released; (d) exhausted budget (pass_budget=1 on a hard system) with fallback returns PDIP bit-exactly (np.array_equal) and without fallback returns certified=False; (e) jax.jit and jax.vmap over B=4 distinct systems produce the per-system scalar results; (f) gradient: jax.grad(lambda q: w @ solve_certified(Q, q)[0]) equals central finite differences on the free set to rtol 1e-6 and agrees with jax.grad through solve_nnls_primal to rtol 1e-4 (PDIP's gradient is a relaxed approximation; tolerance declared before running, recorded in the test docstring); (g) module never imports jax at import time (mirror test_jax_nnls.py:10).
  • test_inversion_util.py: reconstruction_positive_only_from(..., solver="certified", xp=jnp) equals solver="pdip" within 1e-8 on the existing fixtures; solver is ignored on NumPy (xp=np result unchanged, fnnls still called — monkeypatch spy).
  • test_factory.py / a new test_positive_only_dispatch.py: with Settings(positive_only_solver="certified") and xp=jnp: mapper-only inversion (MockMapper) records positive_only_solver_used == "certified" and matches PDIP; an inversion containing MockLinearObjFuncList records "pdip"; xp=np records "pdip"; edge-zeroed subset branch still zeroes edge pixels (compare with use_edge_zeroed_pixels PDIP result).
  • test_settings_dict.py: the four new keys round-trip; test_jax_nnls.py-style default test: knobs default to pdip/16/pdip/1e-9 from the packaged config.

Step 5 — downstream validation (evidence for the PR body; no workspace commits)

  • In the task worktree's autolens_workspace_test (attach via worktree_add_repo or use its config override): set general.yaml inversion.positive_only_solver: certified in a scratch config dir passed by --config / WORKSPACE-config mechanism the scripts already use, run scripts/imaging/jax_likelihood/{lp,rectangular_mge,potential_correction}.py, multi_dataset/jax_likelihood/delaunay.py, weak/jax_grad.py and one Delaunay pixelization cell; assert each printed likelihood/gradient matches the run with pdip (record the relative differences) and prove the JAX path was taken (e.g. positive_only_solver_used printed or a jax.debug.callback counter). rectangular_mge must dispatch to pdip (MGE present) — assert it.
  • RTX 2060 scalar-jit timing on the fixed-light system through the library setting (not the monkeypatch): reuse autolens_profiling fixed_light_trace.py route b with a config override selecting certified, N=1500 Delaunay: expect the route-d-like whole-call number and a 1e-9 pin against route b/pdip. A100 timing and the vmap policy are phase B.

Step 6 — docs

  • Docstrings in the new module and in reconstruction_positive_only_from (algorithm, certification, budgets, fallback, gradient contract, the vmap caveat, why NumPy is unchanged, phase-B pointer). autoarray/config/general.yaml comments. A short entry in PyAutoArray's changelog/release notes file if one exists (check docs/ — release_notes heading convention per Mind memory: breaking changes need an "API changes" heading; this is additive).

Delegation

  • Steps 1-4 + 6: one Opus subagent (worktree, progress file ~/Code/PyAutoLabs-wt/certified-positive-solver/.progress_cps.md, Monitor). Gates: pytest test_autoarray -q (record baseline count first), black --check/ruff per repo convention, the new JAX tests pass under ~/venv/PyAuto.
  • Step 5: a second Opus subagent after I review the diff (RTX via ~/venv/PyAutoGPU with JAX_PLATFORM_NAME=cuda).
  • Ship: /ship_library (Heart RED → development-only human override will be needed again; merge stays /prm with the library freeze gate).

Verification

  • PyAutoArray suite green with the new tests (count reported vs baseline); new module import-safe without jax.
  • Numerics: certified == SciPy NNLS (1e-10) and == PDIP (1e-8) on seeded systems; exhausted budget == PDIP bit-exact; edge-zero subsetting preserved; jit + vmap OK; gradient FD parity.
  • Dispatch: certified only on (JAX, mapper-only, setting on); default config leaves every existing test byte-identical (run the suite once with the default and once with positive_only_solver: certified forced — the latter may only differ on JAX mapper-only tests).
  • Downstream parity scripts equal under both solvers with the JAX path proven; RTX scalar timing recorded.

Out of scope (phase B and later)

Production default flip and the vmap/batched policy (PDIP-under-vmap baseline never measured; lax.cond both-branch cost), any PyAutoFit Fitness._vmap change, a cond-free batched fallback, NumPy certified/factor-reuse solvers (measured no-win), MGE-inclusive certified solves (never certify), mixed-precision solves, Euclid rectangular budget sweeps, the fixed-lens-light workflow (S3 conversion of MGE to regular profiles) as a production stage.

Original Prompt

Click to expand starting prompt

Implement and optimize certified positive solver with structure-aware CPU and JAX dispatch

Type: feature
Target: PyAutoArray
Repos:

  • PyAutoArray
    Difficulty: large
    Autonomy: supervised
    Priority: normal
    Status: scoped — phase A issued 2026-09-23 (Fable start_dev); phase B filed as draft/research/autolens_profiling/certified_solver_production_default.md
    Epic: certified-positive-solver
    Phase: A
    Consequence: judge
    Witness: Representative mapper-only and MGE-inclusive fits preserve constrained reconstruction/evidence within declared tolerances, including forced fallback and JAX jit/vmap; benchmark records justify every automatically selected solver against its backend's current baseline.
    Review-minutes: 20
    Unattended: needs-slicing
    Filed: 2026-09-15

Scoped 2026-09-23 (Fable start_dev) — phase A is THIS issue, phase B is filed separately

Phase A (PyAutoArray only, opt-in): a JAX certified active-set positive solver in a new
autoarray/util/jax_active_set.py (budgeted lax.while_loop that stops at certification; primal +
dual KKT certification; permanent fixed set; lax.cond PDIP fallback on an exhausted budget); the
active set is searched under stop_gradient and the final masked Cholesky solve is autodiffed, so
the gradient is the exact implicit active-set derivative (tested against finite differences and
PDIP's custom_vjp). Wired behind Settings / general.yaml keys positive_only_solver
(pdip default | certified), certified_pass_budget (16), certified_fallback (pdip | none),
certified_tau_rel (1e-9); selected only on the JAX backend for mapper-only inversions
(has(Mapper) and not has(AbstractLinearObjFuncList)) — MGE-inclusive systems never certified
(fixed_light_probe: 40 passes, cond 4e10) and keep PDIP; the NumPy path is untouched (the NumPy
certified scheme measured 3-7 % slower than fnnls, factor-reuse lost to the memo).

Phase B (autolens_profiling, maybe PyAutoFit): measure the production composition
jax.jit(jax.vmap(fn)) with PDIP (the never-measured baseline) against the shipped certified solver
with fallback pdip and none, decide the production default and the batched policy (the
lax.cond fallback runs both branches under vmap: 44.6 vs 31.1 ms/lane at B=16 in residue phase 2).
Only phase B may flip the default.

Evidence added since filing: residue phases 1-3 (autolens_profiling#268/#273/#295) — the certified
solve at budget 7 is 10.2 ms of a 31.6 ms A100 call; whole-call 31.7 vs 50.8 ms PDIP; the harness
kernel is active_set_steps.active_set_masked_jax (static lax.scan, no gradient rule) and
library_solver_injection.certified_reconstruction_from (Jacobi scaling + lax.cond fallback).

Original request

Ok, first intake an issue which is to implement the cerified solve in the source code, have a final stab at speeding it up and optimizing it, and also doing this for CPU. That is, I guess the method should ask if the system it is solving is a Mapper only (e.g. sparse struture) or has dense MGE like structures, and thus chooses the solver which is most efficient for the task. Is it true that the cerified solver works less well than the original code when MGE is included?

Scope

Promote the experimental certified active-set positive solver from the profiling harness into @PyAutoArray, with a bounded final optimization pass on CPU and JAX/GPU. One production-library task and one PR; other repositories below are evidence and validation dependencies only.

Choose the efficient supported solver using backend and actual inversion composition: mapper-only versus systems containing linear-light/MGE coefficients. Use existing inversion metadata (linear_obj_list / has(cls)) before the numeric utility, without lens-profile imports into PyAutoArray. Mapper-only does not imply sparse assembled matrices: the cited benchmarks are dense. Preserve existing sparse paths/fallback unless explicitly supported.

Preserve positivity, fixed/edge-zero constraints, API/return behavior, NumPy fingerprint/memo warm starts, and existing backend/differentiation contracts. Certify primal and dual/KKT conditions, including inactive coefficients with invalid negative dual values. Fall back to the existing robust solver on non-certification, exhausted budgets or numerical failure. Fit settings/override and solver observability into existing APIs.

Use conservative bounded budgets: measured HST 7 Delaunay / 11 rectangular are empirical, not universal; Euclid rectangular already reaches 11 in a fiducial case. Ensure correct, efficient jit/vmap operation: lax.cond inside vmap can execute both solvers, so provide an explicit batching strategy and measure its actual cost, including failed cases.

Make a finite profiling-guided optimization pass. For native CPU, compare against production NumPy/SciPy fnnls separately from JAX-CPU versus PDIP. The existing NumPy certified prototype did not reliably beat fnnls: retain fnnls or keep the new CPU route opt-in unless measurements justify automatic selection. Likewise retain the existing MGE-inclusive solver until matched evidence supports changing it. No speedup is promised on every backend.

Acceptance

Reuse existing fixtures for a finite representative subset: mapper-only rectangular/Delaunay and MGE-inclusive systems, good/poor fits, small/large source sizes (around 500/1500/2500/4000 where useful), available CPU/GPU, fp64 reference. Record solve-only and whole-likelihood times separately, fallback frequency/cost, correctness tolerances and dispatch decisions. Automatically selected routes must preserve constrained reconstruction/evidence and avoid material runtime regression against their production baseline; unproven cases retain the baseline.

Cover positivity/KKT, forced fallback, edge/failure cases and existing derivative contracts. Run appropriate source tests and downstream JAX jit/vmap parity in autogalaxy_workspace_test/autolens_workspace_test; PyAutoArray's own suite is NumPy-only. Document configuration, backend/composition rules, measurements and limitations.

Evidence and boundaries

Formalizes the production-solver seed in PyAutoMind/ideas.md from autolens_profiling#259. In autolens_profiling:

  • results/notes/fixed_lens_light_source_only_2026_09.md
  • results/notes/fixed_lens_light_library_path_2026_09.md
  • results/notes/fixed_lens_light_hardware_2026_09.md
  • results/notes/fixed_lens_light_low_likelihood_draws_2026_09.md
  • results/notes/fixed_lens_light_source_pixel_scaling_2026_09.md
  • results/notes/fixed_lens_light_verdict_2026_09.md
  • results/misc/fixed_light_probe/{delaunay_1500,rectangular_1521}.{md,json}

Joint 60-MGE+Delaunay1500 free_all failed certification after 40 passes (41 factorizations), while free_one certified at 32; PDIP converged in 22 iterations. Joint rectangular1521 variants failed after 40; PDIP took 21. Source-only free_all took 2 Delaunay / 7 rectangular passes. These demonstrate poorer convergence with MGE, not a matched joint-GPU timing comparison.

Reuse the completed CPU decomposition records PyAutoMind/complete/2026/09/fixed-light-numba-phase1.md and complete/2026/09/fixed-light-numba-solver.md, with the final campaign verdict in complete/2026/09/fixed-lens-light-numba-cpu.md; do not duplicate that evidence or the separate hst_gpu_non_solver_residue_programme.md.

Fixed-light measurements use pre-solved intensities and exclude preparation; they do not prove that freezing one estimate throughout a search preserves every likelihood. Out of scope: new sparse operators, changes to statistical modeling/positivity, and broad profiling campaigns.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions