Overview
Phase 4b of the point-source CPU speed-up epic (point-source-cpu-speed). Phase 4a (autolens_profiling#314, PR #321) measured step-0 containment at 1.21 ms of a 1.82 ms single-source solved likelihood on a RAL Xeon 8490H, of which ~0.90 ms is materialising the (23283, 3, 2) triangle array in ArrayTriangles.triangles. This task computes step-0 containment on the static lattice without that gather, keeping the result bit-identical (kept indices, image sets, fiducial log L 7.743201200876812, grad, vmap). No geometry, default or completeness change.
Plan
- Prototype three candidates behind a private switch and measure before choosing: (A) structured strided slices of the traced static vertex table, (B) per-component 1-D gathers, (C) drop only the no-op NaN pad/where.
- Plumb the chosen route:
CoordinateArrayTriangles.with_vertices marks the returned ArrayTriangles (static pytree aux data); containing_indices branches on the marker; refinement steps 1-7 keep the existing path; ArrayTriangles.triangles unchanged.
- Tests: bit-identity fuzz vs the gather path (jit + vmap, several geometries), refinement-path-unchanged, and an HLO guard that the compiled step-0 containment has no
f64[23283,3,2] gather (red on main, green on branch). PyAutoLens suite incl. the static-lattice tie test unchanged.
- Measure with the phase-3 interleaved A/B protocol in
solver_config_sweep.py (new branch feature/point-source-cpu-p4b in autolens_profiling), RAL CPU 8490H + A100 no-regression row.
- Ship library-first: PyAutoArray PR, PyAutoLens only if changed, then the autolens_profiling data PR. Stop rule: if no candidate beats control by >= 1.3x on RAL with bit-identity, ship the data PR only as a documented no-go.
Detailed implementation plan
What the code does today (traced)
ArrayTriangles.triangles (array/PyAutoArray/autoarray/structures/triangles/array.py:121-136) builds the triangle array. It pads the indices, gathers self.vertices[safe_indices] → (N,3,2), then applies a where against NaN. At step 0 no index is −1, so that where does nothing, but it is still traced.
containing_indices (array.py:147-168) is its only consumer: shape.mask(self.triangles) → jnp.where(inside, size=15, fill=-1).
Point.mask (shape.py:164-185 → _barycentric_contains :122-139) uses only six (N,) component vectors: a0 a1 b0 b1 c0 c1.
- The step-0 index map is closed-form, not scrambled.
static_vertex_table (coordinate_array.py:31-103) orders the vertices by lexicographic integer key.
- Every key row has the same width W (59 for ±9.9/0.2).
vid = (ky-ky_min)*W + (kx-kx_min(ky))//4, where kx_min alternates with the parity of ky.
- Triangles are row-major, with up and down interleaved by
(cy+cx) parity.
- Refinement steps 1–7 use the same
containing_indices, but on derived lattices with no vertex_table, where the indices are arange(3N). They must stay untouched.
Approach
When the triangle object came from a step-0 vertex_table, compute the barycentric test on the six component vectors taken by slicing and reshaping the traced vertex table, not by the general (N,3,2) gather. The values in are identical, so the booleans out are identical. Everything else keeps the current path.
-
Worktrees / claims: start_library for PyAutoArray (primary). Add PyAutoLens only if the solver needs a change; the expectation is none, since with_vertices already runs through PyAutoArray. autolens_profiling gets a data PR afterwards on a new branch feature/point-source-cpu-p4b. Register a new issue on PyAutoArray.
-
Prototype three candidates behind a private switch, measured before choosing:
- (A) Structured slices. At
static_vertex_table build time (NumPy, cached), precompute a small layout: W, the parity row offsets, and per-corner slice or stride descriptors. Then (V,2) → reshape (rows, W, 2), take strided slices per corner for the up/down classes, and evaluate _barycentric_contains on the lattice-shaped arrays. Reorder only the boolean mask back to triangle order (row-major, parity-interleaved) before jnp.where.
- (B) Per-component 1-D gathers.
vertices[:,0][idx[:,k]] for k in 0..2: six (N,) gathers, no (N,3,2) array, no NaN where. This is the simplest option.
- (C) Drop the no-op NaN
where/pad when a table is present, keeping the gather. This is the control-minus-overhead floor.
Pick the fastest candidate that passes every gate. Prefer (B) if it is within ~10 % of (A), because it is simpler.
-
Plumb it:
CoordinateArrayTriangles.with_vertices passes the layout, or a structured=True marker, into the ArrayTriangles it returns.
ArrayTriangles.containing_indices branches on that marker. The attribute must be static in the pytree (aux data), not a traced child, so jit and vmap see a Python constant.
for_indexes, neighborhood and up_sample do not propagate the marker, matching how vertex_table is dropped today.
- Keep
ArrayTriangles.triangles itself unchanged for any other caller.
-
Tests (PyAutoArray, alongside test_coordinate_jax.py):
- Bit-identity fuzz. Structured vs gather containment on about 10⁵ points: random, exactly on vertices, edge midpoints and centroids, plus the NaN-padded case. Compare kept index arrays with
array_equal, under both jit and vmap.
- HLO guard. The compiled step-0 containment has no
gather producing f64[23283,3,2]. Copy the phase-2 no-sort guard pattern. The guard must fail on main.
- Refinement untouched. A derived lattice takes the old path.
- PyAutoLens. Run the whole suite, including
test_static_lattice_jax.py and its tie test test__source_on_a_step_0_vertex_returns_the_two_true_images, unchanged.
-
Measure (autolens_profiling, scripts/point_source_image/likelihood_breakdown/solver_config_sweep.py):
- Add a
library route beside control.
- Use the phase-3 interleaved A/B protocol: fresh closures +
jax.clear_caches(), 20×20, bootstrap 90 % CI.
- Record the
step0_split before/after, the 200 prior + 200 stress completeness draws against the control (they must be identical), vmap 1/4/16, compile time (≤ +20 %) and XLA memory.
- Run on RAL CPU pinned to an 8490H node, plus an A100 no-regression row, using the phase-4a submits.
- The library route runs from branch clones on RAL; the mirror is not touched while Euclid arrays run. Assert
source_revisions in the job.
-
Ship library-first.
- PyAutoArray PR, then PyAutoLens (only if changed), then the autolens_profiling data PR with the campaign-note "Phase 4b" section.
- The human runs
/prm on each.
- Release-gated follow-up: after the next release, the library route should read
library_matches.
Execution
This Opus session plans and judges; the prototype, tests, RAL runs and ship run in model: "opus" subagents with progress files + Monitors. Nothing is armed beyond the turn. Stop rule: if no candidate beats the control by ≥ 1.3× on RAL while keeping bit-identity, ship the data PR only, as a documented no-go, with no library change.
Verification
- Kept indices match bit for bit on the fuzz set and on the 400 completeness draws.
- The fiducial
7.743201200876812 holds bit-exactly.
- grad and vmap equal control.
- The HLO guard is red on main and green on the branch.
- The PyAutoArray and PyAutoLens suites pass.
- autolens_workspace_test point-source
jax_likelihood ×4 + jax_grad are identical to main.
- The RAL CPU speed-up has a 90 % CI; the A100 shows ≈1×, with no regression.
Affected Repositories
- PyAutoArray (primary)
- autolens_profiling (measurement / data PR)
- PyAutoLens (only if strictly needed; expected none)
Branch Survey
| Repository |
Current Branch |
Dirty? |
| ./PyAutoArray |
main |
clean |
| ./PyAutoLens |
main |
clean |
| ./autolens_profiling |
main |
1 untracked/modified path (not this task) |
Suggested branch: feature/pointsolver-step0-gather (PyAutoArray); feature/point-source-cpu-p4b (autolens_profiling, based on feature/point-source-cpu-p4 until #321 merges)
Key Files
autoarray/structures/triangles/array.py — ArrayTriangles.triangles, containing_indices
autoarray/structures/triangles/coordinate_array.py — static_vertex_table, with_vertices
autoarray/structures/triangles/shape.py — Point.mask, _barycentric_contains
test_autoarray/structures/triangles/test_coordinate_jax.py — tests
Original Prompt
Prompt: https://github.com/PyAutoLabs/PyAutoMind/blob/main/active/pointsolver_step0_gather_containment.md
Click to expand starting prompt
Point-source CPU speed-up phase 4b — cut the step-0 triangle gather / containment in the JAX PointSolver
Type: feature
Target: autoarray
Repos:
- PyAutoArray
- PyAutoLens
- autolens_profiling
Themes:
- point-source
- profiling
- jax-compile
Difficulty: medium
Autonomy: supervised
Priority: high
Status: formalised
Consequence: judge
Review-minutes: 20
Unattended: ready
Epic: point-source-cpu-speed
Filed: 2026-09-26
Parent: active/pointsolver_cpu_speed_phase_4.md (issue autolens_profiling#314)
Goal
Remove, or sharply cut, the ≈ 0.9 ms vertices[indices] materialisation at step 0 of the JAX
PointSolver. On the released code it is ≈ 49 % of the single-source likelihood. This is a pure
code lever. The solver geometry (±9.9″ / 0.2″ / 1e-3, MAX_CONTAINING_SIZE 15) does not
change, and the result must stay bit-identical:
- image sets and counts;
- the fiducial simple solved log L
7.743201200876812;
jax.grad;
vmap.
Human decision (2026-09-26): this is phase 4b, first of the phase-4a follow-ups.
Evidence (phase 4a, RAL CPU, Xeon 8490H euclid-ral-compute-10-4, 8 CPUs, fp64, JAX 0.10.2)
Ledger: lens/autolens_profiling/results/notes/point_source_cpu_campaign.md, section "Phase 4a".
- Re-baseline, job 356365 (
results/breakdown/point_source_image/image_plane_hpc_ral_cpu_fp64_p4.json,
5 runs): fused solved median 2.095 ms. Step 0 is 66 % of the call. Refinement steps 1–7 are 15 %,
β* 8 %, magnification 7 % and χ² 4 %. The phase-3 FLOP estimate (refinement ≈ 60 %) was wrong in
wall time. That cell's own step-0 "ray trace vs containment" split counts the gather as trace, so
do not quote it.
- Sweep, job 356367 (
results/breakdown/point_source_image/solver_config_sweep_hpc_ral_cpu_fp64.json,
key step0_split.control): the control is 1.824 ms and step 0 is 1.48 ms (81 %). It divides into:
- ray trace of the 11 859 static-lattice vertices: 0.27 ms;
- containment: 1.21 ms (66 % of the likelihood). This is
containing_indices: the gather +
Point.mask + jnp.where.
- Of the containment, the gather that materialises the
(23 283, 3, 2) triangle array alone is
≈ 0.90 ms (triangle_materialisation_ms).
- Under ±2.5″/0.4 (243 rows), containment is 0.06 ms, which shows how much of the cost is the
23 283-triangle size of the lattice.
- Why this lever over shrinking the grid: it recovers most of the extent/scale speed-up (the
best admissible config, ±2.5/0.4, was 2.37× and 5.55× at vmap-16) with no completeness risk and no
default change. The extent became a workspace choice, per draft/feature/autolens/pointsolver_extent_sanity_check.md
and draft/feature/autolens_workspace/pointsolver_grid_extent_per_package.md.
- Revisions measured: PyAutoArray
3de624b5, PyAutoLens 86054bbc, autolens_profiling 6c45fec.
The point-source path is byte-identical to 2026.9.26.1.
Candidate mechanisms (choose by measurement)
- Exploit the regular step-0 lattice. Step 0 is the static phase-3 lattice
(static_vertex_table, PyAutoArray autoarray/structures/triangles/coordinate_array.py; wired in
PyAutoLens AbstractSolver._initial_triangles). Each triangle's vertices are fixed offsets into
the lattice. Compute containment by row/column arithmetic or strided slicing of the traced
vertex table instead of the fancy-index gather. For example, evaluate the barycentric / sign test
on the up- and down-triangles of each lattice row as two dense slices.
- A fused containment kernel. Do the sign test on the index arrays directly, so XLA never
materialises the (N, 3, 2) array. Check the optimised HLO for the gather's disappearance.
- Anything else that removes the materialisation while keeping the kept set bit-identical.
Beware: phase 3's tie study showed that a changed float path at step-0 vertices moves the kept set.
Pin the tie test test_static_lattice_jax.py::test__source_on_a_step_0_vertex_returns_the_two_true_images.
Protocol (as phase 3)
- A red control on main, then an interleaved A/B: fresh closures +
jax.clear_caches(),
20 rounds × 20 calls, rotated round-robin, bootstrap 90 % CI.
- The harness is
lens/autolens_profiling/scripts/point_source_image/likelihood_breakdown/solver_config_sweep.py.
Its step0_split prefixes (source centre / jnp.sum(plane.vertices) / materialised triangles /
containing_indices) are the before/after instrument, and its 200 prior + 200 stress completeness
draws are the regression set. Add a library route beside control.
- Gates: bit-identical log L on the stream and on the fiducial, positions, image counts, grad, and
vmap 1 / 4 / 16. Also compile no worse than +20 %, and no memory regression.
- RAL CPU pinned to an 8490H node (check
sinfo -p ral -N -o "%N %T %C"; idle* nodes are
unreachable), plus an A100 no-regression row. The A100 is launch-bound (phase 3), so expect ~1×.
- Library-first ship: PyAutoArray → PyAutoLens → autolens_profiling data PR.
Traps
- PyAutoArray is currently claimed by task
interferometer-transform-real-scatter. Check the
claim at start_dev (run worktree_check_conflict, not a grep). A parallel claim is fine only if
the file sets are disjoint.
- JAX caches jaxprs on function identity. A monkeypatch A/B needs a distinct function object and
jax.clear_caches() before each compile.
jax.grad through an AnalysisPoint needs autofit.jax.register_model(model); without it the
gradient is silently all-zero.
- Fix, or at least do not quote,
image_plane.py's step-0 prefix split. It counts the gather as ray
trace.
Overview
Phase 4b of the point-source CPU speed-up epic (
point-source-cpu-speed). Phase 4a (autolens_profiling#314, PR #321) measured step-0 containment at 1.21 ms of a 1.82 ms single-source solved likelihood on a RAL Xeon 8490H, of which ~0.90 ms is materialising the(23283, 3, 2)triangle array inArrayTriangles.triangles. This task computes step-0 containment on the static lattice without that gather, keeping the result bit-identical (kept indices, image sets, fiducial log L7.743201200876812, grad, vmap). No geometry, default or completeness change.Plan
CoordinateArrayTriangles.with_verticesmarks the returnedArrayTriangles(static pytree aux data);containing_indicesbranches on the marker; refinement steps 1-7 keep the existing path;ArrayTriangles.trianglesunchanged.f64[23283,3,2]gather (red on main, green on branch). PyAutoLens suite incl. the static-lattice tie test unchanged.solver_config_sweep.py(new branchfeature/point-source-cpu-p4bin autolens_profiling), RAL CPU 8490H + A100 no-regression row.Detailed implementation plan
What the code does today (traced)
ArrayTriangles.triangles(array/PyAutoArray/autoarray/structures/triangles/array.py:121-136) builds the triangle array. It pads the indices, gathersself.vertices[safe_indices]→(N,3,2), then applies awhereagainst NaN. At step 0 no index is −1, so thatwheredoes nothing, but it is still traced.containing_indices(array.py:147-168) is its only consumer:shape.mask(self.triangles)→jnp.where(inside, size=15, fill=-1).Point.mask(shape.py:164-185→_barycentric_contains:122-139) uses only six(N,)component vectors:a0 a1 b0 b1 c0 c1.static_vertex_table(coordinate_array.py:31-103) orders the vertices by lexicographic integer key.vid = (ky-ky_min)*W + (kx-kx_min(ky))//4, wherekx_minalternates with the parity ofky.(cy+cx)parity.containing_indices, but on derived lattices with novertex_table, where the indices arearange(3N). They must stay untouched.Approach
When the triangle object came from a step-0
vertex_table, compute the barycentric test on the six component vectors taken by slicing and reshaping the traced vertex table, not by the general(N,3,2)gather. The values in are identical, so the booleans out are identical. Everything else keeps the current path.Worktrees / claims:
start_libraryfor PyAutoArray (primary). Add PyAutoLens only if the solver needs a change; the expectation is none, sincewith_verticesalready runs through PyAutoArray. autolens_profiling gets a data PR afterwards on a new branchfeature/point-source-cpu-p4b. Register a new issue on PyAutoArray.Prototype three candidates behind a private switch, measured before choosing:
static_vertex_tablebuild time (NumPy, cached), precompute a small layout: W, the parity row offsets, and per-corner slice or stride descriptors. Then(V,2)→ reshape(rows, W, 2), take strided slices per corner for the up/down classes, and evaluate_barycentric_containson the lattice-shaped arrays. Reorder only the boolean mask back to triangle order (row-major, parity-interleaved) beforejnp.where.vertices[:,0][idx[:,k]]for k in 0..2: six(N,)gathers, no(N,3,2)array, no NaNwhere. This is the simplest option.where/pad when a table is present, keeping the gather. This is the control-minus-overhead floor.Pick the fastest candidate that passes every gate. Prefer (B) if it is within ~10 % of (A), because it is simpler.
Plumb it:
CoordinateArrayTriangles.with_verticespasses the layout, or astructured=Truemarker, into theArrayTrianglesit returns.ArrayTriangles.containing_indicesbranches on that marker. The attribute must be static in the pytree (aux data), not a traced child, sojitandvmapsee a Python constant.for_indexes,neighborhoodandup_sampledo not propagate the marker, matching howvertex_tableis dropped today.ArrayTriangles.trianglesitself unchanged for any other caller.Tests (PyAutoArray, alongside
test_coordinate_jax.py):array_equal, under bothjitandvmap.gatherproducingf64[23283,3,2]. Copy the phase-2 no-sort guard pattern. The guard must fail on main.test_static_lattice_jax.pyand its tie testtest__source_on_a_step_0_vertex_returns_the_two_true_images, unchanged.Measure (autolens_profiling,
scripts/point_source_image/likelihood_breakdown/solver_config_sweep.py):libraryroute besidecontrol.jax.clear_caches(), 20×20, bootstrap 90 % CI.step0_splitbefore/after, the 200 prior + 200 stress completeness draws against the control (they must be identical), vmap 1/4/16, compile time (≤ +20 %) and XLA memory.source_revisionsin the job.Ship library-first.
/prmon each.library_matches.Execution
This Opus session plans and judges; the prototype, tests, RAL runs and ship run in
model: "opus"subagents with progress files + Monitors. Nothing is armed beyond the turn. Stop rule: if no candidate beats the control by ≥ 1.3× on RAL while keeping bit-identity, ship the data PR only, as a documented no-go, with no library change.Verification
7.743201200876812holds bit-exactly.jax_likelihood×4 +jax_gradare identical to main.Affected Repositories
Branch Survey
Suggested branch:
feature/pointsolver-step0-gather(PyAutoArray);feature/point-source-cpu-p4b(autolens_profiling, based onfeature/point-source-cpu-p4until #321 merges)Key Files
autoarray/structures/triangles/array.py—ArrayTriangles.triangles,containing_indicesautoarray/structures/triangles/coordinate_array.py—static_vertex_table,with_verticesautoarray/structures/triangles/shape.py—Point.mask,_barycentric_containstest_autoarray/structures/triangles/test_coordinate_jax.py— testsOriginal Prompt
Prompt: https://github.com/PyAutoLabs/PyAutoMind/blob/main/active/pointsolver_step0_gather_containment.md
Click to expand starting prompt
Point-source CPU speed-up phase 4b — cut the step-0 triangle gather / containment in the JAX PointSolver
Type: feature
Target: autoarray
Repos:
Themes:
Difficulty: medium
Autonomy: supervised
Priority: high
Status: formalised
Consequence: judge
Review-minutes: 20
Unattended: ready
Epic: point-source-cpu-speed
Filed: 2026-09-26
Parent: active/pointsolver_cpu_speed_phase_4.md (issue autolens_profiling#314)
Goal
Remove, or sharply cut, the ≈ 0.9 ms
vertices[indices]materialisation at step 0 of the JAXPointSolver. On the released code it is ≈ 49 % of the single-source likelihood. This is a purecode lever. The solver geometry (±9.9″ / 0.2″ / 1e-3,
MAX_CONTAINING_SIZE15) does notchange, and the result must stay bit-identical:
7.743201200876812;jax.grad;vmap.Human decision (2026-09-26): this is phase 4b, first of the phase-4a follow-ups.
Evidence (phase 4a, RAL CPU, Xeon 8490H
euclid-ral-compute-10-4, 8 CPUs, fp64, JAX 0.10.2)Ledger:
lens/autolens_profiling/results/notes/point_source_cpu_campaign.md, section "Phase 4a".results/breakdown/point_source_image/image_plane_hpc_ral_cpu_fp64_p4.json,5 runs): fused solved median 2.095 ms. Step 0 is 66 % of the call. Refinement steps 1–7 are 15 %,
β* 8 %, magnification 7 % and χ² 4 %. The phase-3 FLOP estimate (refinement ≈ 60 %) was wrong in
wall time. That cell's own step-0 "ray trace vs containment" split counts the gather as trace, so
do not quote it.
results/breakdown/point_source_image/solver_config_sweep_hpc_ral_cpu_fp64.json,key
step0_split.control): the control is 1.824 ms and step 0 is 1.48 ms (81 %). It divides into:containing_indices: the gather +Point.mask+jnp.where.(23 283, 3, 2)triangle array alone is≈ 0.90 ms (
triangle_materialisation_ms).23 283-triangle size of the lattice.
best admissible config, ±2.5/0.4, was 2.37× and 5.55× at vmap-16) with no completeness risk and no
default change. The extent became a workspace choice, per
draft/feature/autolens/pointsolver_extent_sanity_check.mdand
draft/feature/autolens_workspace/pointsolver_grid_extent_per_package.md.3de624b5, PyAutoLens86054bbc, autolens_profiling6c45fec.The point-source path is byte-identical to 2026.9.26.1.
Candidate mechanisms (choose by measurement)
(
static_vertex_table, PyAutoArrayautoarray/structures/triangles/coordinate_array.py; wired inPyAutoLens
AbstractSolver._initial_triangles). Each triangle's vertices are fixed offsets intothe lattice. Compute containment by row/column arithmetic or strided slicing of the traced
vertex table instead of the fancy-index gather. For example, evaluate the barycentric / sign test
on the up- and down-triangles of each lattice row as two dense slices.
materialises the
(N, 3, 2)array. Check the optimised HLO for the gather's disappearance.Beware: phase 3's tie study showed that a changed float path at step-0 vertices moves the kept set.
Pin the tie test
test_static_lattice_jax.py::test__source_on_a_step_0_vertex_returns_the_two_true_images.Protocol (as phase 3)
jax.clear_caches(),20 rounds × 20 calls, rotated round-robin, bootstrap 90 % CI.
lens/autolens_profiling/scripts/point_source_image/likelihood_breakdown/solver_config_sweep.py.Its
step0_splitprefixes (source centre /jnp.sum(plane.vertices)/ materialised triangles /containing_indices) are the before/after instrument, and its 200 prior + 200 stress completenessdraws are the regression set. Add a library route beside
control.vmap 1 / 4 / 16. Also compile no worse than +20 %, and no memory regression.
sinfo -p ral -N -o "%N %T %C";idle*nodes areunreachable), plus an A100 no-regression row. The A100 is launch-bound (phase 3), so expect ~1×.
Traps
interferometer-transform-real-scatter. Check theclaim at
start_dev(runworktree_check_conflict, not a grep). A parallel claim is fine only ifthe file sets are disjoint.
jax.clear_caches()before each compile.jax.gradthrough anAnalysisPointneedsautofit.jax.register_model(model); without it thegradient is silently all-zero.
image_plane.py's step-0 prefix split. It counts the gather as raytrace.