Skip to content

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

Merged
Jammy2211 merged 3 commits into
mainfrom
feature/certified-positive-solver
Sep 23, 2026
Merged

Jammy2211 merged 3 commits into
mainfrom
feature/certified-positive-solver

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Summary

Promotes the certified active-set positive-only (NNLS) solver, prototyped and measured in the autolens_profiling fixed-lens-light campaign, into PyAutoArray as an opt-in alternative to the jaxnnls PDIP solve on the JAX path.

  • New module autoarray.util.jax_active_set: a budgeted lax.while_loop of masked full-size Cholesky solves that stops once the iterate satisfies the KKT conditions (primal + dual, relative tolerance tau_rel), with the library's PDIP solve as a lax.cond fallback when the pass budget is exhausted.
  • Exact implicit gradient: the active-set search runs under stop_gradient and returns only the fixed set; the solution is one final differentiable masked solve on that set, so jax.grad gives the exact active-set derivative (PDIP's custom_vjp differentiates a relaxed central-path system).
  • Mapper-only JAX dispatch: AbstractInversion.positive_only_solver_used selects "certified" only when Settings.positive_only_solver == "certified", the inversion is on JAX, and it contains a Mapper and no linear light profiles / MGE (dense MGE blocks converge poorly: 60-MGE + Delaunay-1500 failed to certify in 40 passes).
  • NumPy path untouched: fnnls + warm-start memo stays (a NumPy port measured 3-7 % slower than fnnls).
  • Default stays pdip until phase B measures the batched (vmap) policy (under vmap the fallback cond runs both branches per lane).

API Changes

Additive only — no removals, no renamed symbols, default behaviour byte-identical (positive_only_solver: pdip).

  • reconstruction_positive_only_from gains keyword args solver="pdip" and stats=None.
  • Settings gains kwargs + properties positive_only_solver, certified_pass_budget, certified_fallback, certified_tau_rel (backed by new general.yaml inversion: keys with packaged fallbacks).
  • New property AbstractInversion.positive_only_solver_used; new module autoarray.util.jax_active_set.
  • Downstream note: harness code that wraps reconstruction_positive_only_from with a strict signature check must accept the new kwargs — the linked autolens_profiling PR does exactly that.
    See full details below.

Test Plan

  • PyAutoArray suite: 1584 baseline → 1616 passed under the default config (32 new tests), re-run at head 233cfc0.
  • Forced-certified sweep (config positive_only_solver: certified): 1614 passed; only the two default-value assertions differ.
  • Numerics: vs SciPy nnls 2.2e-16; vs PDIP 7.6e-12; exhausted budget == PDIP bit-exact; jit/vmap (B=4) match.
  • Gradients: vs central finite differences 3.0e-9 rel; vs PDIP custom_vjp 1.5e-5.
  • GPU (RTX 2060, CUDA): the new tests, 24 passed.
  • Downstream autolens_workspace_test JAX parity scripts identical under both solvers (Delaunay pixelization cell 2.3e-14, dispatched certified 3x; rectangular_mge stayed pdip; lp / potential_correction / weak jax_grad identical). multi_dataset/jax_likelihood/delaunay.py fails on main with both solvers (pre-existing, filed as a Mind draft, not caused by this PR).
  • RTX whole-call timing through the library setting: 734.9 ms (PDIP) → 566.1 ms (certified, budget 16) vs 618.2 ms for the harness monkeypatch (budget 7); logL rel ≤ 1.2e-10.
  • autolens_profiling suite against this branch: 795 passed / 5 skipped / 0 failed (with the linked harness fix).
  • Independent review (Brain review faculty): FINDINGS at 8fe4343 (one minor: the module-level jax skip in test_jax_active_set.py also skipped the no-module-level-jax import guard it contained) → fixed in 233cfc0 → CLEAN at 233cfc0. Claim "NumPy fnnls path unchanged" → basis-cited: the NumPy branch of reconstruction_positive_only_from is unchanged apart from a comment and a solver validation that only fires on an invalid value; the pre-existing fnnls tests pass in the 1616.
Full API Changes (for automation & release notes)

Added

  • autoarray.util.jax_active_set — masked_solve(Q, q, fixed), certify(Q, q, x, fixed, permanent, tau_rel), active_set_search(Q, q, permanent=None, pass_budget=16, tau_rel=1e-9), solve_certified(...), solve_certified_with_fallback(Q, q, pdip_fn, fallback=True, pass_budget=16, tau_rel=1e-9, permanent=None).
  • Settings(positive_only_solver=None, certified_pass_budget=None, certified_fallback=None, certified_tau_rel=None) and matching properties ("pdip"|"certified", 16, "pdip"|"none", 1.0e-9); invalid explicit values raise ValueError at construction.
  • general.yaml inversion: keys positive_only_solver: pdip, certified_pass_budget: 16, certified_fallback: pdip, certified_tau_rel: 1.0e-9.
  • AbstractInversion.positive_only_solver_used -> str.

Changed Signature

  • inversion_util.reconstruction_positive_only_from(..., solver: str = "pdip", stats: Optional[dict] = None) — both optional; invalid solver raises ValueError. stats receives solver (and, for certified, traced certified / passes).

Changed Behaviour

  • None by default. With positive_only_solver: certified, mapper-only JAX inversions use the certified solver; all others unchanged.

Migration

  • None required for users. Wrappers that monkeypatch reconstruction_positive_only_from and assert signature coverage must add solver="pdip", stats=None (autolens_profiling: linked PR).

Heart RED override

  • Heart verdict (re-read immediately before this PR, verbatim): RED release validation FAILED (stage integrate); YELLOW workspace validation not passing (4 failed, cloud#35579888156: autolens notebooks/cluster/modeling.ipynb, autolens notebooks/weak/a2744.ipynb, autolens scripts/cluster/modeling.py, +1 more); YELLOW manifest drift: hub organism blurb (organs present) — 7 mismatch(es) vs PyAutoMind/repos.yaml. Freeze: not frozen.
  • Branch gates passed: PyAutoArray tests 1616 passed at 233cfc0; downstream JAX parity scripts identical; GPU 24 passed; independent review CLEAN at 233cfc0.
  • Authorization: live user message i authorize, to the named issue-feat(inversion): certified active-set positive solver on the JAX path, opt-in, mapper-only dispatch (phase A) #566 (+ linked autolens_profiling fix) development override, per AUTONOMY.md "Human override for Heart RED (development only)".
  • Scope: commit / push / pending-release PRs only. No merge, no release, no CI bypass; this branch does not fix Heart, which stays RED for release purposes. Merge needs a separate human /prm with all required checks green.

Closes #566.

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

Jammy2211 and others added 3 commits September 23, 2026 14:53
…, opt-in, mapper-only dispatch

New autoarray/util/jax_active_set.py: budgeted lax.while_loop free-all
active-set search with primal + dual (KKT) certification and a permanent
fixed set, run under stop_gradient; the returned solution is one final
differentiable masked Cholesky solve (exact implicit active-set gradient);
lax.cond fallback to the PDIP solve on an exhausted budget.

reconstruction_positive_only_from gains solver="pdip"|"certified" and an
optional stats out-dict (traced certified/passes); the pdip path and the
NumPy fnnls path are unchanged. Settings gains positive_only_solver,
certified_pass_budget, certified_fallback, certified_tau_rel (packaged
defaults pdip/16/pdip/1e-9). AbstractInversion.positive_only_solver_used
selects certified only on JAX for mapper-only inversions (no
AbstractLinearObjFuncList).

Refs #566

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
…vmap, gradients, dispatch, settings

Refs #566

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Move the no-module-level-jax import guard into its own test module: the
module-level skip in test_jax_active_set.py (applied when jax is absent)
also skipped the guard, contradicting its comment, so the guard never ran
on the NumPy-only env it exists for.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
@Jammy2211

Copy link
Copy Markdown
Collaborator Author

Workspace PR: PyAutoLabs/autolens_profiling#299

@Jammy2211
Jammy2211 merged commit 11b9347 into main Sep 23, 2026
3 checks passed
@Jammy2211
Jammy2211 deleted the feature/certified-positive-solver branch September 23, 2026 14:20
@Jammy2211 Jammy2211 removed the pending-release PR queued for the next release build label Sep 26, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

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

1 participant