Skip to content

fix: mixed-precision inversion JAX/NumPy gap on low-pixel-count data #552

Description

@Jammy2211

Overview

autogalaxy_workspace_test/scripts/imaging/jax_likelihood/rectangular.py (a RectangularBilinearAdaptImage + Adapt inversion with use_mixed_precision=True) asserts jax.jit(analysis.fit_from).log_likelihood against a NumPy fit at rtol=2e-2. After autogalaxy_workspace_test#117 coarsened the shared dataset (180×180 @ 0.2" → 100×100 @ 0.3", 716 → 316 masked pixels) the absolute JAX-vs-NumPy gap grew from 5.8 to ~15 nats and is essentially mesh-independent across eight mesh sizes. The prompt asks where the fp32 part of the mixed-precision inversion loses precision as the pixel count falls, whether 2 % relative is the right contract, and to fix the source if there is a genuine loss.

A read-only code survey (2026-09-15) reframed the bug before any code was touched:

  1. The asserted number is log_likelihood = -0.5·(χ² + noise_normalization) (autoarray/fit/fit_util.py:360). noise_normalization is model-independent and identical on both paths, so the whole gap is Δχ² ≈ 30. The regularization term and both log-dets never enter the assertion — they only enter log_evidence, which the script prints but does not assert.
  2. The "float64 NumPy reference" is not float64. mapper_util.mapping_matrix_from (autoarray/inversion/mappers/mapper_util.py:311) honours use_mixed_precision regardless of xp, and the script passes use_mixed_precision=True to the NumPy analysis too. This contradicts the Settings docstring (autoarray/settings.py:38-39: "only the JAX paths honor the flag — the NumPy backend always runs in fp64"). No mixed-precision parity script in either test workspace has a true fp64 oracle today.
  3. The documented fp32 curvature accumulation is inert, but leaves an inconsistency. In curvature_matrix_via_mapping_matrix_from (autoarray/inversion/inversion/inversion_util.py:130-140) the blurred mapping matrix arrives fp64 (the FFT kernel multiply deliberately upcasts, convolver.py:1190-1205), so A * w.astype(float32) promotes to fp64 and A.T@A accumulates in fp64. The only effect is that 1/σ is rounded to fp32 in F, while data_vector uses exact fp64 1/σ² (inversion_imaging_util.py:123). F and D are weighted inconsistently — a small systematic bias, not round-off.
  4. The two sides run different NNLS algorithms: JAX uses the jaxnnls interior-point solver (inversion_util.py:320+, relaxed KKT, ≤50 iterations, target_kappa=1e-11), NumPy uses exact active-set fnnls. They differ even in exact fp64. χ² is not stationary in the reconstruction at the regularised optimum (∂χ²/∂s = 2Hs ≠ 0), so Δχ² is first order in the reconstruction difference. With use_edge_zeroed_pixels: true the 17×17 mesh solves 225 interior pixels against 316 data pixels — a near-square, degenerate active set where the two solvers disagree most.

Control: autolens_workspace_test/scripts/imaging/jax_likelihood/rectangular.py runs the identical construction (positive-only + mixed precision on both sides) on ~430 pixels at 0.3" and holds rtol=1e-4. The gap is regime-specific, not a general mixed-precision property.

Working hypothesis (to be tested, not assumed): the dominant term is the solver-algorithm difference amplified by conditioning as the system approaches square; the genuine fp32 contributions (complex64 forward FFT noise floor into should-be-zero entries of the blurred matrix; the F/D weighting inconsistency) are second-order and may flip active-set membership. Independently of the root cause, the tolerance contract is the wrong shape: rtol on a quantity dominated by an additive constant ∝ N shrinks the budget linearly with pixel count while the error (a Δχ²) does not.

Reproduction hazard: the on-disk autogalaxy_workspace_test/dataset/imaging/jax_test/ on the dev machine is stale (180×180, 21×21 PSF, pre-#117) and should_simulate does not detect the resolution change. Every measurement must start by moving that folder aside so the simulator regenerates it. The reported 15 nats may itself be a stale-hybrid measurement. There is no test-results/ directory in the repo holding the numbers from the prompt.

Prior art (Memory): PyAutoArray#369 (fp32 near F / PDIP diverges, cond·eps32 ≈ 600), #391 (log_det_method opt-in slogdet), regularization-jax-gradient-gaps record. None covers this defect. There are no unit tests for use_mixed_precision in test_autoarray (0 hits).

Plan

  • Regenerate the coarsened dataset from a clean slate and measure the gap in a matrix {JAX, NumPy} × {mixed, fp64} × {positive-only on, off} to separate the solver effect from the fp32 effect. Record χ², Δs, active-set overlap and cond(F+H) per cell; repeat at the old 716-pixel geometry and the autolens ~430-pixel control.
  • Fix the two source defects the survey established from code alone, whatever the matrix says: make the NumPy backend a true fp64 oracle (honour the Settings docstring), and remove the inert fp32 curvature branch so F and D are weighted consistently.
  • If the matrix shows a JAX-only fp32 loss above the solver floor, fix that path too (candidate: the complex64 forward FFT of the mapping cube); otherwise document the solver floor as the residual.
  • Add the first unit tests for use_mixed_precision in test_autoarray (dtype contract per backend, F/D consistency, a small deterministic JAX-vs-NumPy parity bound).
  • Ship the PyAutoArray PR first (library-first gate). Then change the smoke script's contract from rtol=2e-2 to an absolute nats bound on Δlog_likelihood sized from the measured post-fix gap, once autogalaxy_workspace_test is released by jax-runtime-and-parity.
  • File follow-ups via /intake rather than widen this task: the other 14 mixed-precision parity scripts' tolerances, and rectangular_mge.py still using a 28×28 mesh on the coarsened data.
Detailed implementation plan

Work Classification

Both (library + workspace). Library first.

Affected Repositories

  • PyAutoArray (primary)
  • autogalaxy_workspace_test (one script, tolerance contract only)

Branch Survey (2026-09-15)

Repository Current Branch Dirty?
./PyAutoArray main clean
./autogalaxy_workspace_test main clean

Recent PyAutoArray branches: feature/delaunay-area-magnification-audit, claude/autonerves-floor-regime-stamp.

Worktree conflict guard: worktree_check_conflict mixed-precision-inversion-gap PyAutoArray → clear. autogalaxy_workspace_test is claimed by jax-runtime-and-parity (autolens_workspace_test#317, issued 2026-09-15, edits only smoke_tests.txt, other session). Routing: register with PyAutoArray only; add autogalaxy_workspace_test via worktree_add_repo when that claim clears. The workspace leg is one tolerance edit and depends on the library result anyway.

Heart at planning time: STALE ("release validation incomplete: no rehearsal for current source") — gates ship_*, re-read at ship time.

Suggested branch: feature/mixed-precision-inversion-gap
Worktree root: ~/Code/PyAutoLabs-wt/mixed-precision-inversion-gap/ (created by /start_library)

Implementation Steps

Phase A — diagnostics (scratch script, not committed)

  1. Move autogalaxy_workspace_test/dataset/imaging/jax_test aside (never delete) so the simulator regenerates the 100×100 @ 0.3" dataset on first load.
  2. Rebuild the script's objects (dataset, circular mask r=3.0", adapt image = data, RectangularBilinearAdaptImage(17,17) + Adapt, prior-median instance). For each cell {jax, numpy} × {mixed, fp64} × {positive_only True, False} record log_likelihood, chi_squared, regularization_term, both log-dets, ‖s_jax − s_np‖∞, positive-pixel counts and overlap, cond(F+H), max|F_mixed − F_fp64|, max|D_mixed − D_fp64|.
  3. Repeat the 2×2 at the old geometry (180×180 @ 0.2", 28×28 mesh) and the autolens control (~430 px).
  4. Post the table as an issue comment and state which of {solver difference, complex64 FFT floor, 1/σ inconsistency, fp32 mapping matrix} carries the Δχ² at 316 pixels. The positive_only=False column is decisive: same xp.linalg.solve both sides collapses the gap to pure precision effects if the hypothesis holds.

Phase B — source fixes + unit tests (PyAutoArray)

  • B1. NumPy backend is fp64 regardless of the flag. In autoarray/inversion/mappers/mapper_util.py mapping_matrix_from (~line 311) gate out_dtype = float32 on xp.__name__.startswith("jax"); same gate in Convolver.mapping_matrix_native_from (autoarray/operators/convolver.py:865,903). Check autoarray/inversion/mappers/abstract.py:291,352 pass-throughs. Rationale: mixed precision exists for GPU throughput; the NumPy path is the parity oracle for 15 smoke scripts and the docstring already promises fp64. Risk: shared fp32 rounding no longer cancels, so parity gaps in mixed-precision scripts may move either way — Phase A's numpy+fp64 column measures this first.
  • B2. Consistent noise weighting of F and D. In autoarray/inversion/inversion/inversion_util.py:130-140 remove the inert compute_dtype branch; weight A by fp64 1/σ on the JAX path as the NumPy branch does (or collapse both into one xp-agnostic expression). Update the Settings docstring bullet claiming fp32 accumulation of A.T@A — it never happens on the imaging mapping path.
  • B3 (conditional, only if Phase A shows a JAX-only fp32 term above the solver floor). Upcast the native cube before rfft2 in convolved_mapping_matrix_from (convolver.py:1214-1220), or drop the complex64 forward FFT for the mapping-matrix path; measure the cost. Decision recorded here with the Phase A numbers.
  • Not touched: the NNLS solver choice (jaxnnls vs fnnls). A solver-algorithm difference is not a precision bug; if it dominates, it is the documented floor the Phase D contract absorbs.
  • Unit tests (new):
    • test_autoarray/inversion/pixelization/mappers/test_mapper_util.py: mapping_matrix_from(use_mixed_precision=True) → float64 under xp=np, float32 under xp=jnp; values equal within fp32 eps.
    • test_autoarray/inversion/inversion/test_inversion_util.py: curvature_matrix_via_mapping_matrix_from with Settings(use_mixed_precision=True) under jnp equals the NumPy fp64 result to ~1e-12 relative.
    • test_autoarray/operators/test_convolver.py: convolved_mapping_matrix_from mixed vs fp64 within a stated bound; NumPy path unchanged by the flag.
    • One small synthetic parity test (rectangular mapper, ~50 data pixels, use_positive_only_solver=False): JAX mixed vs NumPy fp64 |Δχ²| below a bound derived from fp32 eps × ‖A‖.
  • Run source activate.sh && pytest test_autoarray/inversion test_autoarray/operators test_autoarray/fit -q, then full test_autoarray (baseline 926+ passed per feat: opt-in gradient-safe log-det via Settings (default unchanged) #391). Format with black.

Phase C — ship library (ship_library). ## API Changes: NumPy backend now ignores use_mixed_precision (was: fp32 mapping matrix); curvature_matrix_via_mapping_matrix_from weights in fp64 under mixed precision (was: fp32-rounded 1/σ). Workspace impact: every use_mixed_precision=True parity script (8 in autogalaxy_workspace_test, 7 in autolens_workspace_test) gets a true fp64 reference.

Phase D — workspace contract (after autogalaxy_workspace_test is released by jax-runtime-and-parity). In scripts/imaging/jax_likelihood/rectangular.py replace rtol=2e-2 with an absolute bound on |logL_jit − logL_np| in nats (atol=X, rtol=0), X = measured post-fix gap × margin, with a short comment citing this issue, the χ² sampling floor √(2N) for scale, and why relative tolerance on a noise-normalization-dominated scalar is the wrong contract. Keep the 17×17 mesh. Ship via ship_workspace behind the library-first gate.

Key Files

  • autoarray/settings.py:33-82 — use_mixed_precision docstring (contract to make true)
  • autoarray/inversion/mappers/mapper_util.py:311 — fp32 mapping-matrix write, xp-agnostic today
  • autoarray/inversion/inversion/inversion_util.py:130-140 — curvature compute dtype branch (inert; rounds 1/σ)
  • autoarray/inversion/inversion/inversion_util.py:320+ — jaxnnls vs fnnls dispatch
  • autoarray/operators/convolver.py:865,903,1190-1230 — native cube dtype, complex64 forward FFT, complex128 upcast
  • autoarray/inversion/inversion/imaging/inversion_imaging_util.py:123 — fp64 1/σ² data vector
  • autoarray/fit/fit_util.py:360,398-437 — log_likelihood vs log_evidence assembly
  • autogalaxy_workspace_test/scripts/imaging/jax_likelihood/rectangular.py:107-172 — settings, reference construction, rtol=2e-2 assertion
  • autogalaxy_workspace_test/scripts/imaging/jax_likelihood/simulator.py — the 100×100 @ 0.3" dataset generator

Follow-ups to file via /intake (not in scope here)

  • Same absolute-nats contract for the other 14 mixed-precision parity scripts across both test workspaces, once this task sets the scale.
  • autogalaxy_workspace_test/scripts/imaging/jax_likelihood/rectangular_mge.py still uses a 28×28 mesh on 0.3" data — the under-determined regime Feature/positions analysis simplify #117 fixed in rectangular.py.
  • use_mixed_precision has no YAML key / config fallback unlike its Settings siblings.

Verification

  1. Phase A table reproduces or corrects the 15-nat figure on a freshly simulated dataset, with the positive_only=False column isolating solver vs precision.
  2. Full test_autoarray green; new tests fail on main for B1/B2 and pass on the branch.
  3. Phase A matrix re-run on the branch: NumPy columns identical mixed vs fp64 (B1); max|F_mixed − F_fp64| at fp64 round-off (B2).
  4. rectangular.py under profile_smoke.yaml (env from autohands.env_config.build_env_for_script, workspace CWD, cleared output/, regenerated dataset) passes the new absolute bound; the autolens control still passes at 1e-4.
  5. Library PR CI green; workspace PR workspace-smoke legs green after the library merge.

Execution model

Fable session plans and judges; every phase is delegated to an Opus subagent with a progress heartbeat. Subagents never edit code to make a test pass. Judgement stays in-session: Phase B fix set from Phase A, PR body, tolerance value X, merge (/prm, human).

Original Prompt

Click to expand starting prompt

Mixed-precision inversion: the JAX-vs-NumPy log-likelihood gap grows on low-pixel-count data

Type: bug
Target: PyAutoArray
Repos:

  • PyAutoArray
  • autogalaxy_workspace_test
    Difficulty: medium
    Autonomy: supervised
    Priority: medium
    Status: formalised
    Filed: 2026-09-06

Found during the phase-5 rebuild of the ci-timing-fast-tests epic
(autogalaxy_workspace_test#117). scripts/imaging/jax_likelihood/rectangular.py —
an adapt-image RectangularBilinearAdaptImage inversion run with
use_mixed_precision=True — asserts jax.jit(analysis.fit_from) against the
float64 NumPy path at rtol=2e-2. On the coarsened shared imaging dataset
(100x100 @ 0.3", 316 masked pixels) the absolute gap between the two paths is
~15 nats, against 5.8 nats on the previous dataset (180x180 @ 0.2", 716 masked
pixels; -3150.83 vs -3145.04 there). Measured at eight mesh sizes on the new
data the gap is essentially mesh-independent (14.99 / 26.73 / 21.32 / 15.47 /
17.64 / 19.06 / 16.80 / 14.98 nats at meshes 12, 14, 16, 17, 19, 20, 24, 28), so
it is a property of the mixed-precision inversion on lower-pixel-count data, not
of the mesh. The rebuild kept the script green by sizing the mesh to the data
(17x17, source pixels <= image pixels — the under-determined 28x28 inversion was
the actual failure) and did not touch the tolerance.

Ask: characterise where the fp32 part of the mixed-precision inversion loses the
precision as the pixel count falls (the data-vector / curvature-matrix products,
the regularization term, the log-det), decide whether the smoke script's 2 %
relative tolerance is the right contract or whether an absolute-nats bound is,
and fix the source if a genuine precision loss is found. Reproducer: the
jax_test imaging dataset of autogalaxy_workspace_test after #117 plus that
script; test-results/ numbers above.

Original observation (verbatim, from the #117 rebuild report)

the JAX-vs-NumPy absolute discrepancy is ~15 nats on the new data against
5.8 nats on main's (-3150.83 vs -3145.04 there), and it is essentially
independent of mesh size … So the mesh fix restores the relative margin
(1.23% against the script's 2%) by increasing |log_L|, not by shrinking the
gap. The growth of that mixed-precision discrepancy on lower-pixel-count data
is a finding for the source repos, not something this change fixes.

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