perf(triangles): cached static step-0 vertex table for CoordinateArrayTriangles (#568 phase 3) - #570
Merged
Merged
Conversation
…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>
This was referenced Sep 24, 2026
…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>
Collaborator
Author
|
CI fix pushed at 🤖 Generated with Claude Code |
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 3 of the point-source CPU campaign (part of #568; phase 4 remains). Step 0 of the JAX
PointSolverdeflects 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): anlru_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, andvertices/indicesthen 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 JAXPointSolverstep 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.triangleswithin 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):simple_solvedsimple_solved_vmap4(per batch)simple_plaincluster_solved(2 sources, 13 components)cluster_plainRAL CPU job 350636 — constant folding on (
--constant-folding, HLO probe confirms folding ran):simple_solvedsimple_solved_vmap4(per batch)simple_plaincluster_solvedcluster_plainRAL A100 job 350637 (A100 80GB PCIe, fp64, folding off):
simple_solvedsimple_solved_vmap4(per batch)simple_plaincluster_solvedcluster_plainRead 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_revisionsPyAutoArrayad0bf97b, PyAutoLensb346b6a0(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
7.743201200876812, unchanged since phase 1); solved positions and image counts identical (simple 4, cluster 3 + 3);jax.gradfinite, non-zero and bit-identical; vmap-4 equals scalar; step-0containing_indicessets identical. The plan's gate was the tolerance gate (logL ≤ 1e-12 rel, positions ≤ 1e-10); the A/B met it at Δ = 0.test__source_on_a_step_0_vertex_returns_the_two_true_images.Caveats
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, avertex_table=Noneconstructor argument onCoordinateArrayTriangles, andfor_limits_and_scale(..., static_vertices=False). Defaults leave every existing result unchanged; withstatic_vertices=Truethe JAXvertices/indicesproperties return the geometrically unique table and its index map. See full details below.Test Plan
pytest test_autoarray— 1645 passed atad0bf97b(re-run at ship)test_coordinate_jax.py(10 → 29)pytest test_autolens756 passed, 1 xfailed on the linked PyAutoLens branch, incl. the shape guard (red on main) and the pinned tie casejax_likelihood×4 andjax_grad/gradient.pyidentical to mainFull API Changes (for automation & release notes)
Removed
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 thatvertices[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 byvertices/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, **_)—Trueattaches the cachedstatic_vertex_table(limits and scale must be concrete numbers; they are the cache key).Changed Behaviour
static_vertices=True,CoordinateArrayTriangles.verticesis(V, 2)(11 859 rows rather than 69 849 for the ±9.9″ / 0.2″ lattice) andindicesmaps into it.Migration
Linked PRs (merge order)
PointSolverstep 0; merge second.Part of #568 (phase 4 remains).
Heart RED development override
Shipped under the
AUTONOMY.md"Human override for Heart RED (development only)"./prmwith every required check green.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.yamlpytest test_autoarray1645 passed atad0bf97b;pytest test_autolens756 passed, 1 xfailed at972d454e(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-sourcejax_likelihood×4 andjax_grad/gradient.pyoutput 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; profilingbuild_readme.py --checkandcheck_submits.py --checkpass.Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code