perf: fold fixed-geometry deflection fields out of the JAX trace (memo phase 2, #604) - #605
Merged
Merged
Conversation
…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
4 tasks done
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
Closes #604 — phase 2 of the
gaussian-deflections-precomputeepic. 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 freemass_to_light_ratiois a tracer. The Faddeeva subgraph is therefore a constant, butjax.jitstages everyjax.numpycall 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 returnsmass_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 ajax.Arraythat is not ajax.core.Tracer(looked up viasys.modules, no import; nevertry: np.asarray)._concrete_scalar_valuereads a concrete 0-d array back as a float so pytree-native instances key identically to float ones._takes_ratio_splitaccepts a traced or 0-d-array ratio on JAX (it is only ever multiplied)._numpy_profile_frombuilds the numpy-evaluated copy and re-checks every constructor argument is concrete before evaluation._grid_fingerprint_and_twinhashesnp.asarray(values)(one device→host copy per grid object per trace, weakref-cached) and rebuilds aGrid2D/Grid2DIrregulartwin. Any tracer among key values, or a traced grid → direct JAX call.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 fixedlmp_linear.Gaussians + one free ratio, NFWSph + shear, rectangular bilinear source, hst,jit(vmap)batch 3):mge._wofzcalls on the compiling call numpy 180 / jnp 0 with the memo on vs numpy 0 / jnp 240 off, steady state 0/0 both;jax_folds90 (30 Gaussians × 3 grids); jaxpr 53,369 → 13,289 equations (−75%); log-likelihood −56107.56407588691 vs −56107.56407588643 (8.6e-15 relative);vmap_first_call10.8 → 5.4 s (2.0×; 1.64× on a second run);vmap_steady_x10unchanged (2.5–2.6 s — the compiled inversion dominates this fit, not the Faddeeva block). Controls: kill switch reproduces memo-off; a freegrid_offsetmakes the grid a tracer → direct path,jax_folds0, 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_mgeamplitude 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.pylikewise); 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_arrayon ndarray / list / a jax-module stand-in; numpy twin equality;jax_foldsstays 0 under numpy); no jax import in testsscripts/imaging/jax_likelihood/— 15/15 pass before and after; every captured vmap pin bit-identical; none editedscripts/lens/deflections/{total,dark,stellar,basis}.pynumpy pins held at rtol 1e-6ruff check/ruff format --checkclean on touched filesFull API Changes (for automation & release notes)
Removed
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']counterChanged Behaviour
deflections_memo.deflections_yx_2d_fromon 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
Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code
https://claude.ai/code/session_01XhnA4pFN2NycuKc8Ni6s2R