Overview
ag.mp.Isothermal(...).convergence_2d_from(grid, xp=jnp) cannot be traced under jax.jit when ell_comps are tracers: it raises TracerArrayConversionError. The cause is that PowerLawCore.convergence_2d_from does not forward xp to convergence_func, so NumPy is used on JAX tracers. This blocks jit/grad of convergence for Isothermal and every PowerLawCore subclass (found in point-source phase 2b).
Plan
- Forward
xp in the one convergence_func(...) call that drops it (PowerLawCore.convergence_2d_from).
- Add a jax-free regression test asserting
convergence_func receives the xp passed to convergence_2d_from.
- Confirm with a scratch jit witness that jit output equals NumPy output.
Detailed implementation plan
Affected Repositories
Branch Survey
| Repository |
Current Branch |
Dirty? |
| ./PyAutoGalaxy |
main |
clean |
Suggested branch: feature/isothermal-convergence-jit
Implementation Steps
autogalaxy/profiles/mass/total/power_law_core.py:111: return self.convergence_func(grid_radius=grid_eta, xp=xp).
test_autogalaxy/profiles/mass/total/test_isothermal.py: pass a numpy-delegating spy module as xp, wrap the instance's convergence_func to capture its xp kwarg, and assert it is the spy. Run red on unfixed main, then green; run pytest test_autogalaxy/profiles/mass.
- Scratch witness:
jax.jit over ell_comps gives the same convergence as NumPy.
Key Files
autogalaxy/profiles/mass/total/power_law_core.py — the call that drops xp
test_autogalaxy/profiles/mass/total/test_isothermal.py — regression test
Original Prompt
Click to expand starting prompt
Isothermal.convergence_2d_from not jit-traceable (PowerLawCore drops xp)
Type: bug
Target: @PyAutoGalaxy
Original request (verbatim)
fix this - Isothermal.convergence_2d_from still can't be traced under jax.jit (the phase 2b bug).
Carried from the active.md row point-source-source-plane-p2c carried: intake list.
Witness (reproduced 2026-09-27 on PyAutoGalaxy main)
With ell_comps as JAX tracers, ag.mp.Isothermal(...).convergence_2d_from(grid, xp=jnp)
under jax.jit raises TracerArrayConversionError.
Root cause
PowerLawCore.convergence_2d_from (autogalaxy/profiles/mass/total/power_law_core.py:111)
calls self.convergence_func(grid_radius=grid_eta) without xp, so the default np is
used and Isothermal.axis_ratio(np) calls np.minimum on a tracer. The shared method means
PowerLaw, PowerLawIntermediate, IsothermalCore and other subclasses are affected too. It is
the only convergence_func(...) call site that omits xp.
Plan
power_law_core.py:111 -> return self.convergence_func(grid_radius=grid_eta, xp=xp).
- jax-free regression test in
test_autogalaxy/profiles/mass/total/test_isothermal.py:
call convergence_2d_from with a numpy-delegating spy xp module and assert
convergence_func receives that same xp. Red on unfixed main, then green.
- Scratch jit witness: jit output equals NumPy output.
Overview
ag.mp.Isothermal(...).convergence_2d_from(grid, xp=jnp)cannot be traced underjax.jitwhenell_compsare tracers: it raisesTracerArrayConversionError. The cause is thatPowerLawCore.convergence_2d_fromdoes not forwardxptoconvergence_func, so NumPy is used on JAX tracers. This blocks jit/grad of convergence for Isothermal and every PowerLawCore subclass (found in point-source phase 2b).Plan
xpin the oneconvergence_func(...)call that drops it (PowerLawCore.convergence_2d_from).convergence_funcreceives thexppassed toconvergence_2d_from.Detailed implementation plan
Affected Repositories
Branch Survey
Suggested branch:
feature/isothermal-convergence-jitImplementation Steps
autogalaxy/profiles/mass/total/power_law_core.py:111:return self.convergence_func(grid_radius=grid_eta, xp=xp).test_autogalaxy/profiles/mass/total/test_isothermal.py: pass a numpy-delegating spy module asxp, wrap the instance'sconvergence_functo capture itsxpkwarg, and assert it is the spy. Run red on unfixed main, then green; runpytest test_autogalaxy/profiles/mass.jax.jitoverell_compsgives the same convergence as NumPy.Key Files
autogalaxy/profiles/mass/total/power_law_core.py— the call that dropsxptest_autogalaxy/profiles/mass/total/test_isothermal.py— regression testOriginal Prompt
Click to expand starting prompt
Isothermal.convergence_2d_from not jit-traceable (PowerLawCore drops xp)
Type: bug
Target: @PyAutoGalaxy
Original request (verbatim)
fix this - Isothermal.convergence_2d_from still can't be traced under jax.jit (the phase 2b bug).
Carried from the active.md row point-source-source-plane-p2c
carried:intake list.Witness (reproduced 2026-09-27 on PyAutoGalaxy main)
With
ell_compsas JAX tracers,ag.mp.Isothermal(...).convergence_2d_from(grid, xp=jnp)under
jax.jitraisesTracerArrayConversionError.Root cause
PowerLawCore.convergence_2d_from(autogalaxy/profiles/mass/total/power_law_core.py:111)calls
self.convergence_func(grid_radius=grid_eta)withoutxp, so the defaultnpisused and
Isothermal.axis_ratio(np)callsnp.minimumon a tracer. The shared method meansPowerLaw, PowerLawIntermediate, IsothermalCore and other subclasses are affected too. It is
the only
convergence_func(...)call site that omitsxp.Plan
power_law_core.py:111->return self.convergence_func(grid_radius=grid_eta, xp=xp).test_autogalaxy/profiles/mass/total/test_isothermal.py:call
convergence_2d_fromwith a numpy-delegating spyxpmodule and assertconvergence_funcreceives that samexp. Red on unfixed main, then green.