Coverage for python/lsst/multiprofit/plotting/plot_catalog_bootstrap.py: 78%

110 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 

22__all__ = ["plot_catalog_bootstrap"] 

23 

24from collections import defaultdict 

25from collections.abc import Iterable 

26from typing import Any 

27 

28import astropy.table 

29import matplotlib.pyplot as plt 

30import numpy as np 

31 

32from ..fitting.fit_source import CatalogSourceFitterConfig, ModelConfig 

33from ..utils import set_config_from_dict 

34 

35ln10 = np.log(10) 

36 

37 

38def plot_catalog_bootstrap( 

39 catalog_bootstrap: astropy.table.Table, 

40 n_bins: int | None = None, 

41 paramvals_ref: Iterable[np.ndarray] | None = None, 

42 plot_total_fluxes: bool = False, 

43 plot_colors: bool = False, 

44 **kwargs: Any, 

45) -> tuple[plt.Figure, plt.Axes]: 

46 """Plot a bootstrap catalog for a single source model. 

47 

48 Parameters 

49 ---------- 

50 catalog_bootstrap 

51 A bootstrap catalog, as returned by 

52 `multiprofit.fit_bootstrap_model.CatalogSourceFitterBootstrap`. 

53 n_bins 

54 The number of bins for parameter value histograms. Default 

55 is sqrt(N) with a minimum of 10. 

56 paramvals_ref 

57 Reference parameter values to plot, if any. 

58 plot_total_fluxes 

59 Whether to plot total fluxes, not just component. 

60 plot_colors 

61 Whether to plot colors in addition to fluxes. 

62 **kwargs 

63 Keyword arguments to pass to matplotlib hist calls. 

64 

65 Returns 

66 ------- 

67 fig, ax 

68 Matplotlib figure and axis handles, as returned by plt.subplots. 

69 """ 

70 n_sources = len(catalog_bootstrap) 

71 if n_bins is None: 71 ↛ 74line 71 didn't jump to line 74 because the condition on line 71 was always true

72 n_bins = np.max([int(np.ceil(np.sqrt(n_sources))), 10]) 

73 

74 config = CatalogSourceFitterConfig() 

75 config_dict = catalog_bootstrap.meta["config"] 

76 # TODO: Figure out if this can be implemented correctly in DM-48911 

77 # In the meantime, we don't need to know the ModelConfig to format columns 

78 # However, it would be useful to get the band and component name (if any) 

79 # given a formatted flux (error) column key 

80 config_dict["config_model"] = ModelConfig().toDict() 

81 set_config_from_dict(config, config_dict) 

82 prefix = config.prefix_column 

83 suffix_err = config.suffix_error 

84 len_suffix_err = len(suffix_err) 

85 # This won't work if the flux format isn't a suffix 

86 # TODO: Consider if this can be fixed in DM-48911 

87 suffix_flux = config.get_key_flux("", "") 

88 

89 # TODO: There are probably better ways of doing this 

90 colnames_err = [col for col in catalog_bootstrap.colnames if col.endswith(suffix_err)] 

91 colnames_meas = [col[:-len_suffix_err] for col in colnames_err] 

92 n_params_init = len(colnames_meas) 

93 if paramvals_ref is not None and (len(paramvals_ref) != n_params_init): 93 ↛ 94line 93 didn't jump to line 94 because the condition on line 93 was never true

94 raise ValueError(f"{len(paramvals_ref)=} != {n_params_init=}") 

95 

96 results_good = catalog_bootstrap[catalog_bootstrap[f"{prefix}n_iter"] > 0] 

97 

98 if plot_total_fluxes or plot_colors: 98 ↛ 172line 98 didn't jump to line 172 because the condition on line 98 was always true

99 if paramvals_ref: 99 ↛ 103line 99 didn't jump to line 103 because the condition on line 99 was always true

100 paramvals_ref = { 

101 colname: paramval_ref for colname, paramval_ref in zip(colnames_meas, paramvals_ref) 

102 } 

103 results_dict = {} 

104 for colname_meas, colname_err in zip(colnames_meas, colnames_err): 

105 results_dict[colname_meas] = results_good[colname_meas] 

106 results_dict[colname_err] = results_good[colname_err] 

107 

108 colnames_flux = [colname for colname in colnames_meas if colname.endswith(suffix_flux)] 

109 

110 colnames_flux_band = defaultdict(list) 

111 colnames_flux_comp = defaultdict(list) 

112 

113 for colname in colnames_flux: 

114 colname_short = colname.partition(prefix)[-1] 

115 comp_band = colname_short.split(suffix_flux)[0] 

116 comp, band = comp_band.split("_") if ("_" in comp_band) else ("", comp_band) 

117 colnames_flux_band[band].append(colname) 

118 colnames_flux_comp[comp].append(colname) 

119 

120 n_comps = len(colnames_flux_comp) 

121 

122 band_prev = None 

123 for band, colnames_band in colnames_flux_band.items(): 

124 # There's no need to make a total flux column with one component 

125 # ... unless there's a component with fixed flux, but that isn't 

126 # supported anyway. 

127 if n_comps >= 2: 127 ↛ 128line 127 didn't jump to line 128 because the condition on line 127 was never true

128 for suffix, target in (("", colnames_meas), (suffix_err, colnames_err)): 

129 is_err = suffix == suffix_err 

130 colname_flux = f"{config.get_key_flux(band=band, label=prefix)}{suffix}" 

131 total = np.sum( 

132 [results_good[f"{colname}{suffix}"] ** (1 + is_err) for colname in colnames_band], 

133 axis=0, 

134 ) 

135 if is_err: 

136 total = np.sqrt(total) 

137 elif paramvals_ref and plot_total_fluxes: 

138 if colname_flux in paramvals_ref: 

139 raise RuntimeError( 

140 f"Tried to set a new total flux column {colname_flux} but it already exists" 

141 ) 

142 paramvals_ref[colname_flux] = sum(paramvals_ref[colname] for colname in colnames_band) 

143 results_dict[colname_flux] = total 

144 if plot_total_fluxes: 

145 target.append(colname_flux) 

146 

147 if band_prev: 

148 flux_prev, flux = (results_dict[f"{prefix}{b}{suffix_flux}"] for b in (band_prev, band)) 

149 mag_prev, mag = (-2.5 * np.log10(flux_b) for flux_b in (flux_prev, flux)) 

150 mag_err_prev, mag_err = ( 

151 results_dict[f"{prefix}{b}{suffix_flux}{suffix_err}"] / (-0.4 * flux_b * ln10) 

152 for b, flux_b in ((band_prev, flux_prev), (band, flux)) 

153 ) 

154 colname_color = f"{prefix}{band_prev}-{band}{suffix_flux}" 

155 colnames_meas.append(colname_color) 

156 colnames_err.append(f"{colname_color}{suffix_err}") 

157 

158 results_dict[colname_color] = mag_prev - mag 

159 results_dict[f"{colname_color}{suffix_err}"] = 2.5 / ln10 * np.hypot(mag_err, mag_err_prev) 

160 if paramvals_ref: 160 ↛ 166line 160 didn't jump to line 166 because the condition on line 160 was always true

161 mag_prev_ref, mag_ref = ( 

162 -2.5 * np.log10(paramvals_ref[f"{prefix}{b}{suffix_flux}"]) for b in (band_prev, band) 

163 ) 

164 paramvals_ref[colname_color] = mag_prev_ref - mag_ref 

165 

166 band_prev = band 

167 

168 results_good = results_dict 

169 if paramvals_ref: 169 ↛ 172line 169 didn't jump to line 172 because the condition on line 169 was always true

170 paramvals_ref = tuple(paramvals_ref.values()) 

171 

172 n_colnames = len(colnames_err) 

173 n_cols = 3 

174 n_rows = int(np.ceil(n_colnames / n_cols)) 

175 

176 fig, ax = plt.subplots(nrows=n_rows, ncols=n_cols, constrained_layout=True) 

177 idx_row, idx_col = 0, 0 

178 

179 for idx_colname in range(n_colnames): 

180 colname_meas = colnames_meas[idx_colname] 

181 colname_short = colname_meas.partition(prefix)[-1] 

182 values = results_good[colname_meas] 

183 errors = results_good[colnames_err[idx_colname]] 

184 median = np.median(values) 

185 std = np.std(values) 

186 

187 median_err = np.median(errors) 

188 

189 axis = ax[idx_row][idx_col] 

190 axis.hist(values, bins=n_bins, color="b", label="fit values", **kwargs) 

191 

192 label = "median +/- stddev" 

193 for offset in (-std, 0, std): 

194 axis.axvline(median + offset, label=label, color="k") 

195 label = None 

196 if paramvals_ref is not None: 196 ↛ 201line 196 didn't jump to line 201 because the condition on line 196 was always true

197 value_ref = paramvals_ref[idx_colname] 

198 label_value = f" {value_ref=:.3e} bias={median - value_ref:.3e}" 

199 axis.axvline(value_ref, label="reference", color="k", linestyle="--") 

200 else: 

201 label_value = f" {median=:.3e}" 

202 axis.hist(median + errors, bins=n_bins, color="r", label="median + error", **kwargs) 

203 axis.set_title(f"{colname_short} {std=:.3e} vs {median_err=:.3e}") 

204 axis.set_xlabel(f"{colname_short} {label_value}") 

205 axis.legend() 

206 

207 idx_col += 1 

208 

209 if idx_col == n_cols: 

210 idx_row += 1 

211 idx_col = 0 

212 

213 return fig, ax