perf(triangles): drop the throwaway jnp.unique vertex dedup on the JAX PointSolver path (#568) - #569
Merged
Merged
Conversation
CoordinateArrayTriangles._vertices_and_indices now returns the flat (3N, 2) vertex table and an arange(3N).reshape(N, 3) index map. Under jit, jnp.unique needs a static size, so the "deduplicated" table was 3N rows anyway (no deflection evaluations saved) while costing a lexicographic sort of 3N fp64 rows twice per solver refinement step. NaN padding rows from for_indexes trace to NaN triangles, which every Shape.mask rejects, so containment is unchanged. The NumPy sibling CoordinateArrayTrianglesNp still deduplicates (dynamic shapes). See autolens_profiling#297. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Round trip of the non-deduplicated vertex table (eager and jit), NaN padding from for_indexes never contained, kept-triangle parity with the deduplicating NumPy sibling on a small lattice (interior, vertex, shared edge and centroid source points), NumPy deduplication pinned, and an HLO guard that the traced containment carries no sort op (fails on main 11b9347 with 39 sort matches). Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
5 tasks
Collaborator
Author
|
Workspace PR: PyAutoLabs/autolens_profiling#301 (merge after this one, library-first). |
Merged
6 tasks
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
Phase 2 of the point-source CPU campaign (part of #568, follows autolens_profiling#297). This removes the throwaway
jnp.uniquevertex dedup from the JAX triangle path that thePointSolveruses.CoordinateArrayTriangles._vertices_and_indicesnow returns the flat(3N, 2)vertex table and anarange(3N).reshape(N, 3)index map. Underjit,jnp.uniqueneeds a staticsize, so the "deduplicated" table was padded back to 3N rows anyway. It saved no deflection evaluations but cost a lexicographic sort of 3N fp64 rows twice per solver refinement step. NaN padding rows fromfor_indexestrace to NaN triangles, which everyShape.maskrejects, so containment is unchanged. The NumPy siblingCoordinateArrayTrianglesNpis untouched and still deduplicates, because its shapes are dynamic.The new
test_autoarray/structures/triangles/test_coordinate_jax.pyhas 10 tests:jit11b93476with 39 sort ops.Measured speed-up (median ms/call, control → no-dedup, interleaved in-process A/B, autolens_profiling linked PR)
Compile time falls about 20% on CPU and about 40% for the simple cell on GPU. The GPU run imported PyAutoArray
25894d10(this head) with 20×20 calls.All correctness gates are bit-identical on both CPU and GPU: log-likelihoods, solved positions, non-zero gradients and vmap parity.
Caveats:
API Changes
None — internal changes only. No signatures, classes or modules change. The contents of the JAX-only
CoordinateArrayTriangles.vertices/.indicesproperties change: the vertex table is no longer deduplicated, and padding triangles keep valid indices with NaN vertices instead of-1entries. Solver results are bit-identical. See full details below.Test Plan
pytest test_autoarray— 1626 passed at25894d10test_coordinate_jax.py(10 tests); the HLO no-sort guard fails on main and passes herejax_likelihood×4 andjax_grad/gradient.pygive identical output to mainimage_plane_solvedpin unchanged (7.743201200876817 / 7.743201200876812)Full API Changes (for automation & release notes)
Removed
Added
Changed Behaviour
CoordinateArrayTriangles.vertices(JAX path) — now the flat(3N, 2)table, where row3*i + kis vertexkof trianglei. It is no longer deduplicated throughjnp.unique, and rows of NaN padding triangles are NaN.CoordinateArrayTriangles.indices(JAX path) — nowarange(3N).reshape(N, 3). Padding triangles keep valid indices, where before they were-1.CoordinateArrayTrianglesNp— unchanged.Migration
PointSolverand containment results are bit-identical.Linked PR
Heart RED development override
Shipped under the
AUTONOMY.md"Human override for Heart RED (development only)"./prmwith every required check green.pyauto-heart readiness, verdict RED, score 45):release validation FAILED (stage integrate)workspace validation not passing (4 failed, cloud#35579888156: autolens notebooks/cluster/modeling.ipynb, autolens notebooks/weak/a2744.ipynb, autolens scripts/cluster/modeling.py, +1 more)manifest drift: hub organism blurb (organs present) — 7 mismatch(es) vs PyAutoMind/repos.yamlpytest test_autoarray1626 passed at25894d10(re-run at ship); +10 new JAX tests, including an HLO sort-count guard that fails on main (39 sort ops) and passes on the branch. Downstream autolens_workspace_test point-sourcejax_likelihood×4 andjax_grad/gradient.pyprint likelihoods and gradients identical to main. The autolens_profilingimage_plane_solvedpin is unchanged. A/B correctness gates are bit-identical on CPU (RAL job 350582) and GPU (RAL job 350587). Profilingbuild_readme.py --checkandcheck_submits.py --checkpass.Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code