Coverage for python/lsst/meas/extensions/scarlet/deconvolveExposureTask.py: 72%
159 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-25 22:52 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-25 22:52 +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/>.
22import dataclasses
23import logging
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
37from . import utils
39log = logging.getLogger(__name__)
41__all__ = [
42 "DeconvolveExposureTask",
43 "DeconvolveExposureConfig",
44 "DeconvolveExposureConnections",
45]
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.
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.
59 Parameters
60 ----------
61 observation :
62 Scarlet lite Observation.
64 minScale :
65 Minimum allowed scale factor.
67 defaultScale :
68 Default scale factor to return if noise level is non-finite.
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
92 if np.any(signalMask):
93 medianSignal = np.median(image[signalMask])
94 snr = medianSignal / noiseLevel
95 else:
96 snr = 1.0
98 # Scale factor that decreases with sparsity and increases with SNR
99 scale = min(1.0, (sparsity * np.sqrt(snr)) / 0.1)
101 return max(minScale, scale)
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 )
123class DeconvolveExposureConnections(
124 pipeBase.PipelineTaskConnections,
125 dimensions=("tract", "patch", "skymap", "band"),
126 defaultTemplates={"inputCoaddName": "deep"},
127):
128 """Connections for DeconvolveExposureTask"""
130 coadd = cT.Input(
131 doc="Exposure to deconvolve",
132 name="{inputCoaddName}Coadd_calexp",
133 storageClass="ExposureF",
134 dimensions=("tract", "patch", "band", "skymap"),
135 )
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 )
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 )
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 )
158 deconvolved = cT.Output(
159 doc="Deconvolved exposure",
160 name="deconvolved_{inputCoaddName}_coadd",
161 storageClass="ExposureF",
162 dimensions=("tract", "patch", "band", "skymap"),
163 )
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
182class DeconvolveExposureConfig(
183 pipeBase.PipelineTaskConfig,
184 pipelineConnections=DeconvolveExposureConnections,
185):
186 """Configuration for DeconvolveExposureTask"""
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 )
232class DeconvolveExposureTask(pipeBase.PipelineTask):
233 """Deconvolve an Exposure using scarlet lite."""
235 ConfigClass = DeconvolveExposureConfig
236 _DefaultName = "deconvolveExposure"
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)
243 def runQuantum(self, butlerQC, inputRefs, outputRefs):
244 inputs = butlerQC.get(inputRefs)
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.")
263 catalog = inputs.pop('catalog', None)
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)
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
281 Parameters
282 ----------
283 coadd :
284 Coadd image to deconvolve
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.
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.
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)
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
333 model, loss = self._deconvolve(observation, footprintImage=footprintImage)
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)
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.
353 We don't actually use scarlet, but the optimized convolutions
354 using scarlet data products are still useful.
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.
365 band :
366 Band of the coadd image.
368 """
369 bands = (band,)
370 model_psf = scl.utils.integrated_circular_gaussian(sigma=0.8)
372 # Give zero weight to non-finite pixels
373 weights = np.ones_like(coadd.image.array)
374 weights[~np.isfinite(coadd.image.array)] = 0
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
390 badPixelMasks = utils.defaultBadPixelMasks
391 badPixels = coadd.mask.getPlaneBitMask(badPixelMasks)
392 mask = coadd.mask.array & badPixels
393 weights[mask > 0] = 0
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
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.
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
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}")
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
453 return model, loss
455 def _modelToExposure(self, model: np.ndarray, coadd: afwImage.Exposure) -> afwImage.Exposure:
456 """Convert a deconvolved image array to an Exposure.
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.
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