From 04152bb776350de589ed4349d07e1a7e2cb1608b Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Tue, 8 Sep 2026 09:00:56 -0400 Subject: [PATCH] perf(reg_split): compact the ConstantSplit split-stencil scatter, wide-row supplement + NaN overflow guard (#536) The JAX path of `pixel_splitted_regularization_matrix_from` scattered the full padded `(4P, K, K)` outer product. For the natural-neighbor `DelaunayNN` mesh `K = 33` while the real occupied width is at most ~11, so ~97% of the 6.5M scattered entries were padding, at a cost quadratic in the padded width. It now scatters only the first `SPLIT_REG_COMPACT_WIDTH = 12` columns of every row and supplements the `SPLIT_REG_WIDE_ROW_BUDGET = 256` widest rows with the head x tail / tail x head / tail x tail blocks the compact pass did not cover, in one masked scatter. The result is exact: padded columns contribute mapping 0 / weight 0, so a row of size <= 12 is reproduced bit-for-bit by the compact pass alone, and wider rows are supplemented at full width rather than dropped. Beyond the budget the matrix is poisoned with NaN, matching the Sibson cap convention (NaN weights -> NaN likelihood -> the sample is discarded), instead of returning a silently truncated H. When `K <= compact_width` (the `Delaunay` mesh's `K = 4`, and the adapt-split family) the compaction is a no-op and the function performs today's single scatter with no supplement and no guard. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_011xsh8KqgiHWTfPGgWEYMQw --- .../regularization/regularization_util.py | 137 +++++++-- .../test_pixel_splitted_jax.py | 271 ++++++++++++++++++ 2 files changed, 391 insertions(+), 17 deletions(-) create mode 100644 test_autoarray/inversion/regularizations/test_pixel_splitted_jax.py diff --git a/autoarray/inversion/regularization/regularization_util.py b/autoarray/inversion/regularization/regularization_util.py index e38197c0b..7d508b6f1 100644 --- a/autoarray/inversion/regularization/regularization_util.py +++ b/autoarray/inversion/regularization/regularization_util.py @@ -30,6 +30,34 @@ from autoarray.inversion.regularization.zeroth import zeroth_regularization_matrix_from +# --------------------------------------------------------------------------- +# Split-regularization stencil compaction (JAX path of +# `pixel_splitted_regularization_matrix_from`) +# --------------------------------------------------------------------------- +# +# The split stencil tables are fixed-shape `(4P, K)` arrays: `K = 4` for the `Delaunay` mesh, but +# `K = 33` for the natural-neighbor `DelaunayNN` mesh (`SIBSON_MAX_NEIGHBORS` 32 + 1 spare column +# for the self insertion). On the real HST `DelaunayNN` cell the occupied post-`reg_split_from` +# width is min 1 / median 5 / p99 9 / max 11, so a full `(4P, K, K)` outer-product scatter spends +# ~97% of its 6.5M entries on padding, at a cost quadratic in the padded width: +# +# compact width | 12 16 20 24 28 32 33 (= today, no compaction) +# A100 fp64 ms/call | 0.58 1.47 2.78 4.53 6.71 9.31 10.03 (vmap 16, real HST tables) +# CPU fp64 ms/call | 12.6 72.8 +# +# so width 12 is a 17x GPU and 5.8x CPU improvement over the uncompacted scatter, with no backend +# gate needed. Width 12 covers the production stencil with margin, but the cap audit over 101 +# ensemble geometries (`autolens_profiling/results/notes/delaunay_nn_cap_audit.md`) saw rare tail +# geometries reach 21 natural neighbors (99.9th pct 11, 99.99th pct 15, max 21), so rows wider than +# the compact width are supplemented exactly rather than dropped: the `SPLIT_REG_WIDE_ROW_BUDGET` +# widest rows get a full-width supplementary scatter. The budget of 256 rows is ~4x the count of +# above-width rows in the worst audited geometry and costs `W * (K**2 - kc**2)` ~ 0.24M entries, +# an order of magnitude below the 6.5M it replaces. Beyond the budget the matrix is NaN (see the +# function docstring), never silently truncated. +SPLIT_REG_COMPACT_WIDTH = 12 +SPLIT_REG_WIDE_ROW_BUDGET = 256 + + def split_points_from(points, area_weights, xp=np): """ points : (N, 2) @@ -286,10 +314,12 @@ def pixel_splitted_regularization_matrix_np_from( def pixel_splitted_regularization_matrix_from( regularization_weights: np.ndarray, # (P,) - splitted_mappings: np.ndarray, # (4P, 4) + splitted_mappings: np.ndarray, # (4P, K) splitted_sizes: np.ndarray, # (4P,) - splitted_weights: np.ndarray, # (4P, 4) + splitted_weights: np.ndarray, # (4P, K) xp=np, + compact_width: int = SPLIT_REG_COMPACT_WIDTH, + wide_row_budget: int = SPLIT_REG_WIDE_ROW_BUDGET, ): """ Returns the regularization matrix for the adaptive split-pixel regularization scheme. @@ -302,6 +332,28 @@ def pixel_splitted_regularization_matrix_from( A visual description and further details are provided in the appendix of He et al. (2024): https://arxiv.org/abs/2403.16253 + JAX path: compact main scatter plus a wide-row supplement + --------------------------------------------------------- + The stencil tables are fixed-shape ``(4P, K)`` arrays whose columns beyond each row's + ``splitted_sizes`` entry are padding (mapping ``-1``, weight ``0``). For the natural-neighbor + (``DelaunayNN``) mesh ``K = 33`` while the real occupied width is at most ~11, so scattering the + full ``(4P, K, K)`` outer product spends ~97% of its traffic on padding and costs + ``O(K**2)``. This function therefore scatters only the first ``kc = min(K, compact_width)`` + columns and supplements the ``wide_row_budget`` widest rows with the blocks the compact pass + did not cover (see :data:`SPLIT_REG_COMPACT_WIDTH`). The result is exact: padded columns + contribute mapping ``0`` / weight ``0``, so a row whose size is ``<= kc`` is reproduced + bit-for-bit by the compact pass alone. + + If more rows exceed ``kc`` than the supplement budget holds, the matrix is poisoned with NaN + rather than silently truncated. This is the same NaN-on-overflow contract the Sibson + natural-neighbor caps already use (``mesh/interpolator/sibson.py``, where ``neighbor_overflow`` + / ``failed`` set the interpolation weights to NaN): the likelihood of an out-of-budget geometry + evaluates to NaN and is discarded by the sampler, never returning a silently wrong ``H``. + + When ``K <= compact_width`` (the ``Delaunay`` mesh's ``K = 4``, and the adapt-split family) the + compaction is a no-op: the function performs today's single scatter with no supplement and no + overflow guard. + Parameters ---------- regularization_weights @@ -314,6 +366,14 @@ def pixel_splitted_regularization_matrix_from( splitted_weights The interpolation weights corresponding to each mapping entry, used to apply regularization between split points. + xp + The array module used, `numpy` or `jax.numpy`. + compact_width + The number of stencil columns scattered for every row on the JAX path (see + :data:`SPLIT_REG_COMPACT_WIDTH`). Ignored on the numpy path. + wide_row_budget + The number of widest rows given a full-width supplementary scatter on the JAX path (see + :data:`SPLIT_REG_WIDE_ROW_BUDGET`). Ignored on the numpy path. Returns ------- @@ -328,10 +388,12 @@ def pixel_splitted_regularization_matrix_from( splitted_weights=splitted_weights, ) + import jax import jax.numpy as jnp # How many real pixels? P = splitted_mappings.shape[0] // 4 + K = splitted_mappings.shape[1] # Square, positive regularization weights reg_w = regularization_weights**2.0 # (P,) @@ -342,31 +404,72 @@ def pixel_splitted_regularization_matrix_from( # ----- Build all 4P contributions at once ----- # Mask away padded entries (where mapping = -1) - valid = splitted_mappings != -1 # (4P, 4) + valid = splitted_mappings != -1 # (4P, K) # Extract valid mapping rows and weights - # BUT keep fixed shape (4) and just zero out invalid ones - map_fixed = jnp.where(valid, splitted_mappings, 0) # (4P, 4) - w_fixed = jnp.where(valid, splitted_weights, 0.0) # (4P, 4) + # BUT keep fixed shape (K) and just zero out invalid ones + map_fixed = jnp.where(valid, splitted_mappings, 0) # (4P, K) + w_fixed = jnp.where(valid, splitted_weights, 0.0) # (4P, K) - # Compute all outer products of weights - # w_fixed[:, :, None] * w_fixed[:, None, :] → (4P, 4, 4) - outer = w_fixed[:, :, None] * w_fixed[:, None, :] # (4P, 4, 4) - - # Build corresponding row and col index grids - rows = map_fixed[:, :, None] # (4P, 4, 1) - cols = map_fixed[:, None, :] # (4P, 1, 4) - - # Multiply each 4x4 block by its pixel’s regularization weight - # Rows 0–3 belong to pixel 0, rows 4–7 to pixel 1, etc. + # Each block is scaled by its pixel's regularization weight. + # Rows 0-3 belong to pixel 0, rows 4-7 to pixel 1, etc. pixel_index = jnp.arange(4 * P) // 4 # (4P,) block_scale = reg_w[pixel_index] # (4P,) + + # ----- Compact main scatter over the first kc columns ----- + + kc = min(K, int(compact_width)) + + map_head = map_fixed[:, :kc] # (4P, kc) + w_head = w_fixed[:, :kc] # (4P, kc) + + outer = w_head[:, :, None] * w_head[:, None, :] # (4P, kc, kc) outer_scaled = outer * block_scale[:, None, None] - # Now scatter-add all entries into the (P,P) matrix + rows = map_head[:, :, None] # (4P, kc, 1) + cols = map_head[:, None, :] # (4P, 1, kc) + reg_mat = reg_mat.at[rows, cols].add(outer_scaled) + # ----- Wide-row supplement (only when the tables are padded wider than kc) ----- + + if K > kc: + W = min(4 * P, int(wide_row_budget)) + + # The widest rows, by their post-split occupied size. `top_k` on the integer sizes is not + # differentiated through, so the weights stay fully differentiable. + _, wide_rows = jax.lax.top_k(splitted_sizes, W) # (W,) + + map_wide = map_fixed[wide_rows] # (W, K) + w_wide = w_fixed[wide_rows] # (W, K) + scale_wide = block_scale[wide_rows] # (W,) + + outer_wide = w_wide[:, :, None] * w_wide[:, None, :] # (W, K, K) + outer_wide = outer_wide * scale_wide[:, None, None] + + # Zero the head x head block, which the compact pass above already scattered, leaving the + # head x tail, tail x head and tail x tail blocks. Scattering the masked (W, K, K) block in + # one `.at[].add` costs a single kernel launch, which on GPU beats three block scatters. + col_index = jnp.arange(K) + head_block = (col_index[:, None] < kc) & (col_index[None, :] < kc) # (K, K) + outer_wide = jnp.where(head_block[None, :, :], 0.0, outer_wide) + + rows_wide = map_wide[:, :, None] # (W, K, 1) + cols_wide = map_wide[:, None, :] # (W, 1, K) + + reg_mat = reg_mat.at[rows_wide, cols_wide].add(outer_wide) + + # Overflow guard: more rows wider than the compact width than the supplement budget holds. + overflow = jnp.sum(splitted_sizes > kc) > W + else: + overflow = None + # Divide diagonal by 2 reg_mat = reg_mat.at[jnp.diag_indices(reg_mat.shape[0])].add(-1e-8) + if overflow is not None: + # Poison the matrix on overflow, matching the Sibson cap convention (NaN weights -> NaN + # likelihood -> the sample is discarded), instead of returning a silently truncated matrix. + reg_mat = jnp.where(overflow, jnp.nan, reg_mat) + return reg_mat diff --git a/test_autoarray/inversion/regularizations/test_pixel_splitted_jax.py b/test_autoarray/inversion/regularizations/test_pixel_splitted_jax.py new file mode 100644 index 000000000..884e56ae2 --- /dev/null +++ b/test_autoarray/inversion/regularizations/test_pixel_splitted_jax.py @@ -0,0 +1,271 @@ +""" +JAX leg of the split-pixel regularization matrix builder: the compacted scatter (a narrow main +pass plus a wide-row supplement, see ``SPLIT_REG_COMPACT_WIDTH``) must reproduce the NumPy +reference exactly, poison the matrix with NaN when the wide-row budget overflows, and stay +jit / vmap / grad friendly. + +Skipped when JAX is absent (it is an optional dependency). +""" + +import numpy as np +import pytest + +jax = pytest.importorskip("jax") +jnp = pytest.importorskip("jax.numpy") + +jax.config.update("jax_enable_x64", True) + +from autoarray.inversion.regularization import regularization_util # noqa: E402 +from autoarray.inversion.regularization.regularization_util import ( # noqa: E402 + SPLIT_REG_COMPACT_WIDTH, +) + + +def tables_from_sizes(sizes, width, seed=1): + """ + A synthetic ``(4P, K)`` split stencil table, padded ``DelaunayNN`` style: every column beyond a + row's size is mapping ``-1`` / weight ``0.0``. The number of pixels ``P`` is ``len(sizes) / 4``, + every mapping is a pixel index and no row repeats a pixel, so a row can be at most ``P`` wide. + """ + rng = np.random.default_rng(seed) + + sizes = np.asarray(sizes, dtype=np.int32) + + assert sizes.shape[0] % 4 == 0 + + total_pixels = sizes.shape[0] // 4 + + assert int(sizes.max()) <= total_pixels + assert int(sizes.max()) <= width + + mappings = -np.ones((sizes.shape[0], width), dtype=np.int32) + weights = np.zeros((sizes.shape[0], width), dtype=np.float64) + + for row, size in enumerate(sizes): + pixels = rng.choice(total_pixels, size=int(size), replace=False) + mappings[row, : int(size)] = pixels + values = rng.uniform(0.05, 1.0, size=int(size)) + weights[row, : int(size)] = values / values.sum() + + regularization_weights = rng.uniform(0.5, 3.0, size=total_pixels) + + return regularization_weights, mappings, sizes, weights + + +NARROW_SIZES = [1, 3, 5, 8, 11, 4, 2, 9, 12, 6, 3, 7] * 4 + + +def wide_sizes(): + """ + ``DelaunayNN``-shaped sizes: 144 rows (36 pixels) of which five exceed the compact width, one + of them the full 33-column table width. + """ + sizes = [1, 3, 5, 8, 11, 4, 2, 9, 12, 6, 3, 7] * 12 + + for row, size in [(5, 21), (40, 13), (77, 33), (100, 14), (131, 20)]: + sizes[row] = size + + return sizes + + +WIDE_SIZES = wide_sizes() + + +def matrix_np_from(regularization_weights, mappings, sizes, weights): + return regularization_util.pixel_splitted_regularization_matrix_np_from( + regularization_weights=np.copy(regularization_weights), + splitted_mappings=np.copy(mappings), + splitted_sizes=np.copy(sizes), + splitted_weights=np.copy(weights), + ) + + +def matrix_jax_from(regularization_weights, mappings, sizes, weights, **kwargs): + return regularization_util.pixel_splitted_regularization_matrix_from( + regularization_weights=jnp.asarray(regularization_weights), + splitted_mappings=jnp.asarray(mappings), + splitted_sizes=jnp.asarray(sizes), + splitted_weights=jnp.asarray(weights), + xp=jnp, + **kwargs, + ) + + +def assert_matrices_equal(matrix_jax, matrix_np): + np.testing.assert_allclose( + np.asarray(matrix_jax, dtype=np.float64), matrix_np, rtol=1.0e-12, atol=1.0e-14 + ) + + +def test__padding_convention__valid_mask_agrees_with_sizes(): + """ + The compaction selects wide rows by ``splitted_sizes`` but masks entries by ``mapping != -1``: + the two must describe the same occupied columns. + """ + _, mappings, sizes, _ = tables_from_sizes(sizes=WIDE_SIZES, width=33) + + valid = mappings != -1 + assert np.array_equal(valid.sum(axis=1).astype(np.int32), sizes) + assert np.array_equal(valid, np.arange(33)[None, :] < sizes[:, None]) + + +def test__all_rows_narrower_than_compact_width__matches_numpy(): + reg_weights, mappings, sizes, weights = tables_from_sizes( + sizes=NARROW_SIZES, width=33 + ) + + assert int(sizes.max()) <= SPLIT_REG_COMPACT_WIDTH + + assert_matrices_equal( + matrix_jax_from(reg_weights, mappings, sizes, weights), + matrix_np_from(reg_weights, mappings, sizes, weights), + ) + + +def test__wide_rows_inside_budget__matches_numpy(): + """ + Rows wider than the compact width exercise the head x tail, tail x head and tail x tail + supplement blocks. One row is exactly the full table width. + """ + reg_weights, mappings, sizes, weights = tables_from_sizes( + sizes=WIDE_SIZES, width=33, seed=3 + ) + + assert int((sizes > SPLIT_REG_COMPACT_WIDTH).sum()) == 5 + assert int(sizes.max()) == 33 + + assert_matrices_equal( + matrix_jax_from(reg_weights, mappings, sizes, weights), + matrix_np_from(reg_weights, mappings, sizes, weights), + ) + + +def test__wide_rows_at_the_budget__matches_numpy(): + reg_weights, mappings, sizes, weights = tables_from_sizes( + sizes=WIDE_SIZES, width=33, seed=5 + ) + + wide = int((sizes > SPLIT_REG_COMPACT_WIDTH).sum()) + + assert_matrices_equal( + matrix_jax_from(reg_weights, mappings, sizes, weights, wide_row_budget=wide), + matrix_np_from(reg_weights, mappings, sizes, weights), + ) + + +def test__more_wide_rows_than_the_budget__matrix_is_nan(): + reg_weights, mappings, sizes, weights = tables_from_sizes( + sizes=WIDE_SIZES, width=33, seed=5 + ) + + wide = int((sizes > SPLIT_REG_COMPACT_WIDTH).sum()) + + matrix = matrix_jax_from( + reg_weights, mappings, sizes, weights, wide_row_budget=wide - 1 + ) + + assert np.all(np.isnan(np.asarray(matrix))) + + +def test__table_narrower_than_compact_width__matches_numpy_with_no_supplement(): + """ + The ``Delaunay`` mesh's ``K = 4`` tables are already narrower than the compact width: the + compaction is a no-op, the supplement is not built and the overflow guard is not applied. + """ + reg_weights, mappings, sizes, weights = tables_from_sizes( + sizes=[1, 2, 3, 4] * 6, width=4, seed=7 + ) + + assert_matrices_equal( + matrix_jax_from(reg_weights, mappings, sizes, weights), + matrix_np_from(reg_weights, mappings, sizes, weights), + ) + + jaxpr = jax.make_jaxpr( + lambda w, m, s, ws: regularization_util.pixel_splitted_regularization_matrix_from( + regularization_weights=w, + splitted_mappings=m, + splitted_sizes=s, + splitted_weights=ws, + xp=jnp, + ) + )( + jnp.asarray(reg_weights), + jnp.asarray(mappings), + jnp.asarray(sizes), + jnp.asarray(weights), + ) + + assert "top_k" not in str(jaxpr) + + +def test__jit_and_vmap__match_the_unbatched_result(): + tables = [ + tables_from_sizes(sizes=WIDE_SIZES, width=33, seed=seed) for seed in range(3) + ] + + def matrix_from(reg_weights, mappings, sizes, weights): + return regularization_util.pixel_splitted_regularization_matrix_from( + regularization_weights=reg_weights, + splitted_mappings=mappings, + splitted_sizes=sizes, + splitted_weights=weights, + xp=jnp, + ) + + unbatched = [ + np.asarray(jax.jit(matrix_from)(*[jnp.asarray(a) for a in t])) for t in tables + ] + + for matrix, table in zip(unbatched, tables): + assert_matrices_equal(matrix, matrix_np_from(*table)) + + batched = jax.jit(jax.vmap(matrix_from))( + jnp.asarray(np.stack([t[0] for t in tables])), + jnp.asarray(np.stack([t[1] for t in tables])), + jnp.asarray(np.stack([t[2] for t in tables])), + jnp.asarray(np.stack([t[3] for t in tables])), + ) + + for index in range(3): + assert_matrices_equal(batched[index], unbatched[index]) + + +def test__gradient_of_the_matrix_sum__is_finite_and_matches_finite_differences(): + reg_weights, mappings, sizes, weights = tables_from_sizes( + sizes=WIDE_SIZES, width=33, seed=11 + ) + + def matrix_sum(splitted_weights): + return jnp.sum( + regularization_util.pixel_splitted_regularization_matrix_from( + regularization_weights=jnp.asarray(reg_weights), + splitted_mappings=jnp.asarray(mappings), + splitted_sizes=jnp.asarray(sizes), + splitted_weights=splitted_weights, + xp=jnp, + ) + ) + + gradient = np.asarray(jax.grad(matrix_sum)(jnp.asarray(weights))) + + assert np.all(np.isfinite(gradient)) + assert np.any(gradient != 0.0) + + # A row and column inside the wide-row supplement's tail block. + row = int(np.argmax(sizes)) + column = 15 + + step = 1.0e-6 + + weights_up = np.copy(weights) + weights_up[row, column] += step + weights_down = np.copy(weights) + weights_down[row, column] -= step + + finite_difference = ( + float(matrix_sum(jnp.asarray(weights_up))) + - float(matrix_sum(jnp.asarray(weights_down))) + ) / (2.0 * step) + + assert gradient[row, column] == pytest.approx(finite_difference, rel=1.0e-6)