Coverage for python/lsst/multiprofit/plotting/plot_model_rgb.py: 73%

229 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-09-14 02:16 -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/>. 

21 

22__all__ = ["plot_model_rgb"] 

23 

24import math 

25from typing import Any 

26 

27import astropy.visualization as apVis 

28import matplotlib as mpl 

29import matplotlib.pyplot as plt 

30import numpy as np 

31 

32import lsst.gauss2d as g2 

33import lsst.gauss2d.fit as g2f 

34 

35from .types import Axes, Figure 

36 

37 

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. 

50 

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. 

79 

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 

112 

113 n_data = len(model.data) 

114 has_model = model is not None 

115 observations = {} 

116 models = {} 

117 

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)") 

131 

132 bands = tuple(weights.keys()) 

133 band_str = ",".join(bands) 

134 n_bands = len(bands) 

135 

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() 

139 

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] 

146 

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 

157 

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 ) 

171 

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) 

196 

197 extent = (x_min, x_max, y_min, y_max) 

198 

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 

205 

206 weights_channel = np.linspace(0, 3, len(weights) + 1)[1:] 

207 idx_channel = 0 

208 weight_channel = 0 

209 

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 

215 

216 chis_unweighted = {} 

217 

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 

265 

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]) 

270 

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) 

282 

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) 

287 

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 

292 

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})") 

308 

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 

312 

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) 

339 

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 ) 

347 

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 

357 

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) 

361 

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() 

387 

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 ) 

396 

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) 

400 

401 return fig_rgb, ax_rgb, fig_gs, ax_gs, mask_inv_highsn