Coverage for python/lsst/multiprofit/plotting/plot_model_rgb.py: 73%
229 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-22 02:34 -0700
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-22 02:34 -0700
1# This file is part of 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/>.
22__all__ = ["plot_model_rgb"]
24import math
25from typing import Any
27import astropy.visualization as apVis
28import matplotlib as mpl
29import matplotlib.pyplot as plt
30import numpy as np
32import lsst.gauss2d as g2
33import lsst.gauss2d.fit as g2f
35from .types import Axes, Figure
38def plot_model_rgb(
39 model: g2f.ModelD | None,
40 weights: dict[str, float] | None = None,
41 high_sn_threshold: float | None = None,
42 plot_singleband: bool = True,
43 plot_chi_hist: bool = True,
44 chi_max: float = 5.0,
45 rgb_min_auto: bool = False,
46 rgb_stretch_auto: bool = False,
47 **kwargs: Any,
48) -> tuple[Figure, Axes, Figure, Axes, np.ndarray]:
49 """Plot RGB images of a model, its data and residuals thereof.
51 Parameters
52 ----------
53 model
54 The model to plot. If None, a dict of observations by band may be
55 passed as an additional kwarg; otherwise, only the data will be
56 plotted.
57 weights
58 Linear weights to multiply each band's image by. The default is a
59 weight of one for each band.
60 high_sn_threshold
61 If non-None and given a model, this will return an image with the
62 pixels having a model S/N above this threshold in every band.
63 plot_singleband
64 Whether to make grayscale plots for each band.
65 plot_chi_hist
66 Whether to plot histograms of the chi (scaled residual) values.
67 chi_max
68 The maximum absolute value of chi in residual plots. Values of 3-5 are
69 suitable for good models while inadequate ones may need larger values.
70 rgb_min_auto
71 Whether to set the minimum in RGB plots automatically. Cannot supply
72 minimum in kwargs if enabled.
73 rgb_stretch_auto
74 Whether to set the stretch in RGB plots automatically. Cannot supply
75 stretch in kwargs if enabled.
76 **kwargs
77 Additional keyword arguments to pass to make_lupton_rgb when creating
78 RGB images.
80 Returns
81 -------
82 fig_rgb
83 The Figure for the RGB plots.
84 ax_rgb
85 The Axes for the RGB plots.
86 fig_gs
87 The Figure for the grayscale plots.
88 ax_gs
89 The Axes for the grayscale plots.
90 mask_inv_highsn
91 The inverse mask (1=selected) if high_sn_threshold was specified.
92 """
93 if rgb_min_auto and "minimum" in kwargs: 93 ↛ 94line 93 didn't jump to line 94 because the condition on line 93 was never true
94 raise ValueError(f"Cannot set rgb_min_auto and pass {kwargs['minimum']=}")
95 if rgb_stretch_auto and "stretch" in kwargs: 95 ↛ 96line 95 didn't jump to line 96 because the condition on line 95 was never true
96 raise ValueError(f"Cannot set rgb_stretch_auto and pass {kwargs['stretch']=}")
97 if not (chi_max > 0): 97 ↛ 98line 97 didn't jump to line 98 because the condition on line 97 was never true
98 raise ValueError(f"{chi_max=} not >0")
99 if weights is None: 99 ↛ 100line 99 didn't jump to line 100 because the condition on line 99 was never true
100 if model is None:
101 weights = {band: 1.0 for band in kwargs["observations"].keys()}
102 else:
103 bands_set = set()
104 bands = []
105 weights = {}
106 for obs in model.data:
107 band = obs.channel.name
108 if band not in bands_set:
109 bands_set.add(band)
110 bands.append(band)
111 weights[band] = 1.0
113 n_data = len(model.data)
114 has_model = model is not None
115 observations = {}
116 models = {}
118 if has_model and (n_data < 3): 118 ↛ 119line 118 didn't jump to line 119 because the condition on line 118 was never true
119 if n_data == 1:
120 # pretend this is three bands
121 obs, output_data = model.data[0], model.outputs[0].data
122 band = obs.channel.name
123 weights = {}
124 for idx in range(1, 4):
125 key = f"{band}{idx}"
126 weights[key] = 1.0
127 observations[key] = obs
128 models[key] = output_data
129 elif n_data == 2:
130 raise NotImplementedError("RGB images for two-band data are not supported (yet)")
132 bands = tuple(weights.keys())
133 band_str = ",".join(bands)
134 n_bands = len(bands)
136 if has_model and (not model.outputs or any([output is None for output in model.outputs])): 136 ↛ 137line 136 didn't jump to line 137 because the condition on line 136 was never true
137 model.setup_evaluators(g2f.EvaluatorMode.image)
138 model.evaluate()
140 if not has_model: 140 ↛ 141line 140 didn't jump to line 141 because the condition on line 140 was never true
141 if plot_chi_hist:
142 raise ValueError("Cannot plot chi histograms without a model")
143 obs_kwarg = kwargs.pop("observations")
144 for band in bands:
145 observations[band] = obs_kwarg[band]
147 x_min, x_max, y_min, y_max = np.inf, -np.inf, np.inf, -np.inf
148 coordsys_last = None
149 if has_model and n_data >= 3: 149 ↛ 158line 149 didn't jump to line 158 because the condition on line 149 was always true
150 for obs, output in zip(model.data, model.outputs):
151 band = obs.channel.name
152 if band in bands: 152 ↛ 150line 152 didn't jump to line 150 because the condition on line 152 was always true
153 if band in observations: 153 ↛ 154line 153 didn't jump to line 154 because the condition on line 153 was never true
154 raise ValueError(f"Cannot plot {model=} because {band=} has multiple observations")
155 observations[band] = obs
156 models[band] = output.data
158 for band, obs in observations.items():
159 coordsys = obs.image.coordsys
160 if coordsys: 160 ↛ 166line 160 didn't jump to line 166 because the condition on line 160 was always true
161 coordsys_last = coordsys
162 x_min = int(round(min(x_min, coordsys.x_min), 0))
163 x_max = int(round(max(x_max, coordsys.x_min + obs.image.n_cols), 0))
164 y_min = int(round(min(y_min, coordsys.y_min), 0))
165 y_max = int(round(max(y_max, coordsys.y_min + obs.image.n_rows), 0))
166 elif coordsys_last is not None:
167 raise ValueError(
168 f"coordinate system for {band=} is None but last was not; they must either "
169 f"all be None or all non-None"
170 )
172 if coordsys_last: 172 ↛ 197line 172 didn't jump to line 197 because the condition on line 172 was always true
173 shape_new = (y_max - y_min, x_max - x_min)
174 keys = ("image", "mask_inv", "sigma_inv")
175 if has_model: 175 ↛ 177line 175 didn't jump to line 177 because the condition on line 175 was always true
176 keys += ("model",)
177 for band, obs in observations.items():
178 coordsys = obs.image.coordsys
179 x_min_c = int(round(coordsys.x_min, 0)) - x_min
180 y_min_c = int(round(coordsys.y_min, 0)) - y_min
181 x_min_o, x_max_o = x_min_c, x_min_c + obs.image.n_cols
182 y_min_o, y_max_o = y_min_c, y_min_c + obs.image.n_rows
183 if x_min_o or x_max_o or y_min_o or y_max_o: 183 ↛ 177line 183 didn't jump to line 177 because the condition on line 183 was always true
184 # zero-pad the relevant images into a new observation
185 data_new = {}
186 for key in keys:
187 img = np.zeros(shape_new)
188 img[y_min_o:y_max_o, x_min_o:x_max_o] = (
189 models[band] if (key == "model") else getattr(obs, key).data
190 )
191 if key == "model":
192 models[band] = img
193 else:
194 data_new[key] = (g2.ImageB if (key == "mask_inv") else g2.ImageD)(img)
195 observations[band] = g2f.ObservationD(channel=obs.channel, **data_new)
197 extent = (x_min, x_max, y_min, y_max)
199 images_data = [None] * 3
200 images_data_unweighted = [None] * 3 if has_model else None
201 images_model = [None] * 3 if has_model else None
202 images_model_unweighted = [None] * 3 if has_model else None
203 images_sigma_inv = [None] * 3 if has_model else None
204 masks_inv_rgb = [None] * 3
206 weights_channel = np.linspace(0, 3, len(weights) + 1)[1:]
207 idx_channel = 0
208 weight_channel = 0
210 def add_if_not_none(array: np.ndarray, index: int, arg: float | None) -> None:
211 if array[index] is not None:
212 array[index] += arg
213 else:
214 array[index] = arg
216 chis_unweighted = {}
218 for idx_band, (band, weight) in enumerate(weights.items()):
219 observation = observations[band]
220 if has_model: 220 ↛ 233line 220 didn't jump to line 233 because the condition on line 220 was always true
221 model_band = models[band]
222 sigma_inv = observation.sigma_inv.data
223 sigma_inv_good = sigma_inv > 0
224 variance_band = np.empty_like(sigma_inv)
225 variance_band[sigma_inv_good] = sigma_inv[sigma_inv_good] ** -2
226 variance_band[~sigma_inv_good] = np.nan
227 if plot_chi_hist:
228 chi_good = (sigma_inv > 0) & np.isfinite(sigma_inv)
229 chi_unweighted = (observation.image.data[chi_good] - model_band[chi_good]) * sigma_inv[
230 chi_good
231 ]
232 chis_unweighted[band] = chi_unweighted
233 weight_channel_new = weights_channel[idx_band]
234 idx_channel_new = int(weight_channel_new // 1)
235 if idx_channel_new == idx_channel:
236 weight_low = weight_channel_new - weight_channel
237 weight_high = 0.0
238 else:
239 weight_low = idx_channel_new - weight_channel
240 weight_high = weight_channel_new - idx_channel_new
241 assert weight_high >= 0
242 assert weight_low >= 0
243 if weight_low > 0: 243 ↛ 253line 243 didn't jump to line 253 because the condition on line 243 was always true
244 data_band = observation.image.data * weight_low
245 add_if_not_none(images_data, idx_channel, data_band * weight)
246 add_if_not_none(masks_inv_rgb, idx_channel, observation.mask_inv.data * weight_low)
247 if has_model: 247 ↛ 253line 247 didn't jump to line 253 because the condition on line 247 was always true
248 add_if_not_none(images_data_unweighted, idx_channel, data_band)
249 model_sub = model_band * weight_low
250 add_if_not_none(images_model, idx_channel, model_sub * weight)
251 add_if_not_none(images_model_unweighted, idx_channel, model_sub)
252 add_if_not_none(images_sigma_inv, idx_channel, variance_band * weight_low)
253 if (idx_channel_new != idx_channel) and (weight_high > 0): 253 ↛ 254line 253 didn't jump to line 254 because the condition on line 253 was never true
254 data_band = observation.image.data * weight_high
255 images_data[idx_channel_new] = data_band * weight
256 masks_inv_rgb[idx_channel_new] = observation.mask_inv.data * weight_low
257 if has_model:
258 images_model_unweighted[idx_channel_new] = data_band
259 model_sub = model_band * weight_high
260 images_model[idx_channel_new] = model_sub * weight
261 images_model_unweighted[idx_channel_new] = model_sub
262 images_sigma_inv[idx_channel_new] = variance_band * weight_high
263 weight_channel = weight_channel_new
264 idx_channel = idx_channel_new
266 # convert variance to 1/sigma
267 if has_model: 267 ↛ 271line 267 didn't jump to line 271 because the condition on line 267 was always true
268 for idx in range(3):
269 images_sigma_inv[idx] = 1 / np.sqrt(images_sigma_inv[idx])
271 if rgb_min_auto or rgb_stretch_auto:
272 # The model won't have negative pixels, so it ought to stretch fine
273 # the max/stretch is not as important anyway
274 rgb_min, rgb_max = np.nanpercentile(
275 np.concatenate([image[mask_inv != 0] for mask_inv, image in zip(masks_inv_rgb, images_data)]),
276 (5, 95),
277 )
278 if rgb_min_auto: 278 ↛ 280line 278 didn't jump to line 280 because the condition on line 278 was always true
279 kwargs["minimum"] = rgb_min
280 if rgb_stretch_auto: 280 ↛ 283line 280 didn't jump to line 283 because the condition on line 280 was always true
281 kwargs["stretch"] = 2 * (rgb_max - rgb_min)
283 img_rgb = apVis.make_lupton_rgb(*images_data, **kwargs)
284 if has_model: 284 ↛ 286line 284 didn't jump to line 286 because the condition on line 284 was always true
285 img_model_rgb = apVis.make_lupton_rgb(*images_model, **kwargs)
286 aspect = np.clip((y_max - y_min) / (x_max - x_min), 0.25, 4)
288 n_rows = 1 + has_model
289 n_cols_gs = 1 + has_model
290 n_cols_rgb = 1 + has_model * (1 + plot_chi_hist)
291 figsize_y = 8 * n_rows * aspect
293 fig_rgb, ax_rgb = plt.subplots(nrows=n_rows, ncols=n_cols_rgb, figsize=(8 * n_cols_rgb, figsize_y))
294 fig_gs, ax_gs = (
295 (None, None)
296 if not plot_singleband
297 else plt.subplots(
298 nrows=n_bands,
299 ncols=n_cols_gs,
300 figsize=(8 * n_cols_gs, 8 * aspect * n_bands),
301 )
302 )
303 (ax_rgb[0][0] if has_model else ax_rgb).imshow(img_rgb, extent=extent, origin="lower")
304 (ax_rgb[0][0] if has_model else ax_rgb).set_title("Data")
305 if has_model: 305 ↛ 309line 305 didn't jump to line 309 because the condition on line 305 was always true
306 ax_rgb[1][0].imshow(img_model_rgb, extent=extent, origin="lower")
307 ax_rgb[1][0].set_title(f"Model ({band_str})")
309 masks_inv = {}
310 # Create a mask of high-sn pixels (based on the model)
311 mask_inv_highsn = np.ones(img_rgb.shape[:1], dtype="bool") if high_sn_threshold else None
313 for idx, band in enumerate(bands):
314 obs = observations[band]
315 mask_inv = obs.mask_inv.data
316 masks_inv[band] = mask_inv
317 img_data = obs.image.data
318 img_sigma_inv = obs.sigma_inv.data
319 if plot_singleband:
320 if has_model: 320 ↛ 337line 320 didn't jump to line 337 because the condition on line 320 was always true
321 img_model = models[band]
322 if mask_inv_highsn: 322 ↛ 323line 322 didn't jump to line 323 because the condition on line 322 was never true
323 mask_inv_highsn *= (img_model * np.nanmedian(img_sigma_inv)) > high_sn_threshold
324 residual = (img_data - img_model) * mask_inv
325 value_max = np.nanpercentile(np.abs(residual), 98)
326 ax_gs[idx][0].imshow(residual, cmap="gray", vmin=-value_max, vmax=value_max, origin="lower")
327 ax_gs[idx][0].tick_params(labelleft=False)
328 ax_gs[idx][0].set_title(f"{band}-band Residual (abs.)")
329 ax_gs[idx][1].imshow(
330 np.clip(residual * img_sigma_inv, -chi_max, chi_max),
331 cmap="gray",
332 origin="lower",
333 )
334 ax_gs[idx][1].tick_params(labelleft=False)
335 ax_gs[idx][1].set_title(f"{band}-band Residual (chi, +/- {chi_max:.2f})")
336 else:
337 ax_gs[idx].imshow(img_data * mask_inv * (img_sigma_inv > 0), cmap="gray", origin="lower")
338 ax_gs[idx].set_title(band)
340 if has_model: 340 ↛ 401line 340 didn't jump to line 401 because the condition on line 340 was always true
341 # TODO: Draw masks in each channel? or draw the combined mask, like:
342 # mask_inv_all = np.prod(list(masks_inv.values()), axis=0)
343 residuals = [(images_model_unweighted[idx] - images_data_unweighted[idx]) for idx in range(3)]
344 resid_max = np.nanpercentile(
345 np.abs(np.concatenate([residual[np.isfinite(residual)] for residual in residuals])), 98
346 )
348 # This may or may not be equivalent to make_lupton_rgb
349 # I just can't figure out how to get that scaled so zero = 50% gray
350 stretch = 3
351 residual_rgb = np.stack(
352 [np.arcsinh(np.clip(residuals[idx], -resid_max, resid_max) * stretch) for idx in range(3)],
353 axis=-1,
354 )
355 residual_rgb /= 2 * np.arcsinh(resid_max * stretch)
356 residual_rgb += 0.5
358 ax_rgb[0][1].imshow(residual_rgb, origin="lower")
359 ax_rgb[0][1].set_title(f"Residual (abs., += {resid_max:.3e})")
360 ax_rgb[0][1].tick_params(labelleft=False)
362 if plot_chi_hist:
363 cmap = mpl.colormaps["coolwarm"]
364 residuals_rgb = np.concatenate(tuple(chis_unweighted.values()))
365 residuals_abs = np.abs(residuals_rgb)
366 n_resid = len(residuals_abs)
367 chi_max = 5 + 2.5 * (
368 (np.sum(residuals_abs > 5) / n_resid > 0.1) + (np.sum(residuals_abs > 7.5) / n_resid > 0.1)
369 )
370 n_bins = int(math.ceil(np.clip(n_resid / 50, 2, 20)) * chi_max)
371 # ax_rgb[0][2].set_adjustable('box')
372 ax_rgb[0][2].hist(
373 np.clip(residuals_rgb, -chi_max, chi_max),
374 bins=n_bins,
375 histtype="step",
376 label="all",
377 )
378 band_colors = cmap(np.linspace(0, 1, n_bands))
379 for band, band_color in zip(bands, band_colors):
380 ax_rgb[0][2].hist(
381 np.clip(residuals_rgb, -chi_max, chi_max),
382 bins=n_bins,
383 histtype="step",
384 label=band,
385 )
386 ax_rgb[0][2].legend()
388 # TODO: Plot unscaled residuals in ax_rgb[1][2]? It's unused now.
389 residual_rgb = np.stack(
390 [
391 (np.clip(residuals[idx] * images_sigma_inv[idx], -chi_max, chi_max) + chi_max) / (2 * chi_max)
392 for idx in range(3)
393 ],
394 axis=-1,
395 )
397 ax_rgb[1][1].imshow(residual_rgb, origin="lower")
398 ax_rgb[1][1].set_title(f"Residual (chi, +/- {chi_max:.2f})")
399 ax_rgb[1][1].tick_params(labelleft=False)
401 return fig_rgb, ax_rgb, fig_gs, ax_gs, mask_inv_highsn