From 2605c92464831d626992e83ad24059a4ed0300e7 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Wed, 2 Sep 2026 13:22:58 -0400 Subject: [PATCH] perf(galaxy): skip the second deflection call for the over-sampled grid at sub-size 1 (PyAutoArray#514) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `Galaxy.traced_grid_2d_from` called `deflections_yx_2d_from` twice — once for the slim grid and once for `grid.over_sampled`. When the grid's over-sampler is uniform at sub-size 1 — the pixelization grid in every CPU likelihood cell — `grid.over_sampled` is the slim grid in the same order, so the second call recomputes a bit-identical answer at full cost and is now reused instead. The sub-size is host numpy, so the guard is a static Python bool and the branch is resolved at JAX trace time rather than becoming a traced `cond`. The equivalent guard in `Tracer.traced_grid_2d_list_from` (PyAutoLens) measures the tracer/raw deflection ratio on hst dropping from 2.52x to 1.41x for Isothermal and 1.94x to 0.88x for PowerLaw, with every pinned value still PASSING — the change is bit-identical. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01HWjPT94MPbEHT45kJmDpDh --- autogalaxy/galaxy/galaxy.py | 23 ++++++++++-- test_autogalaxy/galaxy/test_galaxy.py | 51 +++++++++++++++++++++++++++ 2 files changed, 71 insertions(+), 3 deletions(-) diff --git a/autogalaxy/galaxy/galaxy.py b/autogalaxy/galaxy/galaxy.py index ca89ccc28..ee5f1df7f 100644 --- a/autogalaxy/galaxy/galaxy.py +++ b/autogalaxy/galaxy/galaxy.py @@ -374,12 +374,29 @@ def traced_grid_2d_from( The source-plane (y, x) coordinates after deflection. """ if isinstance(grid, aa.Grid2D): + traced_grid = grid - self.deflections_yx_2d_from(grid=grid, xp=xp) + + # At a uniform over sample size of 1 every sub pixel is the pixel itself, so + # `grid.over_sampled` is the slim grid in the same order and tracing it again + # doubles the deflection angle calculation for a bit-identical result. The + # over sample size is host numpy, so this branch is static at JAX trace time. + + sub_size_all_one = bool(np.all(np.asarray(grid.over_sample_size.array) == 1)) + + if sub_size_all_one: + traced_grid_over_sampled = aa.Grid2DIrregular( + values=traced_grid.array, xp=xp + ) + else: + traced_grid_over_sampled = grid.over_sampled - self.deflections_yx_2d_from( + grid=grid.over_sampled, xp=xp + ) + return aa.Grid2D( - values=grid - self.deflections_yx_2d_from(grid=grid, xp=xp), + values=traced_grid, mask=grid.mask, over_sample_size=grid.over_sample_size, - over_sampled=grid.over_sampled - - self.deflections_yx_2d_from(grid=grid.over_sampled, xp=xp), + over_sampled=traced_grid_over_sampled, over_sampler=grid.over_sampler, ) diff --git a/test_autogalaxy/galaxy/test_galaxy.py b/test_autogalaxy/galaxy/test_galaxy.py index 03f661015..ddd0962f1 100644 --- a/test_autogalaxy/galaxy/test_galaxy.py +++ b/test_autogalaxy/galaxy/test_galaxy.py @@ -192,6 +192,57 @@ def test__deflections_yx_2d_from__two_mass_profiles__matches_sum_of_individual_d assert gal_deflections[1, 0] == mp_deflections[1, 0] +@pytest.mark.parametrize("over_sample_size", [1, 4]) +def test__traced_grid_2d_from__over_sampled_grid_is_the_traced_over_sampled_grid( + over_sample_size, +): + """ + The `over_sampled` attribute of the traced grid must be the galaxy's deflections applied to the + input grid's `over_sampled` grid. At a uniform over sample size of 1 the second ray-trace is + skipped as an optimization, so this is pinned for a sub size of 1 and a sub size above 1. + """ + galaxy = ag.Galaxy( + redshift=0.5, mass_profile=ag.mp.IsothermalSph(einstein_radius=1.0) + ) + + grid = ag.Grid2D.uniform( + shape_native=(7, 7), pixel_scales=0.1, over_sample_size=over_sample_size + ) + + traced_grid = galaxy.traced_grid_2d_from(grid=grid) + + traced_over_sampled = galaxy.traced_grid_2d_from(grid=grid.over_sampled) + + assert np.allclose( + np.asarray(traced_grid.over_sampled.array), + np.asarray(traced_over_sampled.array), + rtol=0.0, + atol=1.0e-12, + ) + + assert not np.allclose( + np.asarray(traced_grid.over_sampled.array), np.asarray(grid.over_sampled.array) + ) + + +def test__traced_grid_2d_from__over_sample_size_1__over_sampled_grid_equals_slim_grid(): + """ + At a uniform over sample size of 1 every sub pixel is the pixel itself, so the traced over + sampled grid is bit identical to the traced slim grid. + """ + galaxy = ag.Galaxy( + redshift=0.5, mass_profile=ag.mp.IsothermalSph(einstein_radius=1.0) + ) + + grid = ag.Grid2D.uniform(shape_native=(7, 7), pixel_scales=0.1, over_sample_size=1) + + traced_grid = galaxy.traced_grid_2d_from(grid=grid) + + assert np.array_equal( + np.asarray(traced_grid.over_sampled.array), np.asarray(traced_grid.array) + ) + + def test__no_mass_profile__potential_2d_from__returns_zeros_of_grid_shape(grid_2d_7x7): galaxy = ag.Galaxy(redshift=0.5)