Coverage for python/lsst/meas/extensions/multiprofit/plots.py: 0%

155 statements  

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

1# This file is part of meas_extensions_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 

23from abc import ABC, abstractmethod 

24from collections.abc import Iterable 

25from typing import Any, Self 

26 

27import astropy.table 

28import astropy.units as u 

29import matplotlib.axes 

30import matplotlib.figure 

31import matplotlib.pyplot as plt 

32import numpy as np 

33import pydantic 

34 

35from lsst.multiprofit.plotting import bands_weights_lsst, plot_model_rgb 

36 

37from .rebuild_coadd_multiband import DataLoader, PatchCoaddRebuilder 

38 

39__all__ = [ 

40 "ObjectTable", 

41 "ObjectTableBase", 

42 "ObjectTableCModel", 

43 "ObjectTableMultiProFit", 

44 "ObjectTablePsf", 

45 "TruthSummaryTable", 

46 "downselect_table", 

47 "downselect_table_axis", 

48 "plot_blend", 

49 "plot_objects", 

50] 

51 

52Figure = matplotlib.figure.Figure 

53Axes = matplotlib.axes.Axes | Iterable[matplotlib.axes.Axes] 

54FigureAxes = tuple[Figure, Axes] 

55 

56 

57class ObjectTableBase(ABC, pydantic.BaseModel): 

58 """Base class for retrieving columns from tract-based object tables.""" 

59 

60 model_config = pydantic.ConfigDict(arbitrary_types_allowed=True, frozen=True) 

61 

62 table: astropy.table.Table = pydantic.Field(doc="The object table") 

63 

64 @abstractmethod 

65 def get_flux(self, band: str) -> np.ndarray: 

66 """Return the flux in a given band. 

67 

68 Parameters 

69 ---------- 

70 band 

71 The name of the band. 

72 

73 Returns 

74 ------- 

75 flux 

76 The configured flux in that band. 

77 """ 

78 

79 @abstractmethod 

80 def get_id(self) -> np.ndarray: 

81 """Return a unique source id.""" 

82 

83 @abstractmethod 

84 def get_is_extended(self) -> np.ndarray: 

85 """Return if the source is extended.""" 

86 

87 @abstractmethod 

88 def get_is_variable(self) -> np.ndarray: 

89 """Return if the source is variable.""" 

90 

91 @abstractmethod 

92 def get_x(self) -> np.ndarray: 

93 """Return the x pixel coordinates.""" 

94 

95 @abstractmethod 

96 def get_y(self) -> np.ndarray: 

97 """Return the y pixel coordinates.""" 

98 

99 def make_subset(self, subset) -> Self: 

100 """Make a new table of the same type as self with a subset of rows. 

101 

102 Parameters 

103 ---------- 

104 subset 

105 An array that can be used to select asubset of the rows in 

106 self.table. 

107 

108 Returns 

109 ------- 

110 table 

111 An object of the same type as self with a subsetted table. 

112 The table will be a copy, as it does not appear to be possible 

113 to return views of slices of astropy Table instances. 

114 """ 

115 kwargs_table = {name: getattr(self, name) for name in self.model_fields if name != "table"} 

116 return type(self)(table=self.table[subset], **kwargs_table) 

117 

118 

119class TruthSummaryTable(ObjectTableBase): 

120 """Class for retrieving columns from DC2 truth tables.""" 

121 

122 def get_flux(self, band: str) -> np.ndarray: 

123 return self.table[f"flux_{band}"] 

124 

125 def get_id(self) -> np.ndarray: 

126 return self.table["id"] 

127 

128 def get_is_extended(self) -> np.ndarray: 

129 return self.table["is_pointsource"] == False # noqa: E712 

130 

131 def get_is_variable(self) -> np.ndarray: 

132 return self.table["is_variable"] == True # noqa: E712 

133 

134 def get_x(self): 

135 return self.table["x"] 

136 

137 def get_y(self): 

138 return self.table["y"] 

139 

140 

141class ObjectTable(ObjectTableBase, ABC): 

142 """Base class for objectTable_tract.""" 

143 

144 def get_id(self) -> np.ndarray: 

145 return self.table["objectId"] 

146 

147 def get_is_extended(self) -> np.ndarray: 

148 return self.table["refExtendedness"] >= 0.5 

149 

150 def get_is_variable(self) -> np.ndarray: 

151 return np.zeros(len(self.table), dtype=bool) 

152 

153 def get_x(self): 

154 return self.table["x"] 

155 

156 def get_y(self): 

157 return self.table["y"] 

158 

159 

160class ObjectTableCModel(ObjectTable): 

161 """Class for retrieving CModel fluxes from objectTable_tract.""" 

162 

163 def get_flux(self, band: str) -> np.ndarray: 

164 return self.table[f"{band}_cModelFlux"] 

165 

166 

167class ObjectTableMultiProFit(ObjectTableBase): 

168 """Class for retrieving fluxes from objectTable_tract_multiprofit.""" 

169 

170 name_model: str = pydantic.Field(doc="The name of the MultiProFit model") 

171 prefix_col: str = pydantic.Field(doc="The prefix for object fit columns", default="mpf_") 

172 

173 def get_flux(self, band: str) -> np.ndarray: 

174 return self.table[f"{self.prefix_col}{self.name_model}_{band}_flux"] 

175 

176 def get_id(self) -> np.ndarray: 

177 return self.table["objectId"] 

178 

179 def get_is_extended(self) -> np.ndarray: 

180 return self.table["refExtendedness"] >= 0.5 

181 

182 def get_is_variable(self) -> np.ndarray: 

183 return np.zeros(len(self.table), dtype=bool) 

184 

185 def get_x(self): 

186 return self.table[f"{self.prefix_col}{self.name_model}_cen_x"] 

187 

188 def get_y(self): 

189 return self.table[f"{self.prefix_col}{self.name_model}_cen_y"] 

190 

191 

192class ObjectTablePsf(ObjectTable): 

193 """Class for retreiving PSF fluxes from objectTable_tract.""" 

194 

195 def get_flux(self, band: str) -> np.ndarray: 

196 return self.table[f"{band}_psfFlux"] 

197 

198 

199def downselect_table( 

200 table: ObjectTableBase, 

201 x_min: float, 

202 x_max: float, 

203 y_min: float, 

204 y_max: float, 

205) -> ObjectTableBase: 

206 """Select points from a table within an x,y extent. 

207 

208 Parameters 

209 ---------- 

210 table 

211 The table to downselect. 

212 x_min 

213 The minimum x value. 

214 x_max 

215 The maximum x value. 

216 y_min 

217 The minimum y value. 

218 y_max 

219 The maximum y value. 

220 

221 Returns 

222 ------- 

223 table 

224 A downselected table of the same class. 

225 """ 

226 x_all = table.get_x() 

227 y_all = table.get_y() 

228 within = (x_all > x_min) & (x_all < x_max) & (y_all > y_min) & (y_all < y_max) 

229 return table.make_subset(within) 

230 

231 

232def downselect_table_axis(table: ObjectTableBase, axis) -> ObjectTableBase: 

233 """Select points from a table within a figure axis. 

234 

235 Parameters 

236 ---------- 

237 table 

238 The table to downselect. 

239 axis 

240 The figure axis to determine the extent from. 

241 

242 Returns 

243 ------- 

244 table 

245 A downselected table of the same class. 

246 """ 

247 extent = np.array(axis.axis()) 

248 return downselect_table(table, extent[0], extent[1], extent[2], extent[3]) 

249 

250 

251def plot_objects( 

252 table: ObjectTableBase, 

253 axes: Axes, 

254 bands: Iterable[str], 

255 table_downselected: bool = False, 

256 kwargs_annotate: dict[str, Any] = None, 

257 kwargs_scatter: dict[str, Any] = None, 

258 labels_extended: tuple[str, str] = ("S", "G"), 

259) -> Axes: 

260 """Plot catalog objects on an existing image. 

261 

262 Parameters 

263 ---------- 

264 table 

265 The object table to plot source from. 

266 axes 

267 The figure axes to plot on. 

268 bands 

269 The bands to sum over fluxes to derive a total mag label. 

270 table_downselected 

271 Whether the table has already been downselected to contain only 

272 points within the bounds of the axes. 

273 kwargs_annotate 

274 Keyword arguments to pass to axes.annotate. 

275 kwargs_scatter 

276 Keyword arguments to pass to axes.scatter. 

277 labels_extended 

278 Label prefixes for non-extended and extended objects, respectively. 

279 

280 Returns 

281 ------- 

282 axes 

283 The input axes with added points and labels. 

284 """ 

285 if kwargs_annotate is None: 

286 kwargs_annotate = dict(color="white", fontsize=14, ha="left", va="bottom") 

287 if kwargs_scatter is None: 

288 kwargs_scatter = dict(c="white", marker="+", s=100) 

289 table_within = table if table_downselected else downselect_table_axis(table, axes) 

290 x = table_within.get_x() 

291 y = table_within.get_y() 

292 axes.scatter(x, y, **kwargs_scatter) 

293 fluxes = [table_within.get_flux(band) for band in bands] 

294 is_extended = table_within.get_is_extended() 

295 is_variable = table_within.get_is_variable() 

296 

297 for idx in range(len(table_within.table)): 

298 mag = u.nanojansky.to(u.ABmag, np.sum([fluxcol[idx] for fluxcol in fluxes])) 

299 type_src = f"{'V' if is_variable[idx] else ''}{labels_extended[1 if is_extended[idx] else 0]}" 

300 axes.annotate(f"{type_src}{mag:.1f}", (x[idx], y[idx]), **kwargs_annotate) 

301 

302 return axes 

303 

304 

305def plot_blend( 

306 rebuilder: PatchCoaddRebuilder, 

307 idx_row_parent: int, 

308 weights: dict[str, float] = None, 

309 table_ref_type: type = TruthSummaryTable, 

310 kwargs_plot_parent: dict[str, Any] = None, 

311 kwargs_plot_children: dict[str, Any] = None, 

312) -> tuple[Figure, Axes, Figure, Axes]: 

313 """Plot an image of an entire blend and its deblended children. 

314 

315 Parameters 

316 ---------- 

317 rebuilder 

318 The patch rebuilder to plot from. 

319 idx_row_parent 

320 The row index of the parent object in the reference SourceCatalog. 

321 weights 

322 Multiplicative weights by band name for RGB plots. 

323 table_ref_type 

324 The type of reference table to construct when downselecting. 

325 kwargs_plot_parent 

326 Keyword arguments to pass to make RGB plots of the parent blend. 

327 kwargs_plot_children 

328 Keyword arguments to pass to make RGB plots of deblended children. 

329 

330 Returns 

331 ------- 

332 fig_rgb 

333 The Figure for the RGB plots of the parent. 

334 ax_rgb 

335 The Axes for the RGB plots of the parent. 

336 fig_gs 

337 The Figure for the grayscale plots of the parent. 

338 ax_gs 

339 The Axes for the grayscale plots of the parent. 

340 """ 

341 if kwargs_plot_parent is None: 

342 kwargs_plot_parent = {} 

343 if kwargs_plot_children is None: 

344 kwargs_plot_children = {} 

345 if weights is None: 

346 weights = bands_weights_lsst 

347 

348 plot_chi_hist = kwargs_plot_children.pop("plot_chi_hist", True) 

349 rebuilder_ref = rebuilder.matches[rebuilder.name_model_ref].rebuilder 

350 observations = { 

351 catexp.band: catexp.get_source_observation(catexp.get_catalog()[idx_row_parent], skip_flags=True) 

352 for catexp in rebuilder_ref.catexps 

353 } 

354 

355 fig_rgb, ax_rgb, fig_gs, ax_gs, *_ = plot_model_rgb( 

356 model=None, 

357 weights=weights, 

358 observations=observations, 

359 plot_singleband=False, 

360 plot_chi_hist=False, 

361 **kwargs_plot_parent, 

362 ) 

363 table_within_ref = downselect_table_axis(table_ref_type(table=rebuilder.reference), ax_rgb) 

364 plot_objects(table_within_ref, ax_rgb, weights, table_downselected=True) 

365 

366 objects_primary = rebuilder.objects[rebuilder.objects["detect_isPrimary"] == True] # noqa: E712 

367 kwargs_annotate_obs = dict(color="white", fontsize=14, ha="right", va="top") 

368 kwargs_scatter_obs = dict(c="white", marker="x", s=70) 

369 table_within_cmodel = downselect_table_axis(ObjectTableCModel(table=objects_primary), ax_rgb) 

370 labels_extended_model = ("C", "E") 

371 plot_objects( 

372 table_within_cmodel, 

373 ax_rgb, 

374 weights, 

375 table_downselected=True, 

376 kwargs_annotate=kwargs_annotate_obs, 

377 kwargs_scatter=kwargs_scatter_obs, 

378 labels_extended=labels_extended_model, 

379 ) 

380 plt.show() 

381 

382 objects_mpf = rebuilder.objects_multiprofit 

383 objects_mpf_within = {} 

384 for name, matched in rebuilder.matches.items(): 

385 if matched.rebuilder and objects_mpf: 

386 objects_mpf_within[name] = downselect_table_axis( 

387 ObjectTableMultiProFit(name_model=name, table=objects_mpf), 

388 ax_rgb, 

389 ) 

390 

391 cat_ref = rebuilder_ref.catalog_multi 

392 row_parent = cat_ref[idx_row_parent] 

393 idx_children = ( 

394 (idx_row_parent,) 

395 if (row_parent["parent"] == 0) 

396 else (np.where(rebuilder_ref.catalog_multi["parent"] == row_parent["id"])[0]) 

397 ) 

398 

399 for idx_child in idx_children: 

400 for name, matched in rebuilder.matches.items(): 

401 print(f"Model: {name}") 

402 rebuilder_child = matched.rebuilder 

403 is_dataloader = isinstance(rebuilder_child, DataLoader) 

404 is_scarlet = is_dataloader and (name == "scarlet") 

405 if is_scarlet or rebuilder_child: 

406 try: 

407 if is_dataloader: 

408 model = None 

409 observations = rebuilder_child.load_deblended_object(idx_child) 

410 else: 

411 model = rebuilder_child.make_model(idx_child) 

412 observations = None 

413 

414 _, ax_rgb_c, *_ = plot_model_rgb( 

415 model=model, 

416 weights=weights, 

417 plot_singleband=False, 

418 plot_chi_hist=(not is_dataloader) and plot_chi_hist, 

419 observations=observations, 

420 **kwargs_plot_children, 

421 ) 

422 ax_rgb_c0 = ax_rgb_c[0][0] 

423 plot_objects(table_within_ref, ax_rgb_c0, weights) 

424 tab_mpf = objects_mpf_within.get(name) 

425 if tab_mpf: 

426 plot_objects( 

427 tab_mpf, 

428 ax_rgb_c0, 

429 weights, 

430 kwargs_annotate=kwargs_annotate_obs, 

431 kwargs_scatter=kwargs_scatter_obs, 

432 labels_extended=labels_extended_model, 

433 ) 

434 plt.show() 

435 except Exception as exc: 

436 print(f"{idx_child=} failed to rebuild due to {exc}") 

437 

438 return fig_rgb, ax_rgb, fig_gs, ax_gs