Coverage for tests/test_fit_bootstrap_model.py: 95%
143 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-22 09:35 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-22 09:35 +0000
1# This file is part of multiprofit.
2#
3# Developed for the LSST Data Management System.
4# This product includes software developed by the LSST Project
5# (https://www.lsst.org).
6# See the COPYRIGHT file at the top-level directory of this distribution
7# for details of code ownership.
8#
9# This program is free software: you can redistribute it and/or modify
10# it under the terms of the GNU General Public License as published by
11# the Free Software Foundation, either version 3 of the License, or
12# (at your option) any later version.
13#
14# This program is distributed in the hope that it will be useful,
15# but WITHOUT ANY WARRANTY; without even the implied warranty of
16# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
17# GNU General Public License for more details.
18#
19# You should have received a copy of the GNU General Public License
20# along with this program. If not, see <https://www.gnu.org/licenses/>.
22import math
24import astropy.table
25import numpy as np
26import pytest
28import lsst.gauss2d.fit as g2f
29from lsst.multiprofit.componentconfig import (
30 CentroidConfig,
31 FluxFractionParameterConfig,
32 FluxParameterConfig,
33 GaussianComponentConfig,
34 ParameterConfig,
35 SersicComponentConfig,
36 SersicIndexParameterConfig,
37)
38from lsst.multiprofit.errors import RaDecConversionNotImplementedError
39from lsst.multiprofit.fitting.fit_bootstrap_model import (
40 CatalogExposurePsfBootstrap,
41 CatalogExposureSourcesBootstrap,
42 CatalogPsfBootstrapConfig,
43 CatalogSourceBootstrapConfig,
44 CatalogSourceFitterBootstrap,
45 NoisyObservationConfig,
46 NoisyPsfObservationConfig,
47)
48from lsst.multiprofit.fitting.fit_psf import (
49 CatalogPsfFitter,
50 CatalogPsfFitterConfig,
51 CatalogPsfFitterConfigData,
52)
53from lsst.multiprofit.fitting.fit_source import CatalogSourceFitterConfig, CatalogSourceFitterConfigData
54from lsst.multiprofit.modelconfig import ModelConfig
55from lsst.multiprofit.modeller import ModelFitConfig
56from lsst.multiprofit.observationconfig import CoordinateSystemConfig
57from lsst.multiprofit.plotting import ErrorValues, plot_catalog_bootstrap, plot_loglike
58from lsst.multiprofit.sourceconfig import ComponentGroupConfig, SourceConfig
59from lsst.multiprofit.utils import get_params_uniq
61shape_img = (23, 27)
62reff_x_src, reff_y_src, rho_src, nser_src = 2.5, 3.6, -0.25, 2.0
64# TODO: These can be parameterized; should they be?
65compute_errors_no_covar = True
66compute_errors_from_jacobian = True
67include_point_source = False
68n_sources = 3
69# Set to True for interactive debugging (but don't commit)
70plot = False
73@pytest.fixture(scope="module")
74def channels():
75 """Return dict of generic RGB channels."""
76 return {band: g2f.Channel.get(band) for band in ("R", "G", "B")}
79@pytest.fixture(scope="module")
80def config_fitter_psfs(channels) -> dict[g2f.Channel, CatalogExposurePsfBootstrap]:
81 """Return dict of bootstrap fitter configs."""
82 config_datas = {}
83 for idx, (band, channel) in enumerate(channels.items()):
84 n_rows = 17 + idx * 2
85 n_cols = 15 + idx * 2
86 config = CatalogPsfFitterConfig(
87 model=SourceConfig(
88 component_groups={
89 "": ComponentGroupConfig(
90 centroids={
91 "default": CentroidConfig(
92 x=ParameterConfig(value_initial=n_cols / 2.0),
93 y=ParameterConfig(value_initial=n_rows / 2.0),
94 ),
95 },
96 components_gauss={
97 "gauss1": GaussianComponentConfig(
98 flux=FluxParameterConfig(value_initial=1.0, fixed=True),
99 fluxfrac=FluxFractionParameterConfig(value_initial=0.5, fixed=False),
100 size_x=ParameterConfig(value_initial=1.5 + 0.1 * idx),
101 size_y=ParameterConfig(value_initial=1.7 + 0.13 * idx),
102 rho=ParameterConfig(value_initial=-0.035 - 0.007 * idx),
103 ),
104 "gauss2": GaussianComponentConfig(
105 size_x=ParameterConfig(value_initial=3.1 + 0.24 * idx),
106 size_y=ParameterConfig(value_initial=2.7 + 0.16 * idx),
107 rho=ParameterConfig(value_initial=0.06 + 0.012 * idx),
108 fluxfrac=FluxFractionParameterConfig(value_initial=1.0, fixed=True),
109 ),
110 },
111 is_fractional=True,
112 )
113 }
114 ),
115 )
116 config_boot = CatalogPsfBootstrapConfig(
117 observation=NoisyPsfObservationConfig(n_rows=n_rows, n_cols=n_cols, gain=1e5),
118 n_sources=n_sources,
119 )
120 config_data = CatalogExposurePsfBootstrap(config=config, config_boot=config_boot)
121 config_datas[channel] = config_data
123 return config_datas
126@pytest.fixture(scope="module")
127def config_fitter_source(channels) -> CatalogSourceFitterConfigData:
128 """Return dict of bootstrap source fitter configs."""
129 config = CatalogSourceFitterConfig(
130 config_fit=ModelFitConfig(fit_linear_iter=3),
131 config_model=ModelConfig(
132 sources={
133 "": SourceConfig(
134 component_groups={
135 "": ComponentGroupConfig(
136 components_gauss=(
137 {
138 "ps": GaussianComponentConfig(
139 flux=FluxParameterConfig(value_initial=1000),
140 rho=ParameterConfig(value_initial=0, fixed=True),
141 size_x=ParameterConfig(value_initial=0, fixed=True),
142 size_y=ParameterConfig(value_initial=0, fixed=True),
143 )
144 }
145 if include_point_source
146 else {}
147 ),
148 components_sersic={
149 "ser": SersicComponentConfig(
150 prior_size_mean=reff_y_src,
151 prior_size_stddev=1.0,
152 prior_axrat_mean=reff_x_src / reff_y_src,
153 prior_axrat_stddev=0.2,
154 flux=FluxParameterConfig(value_initial=5000),
155 rho=ParameterConfig(value_initial=rho_src),
156 size_x=ParameterConfig(value_initial=reff_x_src),
157 size_y=ParameterConfig(value_initial=reff_y_src),
158 sersic_index=SersicIndexParameterConfig(fixed=False, value_initial=1.0),
159 ),
160 },
161 )
162 }
163 ),
164 },
165 ),
166 convert_cen_xy_to_radec=False,
167 compute_errors_no_covar=compute_errors_no_covar,
168 compute_errors_from_jacobian=compute_errors_from_jacobian,
169 )
170 config_data = CatalogSourceFitterConfigData(
171 channels=tuple(channels.values()),
172 config=config,
173 )
174 return config_data
177@pytest.fixture(scope="module")
178def tables_psf_fits(config_fitter_psfs) -> dict[g2f.Channel, astropy.table.Table]:
179 """Return fits to bootstrapped PSF."""
180 fitter = CatalogPsfFitter()
181 fits = {
182 channel: fitter.fit(
183 catexp=config_fitter_psf,
184 config_data=config_fitter_psf,
185 )
186 for channel, config_fitter_psf in config_fitter_psfs.items()
187 }
188 return fits
191@pytest.fixture(scope="module")
192def config_data_sources(
193 config_fitter_psfs,
194 tables_psf_fits,
195) -> dict[g2f.Channel, CatalogExposureSourcesBootstrap]:
196 """Return data and configs for bootstrap source fitting."""
197 config_datas = {}
198 for idx, (channel, config_fitter_psf) in enumerate(config_fitter_psfs.items()):
199 table_psf_fits = tables_psf_fits[channel]
200 n_rows = shape_img[0] + idx * 2
201 n_cols = shape_img[1] + idx * 2
202 config_boot = CatalogSourceBootstrapConfig(
203 observation=NoisyObservationConfig(
204 n_rows=n_rows,
205 n_cols=n_cols,
206 band=channel.name,
207 background=100,
208 coordsys=CoordinateSystemConfig(x_min=-2 + 3 * idx, y_min=5 - 4 * idx),
209 ),
210 n_sources=n_sources,
211 )
212 config_data = CatalogExposureSourcesBootstrap(
213 config_boot=config_boot,
214 table_psf_fits=table_psf_fits,
215 )
216 config_datas[channel] = config_data
218 return config_datas
221def test_fit_psf(config_fitter_psfs, tables_psf_fits):
222 """Check that the bootstrap PSF fits are sensible."""
223 for band, results in tables_psf_fits.items():
224 assert len(results) == n_sources
225 assert np.sum(results["mpf_psf_unknown_flag"]) == 0
226 assert all(np.isfinite(list(results[0].values())))
227 config_data_psf = config_fitter_psfs[band]
228 psf_model_init = config_data_psf.config.make_psf_model()
229 psfdata = CatalogPsfFitterConfigData(config=config_data_psf.config)
230 psf_model_fit = psfdata.psf_model
231 psfdata.init_psf_model(results[0])
232 assert len(psf_model_init.components) == len(psf_model_fit.components)
233 params_init = psf_model_init.parameters()
234 params_fit = psf_model_fit.parameters()
235 assert len(params_init) == len(params_fit)
236 sigma_min_sq = config_data_psf.config.sigma_min**2
237 for p_init, p_meas in zip(params_init, params_fit):
238 assert p_meas.fixed == p_init.fixed
239 if p_meas.fixed:
240 assert p_init.value == p_meas.value
241 else:
242 value = p_meas.value
243 # TODO: come up with better (noise-dependent) thresholds here
244 if isinstance(p_init, g2f.IntegralParameterD): 244 ↛ 245line 244 didn't jump to line 245 because the condition on line 244 was never true
245 atol, rtol = 0, 0.02
246 elif isinstance(p_init, g2f.ProperFractionParameterD):
247 atol, rtol = 0.1, 0.01
248 elif isinstance(p_init, g2f.RhoParameterD):
249 atol, rtol = 0.05, 0.1
250 elif isinstance(p_init, g2f.SigmaXParameterD) or isinstance(p_init, g2f.SigmaYParameterD):
251 value = math.sqrt(value**2 + sigma_min_sq)
252 else:
253 atol, rtol = 0.01, 0.1
254 assert np.isclose(p_init.value, value, atol=atol, rtol=rtol)
257def test_fit_source(config_fitter_source, config_data_sources):
258 """Test bootstrap source fitting."""
259 fitter = CatalogSourceFitterBootstrap()
260 # We don't have or need a multiband input catalog - just use the first one
261 catalog_multi = next(iter(config_data_sources.values())).get_catalog()
262 catexps = list(config_data_sources.values())
264 defer_conversion = config_fitter_source.config.defer_radec_conversion
265 config_fitter_source.config.convert_cen_xy_to_radec = True
267 conversion_error_cls = RaDecConversionNotImplementedError
268 conversion_error_key = conversion_error_cls.column_name()
269 fitter.errors_expected[conversion_error_cls] = conversion_error_key
270 config_fitter_source.config.flag_errors[conversion_error_key] = conversion_error_cls.__name__
272 # Test both code paths for failure to convert RA/Dec, returning to original
273 for value in (not defer_conversion, defer_conversion):
274 config_fitter_source.config.defer_radec_conversion = value
275 results = fitter.fit(catalog_multi=catalog_multi, catexps=catexps, config_data=config_fitter_source)
276 assert np.all(results[f"mpf_{conversion_error_key}"] == 1)
278 config_fitter_source.config.convert_cen_xy_to_radec = False
279 results = fitter.fit(catalog_multi=catalog_multi, catexps=catexps, config_data=config_fitter_source)
280 assert len(results) == n_sources
281 assert np.sum(results["mpf_unknown_flag"]) == 0
282 assert all(np.isfinite(list(results[0].values())))
284 model = fitter.get_model(
285 0,
286 catalog_multi=catalog_multi,
287 catexps=catexps,
288 config_data=config_fitter_source,
289 results=results,
290 )
292 model_sources, priors = config_fitter_source.config.make_sources(
293 channels=list(config_data_sources.keys())
294 )
295 model_true = g2f.ModelD(data=model.data, psfmodels=model.psfmodels, sources=model_sources)
296 fitter.initialize_model(model_true, catalog_multi[0], catexps=catexps)
297 params_true = tuple(param.value for param in get_params_uniq(model_true, fixed=False))
298 plot_catalog_bootstrap(
299 results, histtype="step", paramvals_ref=params_true, plot_total_fluxes=True, plot_colors=True
300 )
301 if plot: 301 ↛ 302line 301 didn't jump to line 302 because the condition on line 301 was never true
302 import matplotlib.pyplot as plt
304 plt.show()
306 variances = []
307 for return_negative in (False, True):
308 variances.append(
309 fitter.modeller.compute_variances(
310 model,
311 transformed=False,
312 options=g2f.HessianOptions(return_negative=return_negative),
313 use_diag_only=True,
314 )
315 )
316 assert np.all(variances[-1] > 0)
317 if return_negative:
318 variances = np.array(variances)
319 variances[variances <= 0] = 0
320 variances = list(variances)
322 # Bootstrap errors
323 model.setup_evaluators(evaluatormode=g2f.EvaluatorMode.image)
324 model.evaluate()
325 img_data_old = []
326 for obs, output in zip(model.data, model.outputs):
327 img_data_old.append(obs.image.data.copy())
328 img = obs.image.data
329 img.flat = output.data.flat
330 options_hessian = g2f.HessianOptions(return_negative=return_negative)
331 variances_bootstrap = fitter.modeller.compute_variances(model, transformed=False, options=options_hessian)
332 variances_bootstrap_diag = fitter.modeller.compute_variances(
333 model, transformed=False, options=options_hessian, use_diag_only=True
334 )
335 for obs, img_datum_old in zip(model.data, img_data_old):
336 obs.image.data.flat = img_datum_old.flat
337 variances_jac = fitter.modeller.compute_variances(model, transformed=False)
338 variances_jac_diag = fitter.modeller.compute_variances(model, transformed=False, use_diag_only=True)
340 errors_plot = {
341 "inv_hess": ErrorValues(values=np.sqrt(variances[0]), kwargs_plot={"linestyle": "-", "color": "r"}),
342 "-inv_hess": ErrorValues(values=np.sqrt(variances[1]), kwargs_plot={"linestyle": "--", "color": "r"}),
343 "inv_jac": ErrorValues(values=np.sqrt(variances_jac), kwargs_plot={"linestyle": "-.", "color": "r"}),
344 "boot_hess": ErrorValues(
345 values=np.sqrt(variances_bootstrap), kwargs_plot={"linestyle": "-", "color": "b"}
346 ),
347 "boot_diag": ErrorValues(
348 values=np.sqrt(variances_bootstrap_diag), kwargs_plot={"linestyle": "--", "color": "b"}
349 ),
350 "boot_jac_diag": ErrorValues(
351 values=np.sqrt(variances_jac_diag), kwargs_plot={"linestyle": "-.", "color": "m"}
352 ),
353 }
354 fig, ax = plot_loglike(model, errors=errors_plot, values_reference=params_true)
355 if plot: 355 ↛ 356line 355 didn't jump to line 356 because the condition on line 355 was never true
356 plt.tight_layout()
357 plt.show()