Overview
Phase 1's fixed-geometry deflection memo (autogalaxy/profiles/mass/abstract/deflections_memo.py, PyAutoGalaxy#602) is a no-op on the JAX path at three independent gates: the xp is not np early return in deflections_yx_2d_from, the isinstance(values, np.ndarray) test in _grid_fingerprint, and _scalar_token returning None for a tracer. Under the JAX likelihood, however, the Gaussian geometry and the grid are concrete while only mass_to_light_ratio is a tracer, so the entire Faddeeva subgraph is a constant that JAX nonetheless re-traces and re-executes every evaluation. XLA cannot fold it away either, because this stack's PyAutoNerves/autonerves/jax_wrapper.py sets XLA_FLAGS to disable XLA constant folding — so this phase folds it explicitly, evaluating the unit-ratio field at trace time with numpy/scipy and embedding it as a constant that the traced ratio multiplies.
Plan
- Add positive concreteness tests to the memo module — a concrete scalar (reusing
mge._is_static_scalar) and a concrete array (an np.ndarray, or a jax.Array that is not a jax.core.Tracer, with jax imported lazily and only when the backend is not numpy). Never a try: np.asarray(...).
- Teach the grid fingerprint to accept a concrete JAX-backed grid (hash
np.asarray(values); a tracer still returns None), and cache a numpy twin of that grid beside the fingerprint so the miss path has something numpy can evaluate on.
- Add the JAX branch to
deflections_yx_2d_from: when the backend is JAX, the memo is enabled, the grid is concrete and the profile token is exact, evaluate the field with numpy/scipy on the twin and return ratio * jnp.asarray(field) (L2) or jnp.asarray(field) (L1). Both backends share one memo dict, so the stored bytes are identical.
- Any tracer among the key values, or a traced grid (free
grid_offset / grid_rotation_angle), falls straight through to the direct JAX call — no data-dependent branching on a traced value anywhere.
- Document the recompilation contract in the module docstring, add a
jax_folds counter to memo_stats(), and prove it: a _wofz backend-split call-count witness, a jaxpr op-count delta, a ≤ 1e-9 likelihood agreement, timings, and the free-grid_offset / kill-switch controls.
Detailed implementation plan
Affected Repositories
- PyAutoGalaxy (primary)
- autolens_profiling
- autolens_workspace_test (read / report only — no edit expected)
Branch Survey
| Repository |
Current Branch |
Dirty? |
| ./PyAutoGalaxy |
main |
clean |
| ./autolens_profiling |
main |
clean (6 untracked A100 result files, unrelated) |
Suggested branch: feature/gaussian-precompute-p2
Implementation Steps
- Concreteness helpers in
autogalaxy/profiles/mass/abstract/deflections_memo.py:
_is_concrete_scalar reuses mge._is_static_scalar (Python / numpy scalar, type module not jax/jaxlib) rather than duplicating it; _is_concrete_array(a, xp) returns True for an np.ndarray, and when xp is not np lazily imports jax and returns isinstance(a, jax.Array) and not isinstance(a, jax.core.Tracer). A 0-d concrete JAX array reaching the memo as a parameter value tokenises via float(np.asarray(v)).
- Grid fingerprint on JAX (
_grid_fingerprint): accept a concrete JAX array by hashing np.asarray(values) (one device→host copy, at trace time only, cached by id(grid) behind the existing weakref as today); a tracer returns None → direct call. Cache the numpy twin of the grid beside the fingerprint — type(grid)(values=np.asarray(grid.array), mask=grid.mask) for Grid2D, Grid2DIrregular(values=...) for the irregular case; a grid type that cannot be twinned falls through to the direct call.
- The JAX branch of
deflections_yx_2d_from, taken when xp is not np, the memo is enabled, the grid is concrete and the profile token is exact: a miss evaluates with numpy on the twin (copy.copy(profile) with mass_to_light_ratio = 1.0 for L2 — the copy must carry no tracer in any constructor argument), stores the numpy array in the same memo dict the numpy path uses, and returns ratio * jnp.asarray(field) (L2) or jnp.asarray(field) (L1) via _wrapped_result(..., xp=xp). _takes_ratio_split accepts a concrete-scalar or tracer / 0-d-array ratio for the Gaussian classes; GaussianGradient stays L1. Any tracer among the key values, or a traced grid, → direct call.
- Docstring + counters: state that recompilation happens only when the embedded constant changes (a new model is a new fit), that the trace-time cost is one numpy evaluation per (profile geometry, grid) per compile, that the JAX path inherits scipy accuracy for fixed geometry, and that
AUTOGALAXY_DEFLECTIONS_MEMO=0 / memo_disabled() govern both backends. memo_stats() gains a jax_folds counter (trace-time numpy evaluations performed for a JAX caller) so the witness can read it.
- Unit tests (numpy-only —
test_autogalaxy never imports jax) in test_autogalaxy/profiles/mass/abstract/test_deflections_memo.py: _is_concrete_array on an ndarray / a list / a stand-in whose type module is jax._src.core; the numpy-twin builder reproducing the same array and mask; jax_folds staying 0 when xp is np.
- JAX validation + timing cell
autolens_profiling/scripts/imaging/likelihood_runtime/mge_mass_jax.py, mirroring likelihood_runtime/mge.py's harness (dataset build, al.AnalysisImaging(dataset, use_jax=True), pytree registration, jax.jit(jax.vmap(...)), vmap_first_call / vmap_steady_x10 sections, artifact JSON under results/runtime/imaging/<cell>/) with the SLaM-shaped model from phase 1's pixelization_numba_mge_mass.py (a Basis of 30 fixed lmp_linear.Gaussians with one free mass_to_light_ratio, NFWSph + shear, rectangular bilinear source, hst). Both legs run in one process — memo-off (memo_disabled()) then memo-on — reporting the _wofz backend-split witness, the jaxpr op-count delta, the ≤ 1e-9 likelihood agreement, the timings and the free-grid_offset / kill-switch controls; recorded in results/notes/numpy_deflections_cpu.md.
- Report-only sweeps: every
autolens_workspace_test/scripts/imaging/jax_likelihood/*.py before and after, pins tabled and none edited (mge.py's bulge is a GaussianGradient basis with free gradient parameters → not memoisable → expected unchanged); autolens_profiling/scripts/lens/deflections/{total,dark,stellar,basis}.py numpy pins hold.
Key Files
autogalaxy/profiles/mass/abstract/deflections_memo.py — the phase-1 memo; all of this phase's library changes land here.
autogalaxy/profiles/mass/abstract/mge.py — _is_static_scalar (the concreteness precedent this phase reuses), _wofz (the witness hook), the Weideman-32 JAX Faddeeva.
PyAutoArray/autoarray/structures/grids/uniform_2d.py:762-810 — Grid2D.subtracted_and_rotated_from, where a JAX grid becomes a concrete jax.Array (and a tracer only when grid_offset / grid_rotation_angle are free).
PyAutoArray/autoarray/fit/fit_dataset.py:186-212 — FitDataset.grids, which rebuilds the grid every call (why the memo is content-keyed, not identity-keyed).
PyAutoArray/autoarray/structures/decorators/to_vector_yx.py — VectorYXMaker, which calls no np. function on the values, so a jnp field re-wraps cleanly with xp=jnp.
PyAutoNerves/autonerves/jax_wrapper.py — the XLA_FLAGS that disable XLA constant folding, which is why the fold has to be explicit.
autolens_profiling/scripts/imaging/likelihood_runtime/mge.py — the JAX likelihood-runtime harness the new cell mirrors.
autolens_workspace_test/scripts/imaging/jax_likelihood/mge.py — the Fitness._vmap standalone validation pattern, and the report-only pin surface.
Original Prompt
Click to expand starting prompt
Gaussian precompute phase 2: JAX trace-time constant — fold the fixed-geometry deflection field out of the jaxpr
Type: feature
Epic: gaussian-deflections-precompute
Phase: 2
Target: autogalaxy
Repos:
- @PyAutoGalaxy
- @autolens_profiling
- @autolens_workspace_test
Themes:
- numba-cpu
- mass-profiles
- jax
- profiling
Difficulty: medium
Autonomy: supervised
Priority: medium
Status: formalised
Filed: 2026-09-03
Parent: draft/feature/autogalaxy/precompute_fixed_geometry_gaussian_deflections.md
Phase 2 of the gaussian-deflections-precompute epic — ledger
draft/feature/autogalaxy/precompute_fixed_geometry_gaussian_deflections.md. Successor work to the
completed numpy-deflections-cpu epic (complete/archive/epics/numpy_deflections_cpu_speedup.md).
Phase 1 is a hard predecessor and is SHIPPED (2026-09-03, record
complete/2026/09/gaussian-precompute-p1.md; PyAutoGalaxy#602 + autolens_profiling#214): it landed
deflections_memo.py, the content-keyed grid fingerprint with its weakref cache, the L1/L2 levels
and the Galaxy / Basis summation-site hooks that this phase extends. This is the "JAX doesn't use this so would help there" half of the user's idea.
Goal
Under JAX the fixed Gaussian geometry is concrete and only mass_to_light_ratio is a tracer, so the
whole Faddeeva subgraph is a constant that JAX nonetheless re-traces and re-executes. Compute the
unit-ratio field at trace time with numpy (scipy wofz, exact) and embed it as a constant: the
Faddeeva subgraph leaves the jaxpr entirely and the JAX path inherits scipy accuracy for fixed geometry.
Steps
- Concrete-value test in
deflections_memo.py: a value is concrete when
type(value).__module__.startswith(("numpy", "jax", "jaxlib")) — the same tracer-detection precedent
used, without importing jax, by autogalaxy/jax/registration.py:93-108 (_is_builtin). Fixed
parameters reach the instance as plain Python floats/tuples and free ones as tracers
(PyAutoFit/autofit/mapper/prior_model/prior_model.py:495-530; Constant subclasses float;
fitness.py:727-731 closes over the model and vmaps only the parameter vector). The instance carries
no free/fixed record — the values do.
- JAX branch of the memo: when
xp is JAX, the grid is concrete and every geometry value is concrete
while m2l is a tracer, evaluate the unit field with numpy/scipy and return m2l * jnp.asarray(field).
If anything in the key is a tracer, fall through unchanged — no data-dependent branching on a traced
value anywhere.
- Recompilation happens only when the embedded constant changes (a new model is a new fit anyway); state
this in the module docstring beside the numpy contract.
- autolens_profiling: JAX before/after in the phase-1 profiling cell (
--xp jax if _driver.py has
it, else a sibling script under scripts/lens/deflections/), recorded in
results/notes/numpy_deflections_cpu.md.
Verification
- Validation is
jax.vmap over the free ratio — never jit-on-concrete, which would fake the win by
constant-folding a fixed-value trace.
- A jaxpr check that no
wofz ops remain in the traced graph for a fixed-geometry Gaussian / MGE stack.
test_autogalaxy green under both backends; ruff check + ruff format --check clean.
autolens_workspace_test JAX likelihood pins for a fixed-MGE lens must hold. They are exact-arithmetic
identical except for the scipy-vs-rational Faddeeva difference (≤ 4e-6) — report any shift, do not edit
the pins in this phase.
- Deflection pins in
scripts/lens/deflections/ unchanged at rtol 1e-6, numpy and JAX.
Ship
Library-first: PyAutoGalaxy PR → autolens_profiling PR. autolens_workspace_test is read/reported only —
no edit expected; if a pin genuinely must move, that is a separate filed finding.
Out of scope
The numpy memo itself (phase 1); the downstream sweep (phase 3); editing JAX likelihood pins; the JAX
20-term omega default or any other JAX numerics decision; convergence_2d_from / potential_2d_from.
Overview
Phase 1's fixed-geometry deflection memo (
autogalaxy/profiles/mass/abstract/deflections_memo.py, PyAutoGalaxy#602) is a no-op on the JAX path at three independent gates: thexp is not npearly return indeflections_yx_2d_from, theisinstance(values, np.ndarray)test in_grid_fingerprint, and_scalar_tokenreturningNonefor a tracer. Under the JAX likelihood, however, the Gaussian geometry and the grid are concrete while onlymass_to_light_ratiois a tracer, so the entire Faddeeva subgraph is a constant that JAX nonetheless re-traces and re-executes every evaluation. XLA cannot fold it away either, because this stack'sPyAutoNerves/autonerves/jax_wrapper.pysetsXLA_FLAGSto disable XLA constant folding — so this phase folds it explicitly, evaluating the unit-ratio field at trace time with numpy/scipy and embedding it as a constant that the traced ratio multiplies.Plan
mge._is_static_scalar) and a concrete array (annp.ndarray, or ajax.Arraythat is not ajax.core.Tracer, withjaximported lazily and only when the backend is not numpy). Never atry: np.asarray(...).np.asarray(values); a tracer still returnsNone), and cache a numpy twin of that grid beside the fingerprint so the miss path has something numpy can evaluate on.deflections_yx_2d_from: when the backend is JAX, the memo is enabled, the grid is concrete and the profile token is exact, evaluate the field with numpy/scipy on the twin and returnratio * jnp.asarray(field)(L2) orjnp.asarray(field)(L1). Both backends share one memo dict, so the stored bytes are identical.grid_offset/grid_rotation_angle), falls straight through to the direct JAX call — no data-dependent branching on a traced value anywhere.jax_foldscounter tomemo_stats(), and prove it: a_wofzbackend-split call-count witness, a jaxpr op-count delta, a ≤ 1e-9 likelihood agreement, timings, and the free-grid_offset/ kill-switch controls.Detailed implementation plan
Affected Repositories
Branch Survey
Suggested branch:
feature/gaussian-precompute-p2Implementation Steps
autogalaxy/profiles/mass/abstract/deflections_memo.py:_is_concrete_scalarreusesmge._is_static_scalar(Python / numpy scalar, type module notjax/jaxlib) rather than duplicating it;_is_concrete_array(a, xp)returnsTruefor annp.ndarray, and whenxp is not nplazily importsjaxand returnsisinstance(a, jax.Array) and not isinstance(a, jax.core.Tracer). A 0-d concrete JAX array reaching the memo as a parameter value tokenises viafloat(np.asarray(v))._grid_fingerprint): accept a concrete JAX array by hashingnp.asarray(values)(one device→host copy, at trace time only, cached byid(grid)behind the existing weakref as today); a tracer returnsNone→ direct call. Cache the numpy twin of the grid beside the fingerprint —type(grid)(values=np.asarray(grid.array), mask=grid.mask)forGrid2D,Grid2DIrregular(values=...)for the irregular case; a grid type that cannot be twinned falls through to the direct call.deflections_yx_2d_from, taken whenxp is not np, the memo is enabled, the grid is concrete and the profile token is exact: a miss evaluates with numpy on the twin (copy.copy(profile)withmass_to_light_ratio = 1.0for L2 — the copy must carry no tracer in any constructor argument), stores the numpy array in the same memo dict the numpy path uses, and returnsratio * jnp.asarray(field)(L2) orjnp.asarray(field)(L1) via_wrapped_result(..., xp=xp)._takes_ratio_splitaccepts a concrete-scalar or tracer / 0-d-array ratio for the Gaussian classes;GaussianGradientstays L1. Any tracer among the key values, or a traced grid, → direct call.AUTOGALAXY_DEFLECTIONS_MEMO=0/memo_disabled()govern both backends.memo_stats()gains ajax_foldscounter (trace-time numpy evaluations performed for a JAX caller) so the witness can read it.test_autogalaxynever imports jax) intest_autogalaxy/profiles/mass/abstract/test_deflections_memo.py:_is_concrete_arrayon an ndarray / a list / a stand-in whose type module isjax._src.core; the numpy-twin builder reproducing the samearrayand mask;jax_foldsstaying 0 whenxp is np.autolens_profiling/scripts/imaging/likelihood_runtime/mge_mass_jax.py, mirroringlikelihood_runtime/mge.py's harness (dataset build,al.AnalysisImaging(dataset, use_jax=True), pytree registration,jax.jit(jax.vmap(...)),vmap_first_call/vmap_steady_x10sections, artifact JSON underresults/runtime/imaging/<cell>/) with the SLaM-shaped model from phase 1'spixelization_numba_mge_mass.py(aBasisof 30 fixedlmp_linear.Gaussians with one freemass_to_light_ratio,NFWSph+ shear, rectangular bilinear source, hst). Both legs run in one process — memo-off (memo_disabled()) then memo-on — reporting the_wofzbackend-split witness, the jaxpr op-count delta, the ≤ 1e-9 likelihood agreement, the timings and the free-grid_offset/ kill-switch controls; recorded inresults/notes/numpy_deflections_cpu.md.autolens_workspace_test/scripts/imaging/jax_likelihood/*.pybefore and after, pins tabled and none edited (mge.py's bulge is aGaussianGradientbasis with free gradient parameters → not memoisable → expected unchanged);autolens_profiling/scripts/lens/deflections/{total,dark,stellar,basis}.pynumpy pins hold.Key Files
autogalaxy/profiles/mass/abstract/deflections_memo.py— the phase-1 memo; all of this phase's library changes land here.autogalaxy/profiles/mass/abstract/mge.py—_is_static_scalar(the concreteness precedent this phase reuses),_wofz(the witness hook), the Weideman-32 JAX Faddeeva.PyAutoArray/autoarray/structures/grids/uniform_2d.py:762-810—Grid2D.subtracted_and_rotated_from, where a JAX grid becomes a concretejax.Array(and a tracer only whengrid_offset/grid_rotation_angleare free).PyAutoArray/autoarray/fit/fit_dataset.py:186-212—FitDataset.grids, which rebuilds the grid every call (why the memo is content-keyed, not identity-keyed).PyAutoArray/autoarray/structures/decorators/to_vector_yx.py—VectorYXMaker, which calls nonp.function on the values, so a jnp field re-wraps cleanly withxp=jnp.PyAutoNerves/autonerves/jax_wrapper.py— theXLA_FLAGSthat disable XLA constant folding, which is why the fold has to be explicit.autolens_profiling/scripts/imaging/likelihood_runtime/mge.py— the JAX likelihood-runtime harness the new cell mirrors.autolens_workspace_test/scripts/imaging/jax_likelihood/mge.py— theFitness._vmapstandalone validation pattern, and the report-only pin surface.Original Prompt
Click to expand starting prompt
Gaussian precompute phase 2: JAX trace-time constant — fold the fixed-geometry deflection field out of the jaxpr
Type: feature
Epic: gaussian-deflections-precompute
Phase: 2
Target: autogalaxy
Repos:
Themes:
Difficulty: medium
Autonomy: supervised
Priority: medium
Status: formalised
Filed: 2026-09-03
Parent: draft/feature/autogalaxy/precompute_fixed_geometry_gaussian_deflections.md
Goal
Under JAX the fixed Gaussian geometry is concrete and only
mass_to_light_ratiois a tracer, so thewhole Faddeeva subgraph is a constant that JAX nonetheless re-traces and re-executes. Compute the
unit-ratio field at trace time with numpy (scipy
wofz, exact) and embed it as a constant: theFaddeeva subgraph leaves the jaxpr entirely and the JAX path inherits scipy accuracy for fixed geometry.
Steps
deflections_memo.py: a value is concrete whentype(value).__module__.startswith(("numpy", "jax", "jaxlib"))— the same tracer-detection precedentused, without importing jax, by
autogalaxy/jax/registration.py:93-108(_is_builtin). Fixedparameters reach the instance as plain Python floats/tuples and free ones as tracers
(
PyAutoFit/autofit/mapper/prior_model/prior_model.py:495-530;Constantsubclassesfloat;fitness.py:727-731closes over the model andvmaps only the parameter vector). The instance carriesno free/fixed record — the values do.
xpis JAX, the grid is concrete and every geometry value is concretewhile
m2lis a tracer, evaluate the unit field with numpy/scipy and returnm2l * jnp.asarray(field).If anything in the key is a tracer, fall through unchanged — no data-dependent branching on a traced
value anywhere.
this in the module docstring beside the numpy contract.
--xp jaxif_driver.pyhasit, else a sibling script under
scripts/lens/deflections/), recorded inresults/notes/numpy_deflections_cpu.md.Verification
jax.vmapover the free ratio — never jit-on-concrete, which would fake the win byconstant-folding a fixed-value trace.
wofzops remain in the traced graph for a fixed-geometry Gaussian / MGE stack.test_autogalaxygreen under both backends;ruff check+ruff format --checkclean.autolens_workspace_testJAX likelihood pins for a fixed-MGE lens must hold. They are exact-arithmeticidentical except for the scipy-vs-rational Faddeeva difference (≤ 4e-6) — report any shift, do not edit
the pins in this phase.
scripts/lens/deflections/unchanged at rtol 1e-6, numpy and JAX.Ship
Library-first: PyAutoGalaxy PR → autolens_profiling PR.
autolens_workspace_testis read/reported only —no edit expected; if a pin genuinely must move, that is a separate filed finding.
Out of scope
The numpy memo itself (phase 1); the downstream sweep (phase 3); editing JAX likelihood pins; the JAX
20-term omega default or any other JAX numerics decision;
convergence_2d_from/potential_2d_from.