Skip to content

feat: analysis-declared gradient_mode (reverse | forward) for gradient searches - #1649

Merged
Jammy2211 merged 1 commit into
mainfrom
feature/point-source-gradient-mode
Sep 27, 2026
Merged

Jammy2211 merged 1 commit into
mainfrom
feature/point-source-gradient-mode

Conversation

@Jammy2211

Copy link
Copy Markdown
Collaborator

Summary

Adds an analysis-declared gradient mode: an af.Analysis can declare that its likelihood is best differentiated in forward mode (jax.jacfwd over the flat parameter vector) instead of reverse mode (jax.value_and_grad), and MultiStartGradient can 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 AnalysisPoint declare "forward".

Closes #1648

API Changes

Additive only; no behaviour change unless an analysis declares gradient_mode = "forward" or a search passes gradient_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. None defers to the analysis; a mistyped value raises at construction.
  • Fitness(..., gradient_mode=None): Fitness.grad honours the resolved mode.
  • New module autofit.jax.gradient with the helpers.
  • The resolved mode is logged once per fit and recorded in samples_info["gradient_mode"].

See full details below.

Test Plan

  • New 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.
  • Full PyAutoFit suite: 2926 passed, 2 skipped (serial; xdist trips on prior-property test IDs).
  • The same 28 tests pass on a RAL A100 (CUDA), job 359057.
  • autofit still 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): returns mode if None or valid, else raises ValueError.
  • autofit.jax.gradient.resolve_gradient_mode(analysis, override=None): the override wins, else analysis.gradient_mode, else "reverse".
  • autofit.jax.gradient.value_and_grad_from(func, mode): (value, grad) with the jax.value_and_grad contract. Forward mode uses jax.jacfwd(..., has_aux=True) and traces the likelihood once.
  • autofit.jax.gradient.grad_from(func, mode): the gradient-only counterpart.
  • af.Analysis.gradient_mode class 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

  • MultiStartGradient builds both of its value_and_grad objectives (physical, and scaler/bijector-stepped) through value_and_grad_from with the resolved mode. In reverse mode this is identical to before.
  • Fitness._grad builds through grad_from with the resolved mode.
  • samples_info gains a "gradient_mode" key; search_internal persists the resolved mode.

Migration

  • None required. To opt in: declare gradient_mode = "forward" on your Analysis subclass, or pass af.MultiStartAdam(gradient_mode="forward"). Forward mode carries one tangent per free parameter; if memory is tight, set batch_size or override to "reverse".

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

…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>
@Jammy2211 Jammy2211 added the pending-release PR queued for the next release build label Sep 27, 2026
@Jammy2211
Jammy2211 merged commit 867af1c into main Sep 27, 2026
4 checks passed
@Jammy2211
Jammy2211 deleted the feature/point-source-gradient-mode branch September 27, 2026 18:33
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

pending-release PR queued for the next release build

Projects

None yet

Development

Successfully merging this pull request may close these issues.

feat: analysis-declared gradient_mode (reverse | forward) for gradient searches

1 participant