Skip to content

perf: scatter the real mapping matrix then cast in TransformerNUFFT.transform_mapping_matrix (#577) - #578

Merged
Jammy2211 merged 1 commit into
mainfrom
feature/interferometer-transform-real-scatter
Sep 26, 2026
Merged

Jammy2211 merged 1 commit into
mainfrom
feature/interferometer-transform-real-scatter

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Summary

TransformerNUFFT.transform_mapping_matrix used to cast the mapping matrix to complex128 before scattering it slim→native into the (n_src, N_y, N_x) stack. It now scatters in the matrix's real dtype, flips, and casts to complex128 just before nufft2d2. A real→complex cast is exact, so the output is bit-identical; a new test asserts np.array_equal against the old formula for NumPy and under jax.jit.

On the A100 a complex128 .at[].set scatter is extremely slow; the #308 probe measured 3422 ms for 20 columns, against 0.57 ms for a float64 scatter. RAL job 356364 ran main and this branch on the same node, sma MGE, dense library path:

per call main this PR
step 3 (transformed mapping matrix) 851.5 ms 0.51 ms
full pipeline, single jit 868.5 ms 3.39 ms
vmap batch 4 225.2 ms/call 1.34 ms/call
log L -3153.942384509246 identical

This removes a fixed ~0.85 s from every dense interferometer transform on GPU, whether MGE or pixelization. sma on the A100 goes from 10x slower than the laptop CPU to about 20x faster. CPU is unchanged: ABBA reruns show no regression at sma or alma under load.

Closes #577.

API Changes

None — internal changes only.

Test Plan

  • New test__nufft__transform_mapping_matrix__real_scatter_matches_complex_scatter_exactly (NumPy + jit, circular mask, exact equality, complex128 output)
  • test_autoarray full suite: 1707 passed
  • Targeted interferometer smoke: 21/21 pass (autolens_workspace 11, autolens_workspace_test 5, autogalaxy_workspace 5)
  • A100 witness: RAL job 356364; results in the autolens_profiling PR
  • Heart at ship: YELLOW, acknowledged by the human on 2026-09-26. Reasons: manifest drift: hub organism blurb (organs present) — 7 mismatch(es) vs PyAutoMind/repos.yaml; manifest drift: organism-map blocks (generated) — 1 mismatch(es) vs PyAutoMind/repos.yaml; manifest drift: workspace checkouts (manifest ↔ disk) — 1 mismatch(es) vs PyAutoMind/repos.yaml; release validation stale: source moved since rehearsal (PyAutoFit, PyAutoArray, PyAutoGalaxy, PyAutoLens).

🤖 Generated with Claude Code

…transform_mapping_matrix

TransformerNUFFT.transform_mapping_matrix cast the mapping matrix to complex128
before scattering it slim->native into the (n_src, N_y, N_x) stack. On the A100 a
complex128 .at[].set scatter is a fixed ~0.85 s (autolens_profiling#308, step 3 of
the sma MGE breakdown), while the same scatter in the real dtype followed by a cast
is sub-millisecond. Both the JAX and NumPy branches now scatter in the mapping
matrix's own dtype, flip, and cast to complex128 just before nufft2d2. A real to
complex cast is exact, so the output is bit-identical; a new test asserts
np.array_equal against the previous formula for NumPy and under jax.jit on a
circular (masked) real-space mask.

Refs PyAutoArray#577.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@Jammy2211 Jammy2211 added the pending-release PR queued for the next release build label Sep 26, 2026
@Jammy2211
Jammy2211 merged commit 14d6336 into main Sep 26, 2026
3 checks passed
@Jammy2211
Jammy2211 deleted the feature/interferometer-transform-real-scatter branch September 26, 2026 18:26
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.

perf: scatter the real mapping matrix then cast in TransformerNUFFT.transform_mapping_matrix

1 participant