Skip to content

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

Merged
Jammy2211 merged 1 commit into
mainfrom
feature/delaunay-walk-early-exit
Sep 7, 2026
Merged

Jammy2211 merged 1 commit into
mainfrom
feature/delaunay-walk-early-exit

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Summary

The JAX Delaunay point locator ran its visibility walk inside the same lax.map chunking as the nearest-vertex seed, and always for the full DELAUNAY_WALK_STEPS (128) fori_loop trip count. The chunking exists as a memory guard for the (chunk, N) squared-distance intermediate of the seed argmin; the walk needs no such guard, and being latency-bound it was paying that chunking as ~Q/1024 serialised runs of a loop that in practice resolves every query in fewer than 10 steps.

pix_indexes_delaunay_walk_from is split into two helpers behind an unchanged wrapper:

  • _nearest_vertex_seed_from — the chunked argmin, now the only chunked stage, still bounding the live intermediate at (chunk, N) under vmap.
  • _walk_from_seed — a lax.while_loop over all Q queries at once, cond (step < DELAUNAY_WALK_STEPS) & any(~done & ~outside), so the walk exits as soon as every query resolves instead of always spending the cap. The NumPy path keeps its break. DELAUNAY_WALK_STEPS is now a safety cap, not a trip count.

lax.while_loop has no reverse-mode rule, so the JAX branch stop_gradients query_points and points before the seed and the walk. Nothing is lost: this function returns only int32 indices, piecewise-constant in the vertex and query positions away from measure-zero re-wiring / triangle-crossing events — the same argument the frozen connectivity tables in _jax_delaunay_tables already rest on. Every differentiable downstream quantity (barycentric weights via pixel_weights_delaunay_from, dual areas, split points, Sibson weights) is recomputed from the traced arrays, so jax.grad through the Delaunay likelihood is unaffected — and is certified below.

jax_delaunay additionally locates the data grid and the split-cross points in one concatenated call rather than two separate walks, since a latency-bound walk over 2Q queries costs roughly one walk rather than two.

sibson.py (DelaunayNN) inherits the gain through the same wrapper. The NumPy path (scipy_delaunay, the numba CPU likelihood) is untouched.

API Changes

None — internal changes only. pix_indexes_delaunay_walk_from keeps its exact signature, its return_simplex_indexes contract and its outside-hull [v, -1, -1] fallback convention; the two new module-private helpers (_nearest_vertex_seed_from, _walk_from_seed) are additions. The one behavioural change is confined to the JAX branch: query_points and points are now stop_gradient-wrapped inside the locator, whose outputs are integer indices whose true derivative is zero.

Test Plan

  • pytest test_autoarray/ — 1452 passed, including 3 new NumPy-only tests in test_autoarray/inversion/pixelization/interpolator/test_delaunay_walk.py (seed matches cKDTree.query off exact ties; the walk converges to find_simplex from a deliberately far seed — the property the split relies on; wrapper shapes/dtypes on an odd query count not a multiple of DELAUNAY_LOCATE_CHUNK).
  • JAX parity vs scipy.spatial.Delaunay.find_simplex + cKDTree outside-hull fallback: exact-row fraction 1.00000 on uniform N=400/Q=3000, blob-ring N=400/Q=3000, blob-ring N=1500/Q=15974; jit == eager; odd Q; jit(vmap) batch of 3 == per-member; jax.grad runs with central-FD rel err 1.8e-11 … 1.6e-9.
  • FD certification autolens_workspace_test/scripts/imaging/jax_grad/delaunay.py — PASSED (exit 0; mass/shear rel err 7.4e-6 … 1.1e-3 under rtol 1e-2).
  • CPU no-regression gate (numba CPU and JAX CPU likelihoods must not slow down). JAX CPU likelihood_breakdown/delaunay.py --split-setup, median of 5 warm reps: params→H prefix 179.31 → 106.45 ms (1.68x) — every after-rep below every before-rep; TOTAL 2932.39 → 2797.60 ms; EXPECTED_LOG_EVIDENCE_HST = 29110.92085793 PASSED 5/5 on both sides and not re-pinned. Numba control delaunay_numba.py: TOTAL 0.325 → 0.297 s, pin PASSED (the numba path never enters the walk).
  • Micro-benchmark, CPU fp64 jitted, median of 10: locator Q=21361 124.15 → 76.70 ms (1.62x); jax_delaunay 131.96 → 68.94 ms (1.91x).
  • A100 verification (issue plan section E) — to be judged on the params→H prefix, see the witness caveat below.

⚠️ Witness caveat — read before judging the breakdown rows

Before this change the breakdown script's "Triangulation + interpolation" prefix returned only the data-grid mappings, so XLA dead-code-eliminated the separate split-point walk and that cost was charged to the H row. With one concatenated walk the split-point walk is live inside that prefix, so the row is now flat (118.92 → 120.11 ms) and the H row goes negative (55.58 → −13.48 ms).

The honest apples-to-apples witness is the params→H prefix (equivalently: Tri+interp + H). The A100 verification will be judged on that prefix, and the task's Witness: line has been amended accordingly. Fixing the profiling script's H-row attribution is a follow-up for the workspace PR, not this one.

Downstream / workspace impact

No public API change, so no migration is needed:

  • sibson.py (DelaunayNN) calls the wrapper unchanged and inherits the speedup.
  • The only workspace caller, autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py, uses the same keyword signature and is unaffected.
  • Follow-up workspace work (a separate PR, not required for this one to merge): a new autolens_workspace_test/scripts/misc/jax_assertions/delaunay_walk.py parity script, the autolens_profiling A100 re-measure, and the H-row attribution note above.

Readiness

Heart at ship time: YELLOW, no RED reasons — "workspace validation not passing (5 failed, 2 timeout, cloud#34099198772: autolens notebooks/multi_dataset/modeling.ipynb, autolens scripts/multi_dataset/modeling.py, autolens_test scripts/imaging/delaunay.py, +4 more)" plus stale "release validation incomplete: no rehearsal for current source". Organism-scope: the failing autolens_test scripts/imaging/delaunay.py legs in that run do not exercise the walk and were already fixed on autolens_workspace_test main today (078e445, 4103234).

Closes #530

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

https://claude.ai/code/session_01B5HT8dp7sWc9qDhZp6moGr

The JAX Delaunay point locator ran its visibility walk inside the same
`lax.map` chunking as the nearest-vertex seed, and always for the full
`DELAUNAY_WALK_STEPS` (128) `fori_loop` trip count. The chunking is a memory
guard for the (chunk, N) squared-distance intermediate of the seed argmin —
the walk needs no such guard, and being latency-bound it was paying that
chunking as ~Q/1024 serialised runs of a loop that in practice resolves every
query in fewer than 10 steps.

`pix_indexes_delaunay_walk_from` is split into two helpers behind an unchanged
wrapper:

- `_nearest_vertex_seed_from` — the chunked argmin, now the ONLY chunked
  stage, still bounding the live intermediate at (chunk, N) under vmap.
- `_walk_from_seed` — a `lax.while_loop` over ALL Q queries at once, cond
  `(step < DELAUNAY_WALK_STEPS) & any(~done & ~outside)`, so the walk exits as
  soon as every query is resolved instead of always spending the cap. The
  NumPy path keeps its `break`. `DELAUNAY_WALK_STEPS` is now a safety cap, not
  a trip count.

`lax.while_loop` has no reverse-mode rule, so the JAX branch `stop_gradient`s
`query_points` and `points` before the seed and the walk. Nothing is lost:
this function returns only int32 indices, piecewise-constant in the vertex and
query positions away from measure-zero re-wiring events — the same argument the
frozen connectivity tables already rest on. Every differentiable downstream
quantity (barycentric weights, dual areas, split points, Sibson weights) is
recomputed from the traced arrays.

`jax_delaunay` also locates the data grid and the split-cross points in ONE
concatenated call rather than two walks, since a latency-bound walk over 2Q
queries costs roughly one walk rather than two.

Public signature, `return_simplex_indexes` contract and the outside-hull
fallback convention are unchanged; `sibson.py` (DelaunayNN) inherits the gain;
the NumPy path (`scipy_delaunay`, the numba CPU likelihood) is untouched.

Verification: test_autoarray 1452 passed (3 new NumPy-only walk tests). JAX
parity vs `scipy find_simplex` + cKDTree fallback: exact-row fraction 1.00000
on uniform N=400/Q=3000, blob-ring N=400/Q=3000 and blob-ring N=1500/Q=15974;
jit == eager; jit(vmap) batch == per-member; `jax.grad` runs with central-FD
rel err 1.8e-11..1.6e-9. FD certification `jax_grad/delaunay.py` PASSED.
CPU no-regression: JAX CPU breakdown params->H prefix 179.31 -> 106.45 ms
(1.68x), `EXPECTED_LOG_EVIDENCE_HST` unchanged 5/5 both sides; numba control
0.325 -> 0.297 s, pin PASSED. Micro-benchmark (CPU fp64 jitted, median of 10):
locator Q=21361 124.15 -> 76.70 ms (1.62x), `jax_delaunay` 131.96 -> 68.94 ms
(1.91x).

Closes #530

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#306

Adds scripts/misc/jax_assertions/delaunay_walk.py (registered in smoke_tests.txt beside the delaunay_nn.py siblings) — parity vs find_simplex + cKDTree, jit == eager across the 1,024 seed-chunk boundary and on odd Q, jit(vmap) == per-member, and jax.grad vs central FD. It imports the rewritten locator and fails against released/main autoarray, so it merges after this PR (library-first gate) and is labelled pending-release.

@Jammy2211

Copy link
Copy Markdown
Collaborator Author

Profiling results PR (A100 A/B, autolens_profiling): PyAutoLabs/autolens_profiling#224

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.

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

1 participant