Skip to content

perf: compact the ConstantSplit split-stencil scatter (DelaunayNN) #536

Description

@Jammy2211

Overview

Phase A of the DelaunayNN speed-up (PyAutoArray#533) cut the params→H prefix from 143.9 to 28.2 ms unbatched (5.1x) but only from 24.3 to 16.4 ms per call at vmap 16, because ~10.0 ms per call is the ConstantSplit regularization assembly, untouched by any Sibson change. That assembly is regularization_util.pixel_splitted_regularization_matrix_from (JAX branch): an outer product of shape (4P, K, K) scattered into (P, P) with .at[rows, cols].add, where DelaunayNN's K = 33 (SIBSON_MAX_NEIGHBORS 32 + 1 spare column).

An A100 investigation (jobs 342331/342332, real HST tables) measured the actual post-reg_split_from stencil size on the production cell: min 1 / median 5 / p99 9 / max 11. So 6,534,000 scatter entries carry only 187,242 real contributions into 29,020 cells — 97 % of the traffic is padding, and the cost is quadratic in the padded width. A compact scatter at width 12 measured 0.86 ms unbatched / 0.58 ms per call at vmap 16 (12x / 17x), agreeing with the current result to <= 2.7e-15 absolute and bit-identically on CPU, where it is also a 6x improvement.

The catch is the tail: the cap audit saw split-point stencils reach 21 natural neighbours in rare ensemble geometries, so a fixed narrow width alone would be wrong there. This task ships an exact hybrid compaction with a wide-row budget — narrow main scatter plus a top_k-selected wide-row supplement, with the existing NaN-on-overflow contract as the guard.

Plan

  • Compact the main scatter in pixel_splitted_regularization_matrix_from (JAX branch only) to the first kc = min(K, 12) columns of each row, since every row whose post-split size is <= kc is already bit-identical there.
  • Add a wide-row supplement: jax.lax.top_k selects the W = 256 rows with the largest split sizes, and only the head x tail, tail x head, tail x tail blocks the main pass missed are scattered for them, so the result stays exact for the tail geometries.
  • Guard the budget with the existing NaN-on-overflow convention already used by the Sibson caps: if more rows exceed kc than the budget holds, poison H with NaN so the sampler discards the sample rather than silently accepting a wrong matrix.
  • Leave the NumPy path, reg_split_from, and the hstack spare-column plumbing in sibson.py untouched — the measurement says they are noise.
  • Add unit tests (NumPy path unchanged; a JAX leg covering all-narrow, wide-inside-budget, and over-budget-to-NaN) and a compaction parity check in the autolens_workspace_test jax_assertions script.
  • Verify with a same-node A100 A/B against the merge base and record the numbers in a new autolens_profiling results note; the investigation bench moves into the repo so the numbers are reproducible.
Detailed implementation plan

Work Classification

Both — library first (PyAutoArray), workspace follow-up (autolens_profiling, autolens_workspace_test) once the API impact is known.

Affected Repositories

  • PyAutoArray (primary, library) — the change + unit tests
  • autolens_profiling (workspace) — A100 A/B submits, results note, bench script
  • autolens_workspace_test (workspace) — jax_assertions parity check

Branch Survey

Repository Current Branch Dirty?
./PyAutoArray main clean
./autolens_profiling main clean
./autolens_workspace_test main clean

Worktree claims: PyAutoArray and autolens_workspace_test are unclaimed in active.md. autolens_profiling is claimed in parallel by retire-gpu1-mig-exclusion (awaiting-merge, 88 MIG-exclusion files) and interferometer-preload-cpu (no commits yet, interferometer preload scope) — file sets are disjoint from this task's (new submit_*assembly* files, a new results note, new breakdown JSON, scripts/misc/delaunay_nn/assembly_bench.py), so this task takes its own worktree under the same parallel-claim precedent already recorded twice in active.md. ~/Code/PyAutoLabs-wt/delaunay-area-magnification-audit/PyAutoArray is an unregistered worktree from a different task — left alone.

Suggested branch: feature/delaunay-nn-constant-split-assembly

Worktree root: ~/Code/PyAutoLabs-wt/delaunay-nn-constant-split-assembly/

Investigation numbers (A100 jobs 342331/342332, real HST tables, fp64)

Production cell: N = 1500, S = 6000 split points, padded width K = 33. Post-reg_split_from stencil size min 1 / median 5 / p99 9 / max 11; 6,534,000 scatter entries carry 187,242 real contributions into 29,020 cells.

variant (A100, fp64) unbatched ms per call @ vmap16 ms agrees with current
current scatter, K = 33 10.45 10.03 (reproduces the production row) —
dense GEMM B^T diag(s) B 9.83 10.11 yes, but 5.2x CPU regression
BCOO sparse 20.1 21.0 yes
dedup sort + segment_sum 31.9 39.1 yes
compact scatter, width 12 0.86 0.58 yes (<= 2.7e-15 abs; bit-identical on CPU)
compact width 16 1.8 1.5 yes
compact width 20 3.1 2.8 yes
compact width 24 4.9 4.5 yes
compact width 28 7.1 6.7 yes

CPU (laptop, same tables): current 58.0 / 72.8 ms, compact-12 9.8 / 12.6 ms — a 6x CPU improvement, so no backend gate is needed.

Tail evidence: autolens_profiling/results/notes/delaunay_nn_cap_audit.md saw split-point stencils reach 21 natural neighbours in rare geometries (99.9th pct 11, 99.99th pct 15, rows above 16 = 28 in the worst ensemble geometry). A cap-safe fixed width of 24 only reaches 4.5 ms per call and misses the < 3 ms witness — hence the hybrid.

Implementation Steps

In pixel_splitted_regularization_matrix_from (JAX branch only; the xp is np path is untouched):

  1. Compact main scatter. Scatter the outer product of only the first kc columns of every row, kc = min(K, SPLIT_REG_COMPACT_WIDTH) (default 12). Cost 4P * kc^2. Bit-identical for every row whose post-split size is <= kc, because columns beyond the size already carry mapping 0 / weight 0 via the existing valid mask.
  2. Wide-row supplement. Select the W = min(4P, SPLIT_REG_WIDE_ROW_BUDGET) rows with the largest splitted_sizes via jax.lax.top_k, gather their full K-wide rows, and scatter only the blocks the main pass did not cover — head x tail, tail x head, tail x tail (columns >= kc). Cost W * (K^2 - kc^2) ~= 256 * 945 = 0.24 M entries against the main pass's 0.86 M, both an order of magnitude below today's 6.5 M. Rows inside the budget whose size is <= kc contribute exact zeros.
  3. Overflow guard, existing convention. overflow = (number of rows with size > kc) > W. On overflow poison the matrix with NaN (jnp.where(overflow, nan, H)) — the same NaN-on-overflow contract the Sibson caps already use (sibson.py:550-555), so an out-of-budget geometry yields a NaN likelihood the sampler discards rather than a silently wrong H. Document that the budget is a soft cap tuned from the audit, with the numbers above.
  4. K <= kc collapses to today. Delaunay's K = 4 and the adapt-split family take the single scatter with no supplement — no change for those callers beyond a trivially-false guard. Constants live in regularization_util.py as module-level values, exposed as kwargs on the function; no env override (the chunk env override exists because it needed sweeping without edits; these do not).
  5. Leave reg_split_from and the hstack spare-column plumbing in sibson.py alone — the measurement says they are noise.

Summation-order note. Rows in the wide budget are added in a different order than today, so GPU results differ at the ~1e-13 relative level (the GPU scatter already reorders between variants). The pin EXPECTED_LOG_EVIDENCE_HST = 29144.581944 is checked at rtol 1e-4 and will hold. On CPU with no wide rows the result is bit-identical.

Tests. New JAX leg in test_autoarray/inversion/regularizations/test_pixel_splitted_jax.py, following the test_adapt_power_jax.py importorskip pattern: synthetic (4P, K) tables with (a) all rows <= kc, (b) a few rows > kc inside the budget, (c) more wide rows than the budget -> NaN; each compared against the NumPy pixel_splitted_regularization_matrix_np_from. Run pytest test_autoarray/inversion/regularizations test_autoarray/inversion/mesh, then full pytest test_autoarray/.

Workspace_test. Extend autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py with a compaction parity check (JAX ConstantSplit H on synthetic production-size tables == NumPy, with a forced wide row) and re-run it plus delaunay_nn_caps.py. Also record, per audit geometry, the count of split rows above kc — this is the evidence that W = 256 has margin; if the worst geometry exceeds ~W/4, raise the default before shipping.

A100 A/B protocol (same as #531 / #533)

  • Private PyAutoArray checkouts at the merge base and at the feature branch under /mnt/ral/jnightin/PyAuto_wt/delaunay-nn-constant-split-assembly/; the shared RAL install is untouched and reached only via PYTHONPATH.
  • Same node / same session for control and feature.
  • Command: scripts/imaging/likelihood_breakdown/delaunay_nn.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16, plus the runtime cell.
  • New submits: submit_breakdown_imaging_delaunay_nn_a100_hst_fp64_{assembly_control,assembly} and the runtime twins.
  • Results note: autolens_profiling/results/notes/delaunay_nn_constant_split_assembly.md; breakdown JSON under results/breakdown/imaging/delaunay_nn_hpc_a100_fp64_assembly*.json.
  • The investigation's assembly_bench.py and its real-table builder move into autolens_profiling/scripts/misc/delaunay_nn/ so the numbers above are reproducible.
  • Judge on regularization_matrix_prefix_s and the "Regularization matrix (H, ConstantSplit assembly)" row.

Witness

On the A100 DelaunayNN breakdown:

  • "H, ConstantSplit assembly" drops from 10.0 ms per call at vmap 16 to under 3 ms (expected ~0.7).
  • regularization_matrix_prefix_s drops from 16.4 to under 11 ms per call (expected ~7).
  • Unbatched params→H 28.2 -> ~19 ms.
  • EXPECTED_LOG_EVIDENCE_HST = 29144.581944 unchanged at rtol 1e-4.
  • delaunay_nn.py and delaunay_nn_caps.py jax_assertions pass.
  • Full pytest test_autoarray/ green.
  • CPU no-regression (an improvement is expected).

Key Files

  • PyAutoArray/autoarray/inversion/regularization/regularization_util.py — the change (pixel_splitted_regularization_matrix_from, JAX branch).
  • PyAutoArray/test_autoarray/inversion/regularizations/test_pixel_splitted_jax.py — new JAX test leg.
  • autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py — compaction parity check.
  • autolens_profiling/hpc/batch_gpu/submit_*assembly* — new A100 submits.
  • autolens_profiling/results/notes/delaunay_nn_constant_split_assembly.md — the results note.
  • autolens_profiling/results/breakdown/imaging/delaunay_nn_hpc_a100_fp64_assembly*.json — the A/B outputs.
  • autolens_profiling/scripts/misc/delaunay_nn/assembly_bench.py — the investigation bench, made reproducible.

Related

Original Prompt

Click to expand starting prompt

DelaunayNN ConstantSplit regularization assembly: the 10 ms per call that Phase A left behind

Type: feature
Target: autoarray
Repos:

  • PyAutoArray
  • autolens_profiling
  • autolens_workspace_test
    Themes:
  • jax-gpu
  • delaunay
  • profiling
  • performance
    Difficulty: medium
    Autonomy: supervised
    Priority: high
    Status: draft
    Consequence: judge
    Witness: on the A100 DelaunayNN breakdown (results/breakdown/imaging/delaunay_nn_hpc_a100_fp64_launch_latency.json is the post-sibson: cut DelaunayNN kernel launches — gated candidate unroll, single concatenated pass, chunk as memory guard #533 baseline) the "Regularization matrix (H, ConstantSplit assembly)" row drops from 10.0 ms per call at vmap 16 to under 3 ms and the params→H prefix (regularization_matrix_prefix_s) from 16.4 ms per call to under 11 ms, with EXPECTED_LOG_EVIDENCE_HST = 29144.581944 unchanged (or, if the assembly is reformulated so the fp summation order changes, matching to a stated relative tolerance with the change justified) and the delaunay_nn.py jax_assertions passing
    Review-minutes: 40
    Unattended: ready
    Filed: 2026-09-08

Original request (verbatim):

i agree with your recommendation but its bed soon so once its a good time to stop do that too, but getting some prm done first is good!

(The recommendation agreed to: after DelaunayNN Phase A shipped as PyAutoArray#533, point the
next prompt at the ConstantSplit assembly rather than the cavity early exit.)

The measurement (A100, post PyAutoArray#533, results/notes/delaunay_nn_launch_latency.md)

Phase A cut the DelaunayNN params→H prefix from 143.9 to 28.2 ms unbatched (5.1×), but only
from 24.3 to 16.4 ms per call at vmap 16 (1.48×). The new split-Sibson breakdown stage
(autolens_profiling#227) attributes the remaining per-call cost: the data-side Sibson pass is
~6.4 ms per call, the split-side Sibson ~1 ms, and the ConstantSplit regularization
assembly ~10.0 ms per call at every chunk size and on the control
— 143× barycentric
Delaunay's equivalent H row (0.07 ms per call) and ~19 % of the 52.6 ms batched whole
likelihood. The assembly is reg_split_from fed by InterpolatorDelaunayNN._mappings_sizes_weights_split:
6,000 split-cross points, each with a 33-wide (32 neighbours + 1 spare column) Sibson stencil,
scattered into the N×N regularization matrix — ~6.5 M scatter entries per lane versus
~96 k for Delaunay's 4-wide stencil.

The planned "Phase B" cavity early exit targets only the ~6.4 ms Sibson share and is worth
~1.3 ms per call; it is deferred in favour of this.

Investigation first (one A100 session, then decide)

  1. Instrument the assembly: time reg_split_from alone under jax.jit and jit(vmap) at
    batch 16 on the production tables (N = 1500, S = 6000, width 33) and identify whether the
    cost is the scatter-add (.at[].add into N×N), the gathers over the padded stencil, or
    the hstack spare-column plumbing. Compare against the 4-wide Delaunay call on the same
    inputs to calibrate.
  2. Candidate reformulations, measured on the same session:
    • Dense matmul: build the split mapping matrix M_s (S × N, 33 non-zeros per row) as a
      dense array and form H = M_sᵀ diag(w) M_s (or the actual ConstantSplit combination) as
      one GEMM: 6000 × 1500 × 1500 ≈ 13.5 GFLOP per lane, ~0.2 TFLOP at batch 16, i.e. ~10 ms
      fp64 on an A100 — no better unless the split combination lets the GEMM shrink, so measure
      before believing.
    • Segment-sum over stencil pairs: sort the (i, j) pairs once per fit (they depend only on
      the frozen tables) and accumulate with segment_sum instead of a random scatter into N×N.
    • Stencil truncation for the regularization only: keep the 32-wide Sibson stencil for the
      data mapping but regularize the split points with their k largest weights (k = 8–12,
      renormalized). This changes the regularization scheme and the pin; it is a science
      decision to be presented, not taken.
  3. Ship the winner that keeps the pin unchanged, or present the pin-changing one with its
    evidence-tolerance argument.

Contracts

  • EXPECTED_LOG_EVIDENCE_HST in scripts/imaging/likelihood_breakdown/delaunay_nn.py stays
    unchanged for any pure-reformulation change; a pin shift is a bug unless the reformulation
    is explicitly a summation-order change, in which case the note states the tolerance.
  • Judge on regularization_matrix_prefix_s and the new "ConstantSplit assembly" row.
  • Gradient: the assembly is differentiable through the Sibson weights; any stop_gradient
    must be justified the way _jax_delaunay_tables and the walk do.
  • SIBSON_MAX_NEIGHBORS / caps unchanged unless the truncation option is chosen.

Verification on the A100

Same-node control (merge base) vs feature A/B with
scripts/imaging/likelihood_breakdown/delaunay_nn.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16 plus the runtime cell; report all rows unbatched and per call at vmap 16,
against the post-#533 baseline.

Related: complete/2026/09/delaunay-nn-launch-latency.md (Phase A record),
results/notes/delaunay_nn_launch_latency.md (numbers), the deferred cavity early-exit
(Phase B of the Phase A prompt) which stays unfiled until this lands.

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