Skip to content

Commit 4927738

Browse files
authored
Merge pull request #686 from PyAutoLabs/feature/point-source-defaults-campaign
feat!: point-source position fit defaults to FitPositionsImagePairAllSolved (#678 phase C)
2 parents d582ecf + d4cd012 commit 4927738

6 files changed

Lines changed: 108 additions & 11 deletions

File tree

autolens/fixtures.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -112,6 +112,9 @@ def make_fit_point_dataset_x2_plane():
112112
dataset=make_point_dataset(),
113113
tracer=make_tracer_x2_plane_point(),
114114
solver=make_solver(),
115+
# The tracer's point source is centre-bearing (`ps.Point`), which the
116+
# solved-centre default fit class rejects — use the free-centre pair fit.
117+
fit_positions_cls=al.FitPositionsImagePair,
115118
)
116119

117120

autolens/jax/registration.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,10 +42,21 @@ def register_tracer_classes(tracer) -> bool:
4242
return False
4343

4444
from autoarray.abstract_ndarray import register_instance_pytree
45+
from autogalaxy.galaxy.galaxy import Galaxy
4546
from autolens.lens.tracer import Tracer
4647

4748
register_instance_pytree(Tracer, no_flatten=("cosmology",))
4849

50+
# ``redshift`` rides as aux data, like the cosmology: plane bookkeeping
51+
# (``Tracer.plane_index_via_redshift_from``, reached by any jitted
52+
# ``PointSolver.solve(..., plane_redshift=...)`` on a multi-plane tracer)
53+
# compares redshifts to derive a static plane index, which is impossible
54+
# if the redshift enters the trace as a leaf. On this hand-built /
55+
# simulator path redshifts are per-fit constants; the model-fit path uses
56+
# ``autofit.jax.register_model``, whose classifier already keeps declared
57+
# redshifts constant.
58+
register_instance_pytree(Galaxy, no_flatten=("redshift",))
59+
4960
for galaxy in tracer.galaxies:
5061
_register_object_classes(galaxy)
5162

autolens/point/fit/dataset.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
(image-plane positions, fluxes, and/or time delays) simultaneously. It creates and
66
stores individual fit objects for each component that is present in the dataset:
77
8-
- ``FitPositionsImagePair`` (or another positions fit class) — fits image-plane positions.
8+
- ``FitPositionsImagePairAllSolved`` (or another positions fit class) — fits image-plane positions.
99
- ``FitFluxes`` — fits flux ratios (if fluxes are in the dataset).
1010
- ``FitTimeDelays`` — fits time delays (if time delays are in the dataset).
1111
@@ -21,7 +21,7 @@ class is used by ``AnalysisPoint`` as the evaluation engine inside the
2121
from autolens.point.fit.times_delays import FitTimeDelays
2222
from autolens.lens.tracer import Tracer
2323

24-
from autolens.point.fit.positions.image.pair import FitPositionsImagePair
24+
from autolens.point.fit.positions.image.pair_all import FitPositionsImagePairAllSolved
2525
from autolens import exc
2626

2727

@@ -31,7 +31,7 @@ def __init__(
3131
dataset: PointDataset,
3232
tracer: Tracer,
3333
solver: PointSolver,
34-
fit_positions_cls=FitPositionsImagePair,
34+
fit_positions_cls=FitPositionsImagePairAllSolved,
3535
xp=np,
3636
fit_flux_cls=FitFluxes,
3737
fit_time_delays_cls=FitTimeDelays,

autolens/point/model/analysis.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@
2323

2424
from autolens.analysis.analysis.lens import AnalysisLens
2525
from autolens.analysis.exceptions import raise_fit_exception
26-
from autolens.point.fit.positions.image.pair_repeat import FitPositionsImagePairRepeat
26+
from autolens.point.fit.positions.image.pair_all import FitPositionsImagePairAllSolved
2727
from autolens.point.fit.dataset import FitPointDataset
2828
from autolens.point.fit.fluxes import FitFluxes
2929
from autolens.point.fit.times_delays import FitTimeDelays
@@ -41,7 +41,7 @@ def __init__(
4141
self,
4242
dataset: PointDataset,
4343
solver: PointSolver,
44-
fit_positions_cls=FitPositionsImagePairRepeat,
44+
fit_positions_cls=FitPositionsImagePairAllSolved,
4545
image=None,
4646
cosmology: ag.cosmo.LensingCosmology = None,
4747
title_prefix: str = None,

test_autolens/point/fit/test_fit_dataset.py

Lines changed: 56 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,10 @@ def test__fit_dataset__matching_point_name__positions_log_likelihood_correct(
3333
)
3434

3535
fit = al.FitPointDataset(
36-
dataset=dataset, tracer=point_source_tracer, solver=mock_solver
36+
dataset=dataset,
37+
tracer=point_source_tracer,
38+
solver=mock_solver,
39+
fit_positions_cls=al.FitPositionsImagePair,
3740
)
3841

3942
assert fit.positions.log_likelihood == pytest.approx(-22.14472, 1.0e-4)
@@ -111,13 +114,64 @@ def test__fit_dataset__positions_and_flux__both_log_likelihoods_correct_and_sum(
111114
fluxes_noise_map=flux_noise_map,
112115
)
113116

114-
fit = al.FitPointDataset(dataset=dataset, tracer=tracer, solver=solver)
117+
fit = al.FitPointDataset(
118+
dataset=dataset,
119+
tracer=tracer,
120+
solver=solver,
121+
fit_positions_cls=al.FitPositionsImagePair,
122+
)
115123

116124
assert fit.positions.log_likelihood == pytest.approx(-22.14472, 1.0e-4)
117125
assert fit.flux.log_likelihood == pytest.approx(-2.9920449, 1.0e-4)
118126
assert fit.log_likelihood == fit.positions.log_likelihood + fit.flux.log_likelihood
119127

120128

129+
def test__fit_dataset__default_positions_fit_is_all_to_all_solved(
130+
positions_and_noise, mock_solver
131+
):
132+
"""
133+
The #678 phase B evidence campaign moved the demonstrated defaults to
134+
solved centres with all-to-all pairing: the missing-image discriminator
135+
showed repeat pairing catastrophically mis-ranks truth when an observed
136+
image is absent, while the all-to-all Occam mixture absorbs it.
137+
"""
138+
positions, noise_map = positions_and_noise
139+
dataset = al.PointDataset(
140+
name="point_0", positions=positions, positions_noise_map=noise_map
141+
)
142+
143+
solved_tracer = al.Tracer(
144+
galaxies=[
145+
al.Galaxy(redshift=0.5, mass=al.mp.IsothermalSph(einstein_radius=1.0)),
146+
al.Galaxy(redshift=1.0, point_0=al.ps.PointSolved()),
147+
]
148+
)
149+
150+
fit = al.FitPointDataset(dataset=dataset, tracer=solved_tracer, solver=mock_solver)
151+
152+
assert fit.fit_positions_cls is al.FitPositionsImagePairAllSolved
153+
assert isinstance(fit.positions, al.FitPositionsImagePairAllSolved)
154+
assert np.isfinite(fit.log_likelihood)
155+
156+
157+
def test__fit_dataset__default_with_centre_bearing_profile_raises_loudly(
158+
point_source_tracer, positions_and_noise, mock_solver
159+
):
160+
# A `ps.Point` centre would be silently ignored by a solved-centre fit, so
161+
# the mismatch must raise, pointing the user at the free-centre class.
162+
positions, noise_map = positions_and_noise
163+
dataset = al.PointDataset(
164+
name="point_0", positions=positions, positions_noise_map=noise_map
165+
)
166+
167+
fit = al.FitPointDataset(
168+
dataset=dataset, tracer=point_source_tracer, solver=mock_solver
169+
)
170+
171+
with pytest.raises(al.exc.PointProfileMismatchException):
172+
fit.positions.log_likelihood
173+
174+
121175
def test__fit_dataset__fit_flux_cls_and_fit_time_delays_cls_hooks_are_forwarded_and_default_unchanged(
122176
point_source_tracer, positions_and_noise, mock_solver
123177
):

test_autolens/point/model/test_analysis_point.py

Lines changed: 33 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -69,6 +69,15 @@ def _test__make_result__result_imaging_is_returned(point_dataset):
6969
assert isinstance(result, ResultPoint)
7070

7171

72+
def test__default_fit_positions_cls_is_all_to_all_solved(point_dataset):
73+
# #678 phase B defaults decision: solved centres + all-to-all pairing.
74+
solver = al.m.MockPointSolver(model_positions=point_dataset.positions)
75+
76+
analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False)
77+
78+
assert analysis.fit_positions_cls is al.FitPositionsImagePairAllSolved
79+
80+
7281
def test__figure_of_merit__matches_correct_fit_given_galaxy_profiles(
7382
positions_x2, positions_x2_noise_map
7483
):
@@ -86,7 +95,12 @@ def test__figure_of_merit__matches_correct_fit_given_galaxy_profiles(
8695

8796
solver = al.m.MockPointSolver(model_positions=positions_x2)
8897

89-
analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False)
98+
analysis = al.AnalysisPoint(
99+
dataset=point_dataset,
100+
solver=solver,
101+
fit_positions_cls=al.FitPositionsImagePairRepeat,
102+
use_jax=False,
103+
)
90104

91105
instance = model.instance_from_unit_vector([])
92106
analysis_log_likelihood = analysis.log_likelihood_function(instance=instance)
@@ -107,7 +121,12 @@ def test__figure_of_merit__matches_correct_fit_given_galaxy_profiles(
107121
model_positions = al.Grid2DIrregular([(0.0, 1.0), (1.0, 2.0)])
108122
solver = al.m.MockPointSolver(model_positions=model_positions)
109123

110-
analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False)
124+
analysis = al.AnalysisPoint(
125+
dataset=point_dataset,
126+
solver=solver,
127+
fit_positions_cls=al.FitPositionsImagePairRepeat,
128+
use_jax=False,
129+
)
111130

112131
analysis_log_likelihood = analysis.log_likelihood_function(instance=instance)
113132

@@ -147,7 +166,12 @@ def test__figure_of_merit__includes_fit_fluxes(
147166

148167
solver = al.m.MockPointSolver(model_positions=positions_x2)
149168

150-
analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False)
169+
analysis = al.AnalysisPoint(
170+
dataset=point_dataset,
171+
solver=solver,
172+
fit_positions_cls=al.FitPositionsImagePairRepeat,
173+
use_jax=False,
174+
)
151175

152176
instance = model.instance_from_unit_vector([])
153177

@@ -179,7 +203,12 @@ def test__figure_of_merit__includes_fit_fluxes(
179203
model_positions = al.Grid2DIrregular([(0.0, 1.0), (1.0, 2.0)])
180204
solver = al.m.MockPointSolver(model_positions=model_positions)
181205

182-
analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False)
206+
analysis = al.AnalysisPoint(
207+
dataset=point_dataset,
208+
solver=solver,
209+
fit_positions_cls=al.FitPositionsImagePairRepeat,
210+
use_jax=False,
211+
)
183212

184213
instance = model.instance_from_unit_vector([])
185214
analysis_log_likelihood = analysis.log_likelihood_function(instance=instance)

0 commit comments

Comments
 (0)