Coverage for python/lsst/meas/extensions/scarlet/deconvolveExposureTask.py: 69%

142 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-14 08:07 +0000

1# This file is part of meas_extensions_scarlet. 

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 

22import dataclasses 

23import logging 

24 

25import lsst.afw.detection as afwDet 

26import lsst.afw.image as afwImage 

27import lsst.afw.table as afwTable 

28import lsst.images as imgs 

29import lsst.pex.config as pexConfig 

30import lsst.pipe.base as pipeBase 

31import lsst.pipe.base.connectionTypes as cT 

32import lsst.scarlet.lite as scl 

33import numpy as np 

34from lsst.images.cells import CellCoadd 

35 

36from . import utils 

37 

38log = logging.getLogger(__name__) 

39 

40__all__ = [ 

41 "DeconvolveExposureTask", 

42 "DeconvolveExposureConfig", 

43 "DeconvolveExposureConnections", 

44] 

45 

46 

47def calculate_update_step( 

48 observation: scl.Observation, 

49 min_scale: float = 0.01, 

50 default_scale: float = 0.1, 

51) -> float: 

52 """Calculate the scale factor for the update step in deconvolution. 

53 

54 For most images this will be 1.0 but for images with low SNR 

55 and/or high sparsity (for example LSST u-band images) the scale 

56 factor will be less than 1.0. 

57 

58 Parameters 

59 ---------- 

60 observation : 

61 Scarlet lite Observation. 

62 

63 min_scale : 

64 Minimum allowed scale factor. 

65 

66 default_scale : 

67 Default scale factor to return if noise level is non-finite. 

68 

69 Returns 

70 ------- 

71 scale : float 

72 Scale factor for the update step. 

73 """ 

74 # Calculate sparsity as fraction of pixels significantly above noise 

75 noise_level = observation.noise_rms[0] 

76 # Guard against non-finite or non-positive noise levels 

77 if noise_level <= 0 or not np.isfinite(noise_level): 77 ↛ 78line 77 didn't jump to line 78 because the condition on line 77 was never true

78 return default_scale 

79 signal_mask = observation.images.data > 3*noise_level 

80 signal_pixels = np.sum(signal_mask) 

81 sparsity = signal_pixels / observation.images.data.size 

82 

83 if np.any(signal_mask): 83 ↛ 87line 83 didn't jump to line 87 because the condition on line 83 was always true

84 median_signal = np.median(observation.images.data[signal_mask]) 

85 snr = median_signal / noise_level 

86 else: 

87 snr = 1.0 

88 

89 # Scale factor that decreases with sparsity and increases with SNR 

90 scale = min(1.0, (sparsity * np.sqrt(snr)) / 0.1) 

91 

92 return max(min_scale, scale) 

93 

94 

95class DeconvolveExposureConnections( 

96 pipeBase.PipelineTaskConnections, 

97 dimensions=("tract", "patch", "skymap", "band"), 

98 defaultTemplates={"inputCoaddName": "deep"}, 

99): 

100 """Connections for DeconvolveExposureTask""" 

101 

102 coadd = cT.Input( 

103 doc="Exposure to deconvolve", 

104 name="{inputCoaddName}Coadd_calexp", 

105 storageClass="ExposureF", 

106 dimensions=("tract", "patch", "band", "skymap"), 

107 ) 

108 

109 coadd_cell = cT.Input( 

110 doc="Exposure on which to run deblending", 

111 name="{inputCoaddName}CoaddCell", 

112 storageClass="MultipleCellCoadd", 

113 dimensions=("tract", "patch", "band", "skymap") 

114 ) 

115 

116 background = cT.Input( 

117 doc="Background model to subtract from the cell-based coadd", 

118 name="{inputCoaddName}Coadd_calexp_background", 

119 storageClass="Background", 

120 dimensions=("tract", "patch", "band", "skymap") 

121 ) 

122 

123 catalog = cT.Input( 

124 doc="Catalog of sources detected in the deconvolved image", 

125 name="{inputCoaddName}Coadd_mergeDet", 

126 storageClass="SourceCatalog", 

127 dimensions=("tract", "patch", "skymap"), 

128 ) 

129 

130 deconvolved = cT.Output( 

131 doc="Deconvolved exposure", 

132 name="deconvolved_{inputCoaddName}_coadd", 

133 storageClass="ExposureF", 

134 dimensions=("tract", "patch", "band", "skymap"), 

135 ) 

136 

137 def __init__(self, *, config=None): 

138 if not config.useFootprints: 

139 # Deconvolution will not use input catalog if 

140 # footprints are not used 

141 self.inputs.remove("catalog") 

142 if config.imageType == "future": 

143 self.coadd = dataclasses.replace(self.coadd, storageClass="CellCoadd") 

144 self.deconvolved = dataclasses.replace(self.deconvolved, storageClass="MaskedImageV2") 

145 del self.coadd_cell 

146 del self.background 

147 elif config.useCellCoadds: 

148 del self.coadd 

149 else: 

150 del self.coadd_cell 

151 del self.background 

152 

153 

154class DeconvolveExposureConfig( 

155 pipeBase.PipelineTaskConfig, 

156 pipelineConnections=DeconvolveExposureConnections, 

157): 

158 """Configuration for DeconvolveExposureTask""" 

159 

160 maxIter = pexConfig.Field[int]( 

161 doc="Maximum number of iterations", 

162 default=100, 

163 ) 

164 minIter = pexConfig.Field[int]( 

165 doc="Minimum number of iterations", 

166 default=10, 

167 ) 

168 eRel = pexConfig.Field[float]( 

169 doc="Relative error threshold", 

170 default=1e-3, 

171 ) 

172 backgroundThreshold = pexConfig.Field[float]( 

173 default=0, 

174 doc="Threshold for background subtraction. " 

175 "Pixels in the fit below this threshold will be set to zero", 

176 ) 

177 useFootprints = pexConfig.Field[bool]( 

178 default=True, 

179 doc="Use footprints to constrain the deconvolved model", 

180 ) 

181 useCellCoadds = pexConfig.Field[bool]( 

182 doc="Use cell-based coadd instead of regular coadd?", 

183 default=False, 

184 ) 

185 imageType = pexConfig.ChoiceField[str]( 

186 "Which image type to read and write. " 

187 "This option only directly affects connection storage classes and hence 'runQuantum'; the 'run' " 

188 "method behavior is determined by which type is actually passed in.", 

189 allowed={ 

190 "legacy": ( 

191 "Read a lsst.cell_coadds.MultipleCellCoadd and restore 'background' (if useCellCoadd) or " 

192 "lsst.afw.image.Exposure (if not useCellCoadd), and write an lsst.afw.image.Exposure." 

193 ), 

194 "future": ( 

195 "Read a lsst.images.cells.CellCoadd via 'connections.coadd' and write an " 

196 "lsst.images.MaskedImage. The useCellCoadd option will be ignored." 

197 ), 

198 }, 

199 optional=False, 

200 default="legacy", 

201 ) 

202 

203 

204class DeconvolveExposureTask(pipeBase.PipelineTask): 

205 """Deconvolve an Exposure using scarlet lite.""" 

206 

207 ConfigClass = DeconvolveExposureConfig 

208 _DefaultName = "deconvolveExposure" 

209 

210 def __init__(self, initInputs=None, **kwargs): 

211 if initInputs is None: 211 ↛ 213line 211 didn't jump to line 213 because the condition on line 211 was always true

212 initInputs = {} 

213 super().__init__(initInputs=initInputs, **kwargs) 

214 

215 def runQuantum(self, butlerQC, inputRefs, outputRefs): 

216 inputs = butlerQC.get(inputRefs) 

217 

218 match self.config.imageType: 

219 case "legacy": 

220 if self.config.useCellCoadds: 

221 band = inputRefs.coadd_cell.dataId['band'] 

222 cellCoadd = inputs.pop('coadd_cell') 

223 background = inputs.pop('background') 

224 coadd = cellCoadd.stitch().asExposure() 

225 coadd.image -= background.getImage() 

226 else: 

227 coadd = inputs.pop("coadd") 

228 band = inputRefs.coadd.dataId['band'] 

229 case "future": 

230 coadd = inputs.pop("coadd") 

231 band = inputRefs.coadd.dataId['band'] 

232 case _: 

233 raise AssertionError(f"Invalid choice {self.config.imageType!r} for imageType.") 

234 

235 catalog = inputs.pop('catalog', None) 

236 

237 assert not inputs, "runQuantum got more inputs than expected." 

238 outputs = self.run( 

239 coadd=coadd, 

240 catalog=catalog, 

241 band=band, 

242 ) 

243 butlerQC.put(outputs, outputRefs) 

244 

245 def run( 

246 self, 

247 coadd: afwImage.Exposure | CellCoadd, 

248 catalog: afwTable.SourceCatalog | None = None, 

249 band: str = 'dummy' 

250 ) -> pipeBase.Struct: 

251 """Deconvolve an Exposure 

252 

253 Parameters 

254 ---------- 

255 coadd : 

256 Coadd image to deconvolve 

257 

258 catalog : 

259 Catalog of sources detected in the merged catalog. 

260 This is used to supress noise in regions with no 

261 significant flux about the noise in the coadds. 

262 

263 band : 

264 Band of the coadd image. 

265 Since this is a single band task the band isn't really necessary 

266 but can be useful for debugging so we keep it as a parameter. 

267 

268 Returns 

269 ------- 

270 deconvolved : `pipeBase.Struct` 

271 Deconvolved exposure (if an `lsst.afw.image.Exposure` is provided; 

272 an `lsst.images.MaskedImage` if an `lsst.images.cells.CellCoadd` 

273 is provided). 

274 """ 

275 futureInputImage = None 

276 if isinstance(coadd, CellCoadd): 276 ↛ 288line 276 didn't jump to line 288 because the condition on line 276 was never true

277 # For now we just convert the future CellCoadd into an Exposure for 

278 # the bulk of the work, and convert the result back at the end (we 

279 # just convert to MaskedImage because we don't need to duplicate 

280 # the structured metadata). Converting the internals to use 

281 # lsst.images types would be disruptive but could take advantage of 

282 # the fact that the lsst.images.CellPointSpreadFunction type knows 

283 # which cells are missing and could probably do a better job of 

284 # picking a decent representative PSF for the full image, but it 

285 # would be cleanest to do that while rewriting some of the utility 

286 # functions to work exclusively with lsst.images types, and that 

287 # looks like it might be disruptive. 

288 futureInputImage = coadd 

289 coadd = coadd.to_legacy() 

290 observation = self._buildObservation(coadd, catalog, band) 

291 self.bbox = coadd.getBBox() 

292 

293 # Deconvolve. 

294 # Store the loss history for debugging purposes. 

295 model, self.loss = self._deconvolve(observation, catalog) 

296 

297 deconvolved = self._modelToExposure(model.data[0], coadd) 

298 if futureInputImage: 298 ↛ 299line 298 didn't jump to line 299 because the condition on line 298 was never true

299 deconvolved = imgs.MaskedImage.from_legacy( 

300 deconvolved.maskedImage, 

301 unit=futureInputImage.unit, 

302 plane_map=imgs.get_legacy_deep_coadd_mask_planes(), 

303 sky_projection=futureInputImage.sky_projection, 

304 ) 

305 return pipeBase.Struct(deconvolved=deconvolved) 

306 

307 def _buildObservation( 

308 self, 

309 coadd: afwImage.Exposure, 

310 catalog: afwTable.SourceCatalog | None = None, 

311 band: str = 'dummy' 

312 ) -> scl.Observation: 

313 """Build a scarlet lite Observation from an Exposure. 

314 

315 We don't actually use scarlet, but the optimized convolutions 

316 using scarlet data products are still useful. 

317 

318 Parameters 

319 ---------- 

320 coadd : 

321 Coadd image to deconvolve. 

322 catalog : 

323 Catalog of sources. 

324 This is used to find a location for the PSF if it cannot be 

325 generated at the center of the coadd. 

326 

327 band : 

328 Band of the coadd image. 

329 

330 """ 

331 bands = (band,) 

332 model_psf = scl.utils.integrated_circular_gaussian(sigma=0.8) 

333 

334 # Give zero weight to non-finite pixels 

335 weights = np.ones_like(coadd.image.array) 

336 weights[~np.isfinite(coadd.image.array)] = 0 

337 

338 image = coadd.image.array.copy() 

339 # Set non-finite pixels to zero 

340 image[~np.isfinite(image)] = 0.0 

341 psfCenter = coadd.getBBox().getCenter() 

342 if catalog is not None: 

343 psf, _, _ = utils.computeNearestPsf(coadd, catalog, band, psfCenter) 

344 if psf is None: 344 ↛ 347line 344 didn't jump to line 347 because the condition on line 344 was never true

345 # There were no valid locations from 

346 # which a PSF could be obtained 

347 raise pipeBase.NoWorkFound("No valid PSF could be obtained for deconvolution") 

348 psf = psf.array 

349 else: 

350 psf = coadd.getPsf().computeKernelImage(psfCenter).array 

351 

352 badPixelMasks = utils.defaultBadPixelMasks 

353 badPixels = coadd.mask.getPlaneBitMask(badPixelMasks) 

354 mask = coadd.mask.array & badPixels 

355 weights[mask > 0] = 0 

356 

357 observation = scl.Observation( 

358 images=image[None], 

359 variance=coadd.variance.array.copy()[None], 

360 weights=weights[None], 

361 psfs=psf[None], 

362 model_psf=model_psf[None], 

363 convolution_mode="fft", 

364 bands=bands, 

365 bbox=utils.bboxToScarletBox(coadd.getBBox()), 

366 ) 

367 return observation 

368 

369 def _deconvolve( 

370 self, 

371 observation: scl.Observation, 

372 catalog: afwTable.SourceCatalog | None = None, 

373 ) -> tuple[scl.Image, list[float]]: 

374 """Deconvolve the observed image. 

375 

376 Parameters 

377 ---------- 

378 observation : 

379 Scarlet lite Observation. 

380 catalog : 

381 Catalog of sources detected in the deconvolved image. 

382 This is used to mask the deconvolved image so that 

383 the deconvolved footprints detected downstream will always 

384 fit inside of the original footprints. 

385 """ 

386 model = observation.images.copy() 

387 loss = [] 

388 step = calculate_update_step(observation) 

389 if catalog is not None: 

390 width, height = self.bbox.getDimensions() 

391 x0, y0 = self.bbox.getMin() 

392 footprintImage = afwDet.footprintsToNumpy(catalog, shape=(height, width), xy0=(x0, y0)) 

393 for n in range(self.config.maxIter): 393 ↛ 415line 393 didn't jump to line 415 because the loop on line 393 didn't complete

394 residual = observation.images - observation.convolve(model) 

395 loss.append(-0.5 * np.sum(residual.data**2)) 

396 update = observation.convolve(residual, grad=True) 

397 update.data[:] *= step 

398 model += update 

399 model.data[(model.data < 0) | ~np.isfinite(model.data)] = 0 

400 if catalog is not None: 

401 # Ensure that the deconvolved model footprints fit 

402 # inside of the original footprints by setting regions 

403 # outside of the original footprints to zero. 

404 model.data[:] *= footprintImage 

405 

406 # Check for a diverging model 

407 if len(loss) > 1 and loss[-1] < loss[-2]: 

408 step = step / 2 

409 self.log.warning(f"Loss increased at iteration {n}, decreasing scale to {step}") 

410 

411 # Check for convergence 

412 if n > self.config.minIter and np.abs(loss[-1] - loss[-2]) < self.config.eRel * np.abs(loss[-1]): 

413 break 

414 

415 return model, loss 

416 

417 def _modelToExposure(self, model: np.ndarray, coadd: afwImage.Exposure) -> afwImage.Exposure: 

418 """Convert a scarlet lite Image to an Exposure. 

419 

420 Parameters 

421 ---------- 

422 image : 

423 Scarlet lite Image. 

424 """ 

425 image = afwImage.Image( 

426 array=model, 

427 xy0=coadd.getBBox().getMin(), 

428 deep=False, 

429 dtype=coadd.image.array.dtype, 

430 ) 

431 maskedImage = afwImage.MaskedImage( 

432 image=image, 

433 mask=coadd.mask, 

434 variance=coadd.variance, 

435 dtype=coadd.image.array.dtype, 

436 ) 

437 exposure = afwImage.Exposure( 

438 maskedImage=maskedImage, 

439 exposureInfo=coadd.getInfo(), 

440 dtype=coadd.image.array.dtype, 

441 ) 

442 return exposure