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

159 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-20 09:37 +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 

35from deprecated.sphinx import deprecated 

36 

37from . import utils 

38 

39log = logging.getLogger(__name__) 

40 

41__all__ = [ 

42 "DeconvolveExposureTask", 

43 "DeconvolveExposureConfig", 

44 "DeconvolveExposureConnections", 

45] 

46 

47 

48def calculateUpdateStep( 

49 observation: scl.Observation, 

50 minScale: float = 0.01, 

51 defaultScale: float = 0.1, 

52) -> float: 

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

54 

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

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

57 factor will be less than 1.0. 

58 

59 Parameters 

60 ---------- 

61 observation : 

62 Scarlet lite Observation. 

63 

64 minScale : 

65 Minimum allowed scale factor. 

66 

67 defaultScale : 

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

69 

70 Returns 

71 ------- 

72 scale : float 

73 Scale factor for the update step. 

74 """ 

75 # Calculate sparsity as fraction of unmasked pixels significantly 

76 # above noise. Pixels with zero weight (border, NO_DATA, BAD) are 

77 # excluded from both numerator and denominator so heavily masked 

78 # inputs are not biased toward a small step. 

79 noiseLevel = observation.noise_rms[0] 

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

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

82 return defaultScale 

83 image = observation.images.data[0] 

84 validMask = observation.weights.data[0] > 0 

85 validPixels = np.sum(validMask) 

86 if validPixels == 0: 86 ↛ 87line 86 didn't jump to line 87 because the condition on line 86 was never true

87 return defaultScale 

88 signalMask = (image > 3*noiseLevel) & validMask 

89 signalPixels = np.sum(signalMask) 

90 sparsity = signalPixels / validPixels 

91 

92 if np.any(signalMask): 

93 medianSignal = np.median(image[signalMask]) 

94 snr = medianSignal / noiseLevel 

95 else: 

96 snr = 1.0 

97 

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

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

100 

101 return max(minScale, scale) 

102 

103 

104@deprecated( 

105 reason=( 

106 "Use `calculateUpdateStep` instead; the snake_case name is kept " 

107 "as a shim. Will be removed after v31." 

108 ), 

109 version="v30.0", 

110 category=FutureWarning, 

111) 

112def calculate_update_step( 

113 observation: scl.Observation, 

114 min_scale: float = 0.01, 

115 default_scale: float = 0.1, 

116) -> float: 

117 """Deprecated snake_case alias for `calculateUpdateStep`.""" 

118 return calculateUpdateStep( 

119 observation, minScale=min_scale, defaultScale=default_scale, 

120 ) 

121 

122 

123class DeconvolveExposureConnections( 

124 pipeBase.PipelineTaskConnections, 

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

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

127): 

128 """Connections for DeconvolveExposureTask""" 

129 

130 coadd = cT.Input( 

131 doc="Exposure to deconvolve", 

132 name="{inputCoaddName}Coadd_calexp", 

133 storageClass="ExposureF", 

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

135 ) 

136 

137 coadd_cell = cT.Input( 

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

139 name="{inputCoaddName}CoaddCell", 

140 storageClass="MultipleCellCoadd", 

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

142 ) 

143 

144 background = cT.Input( 

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

146 name="{inputCoaddName}Coadd_calexp_background", 

147 storageClass="Background", 

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

149 ) 

150 

151 catalog = cT.Input( 

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

153 name="{inputCoaddName}Coadd_mergeDet", 

154 storageClass="SourceCatalog", 

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

156 ) 

157 

158 deconvolved = cT.Output( 

159 doc="Deconvolved exposure", 

160 name="deconvolved_{inputCoaddName}_coadd", 

161 storageClass="ExposureF", 

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

163 ) 

164 

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

166 if not config.useFootprints: 

167 # Deconvolution will not use input catalog if 

168 # footprints are not used 

169 self.inputs.remove("catalog") 

170 if config.imageType == "future": 

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

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

173 del self.coadd_cell 

174 del self.background 

175 elif config.useCellCoadds: 

176 del self.coadd 

177 else: 

178 del self.coadd_cell 

179 del self.background 

180 

181 

182class DeconvolveExposureConfig( 

183 pipeBase.PipelineTaskConfig, 

184 pipelineConnections=DeconvolveExposureConnections, 

185): 

186 """Configuration for DeconvolveExposureTask""" 

187 

188 maxIter = pexConfig.Field[int]( 

189 doc="Maximum number of iterations", 

190 default=100, 

191 ) 

192 minIter = pexConfig.Field[int]( 

193 doc="Minimum number of iterations", 

194 default=10, 

195 ) 

196 eRel = pexConfig.Field[float]( 

197 doc="Relative error threshold", 

198 default=1e-3, 

199 ) 

200 backgroundThreshold = pexConfig.Field[float]( 

201 default=0, 

202 doc="Threshold for background subtraction. " 

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

204 ) 

205 useFootprints = pexConfig.Field[bool]( 

206 default=True, 

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

208 ) 

209 useCellCoadds = pexConfig.Field[bool]( 

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

211 default=False, 

212 ) 

213 imageType = pexConfig.ChoiceField[str]( 

214 "Which image type to read and write. " 

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

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

217 allowed={ 

218 "legacy": ( 

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

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

221 ), 

222 "future": ( 

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

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

225 ), 

226 }, 

227 optional=False, 

228 default="legacy", 

229 ) 

230 

231 

232class DeconvolveExposureTask(pipeBase.PipelineTask): 

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

234 

235 ConfigClass = DeconvolveExposureConfig 

236 _DefaultName = "deconvolveExposure" 

237 

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

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

240 initInputs = {} 

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

242 

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

244 inputs = butlerQC.get(inputRefs) 

245 

246 match self.config.imageType: 

247 case "legacy": 

248 if self.config.useCellCoadds: 

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

250 cellCoadd = inputs.pop('coadd_cell') 

251 background = inputs.pop('background') 

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

253 coadd.image -= background.getImage() 

254 else: 

255 coadd = inputs.pop("coadd") 

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

257 case "future": 

258 coadd = inputs.pop("coadd") 

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

260 case _: 

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

262 

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

264 

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

266 outputs = self.run( 

267 coadd=coadd, 

268 catalog=catalog, 

269 band=band, 

270 ) 

271 butlerQC.put(outputs, outputRefs) 

272 

273 def run( 

274 self, 

275 coadd: afwImage.Exposure | CellCoadd, 

276 catalog: afwTable.SourceCatalog | None = None, 

277 band: str = 'dummy' 

278 ) -> pipeBase.Struct: 

279 """Deconvolve an Exposure 

280 

281 Parameters 

282 ---------- 

283 coadd : 

284 Coadd image to deconvolve 

285 

286 catalog : 

287 Catalog of sources detected in the merged catalog. 

288 This is used to supress noise in regions with no 

289 significant flux about the noise in the coadds. 

290 

291 band : 

292 Band of the coadd image. 

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

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

295 

296 Returns 

297 ------- 

298 deconvolved : `pipeBase.Struct` 

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

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

301 is provided). 

302 """ 

303 futureInputImage = None 

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

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

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

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

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

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

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

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

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

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

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

315 # looks like it might be disruptive. 

316 futureInputImage = coadd 

317 coadd = coadd.to_legacy() 

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

319 

320 # Build the per-pixel footprint mask from the catalog, if one 

321 # was supplied, so the deconvolution loop only needs to know 

322 # about the mask itself rather than how it was derived. 

323 if catalog is not None: 

324 bbox = coadd.getBBox() 

325 width, height = bbox.getDimensions() 

326 x0, y0 = bbox.getMin() 

327 footprintImage = afwDet.footprintsToNumpy( 

328 catalog, shape=(height, width), xy0=(x0, y0) 

329 ) 

330 else: 

331 footprintImage = None 

332 

333 model, loss = self._deconvolve(observation, footprintImage=footprintImage) 

334 

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

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

337 deconvolved = imgs.MaskedImage.from_legacy( 

338 deconvolved.maskedImage, 

339 unit=futureInputImage.unit, 

340 plane_map=imgs.get_legacy_deep_coadd_mask_planes(), 

341 sky_projection=futureInputImage.sky_projection, 

342 ) 

343 return pipeBase.Struct(deconvolved=deconvolved, loss=loss) 

344 

345 def _buildObservation( 

346 self, 

347 coadd: afwImage.Exposure, 

348 catalog: afwTable.SourceCatalog | None = None, 

349 band: str = 'dummy' 

350 ) -> scl.Observation: 

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

352 

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

354 using scarlet data products are still useful. 

355 

356 Parameters 

357 ---------- 

358 coadd : 

359 Coadd image to deconvolve. 

360 catalog : 

361 Catalog of sources. 

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

363 generated at the center of the coadd. 

364 

365 band : 

366 Band of the coadd image. 

367 

368 """ 

369 bands = (band,) 

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

371 

372 # Give zero weight to non-finite pixels 

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

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

375 

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

377 # Set non-finite pixels to zero 

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

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

380 if catalog is not None: 

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

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

383 # There were no valid locations from 

384 # which a PSF could be obtained 

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

386 psf = psf.array 

387 else: 

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

389 

390 badPixelMasks = utils.defaultBadPixelMasks 

391 badPixels = coadd.mask.getPlaneBitMask(badPixelMasks) 

392 mask = coadd.mask.array & badPixels 

393 weights[mask > 0] = 0 

394 

395 observation = scl.Observation( 

396 images=image[None], 

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

398 weights=weights[None], 

399 psfs=psf[None], 

400 model_psf=model_psf[None], 

401 convolution_mode="fft", 

402 bands=bands, 

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

404 ) 

405 return observation 

406 

407 def _deconvolve( 

408 self, 

409 observation: scl.Observation, 

410 footprintImage: np.ndarray | None = None, 

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

412 """Deconvolve the observed image. 

413 

414 Parameters 

415 ---------- 

416 observation : 

417 Scarlet lite Observation. 

418 footprintImage : 

419 Per-pixel mask matching ``observation.images.shape[1:]``. 

420 When supplied, the deconvolved model is multiplied by this 

421 mask after each iteration so the recovered footprints stay 

422 inside the input footprints. 

423 """ 

424 model = observation.images.copy() 

425 loss = [] 

426 step = calculateUpdateStep(observation) 

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

428 # cache=True reuses the FFT plan across iterations; the 

429 # image shape is stable inside the loop so this is a free 

430 # speedup at zero correctness cost. 

431 residual = observation.images - observation.convolve(model, cache=True) 

432 if np.all(~np.isfinite(residual.data)): 

433 self.log.warning(f"Residual is non-finite at iteration {n}, stopping deconvolution") 

434 loss.append(-np.inf) 

435 break 

436 loss.append(-0.5 * np.nansum(residual.data**2)) 

437 update = observation.convolve(residual, grad=True, cache=True) 

438 update.data[:] *= step 

439 model += update 

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

441 if footprintImage is not None: 

442 model.data[:] *= footprintImage 

443 

444 # Check for a diverging model 

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

446 step = step / 2 

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

448 

449 # Check for convergence 

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

451 break 

452 

453 return model, loss 

454 

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

456 """Convert a deconvolved image array to an Exposure. 

457 

458 The output exposure's mask is a deep copy of the input coadd's 

459 mask, and its variance plane is fresh and filled with ``inf``. 

460 Convolution-then-deconvolution alters the per-pixel noise 

461 covariance, so the input coadd's variance no longer describes 

462 the deconvolved pixel values; the infinite variance signals 

463 "no information about the noise here" and naturally zero-weights 

464 these pixels under any inverse-variance scheme. Downstream 

465 consumers that need a variance plane must supply their own. 

466 

467 Parameters 

468 ---------- 

469 model : 

470 Deconvolved image array. 

471 coadd : 

472 Input coadd exposure; its image dtype, bbox, ``ExposureInfo``, 

473 and mask contents are reused. 

474 """ 

475 image = afwImage.Image( 

476 array=model, 

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

478 deep=False, 

479 dtype=coadd.image.array.dtype, 

480 ) 

481 # Deep-copy the mask and build a fresh inf-filled variance so 

482 # the output exposure doesn't alias the input coadd's planes. 

483 # The variance is deliberately invalidated because the input's 

484 # variance does not describe the deconvolved pixel values. 

485 mask = coadd.mask.clone() 

486 variance = coadd.variance.Factory(coadd.variance.getBBox()) 

487 variance.array[:] = np.inf 

488 maskedImage = afwImage.MaskedImage( 

489 image=image, 

490 mask=mask, 

491 variance=variance, 

492 dtype=coadd.image.array.dtype, 

493 ) 

494 exposure = afwImage.Exposure( 

495 maskedImage=maskedImage, 

496 exposureInfo=coadd.getInfo(), 

497 dtype=coadd.image.array.dtype, 

498 ) 

499 return exposure