Skip to content

fix: return the polished raw-forward PDIP iterate as the forward value (#594) - #595

Merged
Jammy2211 merged 3 commits into
mainfrom
feature/raw-pdip-forward-polish
Sep 30, 2026
Merged

Jammy2211 merged 3 commits into
mainfrom
feature/raw-pdip-forward-polish

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Summary

Closes #594 (epic linear-solver-programme, phase 2). The mapper-less positive-only solver (nnls_preconditioning_no_mapper: raw, #572) stops on an absolute infinity-norm KKT residual with a data-scaled tolerance, so complementarity s·z is judged against a threshold scaled by max|q| and a column fnnls holds at zero can keep x ≈ tol/z while reporting converged. Phase 1 (autolens_profiling#355) measured this on 81 captured systems: raw PDIP is converged 81/81 with objective gap ~1e-12, yet its source-column flux is >1e-3 off fnnls on 53/81 systems and on the euclid vis_lp system 11.5 % of the reference amplitude sits on reference-inactive columns (total_source_flux +5.76 %, the red euclid latent test since 2026-09-25). Log-likelihood, objective and KKT are all blind to it.

The fix returns the #573 (PR #574) polished iterate — ≤ 10 tight-tolerance PDIP iterations on the Jacobi-preconditioned system, warm-started from the raw iterate — as the forward value, not only as the start of the backward relaxed-KKT solve. _raw_forward_polished() is shared by the custom_vjp primal and forward so jit, grad and eager return the identical value; the backward pass is unchanged; the polish falls back to the raw iterate under the same ok rule as before. Measured on the 81-system corpus with this branch: pdip_raw is now bit-identical to phase 1's pdip_raw_polish candidate on 81/81 (all 48 SLaM points converged), worst inactive-column flux 3.3e-4 (was 0.115), source-flux misses 1/81 (was 53/81; the remaining one is the euclid source-column proxy, see below), polish accepted every time in 1–7 iterations (mostly 5). Euclid total_source_flux jit-vs-eager is now +7.47e-5 (was +5.76e-2) and tests/test_compute_latent_variable.py passes 19/19 on the pipeline's main with no config override.

Not admissible under phase 1's pre-registered rule as written (criteria 2–4 still fail: the euclid source-column proxy at 4.96e-2 because its single active source column carries ~0.4 % of the reference flux, one flat-direction system at amp_rel_max_sig 0.219 shared by every accurate candidate incl. tight solves, and 5/52 KKT ratios at the 3e-16 floating-point floor). Those are recorded rule weaknesses, not solver defects; the regression test gates on the flux metrics that the phase-1 evidence supports.

API Changes

No public signature changes. The raw-mode forward reconstruction (reconstruction_positive_only_from(..., preconditioning="raw") and solve_nnls_primal_raw_forward) now returns the polished iterate, so mapper-less (MGE / linear light profile) JAX reconstructions and every latent derived from them move at the ~1e-4 relative level on affected systems (euclid total_source_flux −5.4 %; log-likelihoods move at ~1e-11 relative — autolens_workspace_test pins unchanged, mge_group.py JIT now equals NumPy to 6e-12). converged / iterations still describe the raw forward solve; raw_forward_backward_status still reports the polish flag and count. RAW_POLISH_MAX_ITER added; RAW_BACKWARD_POLISH_MAX_ITER kept as an alias.
See full details below.

Test Plan

  • Red-first: test_nnls_raw_forward_amplitude.py fails 12/107 on the base bd03e09e (euclid inactive-column + total flux 1.152e-1; source flux k1 1.638e-3, k2 1.220e-3, k3 4.044e-3, k5 6.442e-3 vs 1e-3), both call paths; green 107/107 after the fix
  • test_nnls_mge_convergence.py + test_jax_nnls.py unchanged green (173 with the new file); full pytest test_autoarray 1890 passed; black clean on changed files
  • jit == eager bit-for-bit through reconstruction_positive_only_from; jax.vjp output == plain call; jax.grad finite and non-zero on all 9 fixture systems
  • euclid_strong_lens_modeling_pipeline tests/test_compute_latent_variable.py 19/19 with this branch (no edit, no override); total_source_flux jit-vs-eager +7.47e-5
  • autolens_profiling scripts/lens/solver/accuracy.py (+--posthoc) on 81 systems: pdip_raw == pdip_raw_polish 81/81, 48/48 SLaM converged
  • autolens_workspace_test smoke: imaging jax_likelihood/{lp,mge_group,rectangular_mge,rectangular_mge_rtu,delaunay_mge}.py, interferometer jax_likelihood/delaunay_mge.py, jax_grad/pixelization.py all pass; no pin moved
  • CI green on the PR
Full API Changes (for automation & release notes)

Changed

  • autoarray.util.jax_nnls.solve_nnls_primal_raw_forward(...) / reconstruction_positive_only_from(..., preconditioning="raw") — the forward value is now the polished iterate (≤ RAW_POLISH_MAX_ITER tight warm-started PDIP iterations on the Jacobi system, falling back to the raw iterate when the polish does not converge or leaves the interior). Reconstructions and derived latents on mapper-less positive-only JAX inversions move at ~1e-4 relative on ill-conditioned MGE systems; log-likelihoods at ~1e-11. converged and iterations unchanged in meaning (raw forward solve).
  • autoarray/config/general.yaml nnls_preconditioning_no_mapper comment, Settings.nnls_preconditioning_no_mapper docstring — describe the polished forward value.

Added

  • autoarray.util.jax_nnls.RAW_POLISH_MAX_ITER = 10 (RAW_BACKWARD_POLISH_MAX_ITER kept as an alias)
  • autoarray.util.jax_nnls._raw_forward_polished(Q_pc, q_pc, Q, q, D, solver_tol, max_iter) (private; shared by the custom_vjp primal and forward)
  • test_autoarray/inversion/inversion/files/mge_solver_reference_systems.npz — 8 fix(inversion): JAX PDIP positive-only solve fails to converge on SLaM MGE systems #571 systems + euclid vis_lp with fnnls x_ref (from the autolens_profiling phase-1 corpus at 3ad68af); test_nnls_raw_forward_amplitude.py

Removed

  • none

Migration

  • none required. Pinned mapper-less reconstruction/latent values captured with raw mode before this change move by up to ~5 % on systems where the raw stop left flux on inactive columns (euclid-like); re-pin against the polished value.

Heart RED override (development only)

  • Authorization: live human in the Claude Code CLI session, 2026-09-30: "ok do phase 2" launched this task after the phase-1 hand-back, and at the ship gate the human pushed feature/raw-pdip-forward-polish directly (~16:35 BST) in response to the reported RED reasons, after the auto-mode classifier denied the agent push. Scope: development shipping (push + this pending-release PR); merge only via a separate /prm on all-green required checks; no release.
  • Heart readiness at the 16:25 BST gate (verbatim): RED release validation FAILED (stage integrate); YELLOW workspace validation not passing (0 failed, 1 timeout, cloud#36404726969: autolens_test scripts/multi_dataset/rectangular.py); YELLOW manifest drift: public front-door organ tables (generated) — 1 mismatch(es) vs PyAutoMind/repos.yaml. None involve PyAutoArray. This PR does not repair Heart; Heart remains RED for release.
  • Passed branch gates at 31b1c2d: red-first regression (12/107 red on base, 107/107 after); test_nnls_mge_convergence.py + test_jax_nnls.py unchanged; full pytest test_autoarray 1890 passed; black clean on changed files; euclid pipeline test_compute_latent_variable.py 19/19 with this branch; corpus 81/81 converged incl. 48/48 SLaM; autolens_workspace_test 7 mapper-less/MGE scripts pass with no pin moved; in-session review (supervised, human-approved plan).

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

Jammy2211 and others added 2 commits September 30, 2026 15:47
…#594)

Fixture mge_solver_reference_systems.npz: the 8 #571 SLaM systems and the
euclid vis_lp system with fnnls x_ref, copied from the autolens_profiling
solver corpus at 3ad68af. The new test gates inactive-column, total and
source-column flux at 1e-3, jit == eager, primal == differentiated forward,
and finite non-zero gradients.

Red on bd03e09 (12/107): euclid inactive + total flux, source flux on
k1, k2, k3, k5, on both the library entry and the dispatch path.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
#594)

The raw mode's data-scaled KKT stop also judges complementarity, so a
column fnnls holds at zero kept x ~ tol / z: 11.5 % of the fnnls flux on
inactive columns of the euclid vis_lp system. The #573 polish (<= 10 tight
warm-started PDIP iterations on the Jacobi system) now feeds the forward
value as well as the backward pass, through one _raw_forward_polished
helper shared by the custom_vjp primal and fwd rule, so plain, jitted and
differentiated calls agree bit-for-bit. converged / iterations remain the
raw forward solve's; RAW_POLISH_MAX_ITER added, RAW_BACKWARD_POLISH_MAX_ITER
kept as an alias.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
…association (#594)

The Python 3.12 Actions leg reproduced eager to 5.4e-11 relative on every
fixture system while 3.13 and local runs were bit-exact (PR #595). The
end-to-end check is now a tight allclose (rtol 1e-9, atol 1e-12 * max|x|);
the solver-alone and primal-vs-vjp checks stay bit-exact.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
@Jammy2211

Copy link
Copy Markdown
Collaborator Author

CI 3.12 leg was red on test__raw_forward_jit_matches_eager only (9/9 systems, max 5.4e-11 relative / 9e-12 absolute jit-vs-eager; 3.13 and no-jax green, same jax 0.11.2 on both legs) — runner-dependent XLA reassociation around the solver, not solver behaviour. Pushed a follow-up commit making the end-to-end check a tight allclose (rtol 1e-9, atol 1e-12·max|x|); the solver-alone and primal-vs-vjp checks stay bit-exact. Local: 107/107.

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

Labels

pending-release PR queued for the next release build

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix: return the polished raw-forward PDIP iterate as the forward value (#571/#573 follow-up)

1 participant