Skip to content

perf(triangles): cached static step-0 vertex table for CoordinateArrayTriangles (#568 phase 3) - #570

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

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

Conversation

@Jammy2211

@Jammy2211 Jammy2211 commented Sep 24, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Phase 3 of the point-source CPU campaign (part of #568; phase 4 remains). Step 0 of the JAX PointSolver deflects the vertices of a triangle lattice fixed by the solver geometry alone. After phase 2 (#569) that is the flat (3N, 2) table — 69 849 rows for the ±9.9″ / 0.2″ lattice of a 100×100, 0.2″ grid — although only 11 859 of them are geometrically distinct (28 665 exact-float distinct: a lattice point computed from neighbouring triangle centres differs by ~1 ulp). The cluster 200×200 @ 0.7″ lattice: 276 507 → 46 516.

This PR adds static_vertex_table(y_min, y_max, x_min, x_max, scale): an lru_cached NumPy builder, keyed on the integer lattice position (cy + f·dy, 2·cx + f·dx), that returns the read-only (V, 2) unique vertices (first-occurrence floats, same arithmetic as .triangles) and the (N, 3) index map. CoordinateArrayTriangles.for_limits_and_scale(..., static_vertices=True) attaches it, and vertices / indices then return it instead of the flat table. The default (static_vertices=False) path is unchanged; derived lattices (for_indexes, up_sample, neighborhood) and pytree round-trips drop the table (falling back to the still-correct flat table). The linked PyAutoLens PR switches it on for the JAX PointSolver step 0 only. Table build 34 ms / 0.75 MB (simple) and 184 ms / 2.96 MB (cluster) on the RAL Xeon; cache hit 1–2 µs.

New tests in test_autoarray/structures/triangles/test_coordinate_jax.py (+19 cases, 10 functions): V = 11 859 for the ±9.9″ / 0.2″ lattice; vertices[indices] equals .triangles within an ulp; kept triangles match NumPy; derived lattices drop the table; the cache key changes with scale; the arrays are read-only; the phase-2 no-sort guard is kept.

Measured speed-up (control = flat step-0 table → library = this change; median ms per likelihood call, 20 rounds × 20 calls, bootstrap 90 % CI)

RAL CPU job 350636 (ral, pinned Xeon Platinum 8490H, 8 CPUs, fp64) — constant folding off (PyAutoNerves default):

Row control ms library ms control / library [90 % CI] MDI compile s ctrl → lib FLOPs ctrl → lib
simple_solved 3.540 1.758 2.01× [1.63, 2.09] 34.7 % 2.91 → 3.14 (+8.0 %) 7.05M → 3.38M
simple_solved_vmap4 (per batch) 8.258 5.784 1.43× [1.42, 1.44] 26.7 % 3.39 → 3.71 (+9.5 %) 23.30M → 11.83M
simple_plain 3.479 1.669 2.08× [1.91, 2.13] 36.0 % 2.58 → 2.67 (+3.3 %) 7.05M → 3.38M
cluster_solved (2 sources, 13 components) 47.36 9.085 5.21× [5.10, 5.30] 39.1 % 22.77 → 24.66 (+8.3 %) 139.44M → 39.47M
cluster_plain 47.60 9.225 5.16× [5.04, 5.26] 27.0 % 16.69 → 18.72 (+12.2 %) 139.43M → 39.46M

RAL CPU job 350636 — constant folding on (--constant-folding, HLO probe confirms folding ran):

Row control ms library ms control / library [90 % CI] MDI compile s ctrl → lib FLOPs ctrl → lib
simple_solved 4.093 1.722 2.38× [2.08, 2.45] 42.9 % 2.85 → 2.91 (+1.9 %) 6.79M → 3.10M
simple_solved_vmap4 (per batch) 7.632 5.499 1.39× [1.31, 1.63] 44.6 % 3.22 → 3.43 (+6.4 %) 23.02M → 11.53M
simple_plain 3.564 1.696 2.10× [1.99, 2.25] 49.5 % 2.44 → 2.57 (+5.2 %) 6.79M → 3.10M
cluster_solved 39.03 8.599 4.54× [4.39, 4.64] 37.4 % 21.93 → 24.14 (+10.1 %) 131.39M → 30.34M
cluster_plain 32.34 8.273 3.91× [3.83, 3.99] 31.1 % 16.28 → 18.97 (+16.5 %) 131.38M → 30.33M

RAL A100 job 350637 (A100 80GB PCIe, fp64, folding off):

Row control ms library ms control / library [90 % CI] MDI compile s ctrl → lib FLOPs ctrl → lib
simple_solved 0.8463 0.8418 1.01× [1.00, 1.01] 5.3 % 5.28 → 5.57 (+5.4 %) 6.60M → 3.07M
simple_solved_vmap4 (per batch) 0.9421 0.9251 1.02× [1.02, 1.02] 4.5 % 6.28 → 6.19 (−1.5 %) 21.98M → 10.82M
simple_plain 0.8244 0.8184 1.01× [1.00, 1.01] 7.0 % 5.19 → 5.10 (−1.6 %) 6.60M → 3.07M
cluster_solved 2.696 2.515 1.07× [1.06, 1.09] 9.2 % 43.13 → 48.84 (+13.2 %) 120.38M → 34.59M
cluster_plain 2.073 1.991 1.04× [1.04, 1.05] 11.8 % 36.11 → 38.76 (+7.4 %) 120.37M → 34.57M

Read from results/breakdown/point_source/static_lattice_ab_{hpc_ral_cpu,constant_folding_hpc_ral_cpu,hpc_ral_a100}_fp64.json (autolens_profiling). Every JSON: source_revisions PyAutoArray ad0bf97b, PyAutoLens b346b6a0 (imported from RAL branch clones), library_matches: lattice, all_gates_pass: true. XLA temp memory falls 4.82 → 1.69 MB (simple) and 91.6 → 22.6 MB (cluster) on CPU, 3.36 → 0.39 / 42.1 → 6.35 MB on the A100.

Gates

  • A/B correctness: 31 / 31 gates pass in all three JSONs with max |Δ| = 0 — log-likelihood bit-identical on every row, route and 16-instance stream (fiducial simple solved 7.743201200876812, unchanged since phase 1); solved positions and image counts identical (simple 4, cluster 3 + 3); jax.grad finite, non-zero and bit-identical; vmap-4 equals scalar; step-0 containing_indices sets identical. The plan's gate was the tolerance gate (logL ≤ 1e-12 rel, positions ≤ 1e-10); the A/B met it at Δ = 0.
  • Tie case pinned — a source placed bit-exactly on traced step-0 vertex (−0.8, −1.99185843) of a simple SIE: the flat control returns 3 images (a duplicate of the vertex root), the static lattice exactly the 2 true images. Human decision 2026-09-24: PASS, pinned as PyAutoLens test__source_on_a_step_0_vertex_returns_the_two_true_images.

Caveats

  • Compile figures are a single cold compile per route, not a median; differences of a few percent are not resolved. Worst: +12.2 % (folding off), +16.5 % (folding on), +13.2 % (A100) — all under the +20 % stop rule.
  • The CPU job ran on a Xeon 8490H node (the phase-1 host); phase 2 ran on EPYC 7763. Only in-job ratios compare across phases.
  • vmap-4 gains 1.43×, below 2× MDI on its own; the change is accepted via the cluster clause of the stop rule (reject only if the simple gain < max(5 %, 2× MDI) with no cluster gain — cluster gains 5.21× against 2× MDI = 78 %).
  • MDI is 27–50 % on this node (per-call dispersion across the parameter stream, not uncertainty of the median; the ratio CIs are a few percent wide).
  • GPU gain is only 1–7 %: the A100 call is launch/latency-bound, not FLOP-bound. No regression.
  • PyAutoLens gains a guarded JAX unit-test file (test_static_lattice_jax.py, ~85–100 s, skipped without jax), departing from its no-JAX-in-unit-tests convention.

API Changes

Additive, opt-in: a new public cached builder static_vertex_table, a vertex_table=None constructor argument on CoordinateArrayTriangles, and for_limits_and_scale(..., static_vertices=False). Defaults leave every existing result unchanged; with static_vertices=True the JAX vertices / indices properties return the geometrically unique table and its index map. See full details below.

Test Plan

  • pytest test_autoarray — 1645 passed at ad0bf97b (re-run at ship)
  • +19 static-table test cases in test_coordinate_jax.py (10 → 29)
  • Downstream: pytest test_autolens 756 passed, 1 xfailed on the linked PyAutoLens branch, incl. the shape guard (red on main) and the pinned tie case
  • autolens_workspace_test point-source jax_likelihood ×4 and jax_grad/gradient.py identical to main
  • RAL CPU (350636, folding off/on) and A100 (350637) A/B: 31 / 31 gates, Δ = 0
  • CI green
Full API Changes (for automation & release notes)

Removed

  • None

Added

  • autoarray.structures.triangles.coordinate_array.static_vertex_table(y_min, y_max, x_min, x_max, scale) — lru_cached (maxsize 32) NumPy builder returning (vertices, indices): the read-only (V, 2) float64 geometrically unique lattice vertices and the (N, 3) int index map such that vertices[indices] is the (N, 3, 2) triangle array.

Changed Signature

  • CoordinateArrayTriangles.__init__(..., vertex_table: Optional[Tuple[np.ndarray, np.ndarray]] = None) — optional precomputed table returned by vertices / indices; not inherited by derived lattices, not part of the pytree.
  • CoordinateArrayTriangles.for_limits_and_scale(y_min, y_max, x_min, x_max, scale=1.0, static_vertices=False, **_) — True attaches the cached static_vertex_table (limits and scale must be concrete numbers; they are the cache key).

Changed Behaviour

  • None at the defaults. With static_vertices=True, CoordinateArrayTriangles.vertices is (V, 2) (11 859 rows rather than 69 849 for the ±9.9″ / 0.2″ lattice) and indices maps into it.

Migration

  • None required.

Linked PRs (merge order)

  1. This PR (PyAutoArray) — merge first.
  2. PyAutoLens: perf(point): JAX PointSolver step 0 deflects only the unique lattice vertices (PyAutoArray#568 phase 3) PyAutoLens#749 — turns the table on for the JAX PointSolver step 0; merge second.
  3. autolens_profiling: profiling(point_source): static-lattice A/B, RAL CPU + A100 rows, campaign phase 3 (PyAutoArray#568) autolens_profiling#305 — A/B cell, RAL CPU + A100 rows, campaign note; merge last.

Part of #568 (phase 4 remains).

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 "When ingestion finishes, how should I ship? Heart is RED…", with the options "Ship with RED override (Recommended): Authorize the development-only Heart RED override for PyAutoArray#568 phase 3…" or hold, the human selected "Ship with RED override (Recommended)".
  • Scope: commit, push and open the pending-release PRs (PyAutoArray, PyAutoLens, then autolens_profiling). No merge, no release. Merge needs its own human /prm with every required check green.
  • Exact current Heart RED reasons (pyauto-heart readiness, re-read immediately before opening these PRs, 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 1645 passed at ad0bf97b; pytest test_autolens 756 passed, 1 xfailed at 972d454e (both re-run at ship); the step-0 shape guard (11 859 traced rows) is red on main (69 849) and green on the branch; autolens_workspace_test point-source jax_likelihood ×4 and jax_grad/gradient.py output identical to main (no rtol-1e-4 pin moved); A/B 31 / 31 gates bit-identical on RAL CPU job 350636 (folding off and on) and A100 job 350637; tie case pinned as a test; 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

…yTriangles (#568 phase 3)

Add `static_vertex_table(y_min, y_max, x_min, x_max, scale)`, a module-level
lru_cache'd NumPy builder of the geometrically unique vertices of the initial
triangle lattice and the (N, 3) index map into them. Vertices are keyed on
their exact integer lattice position (cy + f*dy, 2*cx + f*dx), because the
same point computed from neighbouring triangle centres differs by ~1 ulp:
the +-9.9"/0.2" PointSolver lattice has 69 849 slots, 28 665 exact-float
distinct rows and 11 859 geometric points. Each vertex keeps the float of
its first occurrence (same arithmetic as `.triangles`); arrays are read-only.

`CoordinateArrayTriangles` gains `vertex_table=None` and
`for_limits_and_scale(..., static_vertices=False)`. When a table is attached
`vertices` / `indices` return it; derived lattices (for_indexes, up_sample,
neighborhood) and pytree round-trips do not inherit it. Default path and the
NumPy sibling are unchanged.

Tests: solver-lattice counts (11 859 / 28 665 / 23 283), gather round trip
within 4 ulp, kept triangles match NumPy and the flat table, jit == eager,
derived lattices drop the table, cache keyed on geometry, read-only, no sort.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…est-nojax collection, #568)

The module-level @pytest.mark.parametrize decorators read these plain-data
constants at collection time; defining them inside the jax-present branch
raised NameError on the NumPy-only CI leg. Verified: no-jax collection skips
cleanly, with-jax collected count unchanged.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@Jammy2211

Copy link
Copy Markdown
Collaborator Author

CI fix pushed at d6c5e524: the new JAX test module defined its @parametrize constants inside the jax-present branch, so the unittest-nojax leg failed at collection with a NameError. The constants are now defined at module level. Checked locally: with jax absent (simulated) the module collects and skips cleanly, and with jax present the collected count is unchanged. No test or gate changed. The 'passed gates' list in the PR body covered local jax-present suites only; this leg is now covered too.

🤖 Generated with Claude Code

@Jammy2211
Jammy2211 merged commit 7fa8d27 into main Sep 24, 2026
3 checks passed
@Jammy2211
Jammy2211 deleted the feature/point-source-cpu-p3 branch September 24, 2026 14:09
@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