Skip to content

feat: AnalysisPoint declares gradient_mode = "forward" - #752

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

AnalysisPoint now declares gradient_mode = "forward", so gradient searches (af.MultiStartAdam etc.) differentiate the point-source likelihood in forward mode. autolens_profiling #327/#331 measured this as 2–4.5× faster per gradient call and up to 8× faster to compile than reverse mode, with no crossover through 24 free parameters. The reason is the likelihood's inner jax.jacfwd lensing Hessian.

Depends on PyAutoFit#1649 (adds gradient_mode); merge that first.

API Changes

  • al.AnalysisPoint.gradient_mode = "forward" (was the implicit reverse default). Values and gradients are unchanged, since both modes compute the same gradient; only speed, compile time and memory differ.
  • Override per search: af.MultiStartAdam(gradient_mode="reverse").

See full details below.

Test Plan

  • New test_autolens/point/model/test_analysis_point_gradient_mode.py (3 tests): the declared mode; forward ≡ reverse gradient of a FitPositionsSourceSolved likelihood through af.Fitness over the prior medians plus PRNGKey 0..15 (worst relative difference 1.8e-12, finite and non-zero); a short MultiStart fit forward vs reverse. 2 of the 3 failed on unfixed PyAutoLens.
  • Full PyAutoLens suite: 762 passed, 1 xfailed.
  • The same tests pass on a RAL A100 (CUDA), job 359057.
  • Note: the two JAX checks run in a subprocess. In-process, building a JAX Fitness registers Galaxy as a pytree through autofit.jax.register_model, and test_static_lattice_jax.py's later register_tracer_classes then raises "Duplicate custom PyTreeDef type registration". This is a pre-existing latent issue, tracked separately.
Full API Changes (for automation & release notes)

Changed Behaviour

  • autolens.point.model.analysis.AnalysisPoint.gradient_mode = "forward": gradient searches use jax.jacfwd over the flat parameter vector for point-source fits unless the search overrides it.

Migration

  • None required. To restore reverse mode: af.MultiStartAdam(gradient_mode="reverse").

Generated by the PyAutoLabs agent workflow.

🤖 Generated with Claude Code

The source-plane point-source likelihood carries an inner forward-mode
lensing Hessian, so reverse mode runs reverse-over-forward through every
mass profile; jax.jacfwd over the flat parameter vector measured 2-4.5x
faster and up to 8x faster to compile through 24 free parameters
(autolens_profiling #327/#331). Gradient searches (Fitness.grad,
af.MultiStartAdam & co.) now use forward mode for AnalysisPoint by
default; override per search with af.MultiStartAdam(gradient_mode="reverse").

Tests: the declaration; forward == reverse value and gradient of a
FitPositionsSourceSolved likelihood through Fitness over prior medians +
PRNGKey 0..15 draws (finite, non-zero); a short real MultiStartAdam
point-source fit in the declared (forward) mode matching the reverse
override. The two JAX checks run in a subprocess because building a JAX
Fitness registers Galaxy via autofit.jax.register_model, after which
autoarray's register_instance_pytree(Galaxy) (used by the PointSolver
lattice tests) raises on the duplicate registration.

Depends on PyAutoFit#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 b3c9b68 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.

1 participant