feat(inversion): certified active-set positive solver on the JAX path, opt-in, mapper-only dispatch (#566) - #567
Merged
Conversation
…, 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>
Collaborator
Author
|
Workspace PR: PyAutoLabs/autolens_profiling#299 |
This was referenced Sep 23, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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.
autoarray.util.jax_active_set: a budgetedlax.while_loopof masked full-size Cholesky solves that stops once the iterate satisfies the KKT conditions (primal + dual, relative tolerancetau_rel), with the library's PDIP solve as alax.condfallback when the pass budget is exhausted.stop_gradientand returns only the fixed set; the solution is one final differentiable masked solve on that set, sojax.gradgives the exact active-set derivative (PDIP'scustom_vjpdifferentiates a relaxed central-path system).AbstractInversion.positive_only_solver_usedselects"certified"only whenSettings.positive_only_solver == "certified", the inversion is on JAX, and it contains aMapperand no linear light profiles / MGE (dense MGE blocks converge poorly: 60-MGE + Delaunay-1500 failed to certify in 40 passes).pdipuntil phase B measures the batched (vmap) policy (undervmapthe fallbackcondruns both branches per lane).API Changes
Additive only — no removals, no renamed symbols, default behaviour byte-identical (
positive_only_solver: pdip).reconstruction_positive_only_fromgains keyword argssolver="pdip"andstats=None.Settingsgains kwargs + propertiespositive_only_solver,certified_pass_budget,certified_fallback,certified_tau_rel(backed by newgeneral.yamlinversion:keys with packaged fallbacks).AbstractInversion.positive_only_solver_used; new moduleautoarray.util.jax_active_set.reconstruction_positive_only_fromwith a strict signature check must accept the new kwargs — the linked autolens_profiling PR does exactly that.See full details below.
Test Plan
positive_only_solver: certified): 1614 passed; only the two default-value assertions differ.nnls2.2e-16; vs PDIP 7.6e-12; exhausted budget == PDIP bit-exact;jit/vmap(B=4) match.custom_vjp1.5e-5.certified3x; rectangular_mge stayedpdip; lp / potential_correction / weak jax_grad identical).multi_dataset/jax_likelihood/delaunay.pyfails on main with both solvers (pre-existing, filed as a Mind draft, not caused by this PR).test_jax_active_set.pyalso 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 ofreconstruction_positive_only_fromis unchanged apart from a comment and asolvervalidation 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 raiseValueErrorat construction.general.yamlinversion:keyspositive_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; invalidsolverraisesValueError.statsreceivessolver(and, for certified, tracedcertified/passes).Changed Behaviour
positive_only_solver: certified, mapper-only JAX inversions use the certified solver; all others unchanged.Migration
reconstruction_positive_only_fromand assert signature coverage must addsolver="pdip", stats=None(autolens_profiling: linked PR).Heart RED override
release validation FAILED (stage integrate); YELLOWworkspace validation not passing (4 failed, cloud#35579888156: autolens notebooks/cluster/modeling.ipynb, autolens notebooks/weak/a2744.ipynb, autolens scripts/cluster/modeling.py, +1 more); YELLOWmanifest drift: hub organism blurb (organs present) — 7 mismatch(es) vs PyAutoMind/repos.yaml. Freeze: not frozen.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)"./prmwith all required checks green.Closes #566.
Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code