Skip to content

perf(point): JAX PointSolver step 0 deflects only the unique lattice vertices (PyAutoArray#568 phase 3) - #749

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

Jammy2211 merged 3 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 PyAutoArray#568; phase 4 remains). On the JAX path AbstractSolver._initial_triangles now builds its step-0 tiling with static_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):

  • shape guard — the first traced deflection grid has 11 859 rows (red on main: 69 849);
  • changing scale changes the step-0 table;
  • positions and image counts vs the flat table (≤ 1e-10) on generic, near-caustic and outer sources, and vs NumPy;
  • log-likelihood (≤ 1e-12 rel) and non-zero jax.grad vs the flat table; vmap equals scalar solves;
  • pinned tie case 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 under jit: 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):

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

No signatures change. Behavioural change on the JAX PointSolver only: step 0's triangles.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 at 972d454e (was 755 + the new tie test), against the linked PyAutoArray branch
  • Shape guard red on main, green here; tie case pinned
  • autolens_workspace_test point-source jax_likelihood ×4 and jax_grad/gradient.py identical to main (no rtol-1e-4 pin moved)
  • RAL CPU (350636, folding off/on) and A100 (350637) A/B: 31 / 31 gates, Δ = 0
  • CI green (needs the PyAutoArray PR merged first)
Full API Changes (for automation & release notes)

Removed

  • None

Added

  • None (tests only: test_autolens/point/triangles/test_static_lattice_jax.py)

Changed Behaviour

  • autolens.point.solver.shape_solver.AbstractSolver._initial_triangles (JAX path) — calls CoordinateArrayTriangles.for_limits_and_scale(..., static_vertices=True); step-0 triangles.vertices shape (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.grad and vmap identical 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).
  • NumPy PointSolver path — unchanged.

Migration

  • None required. Requires the linked PyAutoArray PR (static_vertices).

Linked PRs (merge order)

  1. PyAutoArray: perf(triangles): cached static step-0 vertex table for CoordinateArrayTriangles (#568 phase 3) PyAutoArray#570 — merge first (adds static_vertices).
  2. This PR (PyAutoLens) — merge second.
  3. autolens_profiling: profiling(point_source): static-lattice A/B, RAL CPU + A100 rows, campaign phase 3 (PyAutoArray#568) autolens_profiling#305 — merge last.

Part of PyAutoLabs/PyAutoArray#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

Jammy2211 and others added 2 commits September 24, 2026 12:55
…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>
…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 c9ba49fb2: 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 86054bb into main Sep 24, 2026
4 checks passed
@Jammy2211
Jammy2211 deleted the feature/point-source-cpu-p3 branch September 24, 2026 14:10
@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