Skip to content

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

Description

@Jammy2211

Overview

The released mapper-less positive-only solver (nnls_preconditioning_no_mapper: raw, #572) stops on a data-scaled absolute KKT residual, so columns fnnls holds at zero keep x ≈ tol / z while reporting "converged": the euclid total_source_flux latent traces +5.76 % off under jax.jit and the euclid latent test has been red since 2026-09-25, invisible to logL, objective and KKT. Phase 1 (autolens_profiling #355) measured this on 81 systems — raw PDIP is converged 81/81 yet its source-column flux is >1e-3 off fnnls on 53/81, and it leaves 11.5 % of the reference amplitude on reference-inactive euclid columns — while applying the #573 polish to the forward value gave 0/81 unconverged, worst inactive-column flux 3.3e-4, source-flux misses 1/81 and euclid latent +7.5e-5 at a median cost of +5 iterations. Decision: return the polished iterate (≤ 10 tight warm-started Jacobi-system PDIP iterations, already in the library from #573 for the backward pass) as the forward value, in both the custom_vjp forward and the primal so jit, grad and eager agree.

Plan

  • Factor the raw forward solve + fix: NaN JAX gradient on MGE positive-only solves after #572 #573 polish into one helper in autoarray/util/jax_nnls.py and return the polished iterate (falling back to the unpolished one when the polish is not ok) as the forward value, from both the custom_vjp forward and the primal; backward pass unchanged.
  • Keep converged / iterations as the raw forward solve's; keep the RAW_BACKWARD_POLISH_MAX_ITER name (imported downstream) and add RAW_POLISH_MAX_ITER as the preferred alias.
  • Update the module and function docstrings, the nnls_preconditioning_no_mapper comment in general.yaml, and the settings.py / inversion_util.py docstrings to describe the polished forward value and why.
  • Add a regression fixture (the 8 fix(inversion): JAX PDIP positive-only solve fails to converge on SLaM MGE systems #571 systems + the euclid vis_lp system, with fnnls x_ref) and a jax-only amplitude test gating on inactive-column flux, total flux and source flux at 1e-3 plus jit==eager and finite non-zero grad — run red-first on unfixed d4298445 and record the failing cases here before the fix.
  • Keep test_nnls_mge_convergence.py and test_jax_nnls.py green unchanged; full PyAutoArray pytest.
  • Verify downstream: the euclid latent jit test passes on the fixed library with no config override; autolens_profiling accuracy re-run (81/81 converged, inactive-column worst ≤ 3.3e-4) becomes the ledger row after the library merge; sweep workspace_test mapper-less likelihood pins for moved values (re-pin is a follow-up).
Detailed implementation plan

Affected Repositories

  • PyAutoArray (primary)
  • autolens_profiling (downstream: ledger row in results/notes/linear_solver_accuracy_2026_09.md + campaign page + regenerated README stats, after the library merge)
  • euclid_strong_lens_modeling_pipeline (verification only, no edit)

Branch Survey

Repository Current Branch Dirty?
./PyAutoArray main (d429844) clean — parallel claim with streaming-p1-array-free-dataset, disjoint files (human-approved 2026-09-30)
./autolens_profiling main (3ad68af) 1 untracked dir (dataset/abell_1201/)
./euclid_strong_lens_modeling_pipeline main (26e4385) clean

Suggested branch: feature/raw-pdip-forward-polish

Deviation from the prompt

The prompt's witness thresholds (amp_rel_max ≤ 1e-3, flux_rel ≤ 1e-4) are not achievable by any candidate: amp_rel_max_sig worst 0.219 is shared by every accurate candidate (fnnls reference noise on near-flat directions), and the euclid "source columns" metric is +4.96e-2 for polish only because its three source columns carry negligible reference flux (the end-to-end latent is +7.5e-5). The regression test therefore gates on inactive-column flux and total/source flux at 1e-3, per system, chosen from the phase-1 rows before the library test is written, and drops the amplitude-max criterion. Recorded as a rule weakness in the phase-1 ledger already.

Implementation Steps

  1. autoarray/util/jax_nnls.py
    • Factor the raw forward + polish out of _raw_forward_backward_point (:196-234) into _raw_forward_polished(Q_pc, q_pc, Q, q, D, solver_tol, max_iter) returning (y_out, converged, pdip_iter, (yp, sp, zp), ok, polish_iter) where y_out = where(ok, yp, y) — the same ok rule as today (polish converged, finite, sp>0, zp>0), falling back to the unpolished iterate.
    • _raw_forward_backward_point calls it, runs solve_relaxed_nnls from (yp, sp, zp) as now, and returns (y_out, converged, pdip_iter) as out — the polished forward value.
    • _solve_nnls_raw_forward_with (:237-296): the primal (:276-281) must call the same _raw_forward_polished and return y_out / … exactly as the fwd does, so jit, grad and eager agree bit-for-bit. Backward unchanged.
    • converged stays the raw forward solve's flag; iterations stays the forward count (status already carries polish_iter, ok). Keep the name RAW_BACKWARD_POLISH_MAX_ITER (imported by autolens_profiling _solvers.py:248); add RAW_POLISH_MAX_ITER = RAW_BACKWARD_POLISH_MAX_ITER and prefer the new name internally.
    • Rewrite the module docstring (:30-34) and the _solve_nnls_raw_forward_with docstring (:262-263): the polish is now forward and backward; say why (complementarity under the data-scaled stop leaves tol/z on inactive columns; phase-1 numbers).
  2. autoarray/config/general.yaml:11 comment for nnls_preconditioning_no_mapper and autoarray/settings.py:227-237 docstring: describe the polished forward value. inversion_util.py:362-364 note on converged/iterations if it mentions the unpolished iterate.
  3. Tests
    • New fixture test_autoarray/inversion/inversion/files/mge_solver_reference_systems.npz (≈180 KB): the 8 fix(inversion): JAX PDIP positive-only solve fails to converge on SLaM MGE systems #571 systems + the euclid vis_lp system copied from autolens_profiling results/lens/solver/corpus/{slam_fixture_571,euclid_vis_lp}.npz with x_ref (fnnls) and a meta JSON carrying source_column_index_list, max_abs_q, group, provenance (corpus commit 3ad68af). Add a paragraph to files/README.md.
    • New test_autoarray/inversion/inversion/test_nnls_raw_forward_amplitude.py (jax-only, mirrors test_nnls_mge_convergence.py skips/EPSILON guard). Per system, through the library entry point solve_nnls_primal_raw_forward (built as inversion_util.py:462-481 builds it) and through reconstruction_positive_only_from with preconditioning="raw":
      • converged == 1, iterations < 50, finite;
      • flux_inactive_rel = Σx[x_ref ≤ 1e-6·max x_ref] / Σx_ref ≤ 1e-3 (raw is red on euclid: 0.115);
      • |Σx − Σx_ref| / Σx_ref ≤ 1e-3 (raw red on euclid);
      • |flux_rel_source| ≤ 1e-3 on the 8 fix(inversion): JAX PDIP positive-only solve fails to converge on SLaM MGE systems #571 systems (raw red on k1, k2, k3, k5: 1.6e-3…6.4e-3; polish ≤ 5e-5) — euclid excluded from this one metric with the reason above;
      • jax.jit(f)(…) equals eager f(…) (assert_array_equal) and jax.grad finite/non-zero on every system (extends the fix: NaN JAX gradient on MGE positive-only solves after #572 #573 tests to the new fixture).
      • Docstring records the thresholds, the phase-1 values they were chosen from, and which cases are red on d4298445.
    • Existing test_nnls_mge_convergence.py and test_jax_nnls.py must stay green unchanged (the objective/control/gradient assertions all tighten under polish; test__control__well_conditioned_pdip_unchanged requires Jacobi mode bit-identical — untouched).
    • Red-first: run the new test on unfixed d4298445 and record the failing cases on the issue before the fix commit.
  4. Full PyAutoArray pytest, jax_grad-style guards: none new.

Verification beyond the library

  • euclid witness (no edit): in the canonical lens/euclid_strong_lens_modeling_pipeline (clean, main), run pytest tests/test_compute_latent_variable.py -k traces_under_jax_jit with the worktree PyAutoArray first on PYTHONPATH; must pass at JIT_VS_EAGER_REL = 1e-3 with no config override (expected +7.5e-5). Also run the full file to confirm the NumPy path is untouched.
  • SLaM 48/48 stays converged + corpus re-measure: in the autolens_profiling worktree, run scripts/lens/solver/accuracy.py --posthoc with the fixed PyAutoArray (its library-entry candidate calls solve_nnls_primal_raw_forward, _solvers.py:163-167); expect 81/81 converged, inactive-column worst ≤ 3.3e-4, euclid latent via euclid_latent.py within 1e-3. This is the downstream row.
  • workspace_test pins: grep lens/autolens_workspace_test for mapper-less (MGE / light_lp linear) likelihood pins and run those scripts under the smoke profile with the fixed library; report any value that moved (re-pin is a follow-up ship_workspace, not this PR).
  • Standard: ruff, black --check per repo convention, full pytest test_autoarray.

Key Files

  • autoarray/util/jax_nnls.py — _raw_forward_backward_point, _solve_nnls_raw_forward_with, new _raw_forward_polished, module docstring
  • autoarray/config/general.yaml — nnls_preconditioning_no_mapper comment
  • autoarray/settings.py — preconditioning docstring
  • autoarray/inversion/inversion/inversion_util.py — raw branch / converged note
  • test_autoarray/inversion/inversion/test_nnls_raw_forward_amplitude.py — new amplitude regression test
  • test_autoarray/inversion/inversion/files/mge_solver_reference_systems.npz + files/README.md — new fixture
  • test_autoarray/inversion/inversion/test_nnls_mge_convergence.py, test_autoarray/util/test_jax_nnls.py — must stay green unchanged

Original Prompt

Click to expand starting prompt

Linear-solver programme phase 2: fix the raw-forward PDIP amplitude bias in PyAutoArray

Type: bug
Target: PyAutoArray
Repos:

  • PyAutoArray
  • euclid_strong_lens_modeling_pipeline
  • autolens_profiling
    Difficulty: medium
    Autonomy: supervised
    Priority: high
    Status: formalised
    Filed: 2026-09-30
    Blocked-by: none (the PyAutoArray claim of sparse-data-none-guard cleared 2026-09-30 — complete/2026/09/sparse-data-none-guard.md, PyAutoArray#591 merged)
    Witness: a PyAutoArray regression test loads the phase-1 corpus npz (copied from autolens_profiling results/lens/solver/corpus/ as a fixture beside test_autoarray/inversion/inversion/files/mge_slam_nnls_systems.npz) and asserts, per system, converged AND amplitude agreement with fnnls (amp_rel_max ≤ 1e-3, source flux_rel ≤ 1e-4) — an amplitude assertion, not logL — red on unfixed main; plus euclid tests/test_compute_latent_variable.py::test_latent_euclid_variables_traces_under_jax_jit passes on library main with no config override.
    Review-minutes: 3
    Consequence: glance
    Unattended: needs-slicing

Epic linear-solver-programme, phase 2. Contract: complete/2026/09/linear-solver-accuracy-study.md (phase-1 record with the original programme prompt folded in; the phase-1 verdict blocker is satisfied — autolens_profiling#355 merged 2026-09-30).

Symptom

Raw PDIP (PyAutoArray#572, nnls_preconditioning_no_mapper: raw, the released mapper-less
default) stops on the objective (gap 4.7e-6 vs fnnls) with amplitudes 4.3 % off fnnls, so
euclid total_source_flux traces to 3.511 under jax.jit vs 3.320 eager. The euclid latent
test has been red since 2026-09-25. logL cannot see it (Δχ² ~1e-5); amplitude latents shift ~6 %.

Fix

Candidate: phase 1 verdict (autolens_profiling#355, complete/2026/09/linear-solver-accuracy-study.md): no candidate passes the pre-registered rule — raw PDIP reports converged on 81/81 (KKT ~3e-14) yet leaves 11.5 % of the reference amplitude on euclid columns inactive in the reference (total_source_flux +5.76 %); jacobi diverges on 29/81. Post-hoc, each of these turns the euclid latent test green: forward polish (+7.5e-5), tol 1e-5 (+5.1e-4), jaxnnls tol with cap > 50 (-3e-8); caps ≤ 16 are unsafe. The binding constraint is the stopping test, not the iteration budget. Choose among these on the post-hoc evidence and gate on inactive-column flux, not logL/KKT.

Implement exactly the candidate the phase-1 pre-registered rule admits (converged 100 %, 48/48
SLaM incl.; worst amp_rel_max ≤ 1e-3; worst source flux_rel ≤ 1e-4; KKT ≤ 10× pdip_jacobi's;
lowest median iterations). The three plausible shapes:

  1. A tighter data-scaled tolerance constant in data_scaled_solver_tol
    (autoarray/util/jax_nnls.py).
  2. Return the PyAutoArray#573 polished iterate (≤ 10 warm-started Jacobi iterations) as the
    forward value in _raw_forward_backward_point / _solve_nnls_raw_forward_with. The
    custom_vjp forward must return the same value the primal (solve_nnls_primal_raw_forward)
    returns, or jit/grad and eager diverge — polish both or neither.
  3. A KKT / solution-based stop instead of the objective-gap stop.

Also touch the inversion_util.py raw branch if the candidate needs it, and update the
nnls_preconditioning_no_mapper comment in general.yaml to describe the new behaviour.

Constraint

The 48/48 SLaM source_lp[1] points must stay converged (#571); the existing
mge_slam_nnls_systems.npz test must stay green.

Rejected

Downstream

  • workspace_test mapper-less likelihood pins: re-pin if the forward value moves.
  • autolens_profiling: re-run scripts/lens/solver/accuracy.py post-fix and append the row to
    results/notes/linear_solver_accuracy_2026_09.md (+ campaign page
    wiki/campaigns/linear_solver_accuracy.md).

Ownership and order

PyAutoArray owns the solver; euclid_strong_lens_modeling_pipeline owns the witness test (no edit
expected there); autolens_profiling is the corpus fixture source. Library-first: ship PyAutoArray,
then re-verify the euclid test on library main.

🤖 Generated with Claude Code

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