diff --git a/autolens/fixtures.py b/autolens/fixtures.py index 7b13a1a26..f6d3ccf57 100644 --- a/autolens/fixtures.py +++ b/autolens/fixtures.py @@ -112,6 +112,9 @@ def make_fit_point_dataset_x2_plane(): dataset=make_point_dataset(), tracer=make_tracer_x2_plane_point(), solver=make_solver(), + # The tracer's point source is centre-bearing (`ps.Point`), which the + # solved-centre default fit class rejects — use the free-centre pair fit. + fit_positions_cls=al.FitPositionsImagePair, ) diff --git a/autolens/jax/registration.py b/autolens/jax/registration.py index a71a4ab33..77b63ed67 100644 --- a/autolens/jax/registration.py +++ b/autolens/jax/registration.py @@ -42,10 +42,21 @@ def register_tracer_classes(tracer) -> bool: return False from autoarray.abstract_ndarray import register_instance_pytree + from autogalaxy.galaxy.galaxy import Galaxy from autolens.lens.tracer import Tracer register_instance_pytree(Tracer, no_flatten=("cosmology",)) + # ``redshift`` rides as aux data, like the cosmology: plane bookkeeping + # (``Tracer.plane_index_via_redshift_from``, reached by any jitted + # ``PointSolver.solve(..., plane_redshift=...)`` on a multi-plane tracer) + # compares redshifts to derive a static plane index, which is impossible + # if the redshift enters the trace as a leaf. On this hand-built / + # simulator path redshifts are per-fit constants; the model-fit path uses + # ``autofit.jax.register_model``, whose classifier already keeps declared + # redshifts constant. + register_instance_pytree(Galaxy, no_flatten=("redshift",)) + for galaxy in tracer.galaxies: _register_object_classes(galaxy) diff --git a/autolens/point/fit/dataset.py b/autolens/point/fit/dataset.py index c295787d4..26052a0a9 100644 --- a/autolens/point/fit/dataset.py +++ b/autolens/point/fit/dataset.py @@ -5,7 +5,7 @@ (image-plane positions, fluxes, and/or time delays) simultaneously. It creates and stores individual fit objects for each component that is present in the dataset: -- ``FitPositionsImagePair`` (or another positions fit class) — fits image-plane positions. +- ``FitPositionsImagePairAllSolved`` (or another positions fit class) — fits image-plane positions. - ``FitFluxes`` — fits flux ratios (if fluxes are in the dataset). - ``FitTimeDelays`` — fits time delays (if time delays are in the dataset). @@ -21,7 +21,7 @@ class is used by ``AnalysisPoint`` as the evaluation engine inside the from autolens.point.fit.times_delays import FitTimeDelays from autolens.lens.tracer import Tracer -from autolens.point.fit.positions.image.pair import FitPositionsImagePair +from autolens.point.fit.positions.image.pair_all import FitPositionsImagePairAllSolved from autolens import exc @@ -31,7 +31,7 @@ def __init__( dataset: PointDataset, tracer: Tracer, solver: PointSolver, - fit_positions_cls=FitPositionsImagePair, + fit_positions_cls=FitPositionsImagePairAllSolved, xp=np, fit_flux_cls=FitFluxes, fit_time_delays_cls=FitTimeDelays, diff --git a/autolens/point/model/analysis.py b/autolens/point/model/analysis.py index 1221575d1..eddae8908 100644 --- a/autolens/point/model/analysis.py +++ b/autolens/point/model/analysis.py @@ -23,7 +23,7 @@ from autolens.analysis.analysis.lens import AnalysisLens from autolens.analysis.exceptions import raise_fit_exception -from autolens.point.fit.positions.image.pair_repeat import FitPositionsImagePairRepeat +from autolens.point.fit.positions.image.pair_all import FitPositionsImagePairAllSolved from autolens.point.fit.dataset import FitPointDataset from autolens.point.fit.fluxes import FitFluxes from autolens.point.fit.times_delays import FitTimeDelays @@ -41,7 +41,7 @@ def __init__( self, dataset: PointDataset, solver: PointSolver, - fit_positions_cls=FitPositionsImagePairRepeat, + fit_positions_cls=FitPositionsImagePairAllSolved, image=None, cosmology: ag.cosmo.LensingCosmology = None, title_prefix: str = None, diff --git a/test_autolens/point/fit/test_fit_dataset.py b/test_autolens/point/fit/test_fit_dataset.py index 395fa313a..7eb7cc532 100644 --- a/test_autolens/point/fit/test_fit_dataset.py +++ b/test_autolens/point/fit/test_fit_dataset.py @@ -33,7 +33,10 @@ def test__fit_dataset__matching_point_name__positions_log_likelihood_correct( ) fit = al.FitPointDataset( - dataset=dataset, tracer=point_source_tracer, solver=mock_solver + dataset=dataset, + tracer=point_source_tracer, + solver=mock_solver, + fit_positions_cls=al.FitPositionsImagePair, ) 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( fluxes_noise_map=flux_noise_map, ) - fit = al.FitPointDataset(dataset=dataset, tracer=tracer, solver=solver) + fit = al.FitPointDataset( + dataset=dataset, + tracer=tracer, + solver=solver, + fit_positions_cls=al.FitPositionsImagePair, + ) assert fit.positions.log_likelihood == pytest.approx(-22.14472, 1.0e-4) assert fit.flux.log_likelihood == pytest.approx(-2.9920449, 1.0e-4) assert fit.log_likelihood == fit.positions.log_likelihood + fit.flux.log_likelihood +def test__fit_dataset__default_positions_fit_is_all_to_all_solved( + positions_and_noise, mock_solver +): + """ + The #678 phase B evidence campaign moved the demonstrated defaults to + solved centres with all-to-all pairing: the missing-image discriminator + showed repeat pairing catastrophically mis-ranks truth when an observed + image is absent, while the all-to-all Occam mixture absorbs it. + """ + positions, noise_map = positions_and_noise + dataset = al.PointDataset( + name="point_0", positions=positions, positions_noise_map=noise_map + ) + + solved_tracer = al.Tracer( + galaxies=[ + al.Galaxy(redshift=0.5, mass=al.mp.IsothermalSph(einstein_radius=1.0)), + al.Galaxy(redshift=1.0, point_0=al.ps.PointSolved()), + ] + ) + + fit = al.FitPointDataset(dataset=dataset, tracer=solved_tracer, solver=mock_solver) + + assert fit.fit_positions_cls is al.FitPositionsImagePairAllSolved + assert isinstance(fit.positions, al.FitPositionsImagePairAllSolved) + assert np.isfinite(fit.log_likelihood) + + +def test__fit_dataset__default_with_centre_bearing_profile_raises_loudly( + point_source_tracer, positions_and_noise, mock_solver +): + # A `ps.Point` centre would be silently ignored by a solved-centre fit, so + # the mismatch must raise, pointing the user at the free-centre class. + positions, noise_map = positions_and_noise + dataset = al.PointDataset( + name="point_0", positions=positions, positions_noise_map=noise_map + ) + + fit = al.FitPointDataset( + dataset=dataset, tracer=point_source_tracer, solver=mock_solver + ) + + with pytest.raises(al.exc.PointProfileMismatchException): + fit.positions.log_likelihood + + def test__fit_dataset__fit_flux_cls_and_fit_time_delays_cls_hooks_are_forwarded_and_default_unchanged( point_source_tracer, positions_and_noise, mock_solver ): diff --git a/test_autolens/point/model/test_analysis_point.py b/test_autolens/point/model/test_analysis_point.py index a9f3dc78c..7bf6d22e2 100644 --- a/test_autolens/point/model/test_analysis_point.py +++ b/test_autolens/point/model/test_analysis_point.py @@ -69,6 +69,15 @@ def _test__make_result__result_imaging_is_returned(point_dataset): assert isinstance(result, ResultPoint) +def test__default_fit_positions_cls_is_all_to_all_solved(point_dataset): + # #678 phase B defaults decision: solved centres + all-to-all pairing. + solver = al.m.MockPointSolver(model_positions=point_dataset.positions) + + analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False) + + assert analysis.fit_positions_cls is al.FitPositionsImagePairAllSolved + + def test__figure_of_merit__matches_correct_fit_given_galaxy_profiles( positions_x2, positions_x2_noise_map ): @@ -86,7 +95,12 @@ def test__figure_of_merit__matches_correct_fit_given_galaxy_profiles( solver = al.m.MockPointSolver(model_positions=positions_x2) - analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False) + analysis = al.AnalysisPoint( + dataset=point_dataset, + solver=solver, + fit_positions_cls=al.FitPositionsImagePairRepeat, + use_jax=False, + ) instance = model.instance_from_unit_vector([]) analysis_log_likelihood = analysis.log_likelihood_function(instance=instance) @@ -107,7 +121,12 @@ def test__figure_of_merit__matches_correct_fit_given_galaxy_profiles( model_positions = al.Grid2DIrregular([(0.0, 1.0), (1.0, 2.0)]) solver = al.m.MockPointSolver(model_positions=model_positions) - analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False) + analysis = al.AnalysisPoint( + dataset=point_dataset, + solver=solver, + fit_positions_cls=al.FitPositionsImagePairRepeat, + use_jax=False, + ) analysis_log_likelihood = analysis.log_likelihood_function(instance=instance) @@ -147,7 +166,12 @@ def test__figure_of_merit__includes_fit_fluxes( solver = al.m.MockPointSolver(model_positions=positions_x2) - analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False) + analysis = al.AnalysisPoint( + dataset=point_dataset, + solver=solver, + fit_positions_cls=al.FitPositionsImagePairRepeat, + use_jax=False, + ) instance = model.instance_from_unit_vector([]) @@ -179,7 +203,12 @@ def test__figure_of_merit__includes_fit_fluxes( model_positions = al.Grid2DIrregular([(0.0, 1.0), (1.0, 2.0)]) solver = al.m.MockPointSolver(model_positions=model_positions) - analysis = al.AnalysisPoint(dataset=point_dataset, solver=solver, use_jax=False) + analysis = al.AnalysisPoint( + dataset=point_dataset, + solver=solver, + fit_positions_cls=al.FitPositionsImagePairRepeat, + use_jax=False, + ) instance = model.instance_from_unit_vector([]) analysis_log_likelihood = analysis.log_likelihood_function(instance=instance)