feat: analysis-declared gradient_mode (reverse | forward) for gradient searches - #1649
Merged
Merged
Conversation
…Gradient
An af.Analysis now declares how gradient searches differentiate its
likelihood via the class attribute `gradient_mode` ("reverse" default,
or "forward" = jax.jacfwd over the flat parameter vector). New helper
module autofit/jax/gradient.py (GRADIENT_MODES, resolve_gradient_mode,
value_and_grad_from, grad_from) resolves the declaration against an
optional override and builds the transform; JAX is imported lazily.
- Fitness: optional `gradient_mode` override; `_grad` built through
grad_from(resolved mode); survives pickling, old pickles default None.
- MultiStartGradient: `gradient_mode` keyword (validated at
construction, dict/pickle round-trip), resolved against the analysis
at fit time; both value_and_grad sites (physical and scaler/bijector
stepped objective) go through value_and_grad_from; resolved mode is
logged at search start, persisted in search_internal and reported in
samples_info.
- Defaults unchanged: every analysis stays "reverse" unless it declares
otherwise (PyAutoLens AnalysisPoint will declare "forward", per
autolens_profiling #327/#331).
Refs #1648
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
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
Adds an analysis-declared gradient mode: an
af.Analysiscan declare that its likelihood is best differentiated in forward mode (jax.jacfwdover the flat parameter vector) instead of reverse mode (jax.value_and_grad), andMultiStartGradientcan override the declaration per search. The default stays"reverse", so nothing changes for any analysis that does not declare otherwise.Motivation (autolens_profiling #327, #331): for the source-plane point-source likelihood, forward mode is 2–4.5× faster, up to 8× faster to compile, and has no crossover through 24 free parameters. That likelihood contains an inner forward-mode lensing Hessian, so reverse mode must run reverse-over-forward through it. The companion PyAutoLens PR makes
AnalysisPointdeclare"forward".Closes #1648
API Changes
Additive only; no behaviour change unless an analysis declares
gradient_mode = "forward"or a search passesgradient_mode=.af.Analysis.gradient_mode = "reverse": a new class attribute that analyses may override.af.MultiStartGradient(..., gradient_mode=None)and its subclasses (MultiStartAdam, …): a new keyword.Nonedefers to the analysis; a mistyped value raises at construction.Fitness(..., gradient_mode=None):Fitness.gradhonours the resolved mode.autofit.jax.gradientwith the helpers.samples_info["gradient_mode"].See full details below.
Test Plan
test_autofit/non_linear/test_gradient_mode.py(28 tests): forward ≡ reverse value and gradient; mode precedence; invalid values raise;Fitness.grad; MultiStart forward vs reverse give the same result to rtol 1e-8 on the plain, scaler and bijector paths; pickle and old-pickle round-trips. Before the fix they failed at import.autofitstill imports without JAX.Full API Changes (for automation & release notes)
Added
autofit.jax.gradient.GRADIENT_MODES = ("reverse", "forward")autofit.jax.gradient.validate_gradient_mode(mode): returnsmodeifNoneor valid, else raisesValueError.autofit.jax.gradient.resolve_gradient_mode(analysis, override=None): the override wins, elseanalysis.gradient_mode, else"reverse".autofit.jax.gradient.value_and_grad_from(func, mode):(value, grad)with thejax.value_and_gradcontract. Forward mode usesjax.jacfwd(..., has_aux=True)and traces the likelihood once.autofit.jax.gradient.grad_from(func, mode): the gradient-only counterpart.af.Analysis.gradient_modeclass attribute (default"reverse").Changed Signature
af.MultiStartGradient.__init__(..., gradient_mode: Optional[str] = None)and its subclasses.Fitness.__init__(..., gradient_mode: Optional[str] = None).Changed Behaviour
MultiStartGradientbuilds both of itsvalue_and_gradobjectives (physical, and scaler/bijector-stepped) throughvalue_and_grad_fromwith the resolved mode. In reverse mode this is identical to before.Fitness._gradbuilds throughgrad_fromwith the resolved mode.samples_infogains a"gradient_mode"key;search_internalpersists the resolved mode.Migration
gradient_mode = "forward"on yourAnalysissubclass, or passaf.MultiStartAdam(gradient_mode="forward"). Forward mode carries one tangent per free parameter; if memory is tight, setbatch_sizeor override to"reverse".Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code