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 07:53 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-14 07:53 +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
36from . import utils
38log = logging.getLogger(__name__)
40__all__ = [
41 "DeconvolveExposureTask",
42 "DeconvolveExposureConfig",
43 "DeconvolveExposureConnections",
44]
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.
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.
58 Parameters
59 ----------
60 observation :
61 Scarlet lite Observation.
63 min_scale :
64 Minimum allowed scale factor.
66 default_scale :
67 Default scale factor to return if noise level is non-finite.
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
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
89 # Scale factor that decreases with sparsity and increases with SNR
90 scale = min(1.0, (sparsity * np.sqrt(snr)) / 0.1)
92 return max(min_scale, scale)
95class DeconvolveExposureConnections(
96 pipeBase.PipelineTaskConnections,
97 dimensions=("tract", "patch", "skymap", "band"),
98 defaultTemplates={"inputCoaddName": "deep"},
99):
100 """Connections for DeconvolveExposureTask"""
102 coadd = cT.Input(
103 doc="Exposure to deconvolve",
104 name="{inputCoaddName}Coadd_calexp",
105 storageClass="ExposureF",
106 dimensions=("tract", "patch", "band", "skymap"),
107 )
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 )
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 )
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 )
130 deconvolved = cT.Output(
131 doc="Deconvolved exposure",
132 name="deconvolved_{inputCoaddName}_coadd",
133 storageClass="ExposureF",
134 dimensions=("tract", "patch", "band", "skymap"),
135 )
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
154class DeconvolveExposureConfig(
155 pipeBase.PipelineTaskConfig,
156 pipelineConnections=DeconvolveExposureConnections,
157):
158 """Configuration for DeconvolveExposureTask"""
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 )
204class DeconvolveExposureTask(pipeBase.PipelineTask):
205 """Deconvolve an Exposure using scarlet lite."""
207 ConfigClass = DeconvolveExposureConfig
208 _DefaultName = "deconvolveExposure"
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)
215 def runQuantum(self, butlerQC, inputRefs, outputRefs):
216 inputs = butlerQC.get(inputRefs)
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.")
235 catalog = inputs.pop('catalog', None)
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)
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
253 Parameters
254 ----------
255 coadd :
256 Coadd image to deconvolve
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.
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.
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()
293 # Deconvolve.
294 # Store the loss history for debugging purposes.
295 model, self.loss = self._deconvolve(observation, catalog)
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)
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.
315 We don't actually use scarlet, but the optimized convolutions
316 using scarlet data products are still useful.
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.
327 band :
328 Band of the coadd image.
330 """
331 bands = (band,)
332 model_psf = scl.utils.integrated_circular_gaussian(sigma=0.8)
334 # Give zero weight to non-finite pixels
335 weights = np.ones_like(coadd.image.array)
336 weights[~np.isfinite(coadd.image.array)] = 0
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
352 badPixelMasks = utils.defaultBadPixelMasks
353 badPixels = coadd.mask.getPlaneBitMask(badPixelMasks)
354 mask = coadd.mask.array & badPixels
355 weights[mask > 0] = 0
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
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.
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
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}")
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
415 return model, loss
417 def _modelToExposure(self, model: np.ndarray, coadd: afwImage.Exposure) -> afwImage.Exposure:
418 """Convert a scarlet lite Image to an Exposure.
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