From f9d3bc671d3c36dc41bf9f1f22b61e1e61e84a07 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Sat, 26 Sep 2026 16:49:40 +0100 Subject: [PATCH] perf(interferometer): scatter the real mapping matrix, then cast, in 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 --- autoarray/operators/transformer.py | 14 ++-- test_autoarray/operators/test_transformer.py | 78 ++++++++++++++++++++ 2 files changed, 86 insertions(+), 6 deletions(-) diff --git a/autoarray/operators/transformer.py b/autoarray/operators/transformer.py index 6ff5a878c..3323aea5b 100644 --- a/autoarray/operators/transformer.py +++ b/autoarray/operators/transformer.py @@ -525,14 +525,15 @@ def transform_mapping_matrix(self, mapping_matrix, xp=np): if xp.__name__.startswith("jax"): import jax.numpy as jnp - mm_T = jnp.asarray(mapping_matrix).T.astype(jnp.complex128) - source_images = jnp.zeros((n_src, n_y, n_x), dtype=jnp.complex128) + # Real scatter, then cast: GPU complex128 scatter floor (autolens_profiling#308). + mm_T = jnp.asarray(mapping_matrix).T + source_images = jnp.zeros((n_src, n_y, n_x), dtype=mm_T.dtype) source_images = source_images.at[ jnp.arange(n_src)[:, None], jnp.asarray(rows)[None, :], jnp.asarray(cols)[None, :], ].set(mm_T) - flipped = source_images[:, ::-1, :] + flipped = source_images[:, ::-1, :].astype(jnp.complex128) x = jnp.asarray(self._x) y = jnp.asarray(self._y) shift = jnp.asarray(self._shift) @@ -542,10 +543,11 @@ def transform_mapping_matrix(self, mapping_matrix, xp=np): ) return vis_batched.T - mm_T = np.asarray(mapping_matrix).T.astype(np.complex128) - source_images = np.zeros((n_src, n_y, n_x), dtype=np.complex128) + # Real scatter, then cast: GPU complex128 scatter floor (autolens_profiling#308). + mm_T = np.asarray(mapping_matrix).T + source_images = np.zeros((n_src, n_y, n_x), dtype=mm_T.dtype) source_images[np.arange(n_src)[:, None], rows[None, :], cols[None, :]] = mm_T - flipped = source_images[:, ::-1, :] + flipped = source_images[:, ::-1, :].astype(np.complex128) vis_batched = ( _nufftax.nufft2d2(self._x, self._y, flipped, self.eps, -1) * self._shift[None, :] diff --git a/test_autoarray/operators/test_transformer.py b/test_autoarray/operators/test_transformer.py index 42fd8e674..5567edfab 100644 --- a/test_autoarray/operators/test_transformer.py +++ b/test_autoarray/operators/test_transformer.py @@ -139,6 +139,84 @@ def test__nufft__transform_mapping_matrix__ones_mapping_matrix__first_element_ma assert transformed_mapping_matrix_nufft[0, 0] == pytest.approx(25.0 + 0.0j, 1.0e-4) +def test__nufft__transform_mapping_matrix__real_scatter_matches_complex_scatter_exactly(): + """Scattering the real mapping matrix and casting afterwards + (autolens_profiling#308) must be bit-identical to the previous + cast-then-scatter formula, for NumPy and under ``jax.jit``.""" + import jax + import jax.numpy as jnp + + from autoarray.operators import transformer as transformer_module + + rng = np.random.default_rng(seed=4) + uv_wavelengths = rng.normal(size=(41, 2)) * 50.0 + real_space_mask = aa.Mask2D.circular( + shape_native=(12, 11), pixel_scales=0.05, radius=0.25 + ) + n_src = 5 + mapping_matrix = rng.normal(size=(real_space_mask.pixels_in_mask, n_src)) + + transformer = aa.TransformerNUFFT( + uv_wavelengths=uv_wavelengths, real_space_mask=real_space_mask + ) + + nufftax = transformer_module._load_nufftax() + rows, cols = real_space_mask.slim_to_native_tuple + n_y, n_x = real_space_mask.shape_native + + source_images = np.zeros((n_src, n_y, n_x), dtype=np.complex128) + source_images[np.arange(n_src)[:, None], rows[None, :], cols[None, :]] = ( + mapping_matrix.T.astype(np.complex128) + ) + expected_numpy = np.array( + np.asarray( + nufftax.nufft2d2( + transformer._x, + transformer._y, + source_images[:, ::-1, :], + transformer.eps, + -1, + ) + * transformer._shift[None, :] + ).T + ) + + result_numpy = transformer.transform_mapping_matrix(mapping_matrix=mapping_matrix) + + assert result_numpy.dtype == np.complex128 + assert np.array_equal(result_numpy, expected_numpy) + + @jax.jit + def old_formula(mm): + images = jnp.zeros((n_src, n_y, n_x), dtype=jnp.complex128) + images = images.at[ + jnp.arange(n_src)[:, None], + jnp.asarray(rows)[None, :], + jnp.asarray(cols)[None, :], + ].set(mm.T.astype(jnp.complex128)) + vis = ( + nufftax.nufft2d2( + jnp.asarray(transformer._x), + jnp.asarray(transformer._y), + images[:, ::-1, :], + transformer.eps, + -1, + ) + * jnp.asarray(transformer._shift)[None, :] + ) + return vis.T + + @jax.jit + def new_formula(mm): + return transformer.transform_mapping_matrix(mapping_matrix=mm, xp=jnp) + + expected_jax = np.asarray(old_formula(jnp.asarray(mapping_matrix))) + result_jax = np.asarray(new_formula(jnp.asarray(mapping_matrix))) + + assert result_jax.dtype == np.complex128 + assert np.array_equal(result_jax, expected_jax) + + def test__nufft__chunk_size__rejects_non_positive(): real_space_mask = aa.Mask2D.all_false(shape_native=(5, 5), pixel_scales=0.005) uv_wavelengths = np.array([[0.2, 1.0], [0.5, 1.1], [0.8, 1.2]])