Skip to content

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

Description

@Jammy2211

Overview

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

  1. 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.
  2. Add a test that the result equals the previous formula exactly, in NumPy and under jax.jit, including a masked-grid case.
  3. Run the PyAutoArray suite plus targeted interferometer smoke.
  4. 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.
  5. Ship the library PR first, then the workspace PR.
Detailed implementation plan

Affected Repositories

  • PyAutoArray (primary)
  • autolens_profiling (A100 witness)

Branch Survey

Repository Current Branch Dirty?
./PyAutoArray main clean
./autolens_profiling main untracked dataset/abell_1201 (unrelated)

Suggested branch: feature/interferometer-transform-real-scatter

Context

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.

  • JAX: :528-535. NumPy: :545-547.
  • On the A100, a complex128 .at[].set scatter 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).
  • 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__.

PyAutoArray, autoarray/operators/transformer.py

  • JAX branch:
    • mm_T = jnp.asarray(mapping_matrix).T
    • source_images = jnp.zeros((n_src, n_y, n_x), dtype=mm_T.dtype).at[...].set(mm_T)
    • flipped = source_images[:, ::-1, :].astype(jnp.complex128)
  • NumPy branch: the same (real zeros, scatter, flip, then .astype(np.complex128)).
  • Keep complex128 as the final dtype, so mixed-precision behaviour is unchanged.
  • Add a one-line comment giving the reason (the GPU complex scatter floor, citing feat: add aa.interp_2d (NumPy + JAX bilinear interpolation) #308).
  • Tests (test_autoarray/operators/test_transformer.py):
    • 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

  1. Exact-equality tests (NumPy and jit) plus the full PyAutoArray suite.
  2. Targeted smoke: interferometer pixelization and MGE scripts in autolens_workspace and autolens_workspace_test.
  3. 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.

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