Coverage for python/lsst/pipe/tasks/fit_coadd_multiband.py: 52%
166 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-20 09:23 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-20 09:23 +0000
1# This file is part of pipe_tasks.
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/>.
22__all__ = [
23 "CoaddMultibandFitConfig", "CoaddMultibandFitConnections", "CoaddMultibandFitSubConfig",
24 "CoaddMultibandFitSubTask", "CoaddMultibandFitTask",
25]
27from .fit_multiband import CatalogExposure, CatalogExposureConfig
29import lsst.afw.table as afwTable
30from lsst.meas.base import SkyMapIdGeneratorConfig
31from lsst.meas.extensions.scarlet.io import updateCatalogFootprints
32import lsst.pex.config as pexConfig
33import lsst.pipe.base as pipeBase
34import lsst.pipe.base.connectionTypes as cT
36import astropy.table
37import dataclasses
38from abc import ABC, abstractmethod
39from pydantic import Field
40from pydantic.dataclasses import dataclass
41from typing import Iterable
43CoaddMultibandFitBaseTemplates = {
44 "name_coadd": "deep",
45 "name_method": "multiprofit",
46 "name_table": "objects",
47}
50@dataclass(frozen=True, kw_only=True, config=CatalogExposureConfig)
51class CatalogExposureInputs(CatalogExposure):
52 table_psf_fits: astropy.table.Table = Field(title="A table of PSF fit parameters for each source")
54 def get_catalog(self):
55 return self.catalog
58class CoaddMultibandFitInputConnections(
59 pipeBase.PipelineTaskConnections,
60 dimensions=("tract", "patch", "skymap"),
61 defaultTemplates=CoaddMultibandFitBaseTemplates,
62):
63 cat_ref = cT.Input(
64 doc="Reference multiband source catalog",
65 name="{name_coadd}Coadd_ref",
66 storageClass="SourceCatalog",
67 dimensions=("tract", "patch", "skymap"),
68 )
69 cats_meas = cT.Input(
70 doc="Deblended single-band source catalogs",
71 name="{name_coadd}Coadd_meas",
72 storageClass="SourceCatalog",
73 dimensions=("tract", "patch", "band", "skymap"),
74 multiple=True,
75 )
76 coadds = cT.Input(
77 doc="Exposures on which to run fits",
78 name="{name_coadd}Coadd_calexp",
79 storageClass="ExposureF",
80 dimensions=("tract", "patch", "band", "skymap"),
81 multiple=True,
82 )
83 coadds_cell = cT.Input(
84 doc="Cell-coadd exposures on which to run fits",
85 name="{name_coadd}CoaddCell",
86 storageClass="MultipleCellCoadd",
87 dimensions=("tract", "patch", "band", "skymap"),
88 multiple=True,
89 )
90 backgrounds = cT.Input(
91 doc="Background models to subtract from the coadds_cell",
92 name="{name_coadd}Coadd_calexp_background",
93 storageClass="Background",
94 dimensions=("tract", "patch", "band", "skymap"),
95 multiple=True,
96 )
97 models_psf = cT.Input(
98 doc="Input PSF model parameter catalog",
99 # Consider allowing independent psf fit method
100 name="{name_coadd}Coadd_psfs_{name_method}",
101 storageClass="ArrowAstropy",
102 dimensions=("tract", "patch", "band", "skymap"),
103 multiple=True,
104 deferLoad=True,
105 )
106 models_scarlet = pipeBase.connectionTypes.Input(
107 doc="Multiband scarlet models produced by the deblender",
108 name="{name_coadd}Coadd_scarletModelData",
109 storageClass="LsstScarletModelData",
110 dimensions=("tract", "patch", "skymap"),
111 )
113 def adjustQuantum(self, inputs, outputs, label, data_id):
114 """Validates the `lsst.daf.butler.DatasetRef` bands against the
115 subtask's list of bands to fit and drops unnecessary bands.
117 Parameters
118 ----------
119 inputs : `dict`
120 Dictionary whose keys are an input (regular or prerequisite)
121 connection name and whose values are a tuple of the connection
122 instance and a collection of associated `DatasetRef` objects.
123 The exact type of the nested collections is unspecified; it can be
124 assumed to be multi-pass iterable and support `len` and ``in``, but
125 it should not be mutated in place. In contrast, the outer
126 dictionaries are guaranteed to be temporary copies that are true
127 `dict` instances, and hence may be modified and even returned; this
128 is especially useful for delegating to `super` (see notes below).
129 outputs : `Mapping`
130 Mapping of output datasets, with the same structure as ``inputs``.
131 label : `str`
132 Label for this task in the pipeline (should be used in all
133 diagnostic messages).
134 data_id : `lsst.daf.butler.DataCoordinate`
135 Data ID for this quantum in the pipeline (should be used in all
136 diagnostic messages).
138 Returns
139 -------
140 adjusted_inputs : `Mapping`
141 Mapping of the same form as ``inputs`` with updated containers of
142 input `DatasetRef` objects. All inputs involving the 'band'
143 dimension are adjusted to put them in consistent order and remove
144 unneeded bands.
145 adjusted_outputs : `Mapping`
146 Mapping of updated output datasets; always empty for this task.
148 Raises
149 ------
150 lsst.pipe.base.NoWorkFound
151 Raised if there are not enough of the right bands to run the task
152 on this quantum.
153 """
154 # Check which bands are going to be fit
155 bands_fit, bands_read_only = self.config.get_band_sets()
156 bands_needed = bands_fit + [band for band in bands_read_only if band not in bands_fit]
157 bands_needed_set = set(bands_needed)
159 adjusted_inputs = {}
160 inputs_to_adjust = {}
161 bands_found = bands_needed_set
162 for connection_name, (connection, dataset_refs) in inputs.items():
163 # Datasets without bands in their dimensions should be fine
164 if 'band' in connection.dimensions: 164 ↛ 162line 164 didn't jump to line 162 because the condition on line 164 was always true
165 datasets_by_band = {dref.dataId['band']: dref for dref in dataset_refs}
166 bands_set = set(datasets_by_band.keys())
167 if self.config.allow_missing_bands:
168 if len(bands_found) == 0: 168 ↛ 169line 168 didn't jump to line 169 because the condition on line 168 was never true
169 raise pipeBase.NoWorkFound(
170 f'DatasetRefs={dataset_refs} for {connection_name=} is empty'
171 )
172 bands_found &= bands_set
173 # All configured bands are treated as necessary
174 elif not bands_needed_set.issubset(bands_set):
175 raise pipeBase.NoWorkFound(
176 f'DatasetRefs={dataset_refs} have data with bands in the'
177 f' set={set(datasets_by_band.keys())},'
178 f' which is not a superset of the required bands={bands_needed} defined by'
179 f' {self.config.__class__}.fit_coadd_multiband='
180 f'{self.config.fit_coadd_multiband._value.__class__}\'s attributes'
181 f' bands_fit={bands_fit} and bands_read_only()={bands_read_only}.'
182 f' Add the required bands={set(bands_needed).difference(datasets_by_band.keys())}.'
183 )
184 # Adjust all datasets with band dimensions to include just
185 # the needed bands, in consistent order.
186 inputs_to_adjust[connection_name] = (connection, datasets_by_band)
188 if self.config.allow_missing_bands: 188 ↛ 196line 188 didn't jump to line 196 because the condition on line 188 was always true
189 bands_needed = [band for band in bands_fit if band in bands_found] + [
190 band for band in bands_read_only if band not in bands_found
191 ]
192 if len(bands_needed) == 0:
193 raise pipeBase.NoWorkFound(
194 f'No common bands remaining for inputs {",".join(inputs_to_adjust.keys())}'
195 )
196 for connection_name, (connection, datasets_by_band) in inputs_to_adjust.items():
197 adjusted_inputs[connection_name] = (
198 connection,
199 [datasets_by_band[band] for band in bands_needed]
200 )
202 # Delegate to super for more checks.
203 inputs.update(adjusted_inputs)
204 super().adjustQuantum(inputs, outputs, label, data_id)
205 return adjusted_inputs, {}
207 def __init__(self, *, config=None):
208 super().__init__(config=config)
209 assert isinstance(config, CoaddMultibandFitBaseConfig)
211 if config.drop_psf_connection: 211 ↛ 212line 211 didn't jump to line 212 because the condition on line 211 was never true
212 del self.models_psf
214 if config.image_type == "future": 214 ↛ 215line 214 didn't jump to line 215 because the condition on line 214 was never true
215 self.coadds = dataclasses.replace(self.coadds, storageClass="CellCoadd")
216 del self.coadds_cell
217 del self.backgrounds
218 elif config.use_cell_coadds: 218 ↛ 219line 218 didn't jump to line 219 because the condition on line 218 was never true
219 del self.coadds
220 else:
221 del self.coadds_cell
222 del self.backgrounds
225class CoaddMultibandFitConnections(CoaddMultibandFitInputConnections):
226 cat_output = cT.Output(
227 doc="Output source model fit parameter catalog",
228 name="{name_coadd}Coadd_{name_table}_{name_method}",
229 storageClass="ArrowTable",
230 dimensions=("tract", "patch", "skymap"),
231 )
234class CoaddMultibandFitSubConfig(pexConfig.Config):
235 """Configuration for implementing fitter subtasks.
236 """
238 bands_fit = pexConfig.ListField[str](
239 default=[],
240 doc="list of bandpass filters to fit",
241 listCheck=lambda x: (len(x) > 0) and (len(set(x)) == len(x)),
242 )
244 @abstractmethod
245 def bands_read_only(self) -> set:
246 """Return the set of bands that the Task needs to read (e.g. for
247 defining priors) but not necessarily fit.
249 Returns
250 -------
251 The set of such bands.
252 """
253 return set()
256class CoaddMultibandFitSubTask(pipeBase.Task, ABC):
257 """Subtask interface for multiband fitting of deblended sources.
259 Parameters
260 ----------
261 **kwargs
262 Additional arguments to be passed to the `lsst.pipe.base.Task`
263 constructor.
264 """
265 ConfigClass = CoaddMultibandFitSubConfig
267 def __init__(self, **kwargs):
268 super().__init__(**kwargs)
270 @abstractmethod
271 def run(
272 self, catexps: Iterable[CatalogExposureInputs], cat_ref: afwTable.SourceCatalog
273 ) -> pipeBase.Struct:
274 """Fit models to deblended sources from multi-band inputs.
276 Parameters
277 ----------
278 catexps : `typing.List [CatalogExposureInputs]`
279 A list of catalog-exposure pairs with metadata in a given band.
280 cat_ref : `lsst.afw.table.SourceCatalog`
281 A reference source catalog to fit.
283 Returns
284 -------
285 retStruct : `lsst.pipe.base.Struct`
286 A struct with a cat_output attribute containing the output
287 measurement catalog.
289 Notes
290 -----
291 Subclasses may have further requirements on the input parameters,
292 including:
293 - Passing only one catexp per band;
294 - Catalogs containing HeavyFootprints with deblended images;
295 - Fitting only a subset of the sources.
296 If any requirements are not met, the subtask should fail as soon as
297 possible.
298 """
301class CoaddMultibandFitBaseConfig(
302 pipeBase.PipelineTaskConfig,
303 pipelineConnections=CoaddMultibandFitInputConnections,
304):
305 """Base class for multiband fitting."""
307 allow_missing_bands = pexConfig.Field[bool](
308 doc="Whether to still fit even if some bands are missing",
309 default=True,
310 )
311 drop_psf_connection = pexConfig.Field[bool](
312 doc="Whether to drop the PSF model connection, e.g. because PSF parameters are in the input catalog",
313 default=False,
314 )
315 fit_coadd_multiband = pexConfig.ConfigurableField(
316 target=CoaddMultibandFitSubTask,
317 doc="Task to fit sources using multiple bands",
318 )
319 use_cell_coadds = pexConfig.Field[bool](
320 doc="Use cell coadd images for object fitting?",
321 default=False,
322 )
323 idGenerator = SkyMapIdGeneratorConfig.make_field()
324 image_type = pexConfig.ChoiceField(
325 "Which image type to expect for the input coadd. "
326 "This option only directly affects connection storage classes and hence 'runQuantum'; the 'run' "
327 "method behavior is determined by which type is actually passed in.",
328 allowed={
329 "legacy": (
330 "Read a lsst.cell_coadds.MultipleCellCoadd via 'coadds_cell` and restore 'background' "
331 "(if use_cell_coadd) or lsst.afw.image.Exposure via `coadds` (if not use_cell_coadd)."
332 ),
333 "future": (
334 "Read lsst.images.cells.CellCoadd via the 'coadds' connection. use_cell_coadd is ignored."
335 ),
336 },
337 dtype=str,
338 optional=False,
339 default="legacy",
340 )
342 def get_band_sets(self):
343 """Get the set of bands required by the fit_coadd_multiband subtask.
345 Returns
346 -------
347 bands_fit : `set`
348 The set of bands that the subtask will fit.
349 bands_read_only : `set`
350 The set of bands that the subtask will only read data
351 (measurement catalog and exposure) for.
352 """
353 try:
354 bands_fit = self.fit_coadd_multiband.bands_fit
355 except AttributeError:
356 raise RuntimeError(f'{__class__}.fit_coadd_multiband must have bands_fit attribute') from None
357 bands_read_only = self.fit_coadd_multiband.bands_read_only()
358 return tuple(list({band: None for band in bands}.keys()) for bands in (bands_fit, bands_read_only))
361class CoaddMultibandFitConfig(
362 CoaddMultibandFitBaseConfig,
363 pipelineConnections=CoaddMultibandFitConnections,
364):
365 """Configuration for a CoaddMultibandFitTask."""
368class CoaddMultibandFitBase:
369 """Base class for tasks that fit or rebuild multiband models.
371 This class only implements data reconstruction.
372 """
374 def build_catexps(self, butlerQC, inputRefs, inputs) -> list[CatalogExposureInputs]:
375 id_tp = self.config.idGenerator.apply(butlerQC.quantum.dataId).catalog_id
376 # This is a roundabout way of ensuring all inputs get sorted and matched
377 if self.config.use_cell_coadds and self.config.image_type == "legacy":
378 keys = ["cats_meas", "coadds_cell", "backgrounds"]
379 else:
380 keys = ["cats_meas", "coadds"]
381 has_psf_models = "models_psf" in inputs
382 if has_psf_models:
383 keys.append("models_psf")
384 input_refs_objs = {key: (getattr(inputRefs, key), inputs[key]) for key in keys}
385 inputs_sorted = {
386 key: {dRef.dataId: obj for dRef, obj in zip(refs, objs, strict=True)}
387 for key, (refs, objs) in input_refs_objs.items()
388 }
389 cats = inputs_sorted["cats_meas"]
390 if self.config.image_type == "future":
391 exps = {data_id: coadd.to_legacy() for data_id, coadd in inputs_sorted["coadds"].items()}
392 elif self.config.use_cell_coadds:
393 exps = {}
394 for data_id, background in inputs_sorted["backgrounds"].items():
395 mcc = inputs_sorted["coadds_cell"][data_id]
396 stitched_coadd = mcc.stitch()
397 exposure = stitched_coadd.asExposure()
398 exposure.image -= background.getImage()
399 exps[data_id] = exposure
400 else:
401 exps = inputs_sorted["coadds"]
403 # Ensure that psf models are loaded with full metadata.
404 if has_psf_models:
405 ref0 = list(inputs_sorted["models_psf"].values())[0]
406 parameters = None
407 if ref0.ref.datasetType.storageClass_name == "ArrowAstropy":
408 parameters = {"strip_astropy_meta_yaml": False}
409 models_psf = {
410 key: ref.get(parameters=parameters)
411 for key, ref in inputs_sorted["models_psf"].items()
412 }
413 else:
414 models_psf = None
416 dataIds = set(cats).union(set(exps))
417 models_scarlet = inputs["models_scarlet"]
418 catexp_dict = {}
419 dataId = None
420 for dataId in dataIds:
421 catalog = cats[dataId]
422 exposure = exps[dataId]
423 updateCatalogFootprints(
424 modelData=models_scarlet,
425 catalog=catalog,
426 band=dataId['band'],
427 imageForRedistribution=exposure,
428 removeScarletData=False,
429 updateFluxColumns=False,
430 )
431 catexp_dict[dataId['band']] = CatalogExposureInputs(
432 catalog=catalog,
433 exposure=exposure,
434 table_psf_fits=models_psf[dataId] if has_psf_models else astropy.table.Table(),
435 dataId=dataId,
436 id_tract_patch=id_tp,
437 )
438 # This shouldn't happen unless this is called with no inputs, but check anyway
439 if dataId is None:
440 raise RuntimeError(f"Did not build any catexps for {inputRefs=}")
441 catexps = []
442 for band in self.config.get_band_sets()[0]:
443 if band in catexp_dict:
444 catexp = catexp_dict[band]
445 else:
446 # Make a dummy catexp with a dataId if there's no data
447 # This should be handled by any subtasks
448 dataId_band = dataId.to_simple(minimal=True)
449 dataId_band.dataId["band"] = band
450 catexp = CatalogExposureInputs(
451 catalog=afwTable.SourceCatalog(),
452 exposure=None,
453 table_psf_fits=astropy.table.Table(),
454 dataId=dataId.from_simple(dataId_band, universe=dataId.universe),
455 id_tract_patch=id_tp,
456 )
457 catexps.append(catexp)
458 return catexps
461class CoaddMultibandFitTask(CoaddMultibandFitBase, pipeBase.PipelineTask):
462 """Fit deblended exposures in multiple bands simultaneously.
464 It is generally assumed but not enforced (except optionally by the
465 configurable `fit_coadd_multiband` subtask) that there is only one exposure
466 per band, presumably a coadd.
467 """
469 ConfigClass = CoaddMultibandFitConfig
470 _DefaultName = "coaddMultibandFit"
472 def __init__(self, initInputs, **kwargs):
473 super().__init__(initInputs=initInputs, **kwargs)
474 self.makeSubtask("fit_coadd_multiband")
476 def make_kwargs(self, butlerQC, inputRefs, inputs):
477 """Make any kwargs needed to be passed to run.
479 This method should be overloaded by subclasses that are configured to
480 use a specific subtask that needs additional arguments derived from
481 the inputs but do not otherwise need to overload runQuantum."""
482 return {}
484 def runQuantum(self, butlerQC, inputRefs, outputRefs):
485 inputs = butlerQC.get(inputRefs)
486 catexps = self.build_catexps(butlerQC, inputRefs, inputs)
487 if not self.config.allow_missing_bands and any([catexp is None for catexp in catexps]):
488 raise RuntimeError(
489 f"Got a None catexp with {self.config.allow_missing_band=}; NoWorkFound should have been"
490 f" raised earlier"
491 )
492 kwargs = self.make_kwargs(butlerQC, inputRefs, inputs)
493 outputs = self.run(catexps=catexps, cat_ref=inputs['cat_ref'], **kwargs)
494 butlerQC.put(outputs, outputRefs)
496 def run(
497 self,
498 catexps: list[CatalogExposure],
499 cat_ref: afwTable.SourceCatalog,
500 **kwargs
501 ) -> pipeBase.Struct:
502 """Fit sources from a reference catalog using data from multiple
503 exposures in the same region (patch).
505 Parameters
506 ----------
507 catexps : `typing.List [CatalogExposure]`
508 A list of catalog-exposure pairs in a given band.
509 cat_ref : `lsst.afw.table.SourceCatalog`
510 A reference source catalog to fit.
512 Returns
513 -------
514 retStruct : `lsst.pipe.base.Struct`
515 A struct with a cat_output attribute containing the output
516 measurement catalog.
518 Notes
519 -----
520 Subtasks may have further requirements; see `CoaddMultibandFitSubTask.run`.
521 """
522 cat_output = self.fit_coadd_multiband.run(catalog_multi=cat_ref, catexps=catexps, **kwargs).output
523 retStruct = pipeBase.Struct(cat_output=cat_output)
524 return retStruct