perf(point): JAX PointSolver step 0 deflects only the unique lattice vertices (PyAutoArray#568 phase 3) - #749
Merged
Merged
Conversation
…vertices (PyAutoArray#568 phase 3) `AbstractSolver._initial_triangles` builds the JAX tiling with `static_vertices=True`, so step 0 ray-traces the cached geometric vertex table (11 859 rows for the 100x100 @ 0.2" grid, not the flat 69 849) and gathers back through `indices`. The lattice depends only on the solver geometry (static pytree aux data), so the table is a compile-time constant and a change of geometry retraces with the matching table. NumPy path unchanged. New test_static_lattice_jax.py (skips without jax): shape guard on the first traced deflection grid (red on main: 69 849), geometry change changes the table, positions/image counts vs the flat table (<= 1e-10, generic and near-caustic sources) and vs NumPy, log L (rel 1e-12) and non-zero jax.grad vs the flat table, vmap vs scalar. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…lattice (PyAutoArray#568 phase 3) A source placed bit-exactly on traced step-0 lattice vertex (-0.8, -1.99185843) of a simple SIE: the static-lattice JAX solver returns exactly the two true images (the vertex root and the counter-image ~(0.41013634, 0.89463539)); the pre-phase-3 flat table returns three (a duplicate of the vertex root). Human gate decision 2026-09-24: PASS, pinned as a test. 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 PyAutoArray#568; phase 4 remains). On the JAX path
AbstractSolver._initial_trianglesnow builds its step-0 tiling withstatic_vertices=True(linked PyAutoArray PR), so step 0 deflects only the 11 859 geometrically unique lattice vertices instead of the 69 849 flat vertex slots (46 516 vs 276 507 on the cluster lattice) and gathers them back to the triangles through the index map. The lattice depends only on the solver geometry, which is static pytree aux data, so the table is a cached compile-time constant and a geometry change retraces with the matching table. Later refinement steps are data-dependent and unchanged; the NumPy path is unchanged.New
test_autolens/point/triangles/test_static_lattice_jax.py(16 tests, skipped without jax):scalechanges the step-0 table;jax.gradvs the flat table;vmapequals scalar solves;test__source_on_a_step_0_vertex_returns_the_two_true_images— source bit-exactly on the image of step-0 vertex (−0.8, −1.99185843) of a simple SIE, source traced underjit: the static lattice returns exactly the 2 true images (the vertex root and the counter-image ~(0.41013634, 0.89463539)); the flat control returns 3 (a duplicate of the vertex root), and main's flat path is eager-vs-jit unstable at such ties. Human gate decision 2026-09-24: PASS, pinned as this test.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
No signatures change. Behavioural change on the JAX
PointSolveronly: step 0'striangles.vertices— the grid every step-0 deflection is evaluated on — is now the(V, 2)unique-vertex table (11 859 rows on the simple lattice) instead of the flat(3N, 2)one (69 849). Solver results are identical within tolerance (bit-identical in the RAL A/B). See full details below.Test Plan
pytest test_autolens— 756 passed, 1 xfailed at972d454e(was 755 + the new tie test), against the linked PyAutoArray branchjax_likelihood×4 andjax_grad/gradient.pyidentical to main (no rtol-1e-4 pin moved)Full API Changes (for automation & release notes)
Removed
Added
test_autolens/point/triangles/test_static_lattice_jax.py)Changed Behaviour
autolens.point.solver.shape_solver.AbstractSolver._initial_triangles(JAX path) — callsCoordinateArrayTriangles.for_limits_and_scale(..., static_vertices=True); step-0triangles.verticesshape(69849, 2)→(11859, 2)for a 100×100, 0.2″ grid (cluster 200×200 @ 0.7″: 276 507 → 46 516). Positions, image counts, log-likelihood,jax.gradandvmapidentical within tolerance. At a measure-zero bit-exact source-on-vertex tie a duplicate image of the flat path is no longer returned (pinned test).PointSolverpath — unchanged.Migration
static_vertices).Linked PRs (merge order)
static_vertices).Part of PyAutoLabs/PyAutoArray#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