perf: scatter the real mapping matrix then cast in TransformerNUFFT.transform_mapping_matrix (#577) - #578
Merged
Conversation
…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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
TransformerNUFFT.transform_mapping_matrixused 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 beforenufft2d2. A real→complex cast is exact, so the output is bit-identical; a new test assertsnp.array_equalagainst the old formula for NumPy and underjax.jit.On the A100 a complex128
.at[].setscatter 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: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
test__nufft__transform_mapping_matrix__real_scatter_matches_complex_scatter_exactly(NumPy + jit, circular mask, exact equality, complex128 output)test_autoarrayfull suite: 1707 passedmanifest 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