You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
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:
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.
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.
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.
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.
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)
Move autogalaxy_workspace_test/dataset/imaging/jax_test aside (never delete) so the simulator regenerates the 100×100 @ 0.3" dataset on first load.
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|.
Repeat the 2×2 at the old geometry (180×180 @ 0.2", 28×28 mesh) and the autolens control (~430 px).
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.pymapping_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‖.
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)
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
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.
Full test_autoarray green; new tests fail on main for B1/B2 and pass on the branch.
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).
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.
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.
Overview
autogalaxy_workspace_test/scripts/imaging/jax_likelihood/rectangular.py(aRectangularBilinearAdaptImage+Adaptinversion withuse_mixed_precision=True) assertsjax.jit(analysis.fit_from).log_likelihoodagainst a NumPy fit atrtol=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:
log_likelihood = -0.5·(χ² + noise_normalization)(autoarray/fit/fit_util.py:360).noise_normalizationis 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 enterlog_evidence, which the script prints but does not assert.mapper_util.mapping_matrix_from(autoarray/inversion/mappers/mapper_util.py:311) honoursuse_mixed_precisionregardless ofxp, and the script passesuse_mixed_precision=Trueto the NumPy analysis too. This contradicts theSettingsdocstring (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.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), soA * w.astype(float32)promotes to fp64 andA.T@Aaccumulates in fp64. The only effect is that1/σis rounded to fp32 in F, whiledata_vectoruses exact fp641/σ²(inversion_imaging_util.py:123). F and D are weighted inconsistently — a small systematic bias, not round-off.inversion_util.py:320+, relaxed KKT, ≤50 iterations,target_kappa=1e-11), NumPy uses exact active-setfnnls. 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. Withuse_edge_zeroed_pixels: truethe 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.pyruns the identical construction (positive-only + mixed precision on both sides) on ~430 pixels at 0.3" and holdsrtol=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:
rtolon 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) andshould_simulatedoes 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 notest-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_methodopt-in slogdet),regularization-jax-gradient-gapsrecord. None covers this defect. There are no unit tests foruse_mixed_precisionintest_autoarray(0 hits).Plan
Settingsdocstring), and remove the inert fp32 curvature branch so F and D are weighted consistently.use_mixed_precisionintest_autoarray(dtype contract per backend, F/D consistency, a small deterministic JAX-vs-NumPy parity bound).rtol=2e-2to an absolute nats bound on Δlog_likelihood sized from the measured post-fix gap, onceautogalaxy_workspace_testis released byjax-runtime-and-parity./intakerather than widen this task: the other 14 mixed-precision parity scripts' tolerances, andrectangular_mge.pystill using a 28×28 mesh on the coarsened data.Detailed implementation plan
Work Classification
Both (library + workspace). Library first.
Affected Repositories
Branch Survey (2026-09-15)
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_testis claimed byjax-runtime-and-parity(autolens_workspace_test#317, issued 2026-09-15, edits onlysmoke_tests.txt, other session). Routing: register with PyAutoArray only; addautogalaxy_workspace_testviaworktree_add_repowhen 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-gapWorktree root:
~/Code/PyAutoLabs-wt/mixed-precision-inversion-gap/(created by/start_library)Implementation Steps
Phase A — diagnostics (scratch script, not committed)
autogalaxy_workspace_test/dataset/imaging/jax_testaside (never delete) so the simulator regenerates the 100×100 @ 0.3" dataset on first load.RectangularBilinearAdaptImage(17,17)+Adapt, prior-median instance). For each cell {jax, numpy} × {mixed, fp64} × {positive_only True, False} recordlog_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|.positive_only=Falsecolumn is decisive: samexp.linalg.solveboth sides collapses the gap to pure precision effects if the hypothesis holds.Phase B — source fixes + unit tests (PyAutoArray)
autoarray/inversion/mappers/mapper_util.pymapping_matrix_from(~line 311) gateout_dtype = float32onxp.__name__.startswith("jax"); same gate inConvolver.mapping_matrix_native_from(autoarray/operators/convolver.py:865,903). Checkautoarray/inversion/mappers/abstract.py:291,352pass-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'snumpy+fp64column measures this first.autoarray/inversion/inversion/inversion_util.py:130-140remove the inertcompute_dtypebranch; weightAby fp641/σon the JAX path as the NumPy branch does (or collapse both into onexp-agnostic expression). Update theSettingsdocstring bullet claiming fp32 accumulation ofA.T@A— it never happens on the imaging mapping path.rfft2inconvolved_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.test_autoarray/inversion/pixelization/mappers/test_mapper_util.py:mapping_matrix_from(use_mixed_precision=True)→ float64 underxp=np, float32 underxp=jnp; values equal within fp32 eps.test_autoarray/inversion/inversion/test_inversion_util.py:curvature_matrix_via_mapping_matrix_fromwithSettings(use_mixed_precision=True)under jnp equals the NumPy fp64 result to ~1e-12 relative.test_autoarray/operators/test_convolver.py:convolved_mapping_matrix_frommixed vs fp64 within a stated bound; NumPy path unchanged by the flag.use_positive_only_solver=False): JAX mixed vs NumPy fp64|Δχ²|below a bound derived from fp32 eps × ‖A‖.source activate.sh && pytest test_autoarray/inversion test_autoarray/operators test_autoarray/fit -q, then fulltest_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 ignoresuse_mixed_precision(was: fp32 mapping matrix);curvature_matrix_via_mapping_matrix_fromweights in fp64 under mixed precision (was: fp32-rounded1/σ). Workspace impact: everyuse_mixed_precision=Trueparity script (8 in autogalaxy_workspace_test, 7 in autolens_workspace_test) gets a true fp64 reference.Phase D — workspace contract (after
autogalaxy_workspace_testis released byjax-runtime-and-parity). Inscripts/imaging/jax_likelihood/rectangular.pyreplacertol=2e-2with 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 viaship_workspacebehind the library-first gate.Key Files
autoarray/settings.py:33-82—use_mixed_precisiondocstring (contract to make true)autoarray/inversion/mappers/mapper_util.py:311— fp32 mapping-matrix write, xp-agnostic todayautoarray/inversion/inversion/inversion_util.py:130-140— curvature compute dtype branch (inert; rounds1/σ)autoarray/inversion/inversion/inversion_util.py:320+— jaxnnls vs fnnls dispatchautoarray/operators/convolver.py:865,903,1190-1230— native cube dtype, complex64 forward FFT, complex128 upcastautoarray/inversion/inversion/imaging/inversion_imaging_util.py:123— fp641/σ²data vectorautoarray/fit/fit_util.py:360,398-437— log_likelihood vs log_evidence assemblyautogalaxy_workspace_test/scripts/imaging/jax_likelihood/rectangular.py:107-172— settings, reference construction,rtol=2e-2assertionautogalaxy_workspace_test/scripts/imaging/jax_likelihood/simulator.py— the 100×100 @ 0.3" dataset generatorFollow-ups to file via
/intake(not in scope here)autogalaxy_workspace_test/scripts/imaging/jax_likelihood/rectangular_mge.pystill uses a 28×28 mesh on 0.3" data — the under-determined regime Feature/positions analysis simplify #117 fixed inrectangular.py.use_mixed_precisionhas no YAML key / config fallback unlike itsSettingssiblings.Verification
positive_only=Falsecolumn isolating solver vs precision.test_autoarraygreen; new tests fail onmainfor B1/B2 and pass on the branch.max|F_mixed − F_fp64|at fp64 round-off (B2).rectangular.pyunderprofile_smoke.yaml(env fromautohands.env_config.build_env_for_script, workspace CWD, clearedoutput/, regenerated dataset) passes the new absolute bound; the autolens control still passes at 1e-4.workspace-smokelegs 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:
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
RectangularBilinearAdaptImageinversion run withuse_mixed_precision=True— assertsjax.jit(analysis.fit_from)against thefloat64 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.83vs-3145.04there). Measured at eight mesh sizes on the newdata 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_testimaging dataset of autogalaxy_workspace_test after #117 plus thatscript;
test-results/numbers above.Original observation (verbatim, from the #117 rebuild report)