Skip to content

perf(triangles): drop the throwaway jnp.unique vertex dedup on the JAX PointSolver path (#568) - #569

Merged
Jammy2211 merged 2 commits into
mainfrom
feature/point-source-cpu-p2
Sep 24, 2026
Merged

Jammy2211 merged 2 commits into
mainfrom
feature/point-source-cpu-p2

Conversation

@Jammy2211

@Jammy2211 Jammy2211 commented Sep 24, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Phase 2 of the point-source CPU campaign (part of #568, follows autolens_profiling#297). This removes the throwaway jnp.unique vertex dedup from the JAX triangle path that the PointSolver uses.

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 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 from for_indexes trace to NaN triangles, which every Shape.mask rejects, so containment is unchanged. The NumPy sibling CoordinateArrayTrianglesNp is untouched and still deduplicates, because its shapes are dynamic.

The new test_autoarray/structures/triangles/test_coordinate_jax.py has 10 tests:

  • the vertex-table round trip, eager and under jit
  • NaN padding is never contained
  • kept-triangle parity with the deduplicating NumPy sibling (interior, vertex, shared-edge and centroid source points)
  • NumPy deduplication pinned
  • an HLO guard that the traced containment carries no sort op. It fails on main 11b93476 with 39 sort ops.

Measured speed-up (median ms/call, control → no-dedup, interleaved in-process A/B, autolens_profiling linked PR)

Cell RAL CPU (AMD EPYC, job 350582) RAL A100 80GB fp64 (job 350587)
simple, solved 24.00 → 5.38 (4.47×) 1.695 → 0.843 (2.01×)
simple, plain 23.76 → 5.39 (4.40×) 1.684 → 0.822 (2.05×)
vmap-4 31.31 → 11.91 (2.63×) 1.865 → 0.958 (1.95×)
cluster, solved 155.2 → 78.3 (1.98×) 4.206 → 2.457 (1.71×)
cluster, plain 153.9 → 77.7 (1.98×) 3.860 → 2.105 (1.83×)

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:

  • The CPU A/B ran on AMD EPYC, while phase 1 ran on Xeon, so only within-host ratios are quoted.
  • The CPU A/B used 5 calls per round, not 20.
  • The GPU protocol resolves about 4%: hash-identical library and no-dedup programs differ by 3–4%.

API Changes

None — internal changes only. No signatures, classes or modules change. The contents of the JAX-only CoordinateArrayTriangles.vertices / .indices properties change: the vertex table is no longer deduplicated, and padding triangles keep valid indices with NaN vertices instead of -1 entries. Solver results are bit-identical. See full details below.

Test Plan

  • pytest test_autoarray — 1626 passed at 25894d10
  • New test_coordinate_jax.py (10 tests); the HLO no-sort guard fails on main and passes here
  • autolens_workspace_test point-source jax_likelihood ×4 and jax_grad/gradient.py give identical output to main
  • autolens_profiling image_plane_solved pin unchanged (7.743201200876817 / 7.743201200876812)
  • RAL CPU (350582) and A100 (350587) A/B: every correctness gate has Δ = 0
  • CI green
Full API Changes (for automation & release notes)

Removed

  • None

Added

  • None

Changed Behaviour

  • CoordinateArrayTriangles.vertices (JAX path) — now the flat (3N, 2) table, where row 3*i + k is vertex k of triangle i. It is no longer deduplicated through jnp.unique, and rows of NaN padding triangles are NaN.
  • CoordinateArrayTriangles.indices (JAX path) — now arange(3N).reshape(N, 3). Padding triangles keep valid indices, where before they were -1.
  • CoordinateArrayTrianglesNp — unchanged.

Migration

  • None required. PointSolver and containment results are bit-identical.

Linked PR

Heart RED development override

Shipped under the AUTONOMY.md "Human override for Heart RED (development only)".

  • Authorization (live human, 2026-09-24, this session), faithfully quoted: asked to authorize the development-only override for PyAutoArray#568 after a GPU check, the user replied "GPU check and then I authorize". The GPU check passed (RAL A100 job 350587), so the authorization is in effect.
  • Scope: commit, push and open the pending-release PRs (PyAutoArray, then the linked autolens_profiling PR). No merge, no release. Merge needs its own human /prm with every required check green.
  • Exact current Heart RED reasons (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.yaml
  • Passed branch gates: pytest test_autoarray 1626 passed at 25894d10 (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-source jax_likelihood ×4 and jax_grad/gradient.py print likelihoods and gradients identical to main. The autolens_profiling image_plane_solved pin is unchanged. A/B correctness gates are bit-identical on CPU (RAL job 350582) and GPU (RAL job 350587). Profiling build_readme.py --check and check_submits.py --check pass.
  • This branch does not fix Heart. Heart stays RED for release purposes.

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

Jammy2211 and others added 2 commits September 23, 2026 16:44
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>
@Jammy2211

Copy link
Copy Markdown
Collaborator Author

Workspace PR: PyAutoLabs/autolens_profiling#301 (merge after this one, library-first).

@Jammy2211
Jammy2211 merged commit 681938a into main Sep 24, 2026
3 checks passed
@Jammy2211
Jammy2211 deleted the feature/point-source-cpu-p2 branch September 24, 2026 08:50
@Jammy2211 Jammy2211 removed the pending-release PR queued for the next release build label Sep 26, 2026
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.

1 participant