feat: JAX Delaunay walk — early-exit while_loop, chunk only the seed argmin - #531
Merged
Merged
Conversation
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
Collaborator
Author
|
Workspace PR: PyAutoLabs/autolens_workspace_test#306 Adds |
Merged
6 tasks
Collaborator
Author
|
Profiling results PR (A100 A/B, autolens_profiling): PyAutoLabs/autolens_profiling#224 |
This was referenced Sep 7, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
The JAX Delaunay point locator ran its visibility walk inside the same
lax.mapchunking as the nearest-vertex seed, and always for the fullDELAUNAY_WALK_STEPS(128)fori_looptrip 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_fromis 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)undervmap._walk_from_seed— alax.while_loopover allQqueries 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 itsbreak.DELAUNAY_WALK_STEPSis now a safety cap, not a trip count.lax.while_loophas no reverse-mode rule, so the JAX branchstop_gradientsquery_pointsandpointsbefore 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_tablesalready rest on. Every differentiable downstream quantity (barycentric weights viapixel_weights_delaunay_from, dual areas, split points, Sibson weights) is recomputed from the traced arrays, sojax.gradthrough the Delaunay likelihood is unaffected — and is certified below.jax_delaunayadditionally locates the data grid and the split-cross points in one concatenated call rather than two separate walks, since a latency-bound walk over2Qqueries 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_fromkeeps its exact signature, itsreturn_simplex_indexescontract 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_pointsandpointsare nowstop_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 intest_autoarray/inversion/pixelization/interpolator/test_delaunay_walk.py(seed matchescKDTree.queryoff exact ties; the walk converges tofind_simplexfrom a deliberately far seed — the property the split relies on; wrapper shapes/dtypes on an odd query count not a multiple ofDELAUNAY_LOCATE_CHUNK).scipy.spatial.Delaunay.find_simplex+ cKDTree outside-hull fallback: exact-row fraction 1.00000 on uniformN=400/Q=3000, blob-ringN=400/Q=3000, blob-ringN=1500/Q=15974;jit== eager; oddQ;jit(vmap)batch of 3 == per-member;jax.gradruns with central-FD rel err 1.8e-11 … 1.6e-9.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).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.92085793PASSED 5/5 on both sides and not re-pinned. Numba controldelaunay_numba.py: TOTAL 0.325 → 0.297 s, pin PASSED (the numba path never enters the walk).Q=21361124.15 → 76.70 ms (1.62x);jax_delaunay131.96 → 68.94 ms (1.91x).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→Hprefix (equivalently: Tri+interp + H). The A100 verification will be judged on that prefix, and the task'sWitness: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.autolens_workspace_test/scripts/misc/jax_assertions/delaunay_nn.py, uses the same keyword signature and is unaffected.autolens_workspace_test/scripts/misc/jax_assertions/delaunay_walk.pyparity script, theautolens_profilingA100 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 failingautolens_test scripts/imaging/delaunay.pylegs in that run do not exercise the walk and were already fixed onautolens_workspace_testmain today (078e445, 4103234).Closes #530
Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code
https://claude.ai/code/session_01B5HT8dp7sWc9qDhZp6moGr