Overview
Add an analysis-declared gradient_mode ("reverse" | "forward") to PyAutoFit, with a MultiStartGradient override, and declare "forward" on PyAutoLens AnalysisPoint. autolens_profiling #327/#331 measured forward-mode gradients 2–4.5× faster (and up to 8× faster to compile) than jax.value_and_grad for the source-plane point-source likelihood, with no crossover through 24 free parameters. The default stays "reverse" for every other analysis.
Plan
- PyAutoFit helper
autofit/jax/gradient.py: mode resolution plus value_and_grad / grad builders for both modes (forward = jacfwd over the flat vector).
af.Analysis.gradient_mode = "reverse"; Fitness and MultiStartGradient build their gradients through the helper; MultiStartGradient(gradient_mode=...) overrides.
- PyAutoLens
AnalysisPoint.gradient_mode = "forward".
- Parity and precedence tests in both libraries, full suites, RAL A100 regression check.
- Workspace follow-up: autolens_profiling cell confirming the speed-up through the real
MultiStartGradient.
Detailed implementation plan
Affected Repositories
- PyAutoFit (primary)
- PyAutoLens
- autolens_profiling (workspace follow-up)
Branch Survey
| Repository |
Current Branch |
Dirty? |
| ./PyAutoFit |
main |
clean |
| ./PyAutoLens |
main |
clean |
| ./autolens_profiling |
main |
untracked dataset/abell_1201/ only |
Suggested branch: feature/point-source-gradient-mode
Context
Phase 2b (#327) and phase 2c (#331) measured the gradient of the source-plane point-source
likelihood. Computing it in forward mode (jax.jacfwd over the flat parameter vector) instead of
reverse mode (jax.value_and_grad):
- is 2–4.5× faster;
- compiles up to 8× faster (80 s → ≤ 11 s at 24 parameters);
- keeps that lead through 24 free parameters, and the lead grows with model size.
The cause is structural: this likelihood contains an inner forward-mode lensing Hessian, so reverse
mode runs reverse-over-forward through every mass profile. You chose (2026-09-27) to make the mode
an analysis-declared default with a search override, not an n_params threshold. This phase puts
that into the libraries. The default stays "reverse" everywhere except AnalysisPoint, so
nothing changes for any other analysis.
Design
PyAutoFit
- New helper module
autofit/jax/gradient.py:
GRADIENT_MODES = ("reverse", "forward");
resolve_gradient_mode(analysis, override=None) -> str. The override wins, otherwise
getattr(analysis, "gradient_mode", "reverse"). An unknown value raises a clear ValueError.
value_and_grad_from(func, mode) -> callable returning (value, grad) with the same contract
as jax.value_and_grad(func):
- reverse →
jax.value_and_grad(func);
- forward →
jax.jacfwd(lambda v: (func(v), func(v)), has_aux=True), swapped to
(value, grad). This is exactly the phase-2b/2c fwd route, over the flat vector.
grad_from(func, mode): the gradient-only counterpart, for Fitness._grad.
af.Analysis: class attribute gradient_mode = "reverse", with a docstring saying when to
declare "forward": when the likelihood contains an inner forward-mode derivative and has few
parameters. The docstring points to the phase-2c result.
Fitness (autofit/non_linear/fitness.py): an optional gradient_mode=None constructor argument.
_grad (:917) builds through grad_from(self.call, resolve_gradient_mode(self.analysis, self.gradient_mode)). The existing log_on_first_compile wrapper and the pickling strip of
_grad (:833) stay unchanged.
MultiStartGradient (autofit/non_linear/search/mle/multi_start_gradient/search.py): a new
keyword gradient_mode: Optional[str] = None, stored and documented, with validation at
construction.
- The mode is resolved against the analysis at fit time.
- Both
jax.value_and_grad sites (:974 physical objective; :1072 scaler/bijector stepped
objective) go through value_and_grad_from(..., mode).
- The downstream code is unchanged:
_value_and_grad_finite, the vmap at :1089, batching,
_broad_starts, value_and_grad_single=jax.jit(_value_and_grad) at :1256.
- The resolved mode is logged once at search start and recorded in the search's output metadata.
- Out of scope: blackjax NUTS/SMC (they differentiate the log-density themselves; a later
step), graphical/EP factor gradients, and changing any default other than AnalysisPoint's.
PyAutoLens
6. AnalysisPoint (autolens/point/model/analysis.py:36): gradient_mode = "forward", with a
docstring citing autolens_profiling #327/#331. Users override per search with
af.MultiStartGradient(gradient_mode="reverse").
Memory caveat (documented): forward mode carries one tangent per free parameter, so under vmap
its memory scales with n_starts × n_params. This is why it is analysis-declared rather than
global. The MultiStartGradient docstring says so next to batch_size.
Tests
- PyAutoFit (
test_autofit/):
- helper parity: forward ≡ reverse value and gradient on a toy JAX analysis (Gaussian model) at
several vectors, rtol 1e-10;
resolve_gradient_mode: analysis default, override precedence, invalid value raises;
Fitness.grad honours the mode;
MultiStartGradient with gradient_mode="forward" vs "reverse" on the toy model with the same
seed gives the same result to tolerance, including the scaler/bijector path;
- invalid keyword raises at construction.
- Existing
test_multi_start_gradient.py stays green.
- PyAutoLens (
test_autolens/point/):
AnalysisPoint.gradient_mode == "forward";
- forward ≡ reverse gradient of a
FitPositionsSourceSolved likelihood through Fitness
(registered model; non-zero, finite; PRNGKey sweep 0..15);
- one short
MultiStartGradient point-source fit that runs in forward mode end-to-end.
- Full PyAutoFit and PyAutoLens suites; GPU regression check on the RAL A100 (the multi-start
point-source test in forward and reverse mode) before ship.
Workspace follow-up (after the library PRs merge)
- autolens_profiling:
scripts/point_source_source/likelihood_breakdown/gradient_mode_library_ab.py.
It runs the real MultiStartGradient step on the phase-2c L5 and L24 rungs with
gradient_mode="forward" vs "reverse", on RAL gpu-node CPU + A100. It confirms the phase-2b/2c
gain arrives through the library path. It also adds a campaign-note section.
- autolens_workspace: no script change needed, because the default flips inside
AnalysisPoint.
If a point-source gradient example exists, add a one-line comment on the new keyword.
Execution
- Task
point-source-gradient-mode; new Mind prompt active/point_source_source_plane_phase_2d.md.
Issue on PyAutoFit, cross-referenced from PyAutoLens.
- Branch
feature/point-source-gradient-mode in PyAutoFit, PyAutoLens and (later) autolens_profiling,
under worktree ~/Code/PyAutoLabs-wt/point-source-gradient-mode.
- Parallel claims (approval requested with this plan):
- PyAutoLens is claimed by
workspace-config-cleanup (its Lens PR is already merged, awaiting
release) and by pointsolver-mcs-headroom (point solver). This phase touches only
autolens/point/model/analysis.py plus a new test.
- autolens_profiling is claimed by
pointsolver-mcs-headroom and interferometer-mesh-numba-p2.
The follow-up touches only a new cell, its results and the source-plane note.
- PyAutoFit is unclaimed.
- Order:
/start_library → implement (Opus subagent, progress file + Monitor) → /ship_library
(PyAutoFit PR first, PyAutoLens PR depends on it) → /start_workspace for the profiling
follow-up → /ship_workspace.
- Related:
jax-grad-nan-zero-components (PyAutoGalaxy#631) is already fixing the NaN-at-zero
gradients found in phase 2c, and is not touched here.
Verification
pytest test_autofit and pytest test_autolens green in the worktree. The new tests fail on
unfixed main first (e.g. AnalysisPoint.gradient_mode missing).
- RAL A100: the point-source multi-start test in both modes, results identical to tolerance.
- The workspace follow-up rows show the forward-mode speed-up through
MultiStartGradient itself.
Original Prompt
Click to expand starting prompt
PyAutoMind campaign prompt draft/research/autolens_profiling/point_source_source_plane_chi_squared_speed.md, "Phase 2c shipped — next is phase 2d": analysis-declared gradient_mode, library-first (human decision 2026-09-27).
Overview
Add an analysis-declared
gradient_mode("reverse" | "forward") to PyAutoFit, with aMultiStartGradientoverride, and declare"forward"on PyAutoLensAnalysisPoint. autolens_profiling #327/#331 measured forward-mode gradients 2–4.5× faster (and up to 8× faster to compile) thanjax.value_and_gradfor the source-plane point-source likelihood, with no crossover through 24 free parameters. The default stays"reverse"for every other analysis.Plan
autofit/jax/gradient.py: mode resolution plusvalue_and_grad/gradbuilders for both modes (forward =jacfwdover the flat vector).af.Analysis.gradient_mode = "reverse";FitnessandMultiStartGradientbuild their gradients through the helper;MultiStartGradient(gradient_mode=...)overrides.AnalysisPoint.gradient_mode = "forward".MultiStartGradient.Detailed implementation plan
Affected Repositories
Branch Survey
dataset/abell_1201/onlySuggested branch:
feature/point-source-gradient-modeContext
Phase 2b (#327) and phase 2c (#331) measured the gradient of the source-plane point-source
likelihood. Computing it in forward mode (
jax.jacfwdover the flat parameter vector) instead ofreverse mode (
jax.value_and_grad):The cause is structural: this likelihood contains an inner forward-mode lensing Hessian, so reverse
mode runs reverse-over-forward through every mass profile. You chose (2026-09-27) to make the mode
an analysis-declared default with a search override, not an
n_paramsthreshold. This phase putsthat into the libraries. The default stays
"reverse"everywhere exceptAnalysisPoint, sonothing changes for any other analysis.
Design
PyAutoFit
autofit/jax/gradient.py:GRADIENT_MODES = ("reverse", "forward");resolve_gradient_mode(analysis, override=None) -> str. The override wins, otherwisegetattr(analysis, "gradient_mode", "reverse"). An unknown value raises a clearValueError.value_and_grad_from(func, mode) -> callablereturning(value, grad)with the same contractas
jax.value_and_grad(func):jax.value_and_grad(func);jax.jacfwd(lambda v: (func(v), func(v)), has_aux=True), swapped to(value, grad). This is exactly the phase-2b/2cfwdroute, over the flat vector.grad_from(func, mode): the gradient-only counterpart, forFitness._grad.af.Analysis: class attributegradient_mode = "reverse", with a docstring saying when todeclare
"forward": when the likelihood contains an inner forward-mode derivative and has fewparameters. The docstring points to the phase-2c result.
Fitness(autofit/non_linear/fitness.py): an optionalgradient_mode=Noneconstructor argument._grad(:917) builds throughgrad_from(self.call, resolve_gradient_mode(self.analysis, self.gradient_mode)). The existinglog_on_first_compilewrapper and the pickling strip of_grad(:833) stay unchanged.MultiStartGradient(autofit/non_linear/search/mle/multi_start_gradient/search.py): a newkeyword
gradient_mode: Optional[str] = None, stored and documented, with validation atconstruction.
jax.value_and_gradsites (:974physical objective;:1072scaler/bijector steppedobjective) go through
value_and_grad_from(..., mode)._value_and_grad_finite, the vmap at:1089, batching,_broad_starts,value_and_grad_single=jax.jit(_value_and_grad)at:1256.step), graphical/EP factor gradients, and changing any default other than
AnalysisPoint's.PyAutoLens
6.
AnalysisPoint(autolens/point/model/analysis.py:36):gradient_mode = "forward", with adocstring citing autolens_profiling #327/#331. Users override per search with
af.MultiStartGradient(gradient_mode="reverse").Memory caveat (documented): forward mode carries one tangent per free parameter, so under vmap
its memory scales with
n_starts × n_params. This is why it is analysis-declared rather thanglobal. The
MultiStartGradientdocstring says so next tobatch_size.Tests
test_autofit/):several vectors, rtol 1e-10;
resolve_gradient_mode: analysis default, override precedence, invalid value raises;Fitness.gradhonours the mode;MultiStartGradientwithgradient_mode="forward"vs"reverse"on the toy model with the sameseed gives the same result to tolerance, including the scaler/bijector path;
test_multi_start_gradient.pystays green.test_autolens/point/):AnalysisPoint.gradient_mode == "forward";FitPositionsSourceSolvedlikelihood throughFitness(registered model; non-zero, finite; PRNGKey sweep 0..15);
MultiStartGradientpoint-source fit that runs in forward mode end-to-end.point-source test in forward and reverse mode) before ship.
Workspace follow-up (after the library PRs merge)
scripts/point_source_source/likelihood_breakdown/gradient_mode_library_ab.py.It runs the real
MultiStartGradientstep on the phase-2c L5 and L24 rungs withgradient_mode="forward"vs"reverse", on RAL gpu-node CPU + A100. It confirms the phase-2b/2cgain arrives through the library path. It also adds a campaign-note section.
AnalysisPoint.If a point-source gradient example exists, add a one-line comment on the new keyword.
Execution
point-source-gradient-mode; new Mind promptactive/point_source_source_plane_phase_2d.md.Issue on PyAutoFit, cross-referenced from PyAutoLens.
feature/point-source-gradient-modein PyAutoFit, PyAutoLens and (later) autolens_profiling,under worktree
~/Code/PyAutoLabs-wt/point-source-gradient-mode.workspace-config-cleanup(its Lens PR is already merged, awaitingrelease) and by
pointsolver-mcs-headroom(point solver). This phase touches onlyautolens/point/model/analysis.pyplus a new test.pointsolver-mcs-headroomandinterferometer-mesh-numba-p2.The follow-up touches only a new cell, its results and the source-plane note.
/start_library→ implement (Opus subagent, progress file + Monitor) →/ship_library(PyAutoFit PR first, PyAutoLens PR depends on it) →
/start_workspacefor the profilingfollow-up →
/ship_workspace.jax-grad-nan-zero-components(PyAutoGalaxy#631) is already fixing the NaN-at-zerogradients found in phase 2c, and is not touched here.
Verification
pytest test_autofitandpytest test_autolensgreen in the worktree. The new tests fail onunfixed main first (e.g.
AnalysisPoint.gradient_modemissing).MultiStartGradientitself.Original Prompt
Click to expand starting prompt
PyAutoMind campaign prompt
draft/research/autolens_profiling/point_source_source_plane_chi_squared_speed.md, "Phase 2c shipped — next is phase 2d": analysis-declaredgradient_mode, library-first (human decision 2026-09-27).