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
TransformerNUFFT.transform_mapping_matrix casts the mapping matrix to complex128 before scattering it slim→native. On the A100 a complex128 scatter of 20 columns takes 3422 ms; scattering float64 and then casting takes 0.57 ms. That scatter adds a fixed ~0.85 s to every dense interferometer transform on the GPU, and it is why sma is 10x slower on the A100 than on CPU (autolens_profiling#308, lever 3). The fix is to scatter in the real dtype and cast afterwards, which gives bit-identical values.
Plan
In both branches of transform_mapping_matrix, scatter the mapping matrix in its own real dtype into a real native stack, flip it, then cast to complex128 just before nufft2d2. The values are bit-identical, since a real→complex cast is exact.
Add a test that the result equals the previous formula exactly, in NumPy and under jax.jit, including a masked-grid case.
Run the PyAutoArray suite plus targeted interferometer smoke.
Workspace witness (autolens_profiling): re-run the A100 sma MGE breakdown leg on the branch and commit the JSON plus a note line. Step 3 should drop from 852.7 ms to under 50 ms with logL unchanged to 1e-9. Also check CPU sma for no regression.
TransformerNUFFT.transform_mapping_matrix (autoarray/operators/transformer.py:508-555) casts the mapping matrix to complex128 before scattering it slim→native into a (n_src, N_y, N_x) stack.
This is a fixed, N_vis-independent cost of about 0.85 s on every dense interferometer transform. It is why sma on the A100 (852.7 ms at step 3) is 10x slower than on the laptop CPU (82 ms).
It hits every dense interferometer inversion that does not use the sparse operator: pixelizations as well as MGE.
The audit found no other slim→native complex scatter:
visibilities_from (:420) scatters in image.dtype (real) and casts later (:363).
:485/:498 and transformer_util.py:129 are accumulators, not scatters.
Detailed plan
Branch feature/interferometer-transform-real-scatter in PyAutoArray and autolens_profiling. Worktree ~/Code/PyAutoLabs-wt/interferometer-transform-real-scatter/. Runs use a private env file pinning PYTHONPATH (not the shared activate.sh), and print __file__.
transform_mapping_matrix output equals the old formula exactly (array_equal, computed inline in the test) for NumPy and for jax.jit, on a masked real-space mask.
Existing transformer and inversion tests stay green.
autolens_profiling, witness only (no code change expected)
RAL: reuse the private branch-libs pattern (/mnt/ral/jnightin/PyAuto_branch/<task>/; the shared mirror is never touched). Adapt the existing submit_breakdown_interferometer_mge_a100_sma_fp64 into a _real_scatter variant, or pass an output suffix.
Commit the new sma A100 JSON/PNG under a distinct name. Add a dated "Lever 3: real scatter (2026-09-26)" subsection to results/notes/interferometer_mge_breakdown_2026_09.md with a before/after row for step 3 and the total.
Include CPU sma before/after (no regression), regenerate the README (--check), and run lint and pytest scripts/misc/test/.
Optional, if cheap in the same job: the alma chunked-arm step 3 (which ran 915 ms).
Verification
Exact-equality tests (NumPy and jit) plus the full PyAutoArray suite.
Targeted smoke: interferometer pixelization and MGE scripts in autolens_workspace and autolens_workspace_test.
A100 sma: step 3 under 50 ms, log L identical to 1e-9. The JSON records the source revisions.
Risks
Mixed precision: if the mapping matrix is float32 under mp, the scatter becomes float32 and the cast becomes complex128. Values are unchanged, because old and new both cast float32 to complex128.
The measured win could be smaller inside the full jit than in the isolated probe (fusion). The A100 witness row settles it.
Original Prompt
Click to expand starting prompt
Interferometer likelihood campaign: scatter the real mapping matrix then cast in transform_mapping_matrix (GPU complex128 scatter floor)
Type: bug
Target: PyAutoArray
Repos:
PyAutoArray
autolens_profiling
Themes:
interferometer
mge
jax-gpu
Difficulty: easy
Autonomy: supervised
Priority: high
Status: draft
Consequence: glance
Witness: on the RAL A100, step 3 ("Transformed mapping matrix (NUFFT)") of results/breakdown/interferometer/sma/mge_hpc_a100_fp64.json drops from 852.7 ms to under 50 ms, with the log_likelihood unchanged to 1e-9 nats.
Review-minutes: 5
Epic: interferometer-likelihood-campaign
Source: autolens_profiling/results/notes/interferometer_mge_breakdown_2026_09.md, lever 3
(autolens_profiling#308); probe numbers in results/breakdown/interferometer/mge_a100_diagnostics_probes_2026_09.json.
Why
TransformerNUFFT.transform_mapping_matrix casts the mapping matrix to complex128 and
scatters it slim->native into a (n_src, N_y, N_x) complex128 stack
(autoarray/operators/transformer.py:525-527) before the NUFFT. On the A100 a complex128
scatter of 20 columns takes 3422 ms (one column 224 ms), while scattering float64 and casting
takes 0.57 ms (~6000x); unique_indices / indices_are_sorted do not help. A bare nufft2d2 of 20 columns on 800x800 is 6.0 ms. The scatter is N_vis-independent: A100 sma
step 3 is 852.7 ms (vs 82.0 ms on the laptop CPU) and alma 915 ms. It is the whole reason
the A100 is ~10x slower than the CPU at sma.
What
In the JAX branch: scatter the real mapping matrix into a float64 native stack, then .astype(complex128) (or pass real input if nufftax accepts it).
Audit the other slim->native complex scatters in transformer.py (e.g. :417) and the
DFT transformer for the same pattern.
Re-run the sma A100 breakdown leg and the imaging/interferometer pixelized dense cells
that share the method.
Watch
CPU is unaffected (82 ms is NUFFT-bound); verify no CPU regression.
Overview
TransformerNUFFT.transform_mapping_matrixcasts the mapping matrix to complex128 before scattering it slim→native. On the A100 a complex128 scatter of 20 columns takes 3422 ms; scattering float64 and then casting takes 0.57 ms. That scatter adds a fixed ~0.85 s to every dense interferometer transform on the GPU, and it is why sma is 10x slower on the A100 than on CPU (autolens_profiling#308, lever 3). The fix is to scatter in the real dtype and cast afterwards, which gives bit-identical values.Plan
transform_mapping_matrix, scatter the mapping matrix in its own real dtype into a real native stack, flip it, then cast to complex128 just beforenufft2d2. The values are bit-identical, since a real→complex cast is exact.jax.jit, including a masked-grid case.Detailed implementation plan
Affected Repositories
Branch Survey
Suggested branch:
feature/interferometer-transform-real-scatterContext
TransformerNUFFT.transform_mapping_matrix(autoarray/operators/transformer.py:508-555) casts the mapping matrix to complex128 before scattering it slim→native into a(n_src, N_y, N_x)stack.:528-535. NumPy::545-547..at[].setscatter of 20 columns takes 3422 ms. The same scatter in float64 followed by a cast takes 0.57 ms (feat: add aa.interp_2d (NumPy + JAX bilinear interpolation) #308 probe JSON).The audit found no other slim→native complex scatter:
visibilities_from(:420) scatters inimage.dtype(real) and casts later (:363).:485/:498andtransformer_util.py:129are accumulators, not scatters.Detailed plan
Branch
feature/interferometer-transform-real-scatterin PyAutoArray and autolens_profiling. Worktree~/Code/PyAutoLabs-wt/interferometer-transform-real-scatter/. Runs use a private env file pinning PYTHONPATH (not the sharedactivate.sh), and print__file__.PyAutoArray,
autoarray/operators/transformer.pymm_T = jnp.asarray(mapping_matrix).Tsource_images = jnp.zeros((n_src, n_y, n_x), dtype=mm_T.dtype).at[...].set(mm_T)flipped = source_images[:, ::-1, :].astype(jnp.complex128).astype(np.complex128)).test_autoarray/operators/test_transformer.py):transform_mapping_matrixoutput equals the old formula exactly (array_equal, computed inline in the test) for NumPy and forjax.jit, on a masked real-space mask.autolens_profiling, witness only (no code change expected)
/mnt/ral/jnightin/PyAuto_branch/<task>/; the shared mirror is never touched). Adapt the existingsubmit_breakdown_interferometer_mge_a100_sma_fp64into a_real_scattervariant, or pass an output suffix.results/notes/interferometer_mge_breakdown_2026_09.mdwith a before/after row for step 3 and the total.--check), and run lint andpytest scripts/misc/test/.Verification
Risks
Original Prompt
Click to expand starting prompt
Interferometer likelihood campaign: scatter the real mapping matrix then cast in transform_mapping_matrix (GPU complex128 scatter floor)
Type: bug
Target: PyAutoArray
Repos:
Themes:
Difficulty: easy
Autonomy: supervised
Priority: high
Status: draft
Consequence: glance
Witness: on the RAL A100, step 3 ("Transformed mapping matrix (NUFFT)") of
results/breakdown/interferometer/sma/mge_hpc_a100_fp64.jsondrops from 852.7 ms to under 50 ms, with the log_likelihood unchanged to 1e-9 nats.Review-minutes: 5
Epic: interferometer-likelihood-campaign
Source:
autolens_profiling/results/notes/interferometer_mge_breakdown_2026_09.md, lever 3(autolens_profiling#308); probe numbers in
results/breakdown/interferometer/mge_a100_diagnostics_probes_2026_09.json.Why
TransformerNUFFT.transform_mapping_matrixcasts the mapping matrix to complex128 andscatters it slim->native into a
(n_src, N_y, N_x)complex128 stack(
autoarray/operators/transformer.py:525-527) before the NUFFT. On the A100 a complex128scatter of 20 columns takes 3422 ms (one column 224 ms), while scattering float64 and casting
takes 0.57 ms (~6000x);
unique_indices/indices_are_sorteddo not help. A barenufft2d2of 20 columns on 800x800 is 6.0 ms. The scatter is N_vis-independent: A100 smastep 3 is 852.7 ms (vs 82.0 ms on the laptop CPU) and alma 915 ms. It is the whole reason
the A100 is ~10x slower than the CPU at sma.
What
.astype(complex128)(or pass real input if nufftax accepts it).transformer.py(e.g.:417) and theDFT transformer for the same pattern.
that share the method.
Watch