Skip to content

perf: evaluate a constant grid shift-and-rotate at compile time under JAX (memo phase 2, PyAutoGalaxy#604) - #520

Merged
Jammy2211 merged 1 commit into
mainfrom
feature/gaussian-precompute-p2
Sep 3, 2026
Merged

Jammy2211 merged 1 commit into
mainfrom
feature/gaussian-precompute-p2

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Summary

Library leg 1 of 2 for PyAutoGalaxy#604 (phase 2 of the gaussian-deflections-precompute epic). Merge first; the PyAutoGalaxy PR depends on it.

Grid2D.subtracted_and_rotated_from is what FitDataset.grids calls on every likelihood evaluation to apply the DatasetModel's grid_offset / grid_rotation_angle. In the overwhelmingly common case both are fixed Python floats, so the result is a compile-time constant — but under jax.jit every jax.numpy call is staged into the jaxpr whatever its operands (jnp.array((0.0, 0.0)) inside a trace is a DynamicJaxprTracer), and this stack disables XLA's constant folding (--xla_disable_hlo_passes=constant_folding, set by autonerves/jax_wrapper.py), so the grid was recomputed every evaluation and reached every downstream consumer as a tracer. That made the phase-1 deflection memo inert on JAX at all 132 profile calls of a SLaM-shaped fit.

This PR wraps the body in jax.ensure_compile_time_eval() on the JAX backend when both offset and angle are concrete (positive test via autoarray.validate.is_concrete_scalar, the same gate the constructor guards use; contextlib.nullcontext on numpy; a traced offset or angle falls back to the staged path). The grid and its over_sampled twin then come out as concrete jax.Arrays. The arithmetic is unchanged — the same operations on the same values, executed now rather than staged — and the log-likelihood of the SLaM-shaped JAX fit is bit-identical (max abs diff 0.0).

API Changes

None — internal changes only. No signature or default changes; numpy behaviour identical; JAX values identical, only their concreteness inside a trace changes.
See full details below.

Test Plan

  • test_autoarray — 1412 passed (-n 4), including 2 new numpy-only tests (_shift_and_rotate_is_constant on floats / tracer stand-in; numpy result unchanged under the context)
  • JAX: inside jax.jit, grid.array and grid.over_sampled.array are concrete ArrayImpl at every memo-hook call with a fixed DatasetModel; with a free grid_offset (6 free parameters) the grid is a tracer and the staged path is taken; log-likelihood identical in both cases
  • autolens_workspace_test scripts/imaging/jax_likelihood/ — all 15 scripts pass before and after under the smoke profile; every captured vmap pin bit-identical; no pin edited
  • ruff check / ruff format --check — no new findings on the touched files (repo baseline has pre-existing ones)
Full API Changes (for automation & release notes)

Removed

  • none

Added

  • autoarray.structures.grids.uniform_2d._shift_and_rotate_is_constant(offset, angle) (private)
  • autoarray.structures.grids.uniform_2d._compile_time_eval_context(offset, angle, xp) (private)

Changed Behaviour

  • Grid2D.subtracted_and_rotated_from(offset, angle, xp) on the JAX backend with concrete offset and angle returns a grid whose arrays are concrete inside a jax.jit trace (values unchanged)

Migration

  • none

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

https://claude.ai/code/session_01XhnA4pFN2NycuKc8Ni6s2R

… JAX (PyAutoGalaxy#604)

Grid2D.subtracted_and_rotated_from: when `offset` and `angle` are concrete
numbers (a fixed or absent DatasetModel — every fit that does not leave
grid_offset / grid_rotation_angle free), the shift-and-rotate is a
compile-time constant. Under jax.jit every jax.numpy call is staged into the
jaxpr regardless of its operands, and this stack disables XLA constant
folding (autonerves/jax_wrapper.py), so the constant grid was rebuilt on
every likelihood evaluation and handed downstream as a tracer. The body now
runs inside jax.ensure_compile_time_eval() on the JAX backend when both
inputs are concrete (validate.is_concrete_scalar), yielding a concrete
jax.Array for the grid and its over-sampled twin. Same arithmetic, same
values (log-likelihood bit-identical); numpy is untouched (nullcontext); a
traced offset or angle takes the staged path as before. This is what lets
autogalaxy's deflections memo fold fixed-geometry fields out of the trace.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01XhnA4pFN2NycuKc8Ni6s2R
@Jammy2211

Copy link
Copy Markdown
Collaborator Author

Phase 2 of the gaussian-deflections-precompute epic (PyAutoGalaxy#604) ships as three PRs:

  1. perf: evaluate a constant grid shift-and-rotate at compile time under JAX (memo phase 2, PyAutoGalaxy#604) #520 (this one) — constant grid shift-and-rotate evaluated at compile time
  2. perf: fold fixed-geometry deflection fields out of the JAX trace (memo phase 2, #604) PyAutoGalaxy#605 — the deflection memo's JAX branch, which depends on 1
  3. likelihood_runtime: JAX validation and before/after for the deflection memo trace-time fold (phase 2, PyAutoGalaxy#604) autolens_profiling#216 — the measurement of record

Merge order: PyAutoArray → PyAutoGalaxy → autolens_profiling.

@Jammy2211
Jammy2211 merged commit e36a5af into main Sep 3, 2026
3 checks passed
@Jammy2211
Jammy2211 deleted the feature/gaussian-precompute-p2 branch September 3, 2026 23:01
@Jammy2211 Jammy2211 removed the pending-release PR queued for the next release build label Sep 4, 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