Skip to content

sibson: cut DelaunayNN kernel launches — gated candidate unroll, single concatenated pass, chunk as memory guard - #533

Merged
Jammy2211 merged 2 commits into
mainfrom
feature/delaunay-nn-launch-latency
Sep 8, 2026
Merged

Jammy2211 merged 2 commits into
mainfrom
feature/delaunay-nn-launch-latency

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

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 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 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:

  • Candidate-edge unroll (backend-gated). The 3-trip fori_loop over a
    cavity triangle's candidate edges in _cavity_triangle_indexes_jax 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. The
    two paths are bit-identical by construction (same add_candidate calls,
    edges 0, 1, 2, same insertion order) and verified so.
  • One concatenated pass. jax_delaunay_nn now locates and interpolates
    the data grid and the 4N split-cross points in a single pass — walk once,
    Sibson once, slice at n_query — the pattern jax_delaunay already uses for
    its 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.
  • The chunk is a memory guard. SIBSON_QUERY_CHUNK (default still 256) can
    now be overridden at import by PYAUTO_SIBSON_QUERY_CHUNK (positive int,
    validated) and is picked up by DelaunayNN.query_chunk, so it can be swept
    on 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.
  • Docstrings covering the loop structure, the launch-count reasoning and the
    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, on
sibson_mappings_weights_from_tables — None recomputes exactly as before, so
every existing call is unaffected. Two new module-level environment overrides
(PYAUTO_SIBSON_QUERY_CHUNK, PYAUTO_SIBSON_UNROLL_CANDIDATES) that default
to today's behaviour when unset. jax_delaunay_nn, jax_sibson,
scipy_delaunay_nn, InterpolatorDelaunayNN and aa.mesh.DelaunayNN keep
their signatures and their outputs. No removals, no renames, no changed
defaults.
See full details below.

Test Plan

  • pytest test_autoarray -q → 1456 passed (59s)
  • 4 new NumPy-only tests in
    test_autoarray/inversion/pixelization/interpolator/test_sibson.py:
    circumcircles= parity against the default call, a fixed-seed
    scipy_delaunay_nn regression digest, env parsing (valid + invalid) for
    both variables, and the unroll gate's module override
  • Bit-identity: all 16 jax_delaunay_nn outputs (data and split
    halves) identical to a frozen main reference — after the unroll, after
    the single pass, on the gated path, and with
    PYAUTO_SIBSON_UNROLL_CANDIDATES forced to 1 and to 0
  • Chunk invariance: outputs identical at query_chunk 64 / 256 / 1024,
    through both the kwarg and PYAUTO_SIBSON_QUERY_CHUNK in fresh
    subprocesses
  • autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py
    exit 0 (parity relative_l2 9.88e-05, corr 0.998595, max_cavity 11,
    max_neighbors 13, flip continuity + gradients green)
  • .../jax_assertions/delaunay_nn_caps.py exit 0 at
    DELAUNAY_NN_CAP_RANDOM_SAMPLES=12 ("cap 32 covers this production-like
    audit")
  • CPU no-regression: interleaved in-process paired A/B of
    jax_delaunay_nn (N=1500, Q=17,980 data + 6,000 split, fp64) reads
    control-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_nn untouched); the local breakdown pin
    EXPECTED_LOG_EVIDENCE_HST = 29144.581944 PASSED on every run and was
    not re-pinned.
  • black --check clean
  • A100 leg (pre-merge): params→H prefix
    (regularization_matrix_prefix_s) ≤ 60 ms unbatched / ≤ 11 ms per call
    at 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 for
    simplices_padded, as returned by delaunay_circumcircles_from. None
    computes them internally, exactly as before.
  • sibson.SIBSON_UNROLL_CANDIDATES / sibson._sibson_unroll_candidates() —
    the candidate-edge loop strategy and its backend gate.
  • Environment overrides: PYAUTO_SIBSON_QUERY_CHUNK (positive int, sets
    SIBSON_QUERY_CHUNK at import; invalid values raise ValueError),
    PYAUTO_SIBSON_UNROLL_CANDIDATES ("0"/"1").

Changed Behaviour

  • jax_delaunay_nn runs one concatenated walk + Sibson pass over
    [query_points, split_points] instead of two passes, slicing at n_query.
    Split points keep their own nearest-vertex seed and outside-hull fallback;
    outputs are bit-identical.
  • On non-CPU backends the cavity walk's candidate-edge loop is unrolled.
    Bit-identical to the rolled loop.

Removed / Renamed / Changed Signature

  • None.

Migration

  • None required.

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

https://claude.ai/code/session_01B5HT8dp7sWc9qDhZp6moGr

…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
@Jammy2211

Copy link
Copy Markdown
Collaborator Author

Workspace PR: PyAutoLabs/autolens_workspace_test#307

Extends scripts/misc/jax_assertions/delaunay_nn.py with the checks that pin this PR's behaviour: single-pass parity (concatenated jax_delaunay_nn vs jax_sibson run separately on the data grid and the split points, both halves exact, NaN-aware), query_chunk invariance at 64/256/1024, and a subprocess check that PYAUTO_SIBSON_QUERY_CHUNK=64 reaches DelaunayNN.query_chunk.

Library-first merge gate: the workspace PR must merge after this one — the script fails against main autoarray until this lands.

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
@Jammy2211

Copy link
Copy Markdown
Collaborator Author

A100 A/B + chunk sweep — Phase A measured, default SIBSON_QUERY_CHUNK set to 4096

RAL euclid-ral-gpu-2, NVIDIA A100 80GB PCIe, fp64 (JAX_ENABLE_X64=True),
xla_flags = --xla_disable_hlo_passes=constant_folding --xla_gpu_autotune_level=0.
All eight job units ran on the same node inside one 20-minute window
(02:16–02:36 UTC+1, 2026-09-08), so this is a same-session A/B, not a
comparison against a recorded baseline. Cell: likelihood_breakdown/imaging/delaunay_nn
HST / Hilbert-1500 / MGE-60 / ConstantSplit, --split-setup --vmap-batch 16.

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.

@Jammy2211

Copy link
Copy Markdown
Collaborator Author

Re-measured at the shipped commit e5d31d74 (default chunk 4096)

Same node, same window, same everything — euclid-ral-gpu-2, jobs 342329 (breakdown,
2:41) and 342330 (runtime, 1:25), PyAutoArray(AB) = e5d31d74, log line
Sibson query chunk: 4096 (PYAUTO_SIBSON_QUERY_CHUNK=unset). The chunk-256 feature rows
from 342322/342325 are kept as ..._launch_latency_chunk256.{json,png}; they are the
same-chunk A/B that isolates the unroll + single concatenated pass from the chunk
change. Pin 29144.581944 PASSED on both new legs.

Headline is now control (d7c96762, chunk 256) vs shipped (e5d31d74, chunk 4096):

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.

@Jammy2211

Copy link
Copy Markdown
Collaborator Author

Profiling PR (A100 results + the new --split-setup split-Sibson stage): PyAutoLabs/autolens_profiling#227

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.

@Jammy2211
Jammy2211 merged commit bc113fb into main Sep 8, 2026
3 checks passed
@Jammy2211
Jammy2211 deleted the feature/delaunay-nn-launch-latency branch September 8, 2026 02:00
@Jammy2211 Jammy2211 removed the pending-release PR queued for the next release build label Sep 26, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant