feat: AnalysisPoint declares gradient_mode = "forward" - #752
Merged
Merged
Conversation
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>
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
AnalysisPointnow declaresgradient_mode = "forward", so gradient searches (af.MultiStartAdametc.) 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 innerjax.jacfwdlensing 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.af.MultiStartAdam(gradient_mode="reverse").See full details below.
Test Plan
test_autolens/point/model/test_analysis_point_gradient_mode.py(3 tests): the declared mode; forward ≡ reverse gradient of aFitPositionsSourceSolvedlikelihood throughaf.Fitnessover 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.FitnessregistersGalaxyas a pytree throughautofit.jax.register_model, andtest_static_lattice_jax.py's laterregister_tracer_classesthen 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 usejax.jacfwdover the flat parameter vector for point-source fits unless the search overrides it.Migration
af.MultiStartAdam(gradient_mode="reverse").Generated by the PyAutoLabs agent workflow.
🤖 Generated with Claude Code