Coverage for tests/test_fit_bootstrap_model.py: 95%

143 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-09-23 09:54 +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/>. 

21 

22import math 

23 

24import astropy.table 

25import numpy as np 

26import pytest 

27 

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 

60 

61shape_img = (23, 27) 

62reff_x_src, reff_y_src, rho_src, nser_src = 2.5, 3.6, -0.25, 2.0 

63 

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 

71 

72 

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")} 

77 

78 

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 

122 

123 return config_datas 

124 

125 

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 

175 

176 

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 

189 

190 

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 

217 

218 return config_datas 

219 

220 

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) 

255 

256 

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()) 

263 

264 defer_conversion = config_fitter_source.config.defer_radec_conversion 

265 config_fitter_source.config.convert_cen_xy_to_radec = True 

266 

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__ 

271 

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) 

277 

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()))) 

283 

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 ) 

291 

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 

303 

304 plt.show() 

305 

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) 

321 

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) 

339 

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()