From b1515e5d3b46462561f78201ccd019c45bc3cab8 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Wed, 2 Sep 2026 13:22:49 -0400 Subject: [PATCH] perf(tracer): skip the second ray-trace of 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 `Tracer.traced_grid_2d_list_from` traced the input `Grid2D` and then traced `grid.over_sampled` a second time to fill the returned grids' `over_sampled` attribute. 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 trace recomputes a bit-identical answer at full cost. 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`. Measured on `autolens_profiling/scripts/lens/deflections/total.py --instrument hst` (the tracer/raw ratio, which normalises out laptop variance): Isothermal 2.52x -> 1.41x, PowerLaw 1.94x -> 0.88x, IsothermalSph 2.78x -> 1.61x, PowerLawSph 2.43x -> 1.38x. Every pinned value still PASSES — the change is bit-identical. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01HWjPT94MPbEHT45kJmDpDh --- autolens/lens/tracer.py | 27 +++++++++++---- test_autolens/lens/test_tracer.py | 55 +++++++++++++++++++++++++++++++ 2 files changed, 75 insertions(+), 7 deletions(-) diff --git a/autolens/lens/tracer.py b/autolens/lens/tracer.py index 84ec1ac1a..da243131d 100644 --- a/autolens/lens/tracer.py +++ b/autolens/lens/tracer.py @@ -496,13 +496,26 @@ def traced_grid_2d_list_from( ) if isinstance(grid, aa.Grid2D): - grid_2d_over_sampled_list = tracer_util.traced_grid_2d_list_from( - planes=self.planes, - grid=grid.over_sampled, - cosmology=self.cosmology, - plane_index_limit=plane_index_limit, - 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 every 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: + grid_2d_over_sampled_list = [ + aa.Grid2DIrregular(values=grid_2d.array, xp=xp) + for grid_2d in grid_2d_list + ] + else: + grid_2d_over_sampled_list = tracer_util.traced_grid_2d_list_from( + planes=self.planes, + grid=grid.over_sampled, + cosmology=self.cosmology, + plane_index_limit=plane_index_limit, + xp=xp, + ) grid_2d_new_list = [] diff --git a/test_autolens/lens/test_tracer.py b/test_autolens/lens/test_tracer.py index ad2b003b2..017d66e48 100644 --- a/test_autolens/lens/test_tracer.py +++ b/test_autolens/lens/test_tracer.py @@ -205,6 +205,61 @@ def test__traced_grid_2d_list_from__plane_index_limit__only_traces_up_to_specifi assert len(traced_grid_list) == 2 +@pytest.mark.parametrize("over_sample_size", [1, 4]) +def test__traced_grid_2d_list_from__over_sampled_grid_is_the_traced_over_sampled_grid( + over_sample_size, +): + """ + The `over_sampled` attribute of each traced grid must be the tracer 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. + """ + g0 = al.Galaxy(redshift=0.5, mass_profile=al.mp.IsothermalSph(einstein_radius=1.0)) + g1 = al.Galaxy(redshift=1.0) + + tracer = al.Tracer(galaxies=[g0, g1]) + + grid = al.Grid2D.uniform( + shape_native=(7, 7), pixel_scales=0.1, over_sample_size=over_sample_size + ) + + traced_grid_list = tracer.traced_grid_2d_list_from(grid=grid) + + traced_over_sampled_list = tracer.traced_grid_2d_list_from(grid=grid.over_sampled) + + assert np.allclose( + np.asarray(traced_grid_list[-1].over_sampled.array), + np.asarray(traced_over_sampled_list[-1].array), + rtol=0.0, + atol=1.0e-12, + ) + + assert not np.allclose( + np.asarray(traced_grid_list[-1].over_sampled.array), + np.asarray(grid.over_sampled.array), + ) + + +def test__traced_grid_2d_list_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. + """ + g0 = al.Galaxy(redshift=0.5, mass_profile=al.mp.IsothermalSph(einstein_radius=1.0)) + g1 = al.Galaxy(redshift=1.0) + + tracer = al.Tracer(galaxies=[g0, g1]) + + grid = al.Grid2D.uniform(shape_native=(7, 7), pixel_scales=0.1, over_sample_size=1) + + traced_grid_list = tracer.traced_grid_2d_list_from(grid=grid) + + for traced_grid in traced_grid_list: + assert np.array_equal( + np.asarray(traced_grid.over_sampled.array), np.asarray(traced_grid.array) + ) + + # ── grid_2d_at_redshift_from ──────────────────────────────────────────────────