Coverage for python/lsst/meas/extensions/multiprofit/plots.py: 0%
155 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 12:36 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 12:36 +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/>.
23from abc import ABC, abstractmethod
24from collections.abc import Iterable
25from typing import Any, Self
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
35from lsst.multiprofit.plotting import bands_weights_lsst, plot_model_rgb
37from .rebuild_coadd_multiband import DataLoader, PatchCoaddRebuilder
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]
52Figure = matplotlib.figure.Figure
53Axes = matplotlib.axes.Axes | Iterable[matplotlib.axes.Axes]
54FigureAxes = tuple[Figure, Axes]
57class ObjectTableBase(ABC, pydantic.BaseModel):
58 """Base class for retrieving columns from tract-based object tables."""
60 model_config = pydantic.ConfigDict(arbitrary_types_allowed=True, frozen=True)
62 table: astropy.table.Table = pydantic.Field(doc="The object table")
64 @abstractmethod
65 def get_flux(self, band: str) -> np.ndarray:
66 """Return the flux in a given band.
68 Parameters
69 ----------
70 band
71 The name of the band.
73 Returns
74 -------
75 flux
76 The configured flux in that band.
77 """
79 @abstractmethod
80 def get_id(self) -> np.ndarray:
81 """Return a unique source id."""
83 @abstractmethod
84 def get_is_extended(self) -> np.ndarray:
85 """Return if the source is extended."""
87 @abstractmethod
88 def get_is_variable(self) -> np.ndarray:
89 """Return if the source is variable."""
91 @abstractmethod
92 def get_x(self) -> np.ndarray:
93 """Return the x pixel coordinates."""
95 @abstractmethod
96 def get_y(self) -> np.ndarray:
97 """Return the y pixel coordinates."""
99 def make_subset(self, subset) -> Self:
100 """Make a new table of the same type as self with a subset of rows.
102 Parameters
103 ----------
104 subset
105 An array that can be used to select asubset of the rows in
106 self.table.
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)
119class TruthSummaryTable(ObjectTableBase):
120 """Class for retrieving columns from DC2 truth tables."""
122 def get_flux(self, band: str) -> np.ndarray:
123 return self.table[f"flux_{band}"]
125 def get_id(self) -> np.ndarray:
126 return self.table["id"]
128 def get_is_extended(self) -> np.ndarray:
129 return self.table["is_pointsource"] == False # noqa: E712
131 def get_is_variable(self) -> np.ndarray:
132 return self.table["is_variable"] == True # noqa: E712
134 def get_x(self):
135 return self.table["x"]
137 def get_y(self):
138 return self.table["y"]
141class ObjectTable(ObjectTableBase, ABC):
142 """Base class for objectTable_tract."""
144 def get_id(self) -> np.ndarray:
145 return self.table["objectId"]
147 def get_is_extended(self) -> np.ndarray:
148 return self.table["refExtendedness"] >= 0.5
150 def get_is_variable(self) -> np.ndarray:
151 return np.zeros(len(self.table), dtype=bool)
153 def get_x(self):
154 return self.table["x"]
156 def get_y(self):
157 return self.table["y"]
160class ObjectTableCModel(ObjectTable):
161 """Class for retrieving CModel fluxes from objectTable_tract."""
163 def get_flux(self, band: str) -> np.ndarray:
164 return self.table[f"{band}_cModelFlux"]
167class ObjectTableMultiProFit(ObjectTableBase):
168 """Class for retrieving fluxes from objectTable_tract_multiprofit."""
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_")
173 def get_flux(self, band: str) -> np.ndarray:
174 return self.table[f"{self.prefix_col}{self.name_model}_{band}_flux"]
176 def get_id(self) -> np.ndarray:
177 return self.table["objectId"]
179 def get_is_extended(self) -> np.ndarray:
180 return self.table["refExtendedness"] >= 0.5
182 def get_is_variable(self) -> np.ndarray:
183 return np.zeros(len(self.table), dtype=bool)
185 def get_x(self):
186 return self.table[f"{self.prefix_col}{self.name_model}_cen_x"]
188 def get_y(self):
189 return self.table[f"{self.prefix_col}{self.name_model}_cen_y"]
192class ObjectTablePsf(ObjectTable):
193 """Class for retreiving PSF fluxes from objectTable_tract."""
195 def get_flux(self, band: str) -> np.ndarray:
196 return self.table[f"{band}_psfFlux"]
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.
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.
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)
232def downselect_table_axis(table: ObjectTableBase, axis) -> ObjectTableBase:
233 """Select points from a table within a figure axis.
235 Parameters
236 ----------
237 table
238 The table to downselect.
239 axis
240 The figure axis to determine the extent from.
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])
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.
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.
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()
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)
302 return axes
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.
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.
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
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 }
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)
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()
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 )
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 )
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
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}")
438 return fig_rgb, ax_rgb, fig_gs, ax_gs