You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
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
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.
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):
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.
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.
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.
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).
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.pyimportorskip 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.
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_profiling/results/notes/delaunay_nn_cap_audit.md — the tail evidence
Phase B (cavity early-exit, ~1.3 ms) stays deferred behind this; re-cost against the new numbers in the completion record.
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)
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.
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.
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.
Overview
Phase A of the DelaunayNN speed-up (PyAutoArray#533) cut the
params→Hprefix 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 isregularization_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'sK = 33(SIBSON_MAX_NEIGHBORS32 + 1 spare column).An A100 investigation (jobs 342331/342332, real HST tables) measured the actual post-
reg_split_fromstencil 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
pixel_splitted_regularization_matrix_from(JAX branch only) to the firstkc = min(K, 12)columns of each row, since every row whose post-split size is<= kcis already bit-identical there.jax.lax.top_kselects theW = 256rows with the largest split sizes, and only thehead x tail,tail x head,tail x tailblocks the main pass missed are scattered for them, so the result stays exact for the tail geometries.kcthan the budget holds, poisonHwith NaN so the sampler discards the sample rather than silently accepting a wrong matrix.reg_split_from, and thehstackspare-column plumbing insibson.pyuntouched — the measurement says they are noise.autolens_workspace_testjax_assertions script.autolens_profilingresults 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
Branch Survey
Worktree claims:
PyAutoArrayandautolens_workspace_testare unclaimed inactive.md.autolens_profilingis claimed in parallel byretire-gpu1-mig-exclusion(awaiting-merge, 88 MIG-exclusion files) andinterferometer-preload-cpu(no commits yet, interferometer preload scope) — file sets are disjoint from this task's (newsubmit_*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 inactive.md.~/Code/PyAutoLabs-wt/delaunay-area-magnification-audit/PyAutoArrayis an unregistered worktree from a different task — left alone.Suggested branch:
feature/delaunay-nn-constant-split-assemblyWorktree root:
~/Code/PyAutoLabs-wt/delaunay-nn-constant-split-assembly/Investigation numbers (A100 jobs 342331/342332, real HST tables, fp64)
Production cell:
N = 1500,S = 6000split points, padded widthK = 33. Post-reg_split_fromstencil size min 1 / median 5 / p99 9 / max 11; 6,534,000 scatter entries carry 187,242 real contributions into 29,020 cells.B^T diag(s) Bsegment_sumCPU (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.mdsaw 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; thexp is nppath is untouched):kccolumns of every row,kc = min(K, SPLIT_REG_COMPACT_WIDTH)(default 12). Cost4P * 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 existingvalidmask.W = min(4P, SPLIT_REG_WIDE_ROW_BUDGET)rows with the largestsplitted_sizesviajax.lax.top_k, gather their fullK-wide rows, and scatter only the blocks the main pass did not cover —head x tail,tail x head,tail x tail(columns>= kc). CostW * (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<= kccontribute exact zeros.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 wrongH. Document that the budget is a soft cap tuned from the audit, with the numbers above.K <= kccollapses to today. Delaunay'sK = 4and the adapt-split family take the single scatter with no supplement — no change for those callers beyond a trivially-false guard. Constants live inregularization_util.pyas 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).reg_split_fromand thehstackspare-column plumbing insibson.pyalone — 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.581944is 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 thetest_adapt_power_jax.pyimportorskippattern: synthetic(4P, K)tables with (a) all rows<= kc, (b) a few rows> kcinside the budget, (c) more wide rows than the budget -> NaN; each compared against the NumPypixel_splitted_regularization_matrix_np_from. Runpytest test_autoarray/inversion/regularizations test_autoarray/inversion/mesh, then fullpytest test_autoarray/.Workspace_test. Extend
autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.pywith a compaction parity check (JAX ConstantSplitHon synthetic production-size tables == NumPy, with a forced wide row) and re-run it plusdelaunay_nn_caps.py. Also record, per audit geometry, the count of split rows abovekc— this is the evidence thatW = 256has margin; if the worst geometry exceeds ~W/4, raise the default before shipping.A100 A/B protocol (same as #531 / #533)
/mnt/ral/jnightin/PyAuto_wt/delaunay-nn-constant-split-assembly/; the shared RAL install is untouched and reached only viaPYTHONPATH.scripts/imaging/likelihood_breakdown/delaunay_nn.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16, plus the runtime cell.submit_breakdown_imaging_delaunay_nn_a100_hst_fp64_{assembly_control,assembly}and the runtime twins.autolens_profiling/results/notes/delaunay_nn_constant_split_assembly.md; breakdown JSON underresults/breakdown/imaging/delaunay_nn_hpc_a100_fp64_assembly*.json.assembly_bench.pyand its real-table builder move intoautolens_profiling/scripts/misc/delaunay_nn/so the numbers above are reproducible.regularization_matrix_prefix_sand the "Regularization matrix (H, ConstantSplit assembly)" row.Witness
On the A100 DelaunayNN breakdown:
regularization_matrix_prefix_sdrops from 16.4 to under 11 ms per call (expected ~7).params→H28.2 -> ~19 ms.EXPECTED_LOG_EVIDENCE_HST = 29144.581944unchanged at rtol 1e-4.delaunay_nn.pyanddelaunay_nn_caps.pyjax_assertions pass.pytest test_autoarray/green.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
autolens_profiling/results/notes/delaunay_nn_launch_latency.md— the post-sibson: cut DelaunayNN kernel launches — gated candidate unroll, single concatenated pass, chunk as memory guard #533 baselineautolens_profiling/results/notes/delaunay_nn_cap_audit.md— the tail evidenceOriginal 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:
Themes:
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.jsonis 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, withEXPECTED_LOG_EVIDENCE_HST = 29144.581944unchanged (or, if the assembly is reformulated so the fp summation order changes, matching to a stated relative tolerance with the change justified) and thedelaunay_nn.pyjax_assertions passingReview-minutes: 40
Unattended: ready
Filed: 2026-09-08
Original request (verbatim):
(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_fromfed byInterpolatorDelaunayNN._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)
reg_split_fromalone underjax.jitandjit(vmap)atbatch 16 on the production tables (N = 1500, S = 6000, width 33) and identify whether the
cost is the scatter-add (
.at[].addinto N×N), the gathers over the padded stencil, orthe
hstackspare-column plumbing. Compare against the 4-wide Delaunay call on the sameinputs to calibrate.
M_s(S × N, 33 non-zeros per row) as adense array and form
H = M_sᵀ diag(w) M_s(or the actual ConstantSplit combination) asone 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.
the frozen tables) and accumulate with
segment_suminstead of a random scatter into N×N.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.
evidence-tolerance argument.
Contracts
EXPECTED_LOG_EVIDENCE_HSTinscripts/imaging/likelihood_breakdown/delaunay_nn.pystaysunchanged 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.
regularization_matrix_prefix_sand the new "ConstantSplit assembly" row.stop_gradientmust be justified the way
_jax_delaunay_tablesand 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 16plus 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.