Coverage for python/lsst/analysis/ap/plotUtils.py: 13%

120 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-11 11:12 +0000

1# This file is part of analysis_ap. 

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"""Visualization helpers for AP analysis. 

23 

24Three tools live here: 

25 

26- `lightcurve` plots a per-band psfFlux light curve for a single diaObject, 

27 optionally overlaying forced photometry. It works with either the APDB 

28 `DbQuery` interface (pandas) or the PPDB `PpdbTap` interface (astropy 

29 Tables). 

30- `cutout_grid` lays out science/template/difference cutouts for many 

31 DiaSources in a single mosaic figure. 

32- `summarize_run` returns a per-visit summary DataFrame of an APDB run 

33 (counts, dipole rate, reliability statistics, etc.) suitable for a quick 

34 health check of a processing run. 

35""" 

36 

37from __future__ import annotations 

38 

39__all__ = ["lightcurve", "cutout_grid", "summarize_run", "BAND_COLORS", "band_color"] 

40 

41import inspect 

42import io 

43 

44import numpy as np 

45import pandas as pd 

46 

47from lsst.utils.plotting import get_multiband_plot_colors 

48 

49# The official Rubin band colors (RTN-045 colorblind-friendly palette), 

50# keyed by lower-case band character. Use `band_color` rather than 

51# indexing this directly, so that unknown bands get a sensible fallback. 

52BAND_COLORS = get_multiband_plot_colors() 

53 

54 

55def band_color(band, default="k"): 

56 """Return the official Rubin plotting color for a band. 

57 

58 Parameters 

59 ---------- 

60 band : `str` 

61 Band name, e.g. ``"g"``. 

62 default : `str`, optional 

63 Color to return for a band with no official color. 

64 

65 Returns 

66 ------- 

67 color : `str` 

68 Matplotlib color specification. 

69 """ 

70 return BAND_COLORS.get(band, default) 

71 

72 

73def _time_column(frame): 

74 """Return the name of the MJD-like time column on a DiaSource frame. 

75 

76 Older APDB schemas used ``midPointTai``; current ones use 

77 ``midpointMjdTai``. This helper accepts either. 

78 """ 

79 for candidate in ("midpointMjdTai", "midPointTai"): 79 ↛ 82line 79 didn't jump to line 82 because the loop on line 79 didn't complete

80 if candidate in frame.columns: 80 ↛ 79line 80 didn't jump to line 79 because the condition on line 80 was always true

81 return candidate 

82 raise KeyError("Expected one of 'midpointMjdTai' or 'midPointTai' " 

83 f"in DataFrame; got columns: {list(frame.columns)}") 

84 

85 

86def _to_dataframe(table): 

87 """Return ``table`` as a `pandas.DataFrame`. 

88 

89 The APDB `DbQuery` loaders already return DataFrames; the PPDB `PpdbTap` 

90 loaders return `astropy.table.Table`. Normalizing here lets the plotting 

91 code use a single pandas (groupby-based) path regardless of which 

92 interface produced the data. Masked astropy values become NaN. 

93 """ 

94 if isinstance(table, pd.DataFrame): 

95 return table 

96 return table.to_pandas() 

97 

98 

99def _load_object_sources(query, dia_object_id, exclude_flagged): 

100 """Load one diaObject's DiaSources as a DataFrame across query interfaces. 

101 """ 

102 method = query.load_sources_for_object 

103 if "exclude_flagged" in inspect.signature(method).parameters: 

104 sources = method(dia_object_id, exclude_flagged=exclude_flagged) 

105 elif exclude_flagged: 

106 raise TypeError( 

107 f"{type(query).__name__}.load_sources_for_object does not support " 

108 "exclude_flagged; the PPDB public interface does not expose " 

109 "diaSource flag filtering. Pass exclude_flagged=False.") 

110 else: 

111 sources = method(dia_object_id) 

112 return _to_dataframe(sources) 

113 

114 

115def lightcurve(query, dia_object_id, ax=None, exclude_flagged=False, 

116 include_forced=True): 

117 """Plot a per-band psfFlux light curve for one diaObject. 

118 

119 Parameters 

120 ---------- 

121 query : `lsst.analysis.ap.apdb.DbQuery` or \ 

122 `lsst.analysis.ap.ppdb.PpdbTap` 

123 dia_object_id : `int` 

124 Object id to load. 

125 ax : `matplotlib.axes.Axes`, optional 

126 Axes to draw into; if None, a new figure is created. 

127 exclude_flagged : `bool`, optional 

128 Forwarded to `load_sources_for_object` when the query supports it. 

129 Defaults to False so the lightcurve matches the row count of a direct 

130 APDB query; pass True to drop diaSources matching the configured 

131 bad-flag list. The PPDB `PpdbTap` interface does not expose flag 

132 filtering, so True is rejected there. DiaForcedSources are always 

133 loaded unfiltered. 

134 include_forced : `bool`, optional 

135 If True, also overlay diaForcedSources as small markers. 

136 

137 Returns 

138 ------- 

139 fig : `matplotlib.figure.Figure` 

140 ax : `matplotlib.axes.Axes` 

141 sources : `pandas.DataFrame` 

142 DiaSources used for the plot. 

143 forced : `pandas.DataFrame` or None 

144 DiaForcedSources used for the plot (None if ``include_forced`` is 

145 False). 

146 """ 

147 import matplotlib.pyplot as plt 

148 

149 if ax is None: 

150 fig, ax = plt.subplots(figsize=(8, 5)) 

151 else: 

152 fig = ax.figure 

153 

154 # diaObjectId is NaN for diaSources not associated with any diaObject; 

155 # short-circuit before hitting the database. 

156 if pd.isna(dia_object_id): 

157 ax.text(0.5, 0.5, "no diaObjectId (NaN)", 

158 ha="center", va="center", transform=ax.transAxes) 

159 return fig, ax, pd.DataFrame(), None 

160 

161 sources = _load_object_sources(query, dia_object_id, exclude_flagged) 

162 # DiaForcedSource has a different (and smaller) flag schema than 

163 # DiaSource: applying the diaSource exclusion list would key into 

164 # columns that don't exist on the forced table. Forced photometry is 

165 # also a measurement at a known location rather than a fresh detection, 

166 # so showing it unfiltered is the right behavior. 

167 forced = (_to_dataframe(query.load_forced_sources_for_object(dia_object_id)) 

168 if include_forced else None) 

169 

170 if len(sources) == 0: 

171 ax.text(0.5, 0.5, f"no sources for diaObjectId={dia_object_id}", 

172 ha="center", va="center", transform=ax.transAxes) 

173 return fig, ax, sources, forced 

174 

175 time_col = _time_column(sources) 

176 for band, group in sources.groupby("band"): 

177 color = band_color(band) 

178 ax.errorbar(group[time_col], group["psfFlux"], yerr=group["psfFluxErr"], 

179 fmt="o", color=color, label=f"{band} (n={len(group)})") 

180 

181 if forced is not None and len(forced): 

182 forced_time_col = _time_column(forced) 

183 # Plot per-band forced points in their band colors, but suppress 

184 # individual legend entries so the forced marker is represented 

185 # once (in black) regardless of how many bands are present. 

186 for band, group in forced.groupby("band"): 

187 color = band_color(band) 

188 ax.errorbar(group[forced_time_col], group["psfFlux"], 

189 yerr=group["psfFluxErr"], fmt=".", ms=4, color=color, 

190 alpha=0.4, label="_nolegend_") 

191 ax.plot([], [], ".", color="black", ms=4, alpha=0.4, 

192 label=f"forced (n={len(forced)})") 

193 

194 ax.axhline(0, color="grey", lw=0.5) 

195 ax.set_xlabel(time_col) 

196 ax.set_ylabel("psfFlux (nJy)") 

197 ax.set_title(f"diaObjectId = {dia_object_id}") 

198 ax.legend(frameon=True) 

199 return fig, ax, sources, forced 

200 

201 

202def cutout_grid(sources, butler, instrument, n_per_row=4, config=None, output=None, 

203 figsize=None, ra_column='ra', dec_column='dec', detector_column='detector', 

204 visit_column='visit', id_column='diaSourceId'): 

205 """Render science/template/difference cutouts for many sources in a grid. 

206 

207 This is a thin wrapper around `PlotImageSubtractionCutoutsTask`: it calls 

208 `generate_image` for each source (which returns a PNG in memory) and 

209 arranges the resulting rasters in a single matplotlib figure. 

210 

211 Parameters 

212 ---------- 

213 sources : `pandas.DataFrame` 

214 DiaSources to cut out. Must contain at least 

215 ``ra, dec, diaSourceId, detector, visit, instrument`` plus whatever 

216 annotation fields the task config requires (see 

217 ``PlotImageSubtractionCutoutsConfig.add_metadata``). 

218 butler : `lsst.daf.butler.Butler` 

219 Butler initialized with the relevant collections. 

220 instrument : `str` 

221 Name of the instrument for the data being plotted. 

222 n_per_row : `int` 

223 Number of cutouts per row in the resulting figure. 

224 config : `PlotImageSubtractionCutoutsConfig`, optional 

225 Cutout config to use (see 

226 ``plotImageSubtractionCutouts.PlotImageSubtractionCutoutsConfig``). 

227 Defaults to a fresh instance with ``add_metadata=False`` 

228 (annotations get cluttered in a grid). 

229 output : `str`, optional 

230 If given, save the figure to this path with ``bbox_inches="tight"``. 

231 figsize : `tuple` [`float`, `float`], optional 

232 Figure size in inches. Defaults to ``(n_per_row*3.5, n_rows*1.7)``. 

233 

234 Returns 

235 ------- 

236 fig : `matplotlib.figure.Figure` 

237 """ 

238 import matplotlib.pyplot as plt 

239 import PIL.Image 

240 

241 import lsst.geom 

242 

243 # Local import to avoid a circular dependency at module load time. 

244 from . import plotImageSubtractionCutouts as cutouts_mod 

245 

246 if config is None: 

247 config = cutouts_mod.PlotImageSubtractionCutoutsConfig() 

248 # Annotations get cramped in a grid; default them off here. 

249 config.add_metadata = False 

250 

251 task = cutouts_mod.PlotImageSubtractionCutoutsTask(config=config, output_path="") 

252 cutouts_mod.butler_cache.set(butler, config) 

253 

254 n_sources = len(sources) 

255 n_rows = max(1, (n_sources + n_per_row - 1) // n_per_row) 

256 if figsize is None: 

257 figsize = (n_per_row * 3.5, n_rows * 1.7) 

258 fig, axes = plt.subplots(n_rows, n_per_row, figsize=figsize, squeeze=False) 

259 

260 records = sources.to_records(index=False) 

261 for i, source in enumerate(records): 

262 row, col = divmod(i, n_per_row) 

263 ax = axes[row][col] 

264 ax.set_axis_off() 

265 try: 

266 sci, tmpl, diff = cutouts_mod.butler_cache.get_exposures( 

267 instrument, source[detector_column], source[visit_column]) 

268 center = lsst.geom.SpherePoint(source[ra_column], source[dec_column], 

269 lsst.geom.degrees) 

270 scale = sci.wcs.getPixelScale(sci.getBBox().getCenter()).asArcseconds() 

271 png = task.generate_image( 

272 sci, tmpl, diff, center, scale, 

273 source=source if config.add_metadata else None, 

274 ) 

275 with PIL.Image.open(io.BytesIO(png.getvalue())) as img: 

276 ax.imshow(np.asarray(img)) 

277 ax.set_title(f"{source[id_column]}", fontsize=7) 

278 except Exception as exc: 

279 ax.text(0.5, 0.5, f"{type(exc).__name__}", ha="center", va="center", 

280 fontsize=8, transform=ax.transAxes) 

281 

282 # Blank out any trailing axes in the last row. 

283 for i in range(n_sources, n_rows * n_per_row): 

284 row, col = divmod(i, n_per_row) 

285 axes[row][col].set_axis_off() 

286 

287 fig.tight_layout() 

288 if output is not None: 

289 fig.savefig(output, bbox_inches="tight") 

290 # Remove the figure from pyplot's figure manager so the Jupyter inline 

291 # backend does not auto-display it at end-of-cell in addition to the 

292 # caller's own rendering of the returned Figure (which would duplicate 

293 # the output the first time the function is run in a notebook). The 

294 # returned Figure object remains valid and renders via its repr hooks. 

295 plt.close(fig) 

296 return fig 

297 

298 

299def summarize_run(query, bad_flag_list=None): 

300 """Return a per-visit summary DataFrame for an APDB run. 

301 

302 Useful as a quick health check after a processing run: one row per visit 

303 with source counts, dipole rate, reliability statistics, and the fraction 

304 of sources that fall in the bad-flag list. 

305 

306 Parameters 

307 ---------- 

308 query : `lsst.analysis.ap.apdb.DbQuery` 

309 APDB query interface (sqlite, postgres, or cassandra). 

310 bad_flag_list : `list` [`str`], optional 

311 Flag column names to count as "bad". If omitted, the query's 

312 currently-configured exclusion list is used. The caller's exclusion 

313 list is restored before returning. 

314 

315 Returns 

316 ------- 

317 summary : `pandas.DataFrame` 

318 Indexed by ``visit``. Columns: 

319 ``n_sources``, ``n_unflagged``, ``bad_flag_fraction``, 

320 ``median_reliability`` (if column present), 

321 ``dipole_fraction`` (if column present), 

322 ``median_psf_chi2_per_dof`` (if columns present). 

323 """ 

324 saved_flags = list(query.diaSource_flags_exclude) 

325 try: 

326 if bad_flag_list is not None: 

327 query.set_excluded_diaSource_flags(bad_flag_list) 

328 sources_all = query.load_sources(limit=None) 

329 sources_clean = query.load_sources(exclude_flagged=True, limit=None) 

330 finally: 

331 query.set_excluded_diaSource_flags(saved_flags) 

332 

333 if len(sources_all) == 0: 

334 return pd.DataFrame() 

335 

336 clean_per_visit = sources_clean.groupby("visit").size() if len(sources_clean) else pd.Series(dtype=int) 

337 

338 rows = [] 

339 for visit, group in sources_all.groupby("visit"): 

340 n_all = len(group) 

341 n_clean = int(clean_per_visit.get(visit, 0)) 

342 row = { 

343 "visit": visit, 

344 "n_sources": n_all, 

345 "n_unflagged": n_clean, 

346 "bad_flag_fraction": 1.0 - n_clean / n_all if n_all else 0.0, 

347 } 

348 if "reliability" in group.columns: 

349 row["median_reliability"] = group["reliability"].median() 

350 if "isDipole" in group.columns: 

351 row["dipole_fraction"] = float(group["isDipole"].mean()) 

352 if "psfChi2" in group.columns and "psfNdata" in group.columns: 

353 with np.errstate(divide="ignore", invalid="ignore"): 

354 ratio = group["psfChi2"] / group["psfNdata"] 

355 row["median_psf_chi2_per_dof"] = float(np.nanmedian(ratio)) 

356 rows.append(row) 

357 return pd.DataFrame(rows).set_index("visit").sort_index()