Overview
On the A100 the HST Delaunay imaging likelihood (1500 Hilbert vertices, 15,361 data pixels, MGE-60 lens, ConstantSplit) costs 97 ms per evaluation, of which the --split-setup decomposition attributes 26.6 ms to "Triangulation + interpolation". Almost all of that is latency in the JAX point locator pix_indexes_delaunay_walk_from (autoarray/inversion/mesh/interpolator/delaunay.py): a fixed-trip fori_loop of 128 steps that cannot exit early (typical resolution < 10 steps), run over 1024-query chunks through a sequential lax.map (~22 chunks x 128 = ~2,800 dependent steps per likelihood, since jax_delaunay locates the data grid and the 6000 split-cross points separately). The chunking exists only to bound the (chunk, N) nearest-vertex argmin intermediate under vmap; it happens to serialise the latency-bound walk as well.
This issue is Phase 1 of the prompt: early-exit while_loop, chunk only the seed argmin, one walk call per likelihood, then re-measure on the A100. Phase 2 (static image-plane seed + one-shot fan test, Mapper/AdaptImages plumbing) is filed separately only if Phase 1 leaves the row above a few ms.
Witness: the "Triangulation + interpolation" row of autolens_profiling/results/breakdown/imaging/delaunay_hpc_a100_fp64.json (--split-setup) drops from 26.6 ms to under 8 ms unbatched, EXPECTED_LOG_EVIDENCE_HST unchanged, walk parity tests pass, FD certification passes.
Constraint (user): no slowdown or behaviour change for the numba CPU likelihood or the JAX CPU likelihood. The numba/NumPy path goes through scipy_delaunay (find_simplex + KDTree) and never enters the walk, so it is untouched by construction; the JAX CPU path is gated by a local before/after breakdown run.
Plan
- Replace the fixed-trip
fori_loop in the JAX walk with a while_loop that exits as soon as every query is located or outside the hull; keep the 128-step cap as the safety bound.
- Split the chunk loop's two jobs: keep only the nearest-vertex argmin chunked (memory guard), then run the walk once over all queries with no
lax.map.
- Locate data-grid and split-cross points in a single walk call in
jax_delaunay (concatenate, locate, slice) so one while_loop runs per likelihood instead of two.
- Freeze the walk's float inputs with
stop_gradient (point location is integer-valued; while_loop has no reverse-mode rule); barycentric weights are still recomputed from traced points, so jax.grad through the likelihood is unchanged.
- Keep the NumPy path (unit tests) behaviourally identical; extend the NumPy unit tests to the refactored seed/walk split.
- Prove JAX parity against
scipy.spatial.Delaunay.find_simplex in a workspace_test jax_assertions script (existing meshes + a 1500-vertex Hilbert mesh traced through a few mass models), plus jit and vmap round-trips; re-run the FD certification.
- Ship the library PR; re-measure on the A100 from a RAL worktree, record the four-way split, ship the profiling/workspace_test PR. Decide Phase 2 from the numbers.
Detailed implementation plan
Affected Repositories
- PyAutoArray (primary, library)
- autolens_workspace_test (JAX parity script — unit tests stay NumPy-only)
- autolens_profiling (A100 re-measurement results + notes)
Branch Survey
| Repository |
Current Branch |
Dirty? |
| ./PyAutoArray |
main |
clean (orphan worktree delaunay-area-magnification-audit on a feature branch — unregistered, not a conflict) |
| ./autolens_workspace_test |
main |
clean |
| ./autolens_profiling |
main |
clean (claimed by retire-gpu1-mig-exclusion, disjoint file set — parallel claim) |
Suggested branch: feature/delaunay-walk-early-exit
Implementation Steps
A. PyAutoArray — autoarray/inversion/mesh/interpolator/delaunay.py
Refactor pix_indexes_delaunay_walk_from into three pieces, public signature unchanged (callers: sibson.py x2 with return_simplex_indexes=True, test_delaunay_walk.py, autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py):
_nearest_vertex_seed_from(query_points, points, xp) → (Q,) int32. NumPy: one brute-force argmin as today. JAX: pad Q to a multiple of DELAUNAY_LOCATE_CHUNK, jax.lax.map over (-1, chunk, 2) computing only the per-chunk argmin, reshape, slice to Q. Pad rows at 1e9 as now. Sole survivor of the chunking; the (chunk, N) intermediate stays bounded under vmap.
_walk_from_seed(query_points, seed, points, simplices_padded, simplex_neighbors, vertex_simplex, xp) → (cur, done, outside) over all Q at once. walk_step / weights_of unchanged in content. NumPy: existing for … if (done | outside).all(): break. JAX: jax.lax.while_loop(cond, body, (step, cur, done, outside)) with cond = (step < DELAUNAY_WALK_STEPS) & jnp.any(~done & ~outside). Under vmap the loop runs to the slowest lane.
- Wrapper builds
mappings / simplex_indexes from (seed, cur, done) exactly as the current locate_chunk tail does, honouring return_simplex_indexes.
Gradient contract: lax.while_loop is not reverse-mode differentiable, so the JAX branch wraps query_points and points in jax.lax.stop_gradient before the seed argmin and the walk (the connectivity tables already are, see _jax_delaunay_tables). Location is integer-valued with zero a.e. derivative; every differentiable quantity (barycentric weights via pixel_weights_delaunay_from, dual areas, split points, Sibson weights) is computed downstream from the traced arrays. Document in the docstring; update the DELAUNAY_LOCATE_CHUNK / DELAUNAY_WALK_STEPS comments to the new roles.
jax_delaunay: locate jnp.concatenate([query_points, split_points]) in one call and slice [:Q] / [Q:]; split points still seed from their own nearest vertex. scipy_delaunay (NumPy path) untouched; sibson.py needs no change.
B. PyAutoArray unit tests (NumPy only, never import jax) — test_autoarray/inversion/pixelization/interpolator/test_delaunay_walk.py: existing _assert_matches_find_simplex tests unchanged; add _nearest_vertex_seed_from vs cKDTree on the blob-ring mesh (ties excluded); _walk_from_seed from a deliberately far seed (vertex 0 for every query) still reaches the find_simplex answer within the cap; wrapper returns (Q, 3) int32 and (Q,) int32 with return_simplex_indexes=True on an odd Q.
C. autolens_workspace_test — scripts/misc/jax_assertions/delaunay_walk.py (sibling of delaunay_nn.py, same __Env__ block and ENV: jax full_datasets declaration, bare """ opener; CI runs every script): parity vs find_simplex + cKDTree fallback on uniform / blob-ring meshes (reuse adaptive_mesh) and on a 1500-vertex Hilbert mesh traced through a handful of the mass models from delaunay_nn_caps.py (lives in this repo, not autolens_profiling), data grid and 4N split points, comparing mapping matrices not triangle ids on shared-edge ties; jax.jit equals eager; jitted vmap over a small batch equals per-member; odd Q; jax.grad of a scalar built from the barycentric weights runs and matches FD; print warm per-call time.
D. FD certification + CPU no-regression gate (before the library PR)
- Numba/NumPy path: grep confirms no NumPy-path call site is added; run
autolens_profiling/scripts/imaging/likelihood_breakdown/delaunay_numba.py locally before/after as a control (same numbers to noise, same pin).
- JAX CPU path: run
autolens_profiling/scripts/imaging/likelihood_breakdown/delaunay.py --split-setup locally on main and on the feature branch, warm, three repeats each. Gate: "Triangulation + interpolation" row and total per-eval time not slower; EXPECTED_LOG_EVIDENCE_HST unchanged. Both tables go in the PR.
autolens_workspace_test/scripts/imaging/jax_grad/delaunay.py must pass.
- Scratchpad micro-benchmark of the locator alone (N=1500, Q=15,361 + 6000 split points) on CPU, before vs after.
- Read the cloud#34099198772 failure for
autolens_test scripts/imaging/delaunay.py and state whether it touches the walk.
E. Ship + A100 verification
ship_library: tests, commit, push, PR on PyAutoArray (pending-release).
- From a RAL worktree with
HPCPullPyAuto on the feature branch, run scripts/imaging/likelihood_breakdown/delaunay.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16 and the delaunay_nn.py sibling. Pins must pass unchanged — a shift is a bug, not a re-pin (PyAutoLens#721 knife-edge lesson: bisect before re-pinning anything).
- Results JSON/PNG under
autolens_profiling/results/breakdown/imaging/ and results/notes/delaunay_walk_early_exit.md with the four-way split unbatched and per call under vmap, the H row and the single-JIT cell against the 2026-09 baselines.
ship_workspace for autolens_workspace_test + autolens_profiling behind the library-first merge gate. Decide whether to file the Phase 2 prompt.
Key Files
PyAutoArray/autoarray/inversion/mesh/interpolator/delaunay.py — pix_indexes_delaunay_walk_from, jax_delaunay, _jax_delaunay_tables, the two module constants
PyAutoArray/autoarray/inversion/mesh/interpolator/sibson.py — DelaunayNN caller (unchanged; inherits the gain)
PyAutoArray/test_autoarray/inversion/pixelization/interpolator/test_delaunay_walk.py — NumPy parity tests
autolens_workspace_test/scripts/misc/jax_assertions/delaunay_walk.py — new JAX parity/jit/vmap/grad script
autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn_caps.py — lensing geometry to reuse
autolens_workspace_test/scripts/imaging/jax_grad/delaunay.py — FD certification
autolens_profiling/scripts/imaging/likelihood_breakdown/delaunay.py, delaunay_nn.py, delaunay_numba.py — breakdown / control scripts
autolens_profiling/results/notes/preopt_breakdown_baseline.md — the 26.6 ms baseline
Original Prompt
Click to expand starting prompt
JAX Delaunay point location: early-exit walk, unchunked loop, static image-plane seed
Type: feature
Target: autoarray
Repos:
- PyAutoArray
- autolens_profiling
Themes:
- jax-gpu
- delaunay
- profiling
- performance
Difficulty: medium
Autonomy: supervised
Priority: high
Status: draft
Consequence: judge
Witness: the A100 Delaunay breakdown's "Triangulation + interpolation" row (26.6 ms, results/breakdown/imaging/delaunay_hpc_a100_fp64.json, --split-setup) drops to under 8 ms unbatched with EXPECTED_LOG_EVIDENCE_HST unchanged, the walk parity tests pass, and the FD certification still passes
Review-minutes: 30
Unattended: ready
Filed: 2026-09-05
Original request (verbatim):
give me a prompt to work on this Target 1: point location is a fixed-trip loop over chunks
The measurement
On the A100 the HST Delaunay imaging likelihood (1500 Hilbert vertices, 15,361 data
pixels, MGE-60 lens, ConstantSplit) costs 97 ms per evaluation, of which the four-way
--split-setup decomposition attributes 26.6 ms to "Triangulation + interpolation"
(autolens_profiling/results/notes/preopt_breakdown_baseline.md). The qhull host callback
is a few ms of that at most (1500 points, pure_callback, vmap_method="sequential"). The
rest is the JAX-side point location in
PyAutoArray/autoarray/inversion/mesh/interpolator/delaunay.py,
pix_indexes_delaunay_walk_from, and it is latency, not FLOPs:
- The walk is
jax.lax.fori_loop(0, DELAUNAY_WALK_STEPS, ...) with
DELAUNAY_WALK_STEPS = 128 (module constant at line 75, loop at line 252). A fixed-trip
fori_loop cannot exit early. Seeded from the nearest vertex the walk resolves in a
handful of steps, so well over 95% of the 128 iterations are no-ops that still launch a
gather, a cross-product batch, an argmin and three wheres each.
- Queries are processed in chunks of
DELAUNAY_LOCATE_CHUNK = 1024 (line 69) through
jax.lax.map (line 279), which is sequential: 16 chunks for the data grid and 6 more for
the 6000 ConstantSplit split-cross points (jax_delaunay calls the walk twice). That is
roughly 22 x 128 = 2,800 dependent walk steps per likelihood.
- The chunking exists only to bound the
(chunk, N) nearest-vertex distance intermediate
under vmap (184 MB per replica at full Q; ~12 GB at batch 64). It is a memory guard
that happens to serialise the latency-bound part.
Nautilus at n_batch=256 gave 51 ms per eval against 62 ms at n_batch=16, so batching
barely amortises this today. DelaunayNN (interpolator/sibson.py) seeds its cavity walk
from pix_indexes_delaunay_walk_from(..., return_simplex_indexes=True), so it inherits
every gain here. interpolator/knn.py is a separate approach and out of scope.
Phase 1: early exit, and chunk only the argmin
- Replace the
fori_loop with a jax.lax.while_loop whose predicate is
(step < DELAUNAY_WALK_STEPS) & jnp.any(~done & ~outside). The 128 cap stays as the
safety bound; the typical exit is under ten steps. Under vmap the loop runs to the
slowest lane, which is fine.
- Split the two jobs the chunk loop currently does. Keep the nearest-vertex argmin
chunked (it is the memory hazard, and a lax.map over 16 single-argmin chunks is cheap),
then run the walk once over all Q queries with no lax.map. The walk's working set is
O(Q), not O(Q x N).
- Keep the NumPy path (
xp is np, used by the unit tests) behaviourally identical; it
already early-exits.
- Prove parity: the JAX walk must return the same
(Q, 3) mappings and simplex indexes as
scipy.spatial.Delaunay.find_simplex on the existing test meshes and on a 1500-vertex
Hilbert mesh traced through a few mass models (reuse the geometry from
autolens_profiling/scripts/misc/jax_assertions/delaunay_nn_caps.py). Points exactly on
a shared edge may legitimately resolve to either adjacent triangle; the barycentric
weights agree because the opposite vertex gets weight zero, so compare mapping matrices,
not triangle ids, in that case.
Phase 2: static image-plane seed and one-shot fan test
The image-plane mesh is fixed per fit by the adapt image, and ray tracing is continuous
away from critical curves, so each data pixel's nearest image-plane mesh vertex is a
near-perfect source-plane seed. Precomputing that index once per fit (in the Mapper /
AdaptImages layer, where image_plane_mesh_grid is known) removes the brute-force
(Q, N) argmin and its memory hazard entirely, which also retires the chunking. Then have
the qhull callback also return a padded vertex-to-incident-triangle table (cap around 12,
audit the actual max the way delaunay_nn_cap_audit.md did) so most queries resolve with
one vectorised barycentric test over the seed's fan; the while_loop walk from Phase 1
handles only the residual near critical curves and outside the hull. Split-cross points
have no image-plane parent pixel; seed them at their parent vertex (they are offsets from
it) and let the fan test cover them.
Phase 2 changes the interpolator's inputs (a seed array), so it touches
InterpolatorDelaunay, InterpolatorDelaunayNN, the Mapper constructor path and
FitImaging plumbing. Land Phase 1 first and re-measure; Phase 2 is only worth its
plumbing if Phase 1 leaves the row above a few ms.
Gradient contract (do not relax)
Point location is integer-valued; its derivative is zero almost everywhere and it is
already wrapped by stop_gradient semantics through the frozen connectivity tables
(_jax_delaunay_tables docstring). The barycentric weights are recomputed from the traced
points after location exactly as now, so jax.grad through the likelihood is unchanged.
Re-run the FD certification autolens_workspace_test/scripts/imaging/jax_grad/delaunay.py
after each phase.
Verification on the A100
Use the delaunay-nn-breakdown tooling (autolens_profiling#219): rerun
scripts/imaging/likelihood_breakdown/delaunay.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16 and the delaunay_nn.py sibling from a RAL worktree with
HPCPullPyAuto pointed at the feature branch. Report the four-way split unbatched and per
call under vmap, the H row, and the single-JIT runtime cell, against the 2026-09 baselines.
Pins (EXPECTED_LOG_EVIDENCE_HST in both scripts) must pass unchanged; a shift means a
mapping changed and is a bug, not a re-pin. Note the symmetric knife-edge lesson from
PyAutoLens#721: if a positions-threshold or point-solver pin elsewhere moves, bisect
before re-pinning.
Overview
On the A100 the HST Delaunay imaging likelihood (1500 Hilbert vertices, 15,361 data pixels, MGE-60 lens, ConstantSplit) costs 97 ms per evaluation, of which the
--split-setupdecomposition attributes 26.6 ms to "Triangulation + interpolation". Almost all of that is latency in the JAX point locatorpix_indexes_delaunay_walk_from(autoarray/inversion/mesh/interpolator/delaunay.py): a fixed-tripfori_loopof 128 steps that cannot exit early (typical resolution < 10 steps), run over 1024-query chunks through a sequentiallax.map(~22 chunks x 128 = ~2,800 dependent steps per likelihood, sincejax_delaunaylocates the data grid and the 6000 split-cross points separately). The chunking exists only to bound the(chunk, N)nearest-vertex argmin intermediate under vmap; it happens to serialise the latency-bound walk as well.This issue is Phase 1 of the prompt: early-exit
while_loop, chunk only the seed argmin, one walk call per likelihood, then re-measure on the A100. Phase 2 (static image-plane seed + one-shot fan test, Mapper/AdaptImages plumbing) is filed separately only if Phase 1 leaves the row above a few ms.Witness: the "Triangulation + interpolation" row of
autolens_profiling/results/breakdown/imaging/delaunay_hpc_a100_fp64.json(--split-setup) drops from 26.6 ms to under 8 ms unbatched,EXPECTED_LOG_EVIDENCE_HSTunchanged, walk parity tests pass, FD certification passes.Constraint (user): no slowdown or behaviour change for the numba CPU likelihood or the JAX CPU likelihood. The numba/NumPy path goes through
scipy_delaunay(find_simplex+ KDTree) and never enters the walk, so it is untouched by construction; the JAX CPU path is gated by a local before/after breakdown run.Plan
fori_loopin the JAX walk with awhile_loopthat exits as soon as every query is located or outside the hull; keep the 128-step cap as the safety bound.lax.map.jax_delaunay(concatenate, locate, slice) so onewhile_loopruns per likelihood instead of two.stop_gradient(point location is integer-valued;while_loophas no reverse-mode rule); barycentric weights are still recomputed from traced points, sojax.gradthrough the likelihood is unchanged.scipy.spatial.Delaunay.find_simplexin a workspace_testjax_assertionsscript (existing meshes + a 1500-vertex Hilbert mesh traced through a few mass models), plus jit and vmap round-trips; re-run the FD certification.Detailed implementation plan
Affected Repositories
Branch Survey
delaunay-area-magnification-auditon a feature branch — unregistered, not a conflict)retire-gpu1-mig-exclusion, disjoint file set — parallel claim)Suggested branch:
feature/delaunay-walk-early-exitImplementation Steps
A. PyAutoArray —
autoarray/inversion/mesh/interpolator/delaunay.pyRefactor
pix_indexes_delaunay_walk_frominto three pieces, public signature unchanged (callers:sibson.pyx2 withreturn_simplex_indexes=True,test_delaunay_walk.py,autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py):_nearest_vertex_seed_from(query_points, points, xp)→(Q,)int32. NumPy: one brute-force argmin as today. JAX: padQto a multiple ofDELAUNAY_LOCATE_CHUNK,jax.lax.mapover(-1, chunk, 2)computing only the per-chunkargmin, reshape, slice toQ. Pad rows at 1e9 as now. Sole survivor of the chunking; the(chunk, N)intermediate stays bounded under vmap._walk_from_seed(query_points, seed, points, simplices_padded, simplex_neighbors, vertex_simplex, xp)→(cur, done, outside)over allQat once.walk_step/weights_ofunchanged in content. NumPy: existingfor … if (done | outside).all(): break. JAX:jax.lax.while_loop(cond, body, (step, cur, done, outside))withcond = (step < DELAUNAY_WALK_STEPS) & jnp.any(~done & ~outside). Under vmap the loop runs to the slowest lane.mappings/simplex_indexesfrom(seed, cur, done)exactly as the currentlocate_chunktail does, honouringreturn_simplex_indexes.Gradient contract:
lax.while_loopis not reverse-mode differentiable, so the JAX branch wrapsquery_pointsandpointsinjax.lax.stop_gradientbefore the seed argmin and the walk (the connectivity tables already are, see_jax_delaunay_tables). Location is integer-valued with zero a.e. derivative; every differentiable quantity (barycentric weights viapixel_weights_delaunay_from, dual areas, split points, Sibson weights) is computed downstream from the traced arrays. Document in the docstring; update theDELAUNAY_LOCATE_CHUNK/DELAUNAY_WALK_STEPScomments to the new roles.jax_delaunay: locatejnp.concatenate([query_points, split_points])in one call and slice[:Q]/[Q:]; split points still seed from their own nearest vertex.scipy_delaunay(NumPy path) untouched;sibson.pyneeds no change.B. PyAutoArray unit tests (NumPy only, never import jax) —
test_autoarray/inversion/pixelization/interpolator/test_delaunay_walk.py: existing_assert_matches_find_simplextests unchanged; add_nearest_vertex_seed_fromvscKDTreeon the blob-ring mesh (ties excluded);_walk_from_seedfrom a deliberately far seed (vertex 0 for every query) still reaches thefind_simplexanswer within the cap; wrapper returns(Q, 3)int32 and(Q,)int32 withreturn_simplex_indexes=Trueon an oddQ.C. autolens_workspace_test —
scripts/misc/jax_assertions/delaunay_walk.py(sibling ofdelaunay_nn.py, same__Env__block andENV: jax full_datasetsdeclaration, bare"""opener; CI runs every script): parity vsfind_simplex+cKDTreefallback on uniform / blob-ring meshes (reuseadaptive_mesh) and on a 1500-vertex Hilbert mesh traced through a handful of the mass models fromdelaunay_nn_caps.py(lives in this repo, not autolens_profiling), data grid and 4N split points, comparing mapping matrices not triangle ids on shared-edge ties;jax.jitequals eager; jittedvmapover a small batch equals per-member; oddQ;jax.gradof a scalar built from the barycentric weights runs and matches FD; print warm per-call time.D. FD certification + CPU no-regression gate (before the library PR)
autolens_profiling/scripts/imaging/likelihood_breakdown/delaunay_numba.pylocally before/after as a control (same numbers to noise, same pin).autolens_profiling/scripts/imaging/likelihood_breakdown/delaunay.py --split-setuplocally onmainand on the feature branch, warm, three repeats each. Gate: "Triangulation + interpolation" row and total per-eval time not slower;EXPECTED_LOG_EVIDENCE_HSTunchanged. Both tables go in the PR.autolens_workspace_test/scripts/imaging/jax_grad/delaunay.pymust pass.autolens_test scripts/imaging/delaunay.pyand state whether it touches the walk.E. Ship + A100 verification
ship_library: tests, commit, push, PR on PyAutoArray (pending-release).HPCPullPyAutoon the feature branch, runscripts/imaging/likelihood_breakdown/delaunay.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16and thedelaunay_nn.pysibling. Pins must pass unchanged — a shift is a bug, not a re-pin (PyAutoLens#721 knife-edge lesson: bisect before re-pinning anything).autolens_profiling/results/breakdown/imaging/andresults/notes/delaunay_walk_early_exit.mdwith the four-way split unbatched and per call under vmap, the H row and the single-JIT cell against the 2026-09 baselines.ship_workspacefor autolens_workspace_test + autolens_profiling behind the library-first merge gate. Decide whether to file the Phase 2 prompt.Key Files
PyAutoArray/autoarray/inversion/mesh/interpolator/delaunay.py—pix_indexes_delaunay_walk_from,jax_delaunay,_jax_delaunay_tables, the two module constantsPyAutoArray/autoarray/inversion/mesh/interpolator/sibson.py— DelaunayNN caller (unchanged; inherits the gain)PyAutoArray/test_autoarray/inversion/pixelization/interpolator/test_delaunay_walk.py— NumPy parity testsautolens_workspace_test/scripts/misc/jax_assertions/delaunay_walk.py— new JAX parity/jit/vmap/grad scriptautolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn_caps.py— lensing geometry to reuseautolens_workspace_test/scripts/imaging/jax_grad/delaunay.py— FD certificationautolens_profiling/scripts/imaging/likelihood_breakdown/delaunay.py,delaunay_nn.py,delaunay_numba.py— breakdown / control scriptsautolens_profiling/results/notes/preopt_breakdown_baseline.md— the 26.6 ms baselineOriginal Prompt
Click to expand starting prompt
JAX Delaunay point location: early-exit walk, unchunked loop, static image-plane seed
Type: feature
Target: autoarray
Repos:
Themes:
Difficulty: medium
Autonomy: supervised
Priority: high
Status: draft
Consequence: judge
Witness: the A100 Delaunay breakdown's "Triangulation + interpolation" row (26.6 ms,
results/breakdown/imaging/delaunay_hpc_a100_fp64.json,--split-setup) drops to under 8 ms unbatched withEXPECTED_LOG_EVIDENCE_HSTunchanged, the walk parity tests pass, and the FD certification still passesReview-minutes: 30
Unattended: ready
Filed: 2026-09-05
Original request (verbatim):
The measurement
On the A100 the HST Delaunay imaging likelihood (1500 Hilbert vertices, 15,361 data
pixels, MGE-60 lens, ConstantSplit) costs 97 ms per evaluation, of which the four-way
--split-setupdecomposition attributes 26.6 ms to "Triangulation + interpolation"(
autolens_profiling/results/notes/preopt_breakdown_baseline.md). The qhull host callbackis a few ms of that at most (1500 points,
pure_callback,vmap_method="sequential"). Therest is the JAX-side point location in
PyAutoArray/autoarray/inversion/mesh/interpolator/delaunay.py,pix_indexes_delaunay_walk_from, and it is latency, not FLOPs:jax.lax.fori_loop(0, DELAUNAY_WALK_STEPS, ...)withDELAUNAY_WALK_STEPS = 128(module constant at line 75, loop at line 252). A fixed-tripfori_loopcannot exit early. Seeded from the nearest vertex the walk resolves in ahandful of steps, so well over 95% of the 128 iterations are no-ops that still launch a
gather, a cross-product batch, an argmin and three
wheres each.DELAUNAY_LOCATE_CHUNK = 1024(line 69) throughjax.lax.map(line 279), which is sequential: 16 chunks for the data grid and 6 more forthe 6000 ConstantSplit split-cross points (
jax_delaunaycalls the walk twice). That isroughly 22 x 128 = 2,800 dependent walk steps per likelihood.
(chunk, N)nearest-vertex distance intermediateunder
vmap(184 MB per replica at full Q; ~12 GB at batch 64). It is a memory guardthat happens to serialise the latency-bound part.
Nautilus at
n_batch=256gave 51 ms per eval against 62 ms atn_batch=16, so batchingbarely amortises this today.
DelaunayNN(interpolator/sibson.py) seeds its cavity walkfrom
pix_indexes_delaunay_walk_from(..., return_simplex_indexes=True), so it inheritsevery gain here.
interpolator/knn.pyis a separate approach and out of scope.Phase 1: early exit, and chunk only the argmin
fori_loopwith ajax.lax.while_loopwhose predicate is(step < DELAUNAY_WALK_STEPS) & jnp.any(~done & ~outside). The 128 cap stays as thesafety bound; the typical exit is under ten steps. Under
vmapthe loop runs to theslowest lane, which is fine.
chunked (it is the memory hazard, and a
lax.mapover 16 single-argmin chunks is cheap),then run the walk once over all Q queries with no
lax.map. The walk's working set isO(Q), not O(Q x N).
xp is np, used by the unit tests) behaviourally identical; italready early-exits.
(Q, 3)mappings and simplex indexes asscipy.spatial.Delaunay.find_simplexon the existing test meshes and on a 1500-vertexHilbert mesh traced through a few mass models (reuse the geometry from
autolens_profiling/scripts/misc/jax_assertions/delaunay_nn_caps.py). Points exactly ona shared edge may legitimately resolve to either adjacent triangle; the barycentric
weights agree because the opposite vertex gets weight zero, so compare mapping matrices,
not triangle ids, in that case.
Phase 2: static image-plane seed and one-shot fan test
The image-plane mesh is fixed per fit by the adapt image, and ray tracing is continuous
away from critical curves, so each data pixel's nearest image-plane mesh vertex is a
near-perfect source-plane seed. Precomputing that index once per fit (in the
Mapper/AdaptImageslayer, whereimage_plane_mesh_gridis known) removes the brute-force(Q, N)argmin and its memory hazard entirely, which also retires the chunking. Then havethe qhull callback also return a padded vertex-to-incident-triangle table (cap around 12,
audit the actual max the way
delaunay_nn_cap_audit.mddid) so most queries resolve withone vectorised barycentric test over the seed's fan; the
while_loopwalk from Phase 1handles only the residual near critical curves and outside the hull. Split-cross points
have no image-plane parent pixel; seed them at their parent vertex (they are offsets from
it) and let the fan test cover them.
Phase 2 changes the interpolator's inputs (a seed array), so it touches
InterpolatorDelaunay,InterpolatorDelaunayNN, theMapperconstructor path andFitImagingplumbing. Land Phase 1 first and re-measure; Phase 2 is only worth itsplumbing if Phase 1 leaves the row above a few ms.
Gradient contract (do not relax)
Point location is integer-valued; its derivative is zero almost everywhere and it is
already wrapped by
stop_gradientsemantics through the frozen connectivity tables(
_jax_delaunay_tablesdocstring). The barycentric weights are recomputed from the tracedpoints after location exactly as now, so
jax.gradthrough the likelihood is unchanged.Re-run the FD certification
autolens_workspace_test/scripts/imaging/jax_grad/delaunay.pyafter each phase.
Verification on the A100
Use the
delaunay-nn-breakdowntooling (autolens_profiling#219): rerunscripts/imaging/likelihood_breakdown/delaunay.py --config-name hpc_a100_fp64 --split-setup --vmap-batch 16and thedelaunay_nn.pysibling from a RAL worktree withHPCPullPyAutopointed at the feature branch. Report the four-way split unbatched and percall under vmap, the H row, and the single-JIT runtime cell, against the 2026-09 baselines.
Pins (
EXPECTED_LOG_EVIDENCE_HSTin both scripts) must pass unchanged; a shift means amapping changed and is a bug, not a re-pin. Note the symmetric knife-edge lesson from
PyAutoLens#721: if a positions-threshold or point-solver pin elsewhere moves, bisect
before re-pinning.