Skip to content

perf: fold fixed-geometry deflection fields out of the JAX trace (memo phase 2, #604) - #605

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

Closes #604 — phase 2 of the gaussian-deflections-precompute epic. Depends on PyAutoLabs/PyAutoArray#520 (constant grid shift-and-rotate evaluated at compile time), which must merge first.

Phase 1's memo was a no-op on JAX at three independent gates. Under the JAX likelihood a fixed Gaussian's geometry reaches the profile as Python floats and — after the PyAutoArray change — the grid as a concrete jax.Array, while the free mass_to_light_ratio is a tracer. The Faddeeva subgraph is therefore a constant, but jax.jit stages every jax.numpy call regardless of operands and this stack disables XLA constant folding (autonerves/jax_wrapper.py), so nothing folded it away. The memo now does so explicitly: on a miss it evaluates the unit-ratio field with numpy and scipy on a numpy twin of the grid, stores it in the same dict the numpy path uses (both backends share entries — the same bytes), and returns mass_to_light_ratio * xp.asarray(field). The Faddeeva block leaves the jaxpr, replaced by one constant and one multiply.

  • _is_concrete_array(array, xp): ndarray → True; on JAX a jax.Array that is not a jax.core.Tracer (looked up via sys.modules, no import; never try: np.asarray). _concrete_scalar_value reads a concrete 0-d array back as a float so pytree-native instances key identically to float ones. _takes_ratio_split accepts a traced or 0-d-array ratio on JAX (it is only ever multiplied). _numpy_profile_from builds the numpy-evaluated copy and re-checks every constructor argument is concrete before evaluation. _grid_fingerprint_and_twin hashes np.asarray(values) (one device→host copy per grid object per trace, weakref-cached) and rebuilds a Grid2D / Grid2DIrregular twin. Any tracer among key values, or a traced grid → direct JAX call.
  • Witness memo_stats()['jax_folds']: trace-time numpy evaluations performed for a JAX caller.

Measured (autolens_profiling scripts/imaging/likelihood_runtime/mge_mass_jax.py, SLaM shape: 30 fixed lmp_linear.Gaussians + one free ratio, NFWSph + shear, rectangular bilinear source, hst, jit(vmap) batch 3): mge._wofz calls on the compiling call numpy 180 / jnp 0 with the memo on vs numpy 0 / jnp 240 off, steady state 0/0 both; jax_folds 90 (30 Gaussians × 3 grids); jaxpr 53,369 → 13,289 equations (−75%); log-likelihood −56107.56407588691 vs −56107.56407588643 (8.6e-15 relative); vmap_first_call 10.8 → 5.4 s (2.0×; 1.64× on a second run); vmap_steady_x10 unchanged (2.5–2.6 s — the compiled inversion dominates this fit, not the Faddeeva block). Controls: kill switch reproduces memo-off; a free grid_offset makes the grid a tracer → direct path, jax_folds 0, value unchanged.

Side effect, reported honestly: JAX fits that evaluate a fully fixed MGE-routed mass profile now get the numpy/scipy field (L1), i.e. the backend floor of ~1e-7 on those profiles (the decompose_convergence_via_mge amplitude sum, recorded in #600) moves onto the numpy value. In autolens_workspace_test two scripts became more accurate against their own numpy references (imaging/jax_likelihood/mge.py: 4.5e-5 → 1e-16; delaunay.py likewise); all 15 pass and no pin moved.

API Changes

None — internal changes only. Private module internals and one new memo_stats() counter. Numpy path unchanged.
See full details below.

Test Plan

  • test_autogalaxy — 1181 passed (-n 4), 4 new numpy-only tests (_is_concrete_array on ndarray / list / a jax-module stand-in; numpy twin equality; jax_folds stays 0 under numpy); no jax import in tests
  • JAX validation cell (companion autolens_profiling PR): vmap traces and runs; witness counts, jaxpr delta, likelihood agreement and both controls as above
  • autolens_workspace_test scripts/imaging/jax_likelihood/ — 15/15 pass before and after; every captured vmap pin bit-identical; none edited
  • autolens_profiling scripts/lens/deflections/{total,dark,stellar,basis}.py numpy pins held at rtol 1e-6
  • ruff check / ruff format --check clean on touched files
Full API Changes (for automation & release notes)

Removed

  • none

Added

  • autogalaxy.profiles.mass.abstract.deflections_memo._is_concrete_array, _concrete_scalar_value, _numpy_twin_from, _grid_fingerprint_and_twin, _numpy_profile_from (private)
  • memo_stats()['jax_folds'] counter

Changed Behaviour

  • deflections_memo.deflections_yx_2d_from on the JAX backend folds a fixed-geometry profile's field into the trace as a constant (values agree with the direct JAX call to ~1e-14 for Gaussians; fully fixed MGE-routed profiles take the numpy/scipy value, ~1e-7 from the previous JAX value)

Migration

  • none

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

https://claude.ai/code/session_01XhnA4pFN2NycuKc8Ni6s2R

…o phase 2, #604)

deflections_memo on the JAX backend is now a trace-time constant fold: when
the grid is a concrete jax.Array (PyAutoArray evaluates a constant
shift-and-rotate at compile time) and every geometry argument is concrete
while mass_to_light_ratio is a tracer, the unit-ratio field is evaluated once
with numpy/scipy on a numpy twin of the grid, stored in the same memo dict the
numpy path uses, and returned as `ratio * xp.asarray(field)`. The Faddeeva
subgraph leaves the jaxpr (53,369 -> 13,289 equations on a SLaM-shaped fit,
-75%), compile time halves, and the JAX path inherits scipy accuracy for fixed
geometry. Concreteness is tested positively (`_is_concrete_array`: ndarray, or
a jax.Array that is not a Tracer via sys.modules — never np.asarray in a try);
a tracer among the key values or a traced grid falls through to the direct
JAX call, so nothing branches on a traced value. `memo_stats()['jax_folds']`
is the witness. Numpy behaviour unchanged.

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 (#604) ships as three PRs:

  1. perf: evaluate a constant grid shift-and-rotate at compile time under JAX (memo phase 2, PyAutoGalaxy#604) PyAutoArray#520 — 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) #605 (this one) — 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 65af112 into main Sep 3, 2026
4 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.

perf: fold fixed-geometry deflections out of the JAX trace (memo phase 2)

1 participant