sibson: cut DelaunayNN kernel launches — gated candidate unroll, single concatenated pass, chunk as memory guard - #533
Conversation
…le concatenated pass, chunk as memory guard Phase A of #532. On the A100 (post-#531) the DelaunayNN params->H prefix costs 144.8 ms unbatched / 24.4 ms per call at vmap 16 against Delaunay's 7.2 / 5.1 ms. The compiled HLO issues 1,244 kernel launches per 256-query chunk over 95 chunks (~118k launches at ~1.1 us each): the cost is a launch floor, not arithmetic, so this cuts launches instead of flops. - _cavity_triangle_indexes_jax: the 3-trip candidate-edge fori_loop is unrolled at trace time when _sibson_unroll_candidates() is true -- PYAUTO_SIBSON_UNROLL_CANDIDATES ("1"/"0") wins, otherwise jax.default_backend() != "cpu". Unrolling removes ~28% of the launches per chunk (1,244 -> 892) on accelerators; on CPU there is no launch cost to remove and the rolled loop measures ~8-9% faster, so CPU stays rolled. Both paths are bit-identical (same add_candidate calls, edges 0,1,2, same insertion order). - jax_delaunay_nn: the data grid and the 4N split-cross points are located and interpolated in ONE concatenated pass (walk once, Sibson once), sliced at n_query -- the pattern jax_delaunay already uses for its walk. The circumcircles depend only on the frozen simplex table, so they are hoisted and handed in through a new optional circumcircles= kwarg on sibson_mappings_weights_from_tables (default None recomputes as before; jax_sibson and scipy_delaunay_nn are unchanged). - SIBSON_QUERY_CHUNK (default still 256) is overridable at import via PYAUTO_SIBSON_QUERY_CHUNK, validated as a positive int, and picked up by DelaunayNN.query_chunk -- so the chunk can be swept on device without editing source. Documented as a memory guard on the (C, 3, 2) per-cavity intermediates, explicitly not a speed knob: it multiplies a latency-bound program. - Docstrings on the loop structure, the launch-count reasoning and the single-pass invariant. Bit-identity: all 16 jax_delaunay_nn outputs (data and split halves) match a frozen main reference after each step, on the gated path and with the env override forced to 1 and to 0. Chunk invariance proven at 64 / 256 / 1024 through both the kwarg and the env var. Tests: pytest test_autoarray -q -> 1456 passed; 4 new NumPy-only tests (circumcircles-kwarg parity, fixed-seed scipy_delaunay_nn regression, env parsing valid/invalid for both variables, unroll-gate override). black clean. Locally: jax_assertions/delaunay_nn.py exit 0 (relative_l2 9.88e-05, corr 0.9986, max_cavity 11, max_neighbors 13, flip continuity + gradients green) and delaunay_nn_caps.py exit 0. CPU no-regression: interleaved in-process paired A/B, control vs feature ratio 0.993 with the gate (0.927 with the unroll forced on). Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01B5HT8dp7sWc9qDhZp6moGr
|
Workspace PR: PyAutoLabs/autolens_workspace_test#307 Extends Library-first merge gate: the workspace PR must merge after this one — the script fails against |
The chunk sweep of #532 on the RAL A100 (NVIDIA A100 80GB PCIe, fp64, euclid-ral-gpu-2, 2026-09-08, jobs 342321/342322 + array 342323_[0-3]) on the HST / Hilbert-1500 / MGE-60 / ConstantSplit DelaunayNN imaging breakdown, --split-setup --vmap-batch 16, params->H prefix and peak sampled nvidia-smi memory: chunk | params->H unbatched | per call @vmap 16 | peak VRAM 256 | 88.29 ms | 23.07 ms | 41,495 MiB 512 | 52.21 ms | 18.24 ms | 41,503 MiB 1024 | 35.32 ms | 19.57 ms | 41,503 MiB 2048 | 27.22 ms | 19.13 ms | 41,503 MiB 4096 | 24.66 ms | 16.45 ms | 41,503 MiB Control (main, d7c9676, chunk 256): 143.90 ms / 24.32 ms / 41,495 MiB. EXPECTED_LOG_EVIDENCE_HST = 29144.581944 held on every row, so the chunk is bit-neutral as designed. 4096 is fastest on both readings. The sweep's VRAM clause turned out uninformative: the ~41.5 GiB plateau is identical at every chunk and on the control, because it is the vmap-16 dense inversion block, not the cavity intermediates -- (C, 3, 2) fp64 intermediates at ~25 kB per query per lane over 4096 queries x 16 lanes is ~1.6 GB, ~2 % of an 80 GB card. PYAUTO_SIBSON_QUERY_CHUNK stays the escape hatch for a smaller GPU or a much larger cell, and the constant's comment now says so. test_delaunay_nn.py's mesh-attribute assertion follows the new default. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01B5HT8dp7sWc9qDhZp6moGr
A100 A/B + chunk sweep — Phase A measured, default
|
| Job | Leg | Elapsed |
|---|---|---|
| 342321 | breakdown control — d7c96762 (merge base = main) |
3:00 |
| 342322 | breakdown feature — 9e9a7d50, chunk 256 |
2:49 |
| 342323_[0-3] | chunk sweep — feature, chunks 512/1024/2048/4096 | 2:26–2:37 |
| 342324 / 342325 | runtime cell, control / feature | 1:22 / 1:18 |
Both legs prepend a private PyAutoArray checkout to PYTHONPATH
(/mnt/ral/jnightin/PyAuto_wt/delaunay-nn-launch-latency/PyAutoArray_{control,feature});
the shared /mnt/ral/jnightin/PyAuto install was not touched (subhalo-validation
and Euclid DR1 CPU runs 342299/342311/342314 were live on it throughout).
backend = gpu, device = cuda:0 on every leg, so the feature took the
unrolled candidate branch.
Control vs feature at the current default (chunk 256)
| Row | control | feature | Δ | control vmap/16 | feature vmap/16 |
|---|---|---|---|---|---|
| Inversion setup (5–8 combined) | 111.552 | 88.270 | −23.282 | 20.364 | 19.239 |
| Triangulation + interpolation | 97.274 | 72.481 | −24.793 | 11.603 | 10.399 |
| Regularization matrix (H) | 45.103 | 14.725 | −30.378 | 12.622 | 12.584 |
| Interpolator prefix (params→step 6) | 98.800 | 73.561 | −25.239 | 11.695 | 10.488 |
| Split-point Sibson | 35.358 | 0.996 | −34.362 | 2.737 | 2.432 |
| H, ConstantSplit assembly | 9.745 | 13.730 | +3.985 | 9.885 | 10.152 |
| params→H prefix (the witness) | 143.903 | 88.287 | −55.616 (1.63×) | 24.317 | 23.072 |
| Total step-by-step (unbatched) | 196.357 | 142.563 | −53.794 | — | — |
Split-point Sibson collapsing 35.4 → 1.0 ms is the attribution shift the
concatenated pass was expected to cause (the split-side Sibson now runs inside
the step-6 prefix), exactly as #531 did for the locate — not a saving on its
own. Judge on regularization_matrix_prefix_s.
Chunk sweep (feature, PYAUTO_SIBSON_QUERY_CHUNK, vmap 16)
| chunk | params→H unbatched | params→H per call @vmap 16 | interp. prefix @vmap 16 | peak sampled VRAM | pin |
|---|---|---|---|---|---|
| 256 (A/B leg) | 88.29 ms | 23.07 ms | 10.49 ms | 41,495 MiB | PASS |
| 512 | 52.21 ms | 18.24 ms | 8.31 ms | 41,503 MiB | PASS |
| 1024 | 35.32 ms | 19.57 ms | 7.06 ms | 41,503 MiB | PASS |
| 2048 | 27.22 ms | 19.13 ms | 6.60 ms | 41,503 MiB | PASS |
| 4096 | 24.66 ms | 16.45 ms | 6.40 ms | 41,503 MiB | PASS |
| control (chunk 256) | 143.90 ms | 24.32 ms | 11.70 ms | 41,495 MiB | PASS |
EXPECTED_LOG_EVIDENCE_HST = 29144.581944 held exactly on every leg and every
chunk, so the chunk is bit-neutral as designed.
Decision — 4096
The plan's rule was "largest of 512/1024/2048/4096 whose peak VRAM at vmap 16
stays under ~50 % of 80 GB and which is fastest per call; if 4096 wins on both,
ship 2048 unless the margin is > 15 %."
4096 is fastest on both readings. The VRAM clause turned out uninformative:
the ~41.5 GiB plateau (50.7 % of 81,920 MiB) is identical at every chunk and
on the control leg — the GPU started each job at 0 MiB with no other tenant, so
the plateau is the vmap-16 dense inversion block, not the cavity intermediates.
The chunk contributes ≤ 8 MiB of difference across a 16× range. The guard
arithmetic says why there is room: (C, 3, 2) fp64 intermediates at ~25 kB per
query per lane over 4096 queries × 16 lanes is ~1.6 GB, ~2 % of an 80 GB card.
With the VRAM leg vacuous and 4096 faster by 14–16 % per call at vmap 16 over
2048 (and 9.4 % unbatched), 4096 ships. PYAUTO_SIBSON_QUERY_CHUNK remains
the escape hatch for a smaller GPU or a much larger cell, and the constant's
comment now records the sweep, the rule and the arithmetic.
Landed as e5d31d74 on this branch (pytest test_autoarray/inversion 499
passed; black --check clean on sibson.py;
autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py exit 0,
its env-override check printing default 4096;
test_delaunay_nn.py's query_chunk assertion updated).
Witness
- Unbatched ≤ 60 ms: MET — 24.66 ms at the shipped default, 5.83× the
143.90 ms control. - Per call at vmap 16 ≤ 11 ms: NOT met — 16.45 ms (1.48× the 24.32 ms
control).
The new "Split-point Sibson / ConstantSplit assembly" stage says why. At vmap 16
and chunk 4096 the params→H prefix decomposes as:
| piece | per call @vmap 16 |
|---|---|
| interpolator prefix (both Sibson passes + locate) | 6.40 ms |
| split-point Sibson (residual after concatenation) | ~1.0 ms |
| H, ConstantSplit assembly | ~10.0 ms |
The assembly row is ~10 ms per call at every chunk and on the control
(Delaunay's equivalent H row is 0.07 ms per call), so after Phase A roughly
60 % of the per-call prefix is the 33-wide split-stencil regularization
assembly, not the Sibson loops. The ≤ 11 ms target could not be reached by any
change to the Sibson chunking.
Whole-likelihood runtime cell
| control | feature | Δ | |
|---|---|---|---|
| Full pipeline (single JIT) | 201.320 | 136.082 | −65.238 (1.48×) |
vmap batch 16, per call |
58.022 | 54.209 | −3.813 (1.07×) |
Both legs ran at chunk 256, so these understate the shipped 4096 default.
Results, PNGs and the full note: autolens_profiling →
results/breakdown/imaging/delaunay_nn_hpc_a100_fp64_{launch_latency,launch_latency_control,chunk512,chunk1024,chunk2048,chunk4096}.{json,png},
results/runtime/imaging/delaunay_nn/delaunay_nn_hpc_a100_fp64_launch_latency{,_control}.json,
results/notes/delaunay_nn_launch_latency.md.
Re-measured at the shipped commit
|
| Row | control | shipped | Δ | control vmap/16 | shipped vmap/16 |
|---|---|---|---|---|---|
| Inversion setup (5–8 combined) | 111.552 | 27.482 | −84.071 | 20.364 | 17.674 |
| Triangulation + interpolation | 97.274 | 13.365 | −83.909 | 11.603 | 8.879 |
| Regularization matrix (H) | 45.103 | 13.768 | −31.335 | 12.622 | 7.491 |
| Interpolator prefix (params→step 6) | 98.800 | 14.444 | −84.356 | 11.695 | 8.944* |
| Split-Sibson prefix | 134.159 | 14.242 | −119.917 | 14.432 | 6.434 |
| params→H prefix (the witness) | 143.903 | 28.212 | −115.691 (5.10×) | 24.317 | 16.435 (1.48×) |
| Total step-by-step (unbatched) | 196.357 | 80.716 | −115.641 | — | — |
* first prefix measured in that run, so it carries warm-up; the split-Sibson prefix in the
same run is a strict superset of it and reads 6.434 ms, matching the sweep's 4096 row
(6.398 ms). Read the Sibson share at vmap 16 as ≈ 6.4 ms per call.
Whole-likelihood runtime cell:
| control | feature @256 | shipped @4096 | Δ (control → shipped) | |
|---|---|---|---|---|
| Full pipeline (single JIT) | 201.320 | 136.082 | 76.322 | −125.0 (2.64×) |
vmap batch 16, per call |
58.022 | 54.209 | 52.618 | −5.4 (1.10×) |
Sweep row for the shipped commit (for the table in my previous comment): chunk 4096,
params→H 28.212 ms unbatched / 16.435 ms per call at vmap 16 — within run-to-run
noise of the 342323_3 sweep row (24.664 / 16.451), which is what the default was chosen on.
Everything else in the previous comment stands: witness unbatched half MET (28.2 ms vs
≤ 60), batched half NOT met (16.44 ms vs ≤ 11) because ~10.0 ms per call of that prefix is
the chunk-independent 33-wide ConstantSplit assembly (Delaunay's equivalent H row is
0.07 ms per call), leaving ≈ 6.4 ms of actual Sibson.
Full note: autolens_profiling/results/notes/delaunay_nn_launch_latency.md.
|
Profiling PR (A100 results + the new Verdict summary on the issue: #532 (comment) — params→H 143.90 → 28.21 ms unbatched (5.10x), 24.32 → 16.44 ms per call at vmap 16; whole likelihood 201.32 → 76.32 ms single-JIT (2.64x); the 29144.581944 pin unchanged on every leg and every chunk. |
Summary
Phase A of #532 — cut the kernel-launch count of the JAX
DelaunayNN(Sibson)interpolation, which is the reason its params→H prefix is ~20× Delaunay's.
Measured on the A100 after #531: the
DelaunayNNparams→H prefix costs144.8 ms unbatched / 24.4 ms per call at vmap 16, against Delaunay's
7.2 / 5.1 ms. The compiled HLO issues 1,244 kernel launches per
256-query chunk across 95 chunks (≈118k launches at ≈1.1 µs each), so the
cost is a launch floor, not arithmetic. This PR removes launches rather than
flops:
fori_loopover acavity triangle's candidate edges in
_cavity_triangle_indexes_jaxisunrolled at trace time when
_sibson_unroll_candidates()is true —PYAUTO_SIBSON_UNROLL_CANDIDATES("1"/"0") wins, otherwisejax.default_backend() != "cpu". Unrolling removes ~28% of the launches perchunk (1,244 → 892) on accelerators; on CPU there is no launch cost to
remove and the rolled loop measures ~8–9% faster, so CPU stays rolled. The
two paths are bit-identical by construction (same
add_candidatecalls,edges 0, 1, 2, same insertion order) and verified so.
jax_delaunay_nnnow locates and interpolatesthe data grid and the
4Nsplit-cross points in a single pass — walk once,Sibson once, slice at
n_query— the patternjax_delaunayalready uses forits walk. Two passes paid two complete sets of launches. The circumcircles
depend only on the frozen simplex table, so they are hoisted out and handed
in through a new optional
circumcircles=kwarg.SIBSON_QUERY_CHUNK(default still 256) cannow be overridden at import by
PYAUTO_SIBSON_QUERY_CHUNK(positive int,validated) and is picked up by
DelaunayNN.query_chunk, so it can be swepton device without editing source. Documented as a bound on the
(C, 3, 2)per-cavity intermediates and explicitly not a speed knob — it multiplies a
latency-bound program.
single-pass invariant.
The new chunk default is not in this PR yet. An A100 sweep (512 / 1024 /
2048 / 4096 at vmap 16, with peak VRAM) runs next and lands as a follow-up
commit on this branch before merge; 256 is unchanged until that number exists.
Phase B (cavity early exit) is a separate library PR on the same issue after
the A100 numbers.
API Changes
One new optional keyword argument,
circumcircles=None, onsibson_mappings_weights_from_tables—Nonerecomputes exactly as before, soevery existing call is unaffected. Two new module-level environment overrides
(
PYAUTO_SIBSON_QUERY_CHUNK,PYAUTO_SIBSON_UNROLL_CANDIDATES) that defaultto today's behaviour when unset.
jax_delaunay_nn,jax_sibson,scipy_delaunay_nn,InterpolatorDelaunayNNandaa.mesh.DelaunayNNkeeptheir signatures and their outputs. No removals, no renames, no changed
defaults.
See full details below.
Test Plan
pytest test_autoarray -q→ 1456 passed (59s)test_autoarray/inversion/pixelization/interpolator/test_sibson.py:circumcircles=parity against the default call, a fixed-seedscipy_delaunay_nnregression digest, env parsing (valid + invalid) forboth variables, and the unroll gate's module override
jax_delaunay_nnoutputs (data and splithalves) identical to a frozen
mainreference — after the unroll, afterthe single pass, on the gated path, and with
PYAUTO_SIBSON_UNROLL_CANDIDATESforced to1and to0query_chunk64 / 256 / 1024,through both the kwarg and
PYAUTO_SIBSON_QUERY_CHUNKin freshsubprocesses
autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.pyexit 0 (parity
relative_l29.88e-05, corr 0.998595,max_cavity11,max_neighbors13, flip continuity + gradients green).../jax_assertions/delaunay_nn_caps.pyexit 0 atDELAUNAY_NN_CAP_RANDOM_SAMPLES=12("cap 32 covers this production-likeaudit")
jax_delaunay_nn(N=1500, Q=17,980 data + 6,000 split, fp64) readscontrol-vs-feature 0.993 with the gate in place (harness null bias
0.998), against 0.927 with the unroll forced on. The numba path never
enters this code (
scipy_delaunay_nnuntouched); the local breakdown pinEXPECTED_LOG_EVIDENCE_HST = 29144.581944PASSED on every run and wasnot re-pinned.
black --checkclean(
regularization_matrix_prefix_s) ≤ 60 ms unbatched / ≤ 11 ms per callat vmap 16 — judged on the prefix, not the Tri+interp or H rows, since
prefix-difference attribution shifts with the fused pass (as it did for
feat: JAX Delaunay walk — early-exit while_loop, chunk only the seed argmin #531). The chunk sweep lands with it.
Full API Changes (for automation & release notes)
Added
autoarray.inversion.mesh.interpolator.sibson.sibson_mappings_weights_from_tables(..., circumcircles=None, ...)—optional precomputed
(centres, radii_squared, valid)triple forsimplices_padded, as returned bydelaunay_circumcircles_from.Nonecomputes them internally, exactly as before.
sibson.SIBSON_UNROLL_CANDIDATES/sibson._sibson_unroll_candidates()—the candidate-edge loop strategy and its backend gate.
PYAUTO_SIBSON_QUERY_CHUNK(positive int, setsSIBSON_QUERY_CHUNKat import; invalid values raiseValueError),PYAUTO_SIBSON_UNROLL_CANDIDATES("0"/"1").Changed Behaviour
jax_delaunay_nnruns one concatenated walk + Sibson pass over[query_points, split_points]instead of two passes, slicing atn_query.Split points keep their own nearest-vertex seed and outside-hull fallback;
outputs are bit-identical.
Bit-identical to the rolled loop.
Removed / Renamed / Changed Signature
Migration
Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code
https://claude.ai/code/session_01B5HT8dp7sWc9qDhZp6moGr