perf: evaluate a constant grid shift-and-rotate at compile time under JAX (memo phase 2, PyAutoGalaxy#604) - #520
Merged
Conversation
… 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
This was referenced Sep 3, 2026
Collaborator
Author
|
Phase 2 of the
Merge order: PyAutoArray → PyAutoGalaxy → autolens_profiling. |
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
Library leg 1 of 2 for PyAutoGalaxy#604 (phase 2 of the
gaussian-deflections-precomputeepic). Merge first; the PyAutoGalaxy PR depends on it.Grid2D.subtracted_and_rotated_fromis whatFitDataset.gridscalls on every likelihood evaluation to apply theDatasetModel'sgrid_offset/grid_rotation_angle. In the overwhelmingly common case both are fixed Python floats, so the result is a compile-time constant — but underjax.jiteveryjax.numpycall is staged into the jaxpr whatever its operands (jnp.array((0.0, 0.0))inside a trace is aDynamicJaxprTracer), and this stack disables XLA's constant folding (--xla_disable_hlo_passes=constant_folding, set byautonerves/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 bothoffsetandangleare concrete (positive test viaautoarray.validate.is_concrete_scalar, the same gate the constructor guards use;contextlib.nullcontexton numpy; a traced offset or angle falls back to the staged path). The grid and itsover_sampledtwin then come out as concretejax.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_constanton floats / tracer stand-in; numpy result unchanged under the context)jax.jit,grid.arrayandgrid.over_sampled.arrayare concreteArrayImplat every memo-hook call with a fixedDatasetModel; with a freegrid_offset(6 free parameters) the grid is a tracer and the staged path is taken; log-likelihood identical in both casesscripts/imaging/jax_likelihood/— all 15 scripts pass before and after under the smoke profile; every captured vmap pin bit-identical; no pin editedruff 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
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 concreteoffsetandanglereturns a grid whose arrays are concrete inside ajax.jittrace (values unchanged)Migration
Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code
https://claude.ai/code/session_01XhnA4pFN2NycuKc8Ni6s2R