Skip to content

feat: JAX Delaunay walk — early-exit while_loop, chunk only the seed argmin #530

Description

@Jammy2211

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

  1. 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.
  2. 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.
  3. 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.
  4. 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.
  5. Keep the NumPy path (unit tests) behaviourally identical; extend the NumPy unit tests to the refactored seed/walk split.
  6. 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.
  7. 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):

  1. _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.
  2. _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.
  3. 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

  1. ship_library: tests, commit, push, PR on PyAutoArray (pending-release).
  2. 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).
  3. 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.
  4. 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

  1. 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.
  2. 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).
  3. Keep the NumPy path (xp is np, used by the unit tests) behaviourally identical; it
    already early-exits.
  4. 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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions