Skip to content

perf: JAX PointSolver step-0 containment without the (N,3,2) triangle gather #579

Description

@Jammy2211

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.

  1. 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.

  2. 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.

  3. 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.
  4. 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.
  5. 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.
  6. 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)

  1. 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.
  2. 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.
  3. 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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions