Coverage for python/lsst/multiprofit/plotting/plot_catalog_bootstrap.py: 78%
110 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-19 09:23 +0000
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-19 09:23 +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/>.
22__all__ = ["plot_catalog_bootstrap"]
24from collections import defaultdict
25from collections.abc import Iterable
26from typing import Any
28import astropy.table
29import matplotlib.pyplot as plt
30import numpy as np
32from ..fitting.fit_source import CatalogSourceFitterConfig, ModelConfig
33from ..utils import set_config_from_dict
35ln10 = np.log(10)
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.
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.
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])
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("", "")
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=}")
96 results_good = catalog_bootstrap[catalog_bootstrap[f"{prefix}n_iter"] > 0]
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]
108 colnames_flux = [colname for colname in colnames_meas if colname.endswith(suffix_flux)]
110 colnames_flux_band = defaultdict(list)
111 colnames_flux_comp = defaultdict(list)
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)
120 n_comps = len(colnames_flux_comp)
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)
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}")
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
166 band_prev = band
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())
172 n_colnames = len(colnames_err)
173 n_cols = 3
174 n_rows = int(np.ceil(n_colnames / n_cols))
176 fig, ax = plt.subplots(nrows=n_rows, ncols=n_cols, constrained_layout=True)
177 idx_row, idx_col = 0, 0
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)
187 median_err = np.median(errors)
189 axis = ax[idx_row][idx_col]
190 axis.hist(values, bins=n_bins, color="b", label="fit values", **kwargs)
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()
207 idx_col += 1
209 if idx_col == n_cols:
210 idx_row += 1
211 idx_col = 0
213 return fig, ax