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

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/>. 

21 

22__all__ = [ 

23 "CoaddMultibandFitConfig", "CoaddMultibandFitConnections", "CoaddMultibandFitSubConfig", 

24 "CoaddMultibandFitSubTask", "CoaddMultibandFitTask", 

25] 

26 

27from .fit_multiband import CatalogExposure, CatalogExposureConfig 

28 

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 

35 

36import astropy.table 

37import dataclasses 

38from abc import ABC, abstractmethod 

39from pydantic import Field 

40from pydantic.dataclasses import dataclass 

41from typing import Iterable 

42 

43CoaddMultibandFitBaseTemplates = { 

44 "name_coadd": "deep", 

45 "name_method": "multiprofit", 

46 "name_table": "objects", 

47} 

48 

49 

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") 

53 

54 def get_catalog(self): 

55 return self.catalog 

56 

57 

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 ) 

112 

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. 

116 

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). 

137 

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. 

147 

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) 

158 

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) 

187 

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 ) 

201 

202 # Delegate to super for more checks. 

203 inputs.update(adjusted_inputs) 

204 super().adjustQuantum(inputs, outputs, label, data_id) 

205 return adjusted_inputs, {} 

206 

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

208 super().__init__(config=config) 

209 assert isinstance(config, CoaddMultibandFitBaseConfig) 

210 

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 

213 

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 

223 

224 

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 ) 

232 

233 

234class CoaddMultibandFitSubConfig(pexConfig.Config): 

235 """Configuration for implementing fitter subtasks. 

236 """ 

237 

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 ) 

243 

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. 

248 

249 Returns 

250 ------- 

251 The set of such bands. 

252 """ 

253 return set() 

254 

255 

256class CoaddMultibandFitSubTask(pipeBase.Task, ABC): 

257 """Subtask interface for multiband fitting of deblended sources. 

258 

259 Parameters 

260 ---------- 

261 **kwargs 

262 Additional arguments to be passed to the `lsst.pipe.base.Task` 

263 constructor. 

264 """ 

265 ConfigClass = CoaddMultibandFitSubConfig 

266 

267 def __init__(self, **kwargs): 

268 super().__init__(**kwargs) 

269 

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. 

275 

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. 

282 

283 Returns 

284 ------- 

285 retStruct : `lsst.pipe.base.Struct` 

286 A struct with a cat_output attribute containing the output 

287 measurement catalog. 

288 

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 """ 

299 

300 

301class CoaddMultibandFitBaseConfig( 

302 pipeBase.PipelineTaskConfig, 

303 pipelineConnections=CoaddMultibandFitInputConnections, 

304): 

305 """Base class for multiband fitting.""" 

306 

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 ) 

341 

342 def get_band_sets(self): 

343 """Get the set of bands required by the fit_coadd_multiband subtask. 

344 

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)) 

359 

360 

361class CoaddMultibandFitConfig( 

362 CoaddMultibandFitBaseConfig, 

363 pipelineConnections=CoaddMultibandFitConnections, 

364): 

365 """Configuration for a CoaddMultibandFitTask.""" 

366 

367 

368class CoaddMultibandFitBase: 

369 """Base class for tasks that fit or rebuild multiband models. 

370 

371 This class only implements data reconstruction. 

372 """ 

373 

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"] 

402 

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 

415 

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 

459 

460 

461class CoaddMultibandFitTask(CoaddMultibandFitBase, pipeBase.PipelineTask): 

462 """Fit deblended exposures in multiple bands simultaneously. 

463 

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 """ 

468 

469 ConfigClass = CoaddMultibandFitConfig 

470 _DefaultName = "coaddMultibandFit" 

471 

472 def __init__(self, initInputs, **kwargs): 

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

474 self.makeSubtask("fit_coadd_multiband") 

475 

476 def make_kwargs(self, butlerQC, inputRefs, inputs): 

477 """Make any kwargs needed to be passed to run. 

478 

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 {} 

483 

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) 

495 

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). 

504 

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. 

511 

512 Returns 

513 ------- 

514 retStruct : `lsst.pipe.base.Struct` 

515 A struct with a cat_output attribute containing the output 

516 measurement catalog. 

517 

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