Coverage for python/lsst/analysis/ap/plotUtils.py: 13%
120 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-22 10:41 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-22 10:41 +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/>.
22"""Visualization helpers for AP analysis.
24Three tools live here:
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"""
37from __future__ import annotations
39__all__ = ["lightcurve", "cutout_grid", "summarize_run", "BAND_COLORS", "band_color"]
41import inspect
42import io
44import numpy as np
45import pandas as pd
47from lsst.utils.plotting import get_multiband_plot_colors
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()
55def band_color(band, default="k"):
56 """Return the official Rubin plotting color for a band.
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.
65 Returns
66 -------
67 color : `str`
68 Matplotlib color specification.
69 """
70 return BAND_COLORS.get(band, default)
73def _time_column(frame):
74 """Return the name of the MJD-like time column on a DiaSource frame.
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)}")
86def _to_dataframe(table):
87 """Return ``table`` as a `pandas.DataFrame`.
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()
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)
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.
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.
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
149 if ax is None:
150 fig, ax = plt.subplots(figsize=(8, 5))
151 else:
152 fig = ax.figure
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
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)
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
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)})")
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)})")
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
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.
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.
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)``.
234 Returns
235 -------
236 fig : `matplotlib.figure.Figure`
237 """
238 import matplotlib.pyplot as plt
239 import PIL.Image
241 import lsst.geom
243 # Local import to avoid a circular dependency at module load time.
244 from . import plotImageSubtractionCutouts as cutouts_mod
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
251 task = cutouts_mod.PlotImageSubtractionCutoutsTask(config=config, output_path="")
252 cutouts_mod.butler_cache.set(butler, config)
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)
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)
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()
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
299def summarize_run(query, bad_flag_list=None):
300 """Return a per-visit summary DataFrame for an APDB run.
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.
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.
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)
333 if len(sources_all) == 0:
334 return pd.DataFrame()
336 clean_per_visit = sources_clean.groupby("visit").size() if len(sources_clean) else pd.Series(dtype=int)
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()