Coverage for python/lsst/analysis/ap/plotDiaSourceLightcurve.py: 70%

177 statements  

« prev     ^ index     » next       coverage.py v7.16.1, created at 2026-09-26 10:29 +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"""Render diaSource cutouts with a per-diaObject lightcurve panel. 

23 

24``PlotDiaSourceLightcurveTask`` extends ``PlotImageSubtractionCutoutsTask`` 

25by adding a lightcurve panel below the science/template/difference cutouts. 

26The panel shows the diaSource history for the associated diaObject and 

27overlays diaForcedSource measurements for any visit that has a forced 

28measurement but no diaSource (drawn with a distinct marker). 

29""" 

30 

31__all__ = ["PlotDiaSourceLightcurveConfig", "PlotDiaSourceLightcurveTask"] 

32 

33import argparse 

34import io 

35import logging 

36import os 

37 

38import lsst.daf.butler 

39import lsst.pex.config as pexConfig 

40 

41import pandas as pd 

42 

43from . import apdb as _apdb_mod 

44from . import plotUtils 

45from .plotImageSubtractionCutouts import ( 

46 PlotImageSubtractionCutoutsConfig, 

47 PlotImageSubtractionCutoutsTask, 

48 _annotate_image, 

49) 

50 

51_log = logging.getLogger(__name__) 

52 

53 

54class PlotDiaSourceLightcurveConfig(PlotImageSubtractionCutoutsConfig): 

55 lightcurve_height = pexConfig.Field( 

56 doc="Height in inches reserved for the lightcurve panel below the cutouts.", 

57 dtype=float, 

58 default=2.5, 

59 ) 

60 lightcurve_exclude_flagged = pexConfig.Field( 

61 doc="Pass exclude_flagged=True to the APDB query when loading " 

62 "diaSources for the lightcurve. Defaults to False so the " 

63 "lightcurve matches the row count of a direct APDB query; " 

64 "set True to drop diaSources matching the configured bad-flag " 

65 "list. DiaForcedSources are always loaded unfiltered.", 

66 dtype=bool, 

67 default=False, 

68 ) 

69 lightcurve_marker_source = pexConfig.Field( 

70 doc="Matplotlib marker style for visits that have a diaSource " 

71 "detection.", 

72 dtype=str, 

73 default="o", 

74 ) 

75 lightcurve_marker_forced_only = pexConfig.Field( 

76 doc="Matplotlib marker style for visits that have a forced " 

77 "measurement but no diaSource.", 

78 dtype=str, 

79 default="v", 

80 ) 

81 highlight_current_source = pexConfig.Field( 

82 doc="Draw a vertical line and open ring on the lightcurve at the " 

83 "MJD/flux of the diaSource being cut out.", 

84 dtype=bool, 

85 default=True, 

86 ) 

87 

88 

89class PlotDiaSourceLightcurveTask(PlotImageSubtractionCutoutsTask): 

90 """Generate cutouts plus a diaObject lightcurve panel for each diaSource. 

91 

92 Parameters 

93 ---------- 

94 output_path : `str` 

95 Path to write outputs to. Same convention as the parent task. 

96 apdb_query : `lsst.analysis.ap.apdb.DbQuery`, optional 

97 Query handle used to load the diaSource and diaForcedSource history 

98 for each diaObject. If None, the lightcurve panel is rendered with 

99 a placeholder message. 

100 

101 Notes 

102 ----- 

103 The input ``data`` DataFrame must include the fields required by the 

104 parent (``ra, dec, diaSourceId, detector, visit, instrument``) plus 

105 ``diaObjectId`` (to load the lightcurve) and ``midpointMjdTai`` (to 

106 highlight the current source on the lightcurve). 

107 """ 

108 ConfigClass = PlotDiaSourceLightcurveConfig 

109 _DefaultName = "plotDiaSourceLightcurve" 

110 

111 def __init__(self, *, output_path, apdb_query=None, **kwargs): 

112 super().__init__(output_path=output_path, **kwargs) 

113 self._apdb_query = apdb_query 

114 # Per-diaObject lightcurve cache so adjacent diaSources on the same 

115 # object don't re-query the APDB. Keyed by diaObjectId. 

116 self._lightcurve_cache = {} 

117 

118 def _reduce_kwargs(self): 

119 kwargs = super()._reduce_kwargs() 

120 kwargs["apdb_query"] = self._apdb_query 

121 return kwargs 

122 

123 def run(self, data, butler, njobs=0): 

124 if njobs > 0: 

125 self.log.warning("njobs=%d ignored; PlotDiaSourceLightcurveTask " 

126 "runs single-process only.", njobs) 

127 return super().run(data, butler, njobs=0) 

128 

129 def write_images(self, data, butler, njobs=0): 

130 if njobs > 0: 130 ↛ 133line 130 didn't jump to line 133 because the condition on line 130 was always true

131 self.log.warning("njobs=%d ignored; PlotDiaSourceLightcurveTask " 

132 "runs single-process only.", njobs) 

133 return super().write_images(data, butler, njobs=0) 

134 

135 def _load_lightcurve_data(self, dia_object_id): 

136 """Return cached (sources, forced) for one diaObject. 

137 

138 Returns ``(None, None)`` if no APDB query handle is configured. 

139 """ 

140 if self._apdb_query is None: 

141 return None, None 

142 if dia_object_id not in self._lightcurve_cache: 

143 sources = self._apdb_query.load_sources_for_object( 

144 dia_object_id, 

145 exclude_flagged=self.config.lightcurve_exclude_flagged, 

146 ) 

147 # Always show all of the forced sources regardless of flags. 

148 forced = self._apdb_query.load_forced_sources_for_object( 

149 dia_object_id, 

150 ) 

151 self._lightcurve_cache[dia_object_id] = (sources, forced) 

152 return self._lightcurve_cache[dia_object_id] 

153 

154 def _plot_cutout(self, science, template, difference, scale, sizes, source=None): 

155 import astropy.visualization as aviz 

156 import matplotlib 

157 matplotlib.use("AGG") 

158 matplotlib.rcParams.update(matplotlib.rcParamsDefault) 

159 import matplotlib.pyplot as plt 

160 from matplotlib import cm 

161 from matplotlib.gridspec import GridSpec 

162 

163 len_sizes = len(sizes) 

164 

165 sources_lc, forced_lc = (None, None) 

166 if source is not None and "diaObjectId" in source.dtype.names: 166 ↛ 179line 166 didn't jump to line 179 because the condition on line 166 was always true

167 dia_object_id = source["diaObjectId"] 

168 # diaObjectId is NaN for diaSources not associated with any 

169 # diaObject (e.g. single unassociated detections). Skip the 

170 # lightcurve query — the panel will render its empty placeholder. 

171 if not pd.isna(dia_object_id): 

172 try: 

173 sources_lc, forced_lc = self._load_lightcurve_data(int(dia_object_id)) 

174 except Exception as e: 

175 self.log.warning("Failed to load lightcurve for diaObjectId=%s: %s. " 

176 "The DiaSource is likely unassociated.", 

177 dia_object_id, e) 

178 

179 cutout_height_in = max(1.7, 1.7 * len_sizes) 

180 lc_height_in = float(self.config.lightcurve_height) 

181 fig_height = cutout_height_in + lc_height_in 

182 fig = plt.figure(figsize=(7, fig_height), constrained_layout=True) 

183 

184 gs = GridSpec(2, 1, height_ratios=[cutout_height_in, lc_height_in], figure=fig) 

185 cutout_gs = gs[0].subgridspec(len_sizes, 3) 

186 cutout_axes = [[fig.add_subplot(cutout_gs[r, c]) for c in range(3)] 

187 for r in range(len_sizes)] 

188 lc_ax = fig.add_subplot(gs[1]) 

189 

190 def plot_one_image(ax, data, size, name=None): 

191 if name == "Difference": 

192 norm = aviz.ImageNormalize( 

193 data[data.shape[0] // 2 - 7:data.shape[0] // 2 + 8, 

194 data.shape[1] // 2 - 7:data.shape[1] // 2 + 8], 

195 interval=aviz.MinMaxInterval(), 

196 stretch=aviz.AsinhStretch(a=0.1), 

197 ) 

198 else: 

199 norm = aviz.ImageNormalize( 

200 data, 

201 interval=aviz.MinMaxInterval(), 

202 stretch=aviz.AsinhStretch(a=0.1), 

203 ) 

204 ax.imshow(data, cmap=cm.bone, interpolation="none", norm=norm, 

205 extent=(0, size, 0, size), origin="lower", aspect="equal") 

206 x_line = 1 

207 y_line = 1 

208 ax.plot((x_line, x_line + 1.0/scale), (y_line, y_line), color="blue", lw=6) 

209 ax.plot((x_line, x_line + 1.0/scale), (y_line, y_line), color="yellow", lw=2) 

210 ax.axis("off") 

211 if name is not None: 211 ↛ exitline 211 didn't return from function 'plot_one_image' because the condition on line 211 was always true

212 ax.set_title(name) 

213 

214 try: 

215 plot_one_image(cutout_axes[0][0], template[0].image.array, sizes[0], "Template") 

216 plot_one_image(cutout_axes[0][1], science[0].image.array, sizes[0], "Science") 

217 plot_one_image(cutout_axes[0][2], difference[0].image.array, sizes[0], "Difference") 

218 for i in range(1, len_sizes): 218 ↛ 219line 218 didn't jump to line 219 because the loop on line 218 never started

219 plot_one_image(cutout_axes[i][0], template[i].image.array, sizes[i], None) 

220 plot_one_image(cutout_axes[i][1], science[i].image.array, sizes[i], None) 

221 plot_one_image(cutout_axes[i][2], difference[i].image.array, sizes[i], None) 

222 

223 self._draw_lightcurve(lc_ax, sources_lc, forced_lc, current_source=source) 

224 

225 if source is not None and self.config.add_metadata: 225 ↛ 233line 225 didn't jump to line 233 because the condition on line 225 was always true

226 # Place metadata text above the figure top, matching the 

227 # multi-size layout in the parent class. 

228 # ``bbox_inches="tight"`` in savefig expands the saved area to 

229 # include them. 

230 _annotate_image(fig, source, len_sizes, 

231 heights=[1.2, 1.15, 1.1, 1.05, 1.0]) 

232 

233 output = io.BytesIO() 

234 plt.savefig(output, bbox_inches="tight", format="png") 

235 output.seek(0) 

236 finally: 

237 plt.close(fig) 

238 return output 

239 

240 def _draw_lightcurve(self, ax, sources, forced, current_source=None): 

241 """Draw the lightcurve panel for a diaObject. 

242 

243 Parameters 

244 ---------- 

245 ax : `matplotlib.axes.Axes` 

246 Axes to draw into. 

247 sources : `pandas.DataFrame` or None 

248 DiaSources for this diaObject, or None if no APDB is configured. 

249 forced : `pandas.DataFrame` or None 

250 DiaForcedSources for this diaObject. Rows whose ``visit`` matches 

251 one in ``sources`` are suppressed; the rest are drawn with the 

252 ``lightcurve_marker_forced_only`` marker. 

253 current_source : `numpy.record`, optional 

254 The diaSource being cut out. If non-None, a vertical line and 

255 open ring mark its MJD/psfFlux on the panel. 

256 """ 

257 if sources is None and forced is None: 

258 ax.text(0.5, 0.5, "no APDB query configured", 

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

260 ax.set_xticks([]) 

261 ax.set_yticks([]) 

262 return 

263 

264 n_src = 0 if sources is None else len(sources) 

265 n_forced = 0 if forced is None else len(forced) 

266 if n_src == 0 and n_forced == 0: 266 ↛ 267line 266 didn't jump to line 267 because the condition on line 266 was never true

267 ax.text(0.5, 0.5, "no lightcurve data", 

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

269 ax.set_xticks([]) 

270 ax.set_yticks([]) 

271 return 

272 

273 # DiaSources are point-like (with the exception of moving objects, 

274 # which appear only once), so a single visit gives at most one 

275 # diaSource per diaObject — dedup on ``visit`` is safe. 

276 if n_forced and n_src: 

277 forced_only = forced[~forced["visit"].isin(sources["visit"])] 

278 elif n_forced: 278 ↛ 279line 278 didn't jump to line 279 because the condition on line 278 was never true

279 forced_only = forced 

280 else: 

281 forced_only = None 

282 

283 def _plot_group(group, color, marker, label): 

284 time_col = plotUtils._time_column(group) 

285 if "psfFluxErr" in group.columns and group["psfFluxErr"].notna().any(): 

286 ax.errorbar(group[time_col], group["psfFlux"], 

287 yerr=group["psfFluxErr"], fmt=marker, 

288 color=color, label=label) 

289 else: 

290 ax.plot(group[time_col], group["psfFlux"], marker, 

291 color=color, label=label) 

292 

293 if n_src: 293 ↛ 300line 293 didn't jump to line 300 because the condition on line 293 was always true

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

295 color = plotUtils.band_color(band) 

296 _plot_group(group, color, 

297 self.config.lightcurve_marker_source, 

298 f"{band} (n={len(group)})") 

299 

300 if forced_only is not None and len(forced_only): 

301 # Plot per-band forced points with their band colors, but 

302 # suppress them from the legend so we only emit one combined 

303 # entry (in black) for the forced marker regardless of how many 

304 # bands are present. 

305 for band, group in forced_only.groupby("band"): 

306 color = plotUtils.band_color(band) 

307 _plot_group(group, color, 

308 self.config.lightcurve_marker_forced_only, 

309 "_nolegend_") 

310 ax.plot([], [], self.config.lightcurve_marker_forced_only, 

311 color="black", 

312 label=f"forced (n={len(forced_only)})") 

313 

314 if self.config.highlight_current_source and current_source is not None: 314 ↛ 325line 314 didn't jump to line 325 because the condition on line 314 was always true

315 try: 

316 x = float(current_source["midpointMjdTai"]) 

317 y = float(current_source["psfFlux"]) 

318 ax.axvline(x, color="grey", lw=0.5, ls="--") 

319 ax.plot([x], [y], "o", markerfacecolor="none", 

320 markeredgecolor="red", markersize=12, 

321 markeredgewidth=1.5) 

322 except (KeyError, ValueError): 

323 pass 

324 

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

326 ax.set_xlabel("MJD (TAI)") 

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

328 ax.legend(frameon=True, fontsize=7, loc="best") 

329 

330 

331def _make_apdbQuery(sqlitefile=None, postgres_url=None, namespace=None): 

332 """Return a query connection to the specified APDB.""" 

333 if sqlitefile is not None: 

334 return _apdb_mod.ApdbSqliteQuery(sqlitefile) 

335 if postgres_url is not None and namespace is not None: 

336 return _apdb_mod.ApdbPostgresQuery(namespace, postgres_url) 

337 raise RuntimeError("Cannot handle database connection args: " 

338 f"sqlitefile={sqlitefile}, postgres_url={postgres_url}, " 

339 f"namespace={namespace}") 

340 

341 

342def build_argparser(): 

343 """Argument parser for the ``plotDiaSourceLightcurve`` command.""" 

344 parser = argparse.ArgumentParser( 

345 description=__doc__, 

346 formatter_class=argparse.RawDescriptionHelpFormatter, 

347 epilog="More information is available at https://pipelines.lsst.io.", 

348 ) 

349 apdbArgs = parser.add_mutually_exclusive_group(required=True) 

350 apdbArgs.add_argument("--sqlitefile", default=None, 

351 help="Path to sqlite APDB file.") 

352 apdbArgs.add_argument("--namespace", default=None, 

353 help="Postgres namespace (schema) to connect to.") 

354 parser.add_argument( 

355 "--postgres_url", 

356 default="rubin@usdf-prompt-processing-dev.slac.stanford.edu/lsst-devl", 

357 help="Postgres connection path.") 

358 parser.add_argument("--limit", default=5, type=int, 

359 help="Number of sources to load (default=5).") 

360 parser.add_argument("-C", "--configFile", 

361 help="PlotDiaSourceLightcurveConfig file to load.") 

362 parser.add_argument("--collections", nargs="*", 

363 help="Butler collection(s) to load image data from.") 

364 parser.add_argument("repo", help="Path to butler repository.") 

365 parser.add_argument("outputPath", 

366 help="Path to write images to (under outputPath/images/).") 

367 parser.add_argument("--reliabilityMin", type=float, default=None, 

368 help="Minimum reliability for diaSource selection.") 

369 parser.add_argument("--reliabilityMax", type=float, default=None, 

370 help="Maximum reliability for diaSource selection.") 

371 return parser 

372 

373 

374def run_lightcurves(args): 

375 """Run PlotDiaSourceLightcurveTask from parsed command-line arguments.""" 

376 logging.basicConfig(level=logging.INFO, 

377 format="{name} {levelname}: {message}", style="{") 

378 

379 butler = lsst.daf.butler.Butler(args.repo, collections=args.collections) 

380 apdb_query = _make_apdbQuery(sqlitefile=args.sqlitefile, 

381 postgres_url=args.postgres_url, 

382 namespace=args.namespace) 

383 

384 config = PlotDiaSourceLightcurveConfig() 

385 if args.configFile is not None: 

386 config.load(os.path.expanduser(args.configFile)) 

387 config.freeze() 

388 task = PlotDiaSourceLightcurveTask(config=config, 

389 output_path=args.outputPath, 

390 apdb_query=apdb_query) 

391 

392 data = next(apdb_query.iter_sources(args.limit, 

393 args.reliabilityMin, 

394 args.reliabilityMax)) 

395 sources = task.run(data, butler) 

396 print(f"Generated {len(sources)} diaSource lightcurve plots to {args.outputPath}.") 

397 

398 

399def main(): 

400 args = build_argparser().parse_args() 

401 run_lightcurves(args)