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
Post-#531 on the A100 (HST / Hilbert-1500 / MGE-60 / ConstantSplit) DelaunayNN costs 197.0 ms per likelihood vs 62.0 ms for barycentric Delaunay; the whole excess is the two Sibson passes (params→H prefix 144.8 ms vs 7.2 ms unbatched; 24.4 vs 5.1 ms per call at vmap 16). The compiled HLO of sibson_mappings_weights_from_tables has three nested loops (a lax.map over SIBSON_QUERY_CHUNK = 256 → 95 sequential chunks, the 32-trip cavity fori_loop, and a 3-trip candidate fori_loop): 1,244 kernel launches per chunk, 95 % in the cavity walk, ≈118k per likelihood at ~1.1 µs each — the A100 launch floor, serialised by the chunk loop. This issue ships Phase A (bit-identical launch-count cuts) and, after measurement, Phase B (cavity early exit). Phase C (loop-free k-ring cavity) is a separate prompt.
Witness: params→H prefix (regularization_matrix_prefix_s) ≤ 60 ms unbatched and ≤ 11 ms per call at vmap 16 after Phase A (baseline 144.789 / 24.424 in results/breakdown/imaging/delaunay_nn_hpc_a100_fp64_walk_early_exit.json), pin EXPECTED_LOG_EVIDENCE_HST = 29144.581944 unchanged, delaunay_nn.py / delaunay_nn_caps.py jax_assertions pass.
Standing constraint: no slowdown for the numba CPU or JAX CPU likelihoods (NumPy Sibson path untouched; JAX CPU gated by a local before/after breakdown run).
Plan
Unroll the 3-trip candidate loop inside the cavity walk (bit-identical; launches per chunk
1,244 → 892).
Locate and interpolate data grid + split points in one concatenated Sibson pass in jax_delaunay_nn, computing the circumcircles once.
Make the query chunk the only memory guard it was meant to be: sweep it on the A100 at
vmap 16 with peak-VRAM readings and ship the largest safe value as the new default.
Add a breakdown stage that separates split-side Sibson from the 33-wide ConstantSplit
assembly, so the remaining ~27 ms residual is attributed.
Prove parity (bit-identical outputs vs main), jit/vmap round-trips and the gradient checks
in the workspace_test jax_assertions; CPU no-regression for the JAX CPU path; numba path is
untouched by construction (NumPy Sibson path scipy_delaunay_nn is not changed).
Ship the library PR; A100 A/B (control at merge base, feature) incl. the chunk sweep; ship
the profiling/workspace_test PR; then Phase B (early-exit cavity while_loop with the stop_gradient boundary) as a second library PR on the same task, measured the same way.
A. PyAutoArray — autoarray/inversion/mesh/interpolator/sibson.py
A.1 Unroll the candidate loop. In _cavity_triangle_indexes_jax (156–208), process_triangle (173) calls jax.lax.fori_loop(0, 3, add_candidate, …) (201). Replace
with a Python for edge in range(3): carry = add_candidate(edge, carry). Insertion order is
unchanged → stencil column order and fp summation order unchanged (the analyst verified
bit-identical mappings/sizes/weights/cavity_sizes/overflow/degenerate on CPU). Keep the
outer cavity fori_loop (203) as is in Phase A.
A.2 One concatenated pass in jax_delaunay_nn (629–720): the local mappings_weights_for is called at 670 (data) and 699 (split). Call it once on jnp.concatenate([query_points, split_points]) and slice all six outputs at n_query. Hoist delaunay_circumcircles_from (computed inside sibson_mappings_weights_from_tables at 451)
so it runs once — add an optional circumcircles= argument to sibson_mappings_weights_from_tables (default None → compute as now) so jax_sibson (505,
used by autolens_workspace_test/.../jax_assertions/delaunay_nn.py) and the NumPy path are
unaffected. Split points keep their own nearest-vertex fallback (the located simplex_indexes/fallback come from the single pix_indexes_delaunay_walk_from call, which
now also runs once on the concatenated queries — the #531 pattern). scipy_delaunay_nn
(546–626) untouched.
A.3 Chunk as a tunable memory guard.SIBSON_QUERY_CHUNK (36) stays the module default
and DelaunayNN.query_chunk (mesh/mesh/delaunay_nn.py:47) the mesh attribute. Add an
environment override read once at import, PYAUTO_SIBSON_QUERY_CHUNK (documented next to the
constant), so the A100 sweep can vary it without code edits. The new default is set from the
sweep result in the same PR (decision rule: largest of 512/1024/2048/4096 whose peak VRAM at
vmap 16 stays under ~50 % of the 80 GB, and which is fastest per call; if 4096 wins on both,
ship 2048 unless the margin is > 15 %, to keep headroom for larger cells). Update the SIBSON_QUERY_CHUNK comment to say what the chunk guards ((C,3,2) intermediates ≈ 15–25 kB
per query per lane) and cite the sweep.
A.4 Docstrings.sibson_mappings_weights_from_tables docstring: the loop structure, why
the chunk exists, launch-count reasoning in one paragraph. jax_delaunay_nn: the single-pass
note mirroring jax_delaunay's.
B. PyAutoArray — unit tests (NumPy only)
test_autoarray/inversion/pixelization/interpolator/test_sibson.py (and the mesh/mapper
DelaunayNN tests): all existing tests must pass unchanged. Add a NumPy-path test that sibson_mappings_weights_from_tables(..., circumcircles=precomputed) equals the default call
bit-for-bit, and that scipy_delaunay_nn output is unchanged (regression on a fixed seed:
mappings, sizes, weights, split outputs). The JAX unroll itself is exercised only on the JAX
path (workspace_test), per the no-JAX-in-unit-tests rule.
C. autolens_workspace_test — scripts/misc/jax_assertions/delaunay_nn.py
Extend (not a new script): (1) a bit-identity check against a frozen reference is not
possible in CI without main, so instead add a parity check jax_delaunay_nn (concatenated)
vs jax_sibson called separately on data and split queries — same tables, same chunk — asserting
exact equality of mappings/sizes/weights/cavity sizes for both halves; (2) a chunk-invariance
check: query_chunk 64 vs 256 vs 1024 give identical outputs (proves the chunk is a pure
memory guard and the env override works); (3) keep the existing gradient/linear-precision/
flip-continuity checks; keep the runtime under the script's current budget (reuse its SIBSON_* overrides). delaunay_nn_caps.py unchanged (caps untouched).
D. autolens_profiling — scripts/imaging/likelihood_breakdown/delaunay_nn.py
D.1 New _setup_prefix_fn stage (819) that stops right after InterpolatorDelaunayNN._mappings_sizes_weights_split (returns the split mappings/weights
alongside step-6 outputs, strict superset like stage 11), timed in --split-setup and
reported as "Split-point Sibson" so the H row becomes "ConstantSplit assembly" alone. Update
the module docstring's row table and the H-row attribution paragraph (prefix-difference
semantics; judge on regularization_matrix_prefix_s).
D.2 Honour PYAUTO_SIBSON_QUERY_CHUNK in the recorded JSON metadata (xla_flags-style
provenance field) so sweep rows are self-describing; --output-dir naming delaunay_nn_hpc_a100_fp64_chunk<N>.json.
D.3 Submit scripts: hpc/batch_gpu/submit_breakdown_imaging_delaunay_nn_a100_hst_fp64_{control,launch_latency}
(the #531 pattern: private checkouts at merge base and feature on PYTHONPATH, provenance
printed in the SLURM log) plus a chunk-sweep array script (one task per chunk, nvidia-smi
peak VRAM sampled to the log).
E. Gates before the library PR
pytest test_autoarray -q green; black --check on touched files.
Workspace_test delaunay_nn.py and delaunay_nn_caps.py pass locally against the worktree
autoarray (CPU); paste PASS lines and wall time.
CPU no-regression (standing user constraint): likelihood_breakdown/delaunay_nn.py --split-setup locally on main vs branch, 3–5 warm reps, params→H not slower, pin
unchanged; delaunay_numba.py control unchanged (NumPy Sibson path untouched; grep confirms
no NumPy-path call site changed).
Local micro-benchmark of jax_delaunay_nn (N=1500, Q=17,980 + 6,000) on CPU before/after
as a fine-grained readout (expect the unroll + single pass to show; the chunk effect is a GPU
story).
Gradient: the Phase A change keeps the fori_loop→scan lowering, so reverse-mode is
unchanged; the jax_assertions gradient checks are the proof.
F. Ship + A100 + Phase B
ship_library (Opus): PR on PyAutoArray, pending-release, issue + Mind updated.
A100 (Opus, RAL): control (merge base) vs feature delaunay_nn.py --split-setup --vmap-batch 16 at the current chunk, then the chunk sweep 512/1024/2048/4096 at vmap 16
with peak VRAM; the runtime cell; DelaunayNN pin unchanged on every leg. Pick the default
per A.3 and push it as a follow-up commit on the same library branch before the PR is
merged (re-run the unit tests + jax_assertions on that commit). Write results/notes/delaunay_nn_launch_latency.md and the JSON/PNGs.
Phase B on the same task after the numbers: cavity fori_loop (203) → lax.while_loop
with cond = (position < max_cavity_triangles) & jnp.any(position < count); overflow
detection (198, 418–423) unchanged; stop_gradient on query, circumcentres, circumradii_squared only where they enter _contains_query inside the loop; the circumcentres[safe_cavity] gather (~259) stays traced. Same gates, second library PR,
same A100 A/B. Expected ~7 ms after Phase A. Decision point for you after Phase A's numbers.
Expected after Phase A: params→H ≈ 50 ms unbatched (2.9×), ≈ 9 ms per call at vmap 16 (2.7×);
with Phase B ≈ 45 / 8. Floor with Phase C ≈ 30 / 6 (Delaunay is 7.2 / 5.1).
none; orphan worktree delaunay-area-magnification-audit still present (unregistered, warn only)
./autolens_workspace_test
main (d818cfb)
clean
none
./autolens_profiling
main (99f4b53)
clean
retire-gpu1-mig-exclusion (hpc/ submits it deletes, README, activate.sh) — disjoint → parallel-claim note as before
Task name delaunay-nn-launch-latency; branch feature/delaunay-nn-launch-latency; worktree ~/Code/PyAutoLabs-wt/delaunay-nn-launch-latency/; active.md status library-dev.
DelaunayNN (Sibson) on the A100: kill the kernel-launch latency in the cavity walk
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: the A100 DelaunayNN breakdown's params→H prefix (regularization_matrix_prefix_s, results/breakdown/imaging/delaunay_nn_hpc_a100_fp64_walk_early_exit.json is the 2026-09-07 baseline: 144.789 ms unbatched / 24.424 ms per call at vmap 16) drops to at most 60 ms unbatched and 11 ms per call at vmap 16 after Phase A, with EXPECTED_LOG_EVIDENCE_HST = 29144.581944 unchanged and the delaunay_nn.py / delaunay_nn_caps.py jax_assertions passing
Review-minutes: 40
Unattended: ready
Filed: 2026-09-07
Supersedes: draft/feature/autoarray/sibson_single_concatenated_walk.md
Original request (verbatim):
ok then prm, then HPCPullPyAuto, then go on to do the work on delaunayNN. We recently did a likelihood_breakdown of delaunay_nn, so check that out and then work out if on the A100 we can make it really fast overall
The measurement (A100, HST / Hilbert-1500 / MGE-60 / ConstantSplit, post PyAutoArray#531)
DelaunayNN costs 197.0 ms per likelihood unbatched against 62.0 ms for barycentric Delaunay,
and the whole excess sits in the two Sibson passes: the params→H prefix is 144.8 ms (NN) vs
7.2 ms (Delaunay); per call at vmap 16 it is 24.4 ms vs 5.1 ms. The four-way split charges
92 ms to "Triangulation + interpolation" (the data-side pass over 17,980 queries) and 45 ms to
the H row (the split-side pass over the 6,000 ConstantSplit points, which re-runs the whole
Sibson pipeline including a second circumcircle computation). The downstream rows are
identical to Delaunay's (blurred mapping matrix 8.36 vs 8.34 ms at vmap 16), so the 32-wide
mapper costs nothing after the setup.
The compiled HLO of sibson_mappings_weights_from_tables
(PyAutoArray/autoarray/inversion/mesh/interpolator/sibson.py) has three nested loops: the lax.map over SIBSON_QUERY_CHUNK = 256 queries (95 sequential chunks: 71 data + 24 split),
the 32-trip cavity fori_loop (MAX_CAVITY_TRIANGLES, the loop always runs to the cap), and a
3-trip candidate fori_loop inside it. That is 32 × (4 + 3 × 11) + 60 = 1,244 kernel launches
per chunk, 95 % of them in the cavity walk, ≈ 118,000 launches per likelihood. Against the
measured 144.8 ms that is ~1.1 µs per launch: the A100 launch floor. Corroboration: 16× the
lanes (vmap 16) cost 2.7× the time. The cost is dispatch latency serialised by the chunk loop,
the same disease DELAUNAY_LOCATE_CHUNK had, one level down. Observed cavity sizes on this
cell are mean 3.5, max 9 (audit maxima across 101 traced meshes: 25 main / 19 split).
Phase A: one PR, all launch-count reductions that keep bit-identical output
Unroll the 3-trip candidate fori_loop in process_triangle (sibson.py ~179–201) into a
Python for edge in range(3). Insertion order is preserved, so stencil column order and
fp summation order are unchanged; verified bit-identical on CPU. Launches per chunk
1,244 → 892.
Raise SIBSON_QUERY_CHUNK (sibson.py:36, bound at mesh/mesh/delaunay_nn.py:47). The
chunk is only a memory guard; sweep 512 / 1024 / 2048 / 4096 on the A100 at vmap 16 and
record peak VRAM (a vmap-64 OOM on this cell is already on record in delaunay_nn_hpc_a100_fp64_vmap64.json). Ship the largest value with comfortable headroom
as the new default; keep the constant overridable.
Locate and interpolate the data grid and the split points in one concatenated mappings_weights_for call in jax_delaunay_nn (sibson.py ~670 and ~699), slicing at n_query, and hoist delaunay_circumcircles_from so the circumcircles are computed once.
Split points keep their own nearest-vertex fallback. Small on its own (~1 chunk), worth it
once the chunk is large.
Add a _setup_prefix_fn stage to autolens_profiling/scripts/imaging/likelihood_breakdown/delaunay_nn.py that stops after _mappings_sizes_weights_split, so the H row separates split-side Sibson from the
33-wide ConstantSplit assembly (the ~27 ms non-chunk residual is currently unattributed).
Expected: params→H ≈ 50 ms unbatched (2.9×) and ≈ 9 ms per call at vmap 16 (2.7×).
Phase B: early-exit the cavity walk (after Phase A is measured)
Convert the cavity fori_loop (sibson.py ~203–208) to a lax.while_loop with cond = (position < cap) & any(position < count), cap kept as the safety bound, overflow
detection (sibson.py ~198, ~418–423) unchanged. Under vmap the exit is at the global max
(~25), so the win is ~20 % of the cavity part: worth ~7 ms after Phase A. Gradient contract:
today the fori_loop lowers to scan and is reverse-mode differentiable; while_loop is
not, so the float inputs to _contains_query (query, circumcentres, circumradii_squared)
must be stop_gradient-wrapped as they enter the loop (integer/bool outputs only, zero a.e.
derivative), while the circumcentres[safe_cavity] gather at ~259 that feeds the Sibson
weights stays traced. Re-run the DelaunayNN gradient checks after.
Phase C (separate prompt, not this task): loop-free cavity via a fixed k-ring gather
Replace the cavity walk with a one-shot containment test over the seed simplex's k-ring
(3-ring 21 / 4-ring 45 candidates), compact, overflow → NaN as now. Removes the loop
entirely (launches per chunk → ~80) for an estimated floor of ~30 ms unbatched / ~6 ms per
call at vmap 16, but changes candidate ordering and therefore fp summation order in the
mapping matrix, so the pin needs a relative tolerance and the cap audit must be re-run.
File after Phase B's numbers.
Contracts (do not relax)
Pin EXPECTED_LOG_EVIDENCE_HST = 29144.581944 in delaunay_nn.py passes unchanged through
Phases A and B; a shift means a mapping or summation order changed and is a bug.
SIBSON_MAX_NEIGHBORS / MAX_CAVITY_TRIANGLES stay at 32 (cap audit results/notes/delaunay_nn_cap_audit.md); any cap change re-runs autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn_caps.py.
jax_sibson (sibson.py ~505) is exercised by autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py (sibson_tables), so
it stays; apply the unroll to the shared sibson_mappings_weights_from_tables so both
entry points benefit.
Unit tests stay NumPy-only; JAX-path parity, jit/vmap and gradient checks live in the
workspace_test jax_assertions scripts.
Verification on the A100
Same-node, same-session A/B as PyAutoArray#531 (control = private checkout at the merge
base, feature = branch; shared /mnt/ral/jnightin/PyAuto untouched while science jobs run): scripts/imaging/likelihood_breakdown/delaunay_nn.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16, plus the chunk sweep for step 2 with peak-VRAM readings, and the runtime
cell. Report every row unbatched and per call at vmap 16, the params→H prefix, the new
ConstantSplit stage, and the single-JIT total, against the 2026-09-07 post-#531 baseline.
Folds in draft/feature/autoarray/sibson_single_concatenated_walk.md (filed 2026-09-07 as the
follow-up of complete/2026/09/delaunay-walk-early-exit.md); the analysis above shows the
concatenation alone is worth ~1 chunk, so it ships as step A.3 rather than on its own.
Overview
Post-#531 on the A100 (HST / Hilbert-1500 / MGE-60 / ConstantSplit) DelaunayNN costs 197.0 ms per likelihood vs 62.0 ms for barycentric Delaunay; the whole excess is the two Sibson passes (params→H prefix 144.8 ms vs 7.2 ms unbatched; 24.4 vs 5.1 ms per call at vmap 16). The compiled HLO of
sibson_mappings_weights_from_tableshas three nested loops (alax.mapoverSIBSON_QUERY_CHUNK = 256→ 95 sequential chunks, the 32-trip cavityfori_loop, and a 3-trip candidatefori_loop): 1,244 kernel launches per chunk, 95 % in the cavity walk, ≈118k per likelihood at ~1.1 µs each — the A100 launch floor, serialised by the chunk loop. This issue ships Phase A (bit-identical launch-count cuts) and, after measurement, Phase B (cavity early exit). Phase C (loop-free k-ring cavity) is a separate prompt.Witness: params→H prefix (
regularization_matrix_prefix_s) ≤ 60 ms unbatched and ≤ 11 ms per call at vmap 16 after Phase A (baseline 144.789 / 24.424 inresults/breakdown/imaging/delaunay_nn_hpc_a100_fp64_walk_early_exit.json), pinEXPECTED_LOG_EVIDENCE_HST = 29144.581944unchanged,delaunay_nn.py/delaunay_nn_caps.pyjax_assertions pass.Standing constraint: no slowdown for the numba CPU or JAX CPU likelihoods (NumPy Sibson path untouched; JAX CPU gated by a local before/after breakdown run).
Plan
1,244 → 892).
jax_delaunay_nn, computing the circumcircles once.vmap 16 with peak-VRAM readings and ship the largest safe value as the new default.
assembly, so the remaining ~27 ms residual is attributed.
in the workspace_test jax_assertions; CPU no-regression for the JAX CPU path; numba path is
untouched by construction (NumPy Sibson path
scipy_delaunay_nnis not changed).the profiling/workspace_test PR; then Phase B (early-exit cavity
while_loopwith thestop_gradientboundary) as a second library PR on the same task, measured the same way.Detailed implementation plan
Affected Repositories
A. PyAutoArray —
autoarray/inversion/mesh/interpolator/sibson.pyA.1 Unroll the candidate loop. In
_cavity_triangle_indexes_jax(156–208),process_triangle(173) callsjax.lax.fori_loop(0, 3, add_candidate, …)(201). Replacewith a Python
for edge in range(3): carry = add_candidate(edge, carry). Insertion order isunchanged → stencil column order and fp summation order unchanged (the analyst verified
bit-identical
mappings/sizes/weights/cavity_sizes/overflow/degenerateon CPU). Keep theouter cavity
fori_loop(203) as is in Phase A.A.2 One concatenated pass in
jax_delaunay_nn(629–720): the localmappings_weights_foris called at 670 (data) and 699 (split). Call it once onjnp.concatenate([query_points, split_points])and slice all six outputs atn_query. Hoistdelaunay_circumcircles_from(computed insidesibson_mappings_weights_from_tablesat 451)so it runs once — add an optional
circumcircles=argument tosibson_mappings_weights_from_tables(default None → compute as now) sojax_sibson(505,used by
autolens_workspace_test/.../jax_assertions/delaunay_nn.py) and the NumPy path areunaffected. Split points keep their own nearest-vertex fallback (the located
simplex_indexes/fallback come from the singlepix_indexes_delaunay_walk_fromcall, whichnow also runs once on the concatenated queries — the #531 pattern).
scipy_delaunay_nn(546–626) untouched.
A.3 Chunk as a tunable memory guard.
SIBSON_QUERY_CHUNK(36) stays the module defaultand
DelaunayNN.query_chunk(mesh/mesh/delaunay_nn.py:47) the mesh attribute. Add anenvironment override read once at import,
PYAUTO_SIBSON_QUERY_CHUNK(documented next to theconstant), so the A100 sweep can vary it without code edits. The new default is set from the
sweep result in the same PR (decision rule: largest of 512/1024/2048/4096 whose peak VRAM at
vmap 16 stays under ~50 % of the 80 GB, and which is fastest per call; if 4096 wins on both,
ship 2048 unless the margin is > 15 %, to keep headroom for larger cells). Update the
SIBSON_QUERY_CHUNKcomment to say what the chunk guards ((C,3,2)intermediates ≈ 15–25 kBper query per lane) and cite the sweep.
A.4 Docstrings.
sibson_mappings_weights_from_tablesdocstring: the loop structure, whythe chunk exists, launch-count reasoning in one paragraph.
jax_delaunay_nn: the single-passnote mirroring
jax_delaunay's.B. PyAutoArray — unit tests (NumPy only)
test_autoarray/inversion/pixelization/interpolator/test_sibson.py(and the mesh/mapperDelaunayNN tests): all existing tests must pass unchanged. Add a NumPy-path test that
sibson_mappings_weights_from_tables(..., circumcircles=precomputed)equals the default callbit-for-bit, and that
scipy_delaunay_nnoutput is unchanged (regression on a fixed seed:mappings, sizes, weights, split outputs). The JAX unroll itself is exercised only on the JAX
path (workspace_test), per the no-JAX-in-unit-tests rule.
C. autolens_workspace_test —
scripts/misc/jax_assertions/delaunay_nn.pyExtend (not a new script): (1) a bit-identity check against a frozen reference is not
possible in CI without main, so instead add a parity check
jax_delaunay_nn(concatenated)vs
jax_sibsoncalled separately on data and split queries — same tables, same chunk — assertingexact equality of mappings/sizes/weights/cavity sizes for both halves; (2) a chunk-invariance
check:
query_chunk64 vs 256 vs 1024 give identical outputs (proves the chunk is a purememory guard and the env override works); (3) keep the existing gradient/linear-precision/
flip-continuity checks; keep the runtime under the script's current budget (reuse its
SIBSON_*overrides).delaunay_nn_caps.pyunchanged (caps untouched).D. autolens_profiling —
scripts/imaging/likelihood_breakdown/delaunay_nn.pyD.1 New
_setup_prefix_fnstage (819) that stops right afterInterpolatorDelaunayNN._mappings_sizes_weights_split(returns the split mappings/weightsalongside step-6 outputs, strict superset like stage 11), timed in
--split-setupandreported as "Split-point Sibson" so the H row becomes "ConstantSplit assembly" alone. Update
the module docstring's row table and the H-row attribution paragraph (prefix-difference
semantics; judge on
regularization_matrix_prefix_s).D.2 Honour
PYAUTO_SIBSON_QUERY_CHUNKin the recorded JSON metadata (xla_flags-styleprovenance field) so sweep rows are self-describing;
--output-dirnamingdelaunay_nn_hpc_a100_fp64_chunk<N>.json.D.3 Submit scripts:
hpc/batch_gpu/submit_breakdown_imaging_delaunay_nn_a100_hst_fp64_{control,launch_latency}(the #531 pattern: private checkouts at merge base and feature on
PYTHONPATH, provenanceprinted in the SLURM log) plus a chunk-sweep array script (one task per chunk,
nvidia-smipeak VRAM sampled to the log).
E. Gates before the library PR
pytest test_autoarray -qgreen;black --checkon touched files.delaunay_nn.pyanddelaunay_nn_caps.pypass locally against the worktreeautoarray (CPU); paste PASS lines and wall time.
likelihood_breakdown/delaunay_nn.py --split-setuplocally on main vs branch, 3–5 warm reps, params→H not slower, pinunchanged;
delaunay_numba.pycontrol unchanged (NumPy Sibson path untouched; grep confirmsno NumPy-path call site changed).
jax_delaunay_nn(N=1500, Q=17,980 + 6,000) on CPU before/afteras a fine-grained readout (expect the unroll + single pass to show; the chunk effect is a GPU
story).
fori_loop→scanlowering, so reverse-mode isunchanged; the jax_assertions gradient checks are the proof.
F. Ship + A100 + Phase B
ship_library(Opus): PR on PyAutoArray, pending-release, issue + Mind updated.delaunay_nn.py --split-setup --vmap-batch 16at the current chunk, then the chunk sweep 512/1024/2048/4096 at vmap 16with peak VRAM; the runtime cell; DelaunayNN pin unchanged on every leg. Pick the default
per A.3 and push it as a follow-up commit on the same library branch before the PR is
merged (re-run the unit tests + jax_assertions on that commit). Write
results/notes/delaunay_nn_launch_latency.mdand the JSON/PNGs.ship_workspace(Opus): workspace_test PR (C) and profiling PR (D + results), library-firstgate.
fori_loop(203) →lax.while_loopwith
cond = (position < max_cavity_triangles) & jnp.any(position < count); overflowdetection (198, 418–423) unchanged;
stop_gradientonquery,circumcentres,circumradii_squaredonly where they enter_contains_queryinside the loop; thecircumcentres[safe_cavity]gather (~259) stays traced. Same gates, second library PR,same A100 A/B. Expected ~7 ms after Phase A. Decision point for you after Phase A's numbers.
Expected after Phase A: params→H ≈ 50 ms unbatched (2.9×), ≈ 9 ms per call at vmap 16 (2.7×);
with Phase B ≈ 45 / 8. Floor with Phase C ≈ 30 / 6 (Delaunay is 7.2 / 5.1).
Branch Survey
d7c96762, #531 merged)delaunay-area-magnification-auditstill present (unregistered, warn only)d818cfb)99f4b53)retire-gpu1-mig-exclusion(hpc/ submits it deletes, README, activate.sh) — disjoint →parallel-claimnote as beforeTask name
delaunay-nn-launch-latency; branchfeature/delaunay-nn-launch-latency; worktree~/Code/PyAutoLabs-wt/delaunay-nn-launch-latency/;active.mdstatuslibrary-dev.Suggested branch:
feature/delaunay-nn-launch-latencyKey Files
PyAutoArray/autoarray/inversion/mesh/interpolator/sibson.py—_cavity_triangle_indexes_jax,sibson_mappings_weights_from_tables,jax_delaunay_nn,SIBSON_QUERY_CHUNKPyAutoArray/autoarray/inversion/mesh/mesh/delaunay_nn.py— mesh-levelquery_chunkbindingPyAutoArray/test_autoarray/inversion/pixelization/interpolator/test_sibson.py— NumPy testsautolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py— JAX parity / gradient gateautolens_profiling/scripts/imaging/likelihood_breakdown/delaunay_nn.py— breakdown stages, pinautolens_profiling/results/notes/delaunay_walk_early_exit.md— the post-feat: JAX Delaunay walk — early-exit while_loop, chunk only the seed argmin #531 baseline and attribution caveatOriginal Prompt
Click to expand starting prompt
DelaunayNN (Sibson) on the A100: kill the kernel-launch latency in the cavity walk
Type: feature
Target: autoarray
Repos:
Themes:
Difficulty: medium
Autonomy: supervised
Priority: high
Status: draft
Consequence: judge
Witness: the A100 DelaunayNN breakdown's params→H prefix (
regularization_matrix_prefix_s,results/breakdown/imaging/delaunay_nn_hpc_a100_fp64_walk_early_exit.jsonis the 2026-09-07 baseline: 144.789 ms unbatched / 24.424 ms per call at vmap 16) drops to at most 60 ms unbatched and 11 ms per call at vmap 16 after Phase A, withEXPECTED_LOG_EVIDENCE_HST = 29144.581944unchanged and thedelaunay_nn.py/delaunay_nn_caps.pyjax_assertions passingReview-minutes: 40
Unattended: ready
Filed: 2026-09-07
Supersedes: draft/feature/autoarray/sibson_single_concatenated_walk.md
Original request (verbatim):
The measurement (A100, HST / Hilbert-1500 / MGE-60 / ConstantSplit, post PyAutoArray#531)
DelaunayNN costs 197.0 ms per likelihood unbatched against 62.0 ms for barycentric Delaunay,
and the whole excess sits in the two Sibson passes: the params→H prefix is 144.8 ms (NN) vs
7.2 ms (Delaunay); per call at vmap 16 it is 24.4 ms vs 5.1 ms. The four-way split charges
92 ms to "Triangulation + interpolation" (the data-side pass over 17,980 queries) and 45 ms to
the H row (the split-side pass over the 6,000 ConstantSplit points, which re-runs the whole
Sibson pipeline including a second circumcircle computation). The downstream rows are
identical to Delaunay's (blurred mapping matrix 8.36 vs 8.34 ms at vmap 16), so the 32-wide
mapper costs nothing after the setup.
The compiled HLO of
sibson_mappings_weights_from_tables(
PyAutoArray/autoarray/inversion/mesh/interpolator/sibson.py) has three nested loops: thelax.mapoverSIBSON_QUERY_CHUNK = 256queries (95 sequential chunks: 71 data + 24 split),the 32-trip cavity
fori_loop(MAX_CAVITY_TRIANGLES, the loop always runs to the cap), and a3-trip candidate
fori_loopinside it. That is 32 × (4 + 3 × 11) + 60 = 1,244 kernel launchesper chunk, 95 % of them in the cavity walk, ≈ 118,000 launches per likelihood. Against the
measured 144.8 ms that is ~1.1 µs per launch: the A100 launch floor. Corroboration: 16× the
lanes (vmap 16) cost 2.7× the time. The cost is dispatch latency serialised by the chunk loop,
the same disease
DELAUNAY_LOCATE_CHUNKhad, one level down. Observed cavity sizes on thiscell are mean 3.5, max 9 (audit maxima across 101 traced meshes: 25 main / 19 split).
Phase A: one PR, all launch-count reductions that keep bit-identical output
fori_loopinprocess_triangle(sibson.py~179–201) into aPython
for edge in range(3). Insertion order is preserved, so stencil column order andfp summation order are unchanged; verified bit-identical on CPU. Launches per chunk
1,244 → 892.
SIBSON_QUERY_CHUNK(sibson.py:36, bound atmesh/mesh/delaunay_nn.py:47). Thechunk is only a memory guard; sweep 512 / 1024 / 2048 / 4096 on the A100 at vmap 16 and
record peak VRAM (a vmap-64 OOM on this cell is already on record in
delaunay_nn_hpc_a100_fp64_vmap64.json). Ship the largest value with comfortable headroomas the new default; keep the constant overridable.
mappings_weights_forcall injax_delaunay_nn(sibson.py~670 and ~699), slicing atn_query, and hoistdelaunay_circumcircles_fromso the circumcircles are computed once.Split points keep their own nearest-vertex fallback. Small on its own (~1 chunk), worth it
once the chunk is large.
_setup_prefix_fnstage toautolens_profiling/scripts/imaging/likelihood_breakdown/delaunay_nn.pythat stops after_mappings_sizes_weights_split, so the H row separates split-side Sibson from the33-wide ConstantSplit assembly (the ~27 ms non-chunk residual is currently unattributed).
Expected: params→H ≈ 50 ms unbatched (2.9×) and ≈ 9 ms per call at vmap 16 (2.7×).
Phase B: early-exit the cavity walk (after Phase A is measured)
Convert the cavity
fori_loop(sibson.py~203–208) to alax.while_loopwithcond = (position < cap) & any(position < count), cap kept as the safety bound, overflowdetection (
sibson.py~198, ~418–423) unchanged. Under vmap the exit is at the global max(~25), so the win is ~20 % of the cavity part: worth ~7 ms after Phase A. Gradient contract:
today the
fori_looplowers toscanand is reverse-mode differentiable;while_loopisnot, so the float inputs to
_contains_query(query,circumcentres,circumradii_squared)must be
stop_gradient-wrapped as they enter the loop (integer/bool outputs only, zero a.e.derivative), while the
circumcentres[safe_cavity]gather at ~259 that feeds the Sibsonweights stays traced. Re-run the DelaunayNN gradient checks after.
Phase C (separate prompt, not this task): loop-free cavity via a fixed k-ring gather
Replace the cavity walk with a one-shot containment test over the seed simplex's k-ring
(3-ring 21 / 4-ring 45 candidates), compact, overflow → NaN as now. Removes the loop
entirely (launches per chunk → ~80) for an estimated floor of ~30 ms unbatched / ~6 ms per
call at vmap 16, but changes candidate ordering and therefore fp summation order in the
mapping matrix, so the pin needs a relative tolerance and the cap audit must be re-run.
File after Phase B's numbers.
Contracts (do not relax)
EXPECTED_LOG_EVIDENCE_HST = 29144.581944indelaunay_nn.pypasses unchanged throughPhases A and B; a shift means a mapping or summation order changed and is a bug.
regularization_matrix_prefix_s), not on the Tri+interp or Hrows: both are prefix differences and step 3 moves work between them exactly as feat: JAX Delaunay walk — early-exit while_loop, chunk only the seed argmin #531 did
for Delaunay (
results/notes/delaunay_walk_early_exit.md).SIBSON_MAX_NEIGHBORS/MAX_CAVITY_TRIANGLESstay at 32 (cap auditresults/notes/delaunay_nn_cap_audit.md); any cap change re-runsautolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn_caps.py.jax_sibson(sibson.py~505) is exercised byautolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py(sibson_tables), soit stays; apply the unroll to the shared
sibson_mappings_weights_from_tablesso bothentry points benefit.
workspace_test jax_assertions scripts.
Verification on the A100
Same-node, same-session A/B as PyAutoArray#531 (control = private checkout at the merge
base, feature = branch; shared
/mnt/ral/jnightin/PyAutountouched while science jobs run):scripts/imaging/likelihood_breakdown/delaunay_nn.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16, plus the chunk sweep for step 2 with peak-VRAM readings, and the runtimecell. Report every row unbatched and per call at vmap 16, the params→H prefix, the new
ConstantSplit stage, and the single-JIT total, against the 2026-09-07 post-#531 baseline.
Folds in
draft/feature/autoarray/sibson_single_concatenated_walk.md(filed 2026-09-07 as thefollow-up of
complete/2026/09/delaunay-walk-early-exit.md); the analysis above shows theconcatenation alone is worth ~1 chunk, so it ships as step A.3 rather than on its own.