Coverage for python/lsst/meas/extensions/multiprofit/rebuild_coadd_multiband.py: 0%

178 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-02 10:24 +0000

1# This file is part of meas_extensions_multiprofit. 

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__ = ["ModelRebuilder", "PatchCoaddRebuilder", "PatchModelMatches"] 

23 

24from collections.abc import Iterable 

25from functools import cached_property 

26 

27import astropy.table 

28import astropy.units as u 

29import numpy as np 

30import pydantic 

31 

32import lsst.afw.table as afwTable 

33import lsst.daf.butler as dafButler 

34import lsst.gauss2d.fit as g2f 

35import lsst.geom as geom 

36from lsst.meas.extensions.scarlet.io import updateCatalogFootprints 

37from lsst.pipe.base import QuantumContext, QuantumGraph 

38from lsst.pipe.tasks.fit_coadd_multiband import ( 

39 CoaddMultibandFitBaseTemplates, 

40 CoaddMultibandFitInputConnections, 

41 CoaddMultibandFitTask, 

42) 

43from lsst.skymap import BaseSkyMap, TractInfo 

44 

45from .fit_coadd_multiband import ( 

46 CatalogExposurePsfs, 

47 CatalogSourceFitterConfigData, 

48 MultiProFitSourceConfig, 

49 MultiProFitSourceTask, 

50) 

51 

52astropy_to_geom_units = { 

53 u.arcmin: geom.arcminutes, 

54 u.arcsec: geom.arcseconds, 

55 u.mas: geom.milliarcseconds, 

56 u.deg: geom.degrees, 

57 u.rad: geom.radians, 

58} 

59 

60 

61def astropy_unit_to_geom(unit: u.Unit, default=None) -> geom.AngleUnit: 

62 """Convert an astropy unit to an lsst.geom unit. 

63 

64 Parameters 

65 ---------- 

66 unit 

67 The astropy unit to convert. 

68 default 

69 The default value to return if no known conversion is found. 

70 

71 Returns 

72 ------- 

73 unit_geom 

74 The equivalent unit, if found. 

75 

76 Raises 

77 ------ 

78 ValueError 

79 Raised if no equivalent unit is found. 

80 """ 

81 unit_geom = astropy_to_geom_units.get(unit, default) 

82 if unit_geom is None: 

83 raise ValueError(f"{unit=} not found in {astropy_to_geom_units=}") 

84 return unit_geom 

85 

86 

87def find_patches(tract_info: TractInfo, ra_array, dec_array, unit: geom.AngleUnit) -> list[int]: 

88 """Find the patches containing a list of ra/dec values within a tract. 

89 

90 Parameters 

91 ---------- 

92 tract_info 

93 The TractInfo object for the tract. 

94 ra_array 

95 The array of right ascension values. 

96 dec_array 

97 The array of declination values (must be same length as ra_array). 

98 unit 

99 The unit of the RA/dec values. 

100 

101 Returns 

102 ------- 

103 patches 

104 A list of patches containing the specified RA/dec values. 

105 """ 

106 radec = [geom.SpherePoint(ra, dec, units=unit) for ra, dec in zip(ra_array, dec_array, strict=True)] 

107 points = np.array([geom.Point2I(tract_info.wcs.skyToPixel(coords)) for coords in radec]) 

108 x_list, y_list = (points[:, idx] // tract_info.patch_inner_dimensions[idx] for idx in range(2)) 

109 patches = [tract_info.getSequentialPatchIndexFromPair((x, y)) for x, y in zip(x_list, y_list)] 

110 return patches 

111 

112 

113def get_radec_unit(table: astropy.table.Table, coord_ra: str, coord_dec: str, default=None): 

114 """Get the RA/dec units for columns in a table. 

115 

116 Parameters 

117 ---------- 

118 table 

119 The table to determine units for. 

120 coord_ra 

121 The key of the right ascension column. 

122 coord_dec 

123 The key of the declination column. 

124 default 

125 The default value to return if no unit is found. 

126 

127 Returns 

128 ------- 

129 unit 

130 The unit of the RA/dec columns or None if none is found. 

131 

132 Raises 

133 ------ 

134 ValueError 

135 Raised if the units are inconsistent. 

136 """ 

137 unit_ra, unit_dec = ( 

138 astropy_unit_to_geom(table[coord].unit, default=default) for coord in (coord_ra, coord_dec) 

139 ) 

140 if unit_ra != unit_dec: 

141 units = {coord: table[coord].unit for coord in (coord_ra, coord_dec)} 

142 raise ValueError(f"Reference table has inconsistent {units=}") 

143 return unit_ra 

144 

145 

146class DataLoader(pydantic.BaseModel): 

147 """A collection of data that can be used to rebuild models.""" 

148 

149 model_config = pydantic.ConfigDict(arbitrary_types_allowed=True, frozen=True) 

150 

151 catexps: list[CatalogExposurePsfs] = pydantic.Field( 

152 doc="List of MultiProFit catalog-exposure-psf objects used to fit PSF-convolved models", 

153 ) 

154 catalog_multi: afwTable.SourceCatalog = pydantic.Field( 

155 doc="Patch-level multiband reference catalog (deepCoadd_ref)", 

156 ) 

157 

158 @cached_property 

159 def channels(self) -> tuple[g2f.Channel]: 

160 channels = tuple(g2f.Channel.get(catexp.band) for catexp in self.catexps) 

161 return channels 

162 

163 @classmethod 

164 def from_butler( 

165 cls, butler: dafButler.Butler, data_id: dict[str], bands: Iterable[str], name_coadd=None, **kwargs 

166 ): 

167 """Construct a DataLoader from a Butler and dataId. 

168 

169 Parameters 

170 ---------- 

171 butler 

172 The butler to load from. 

173 data_id 

174 Key-value pairs for the {name_coadd}Coadd_* dataId. 

175 bands 

176 The list of bands to load. 

177 name_coadd 

178 The prefix of the Coadd datasettype name. 

179 **kwargs 

180 Additional keyword arguments to pass to the init method for 

181 `CoaddMultibandFitInputConnections`. 

182 

183 Returns 

184 ------- 

185 data_loader 

186 An initialized DataLoader. 

187 """ 

188 bands = tuple(bands) 

189 if len(set(bands)) != len(bands): 

190 raise ValueError(f"{bands=} is not a set") 

191 if name_coadd is None: 

192 name_coadd = CoaddMultibandFitBaseTemplates["name_coadd"] 

193 

194 catalog_multi = butler.get( 

195 CoaddMultibandFitInputConnections.cat_ref.name.format(name_coadd=name_coadd), **data_id, **kwargs 

196 ) 

197 

198 catexps = {} 

199 for band in bands: 

200 data_id["band"] = band 

201 catalog = butler.get( 

202 CoaddMultibandFitInputConnections.cats_meas.name.format(name_coadd=name_coadd), 

203 **data_id, 

204 **kwargs, 

205 ) 

206 exposure = butler.get( 

207 CoaddMultibandFitInputConnections.coadds.name.format(name_coadd=name_coadd), 

208 **data_id, 

209 **kwargs, 

210 ) 

211 models_scarlet = butler.get( 

212 CoaddMultibandFitInputConnections.models_scarlet.name.format(name_coadd=name_coadd), 

213 **data_id, 

214 **kwargs, 

215 ) 

216 updateCatalogFootprints( 

217 modelData=models_scarlet, 

218 catalog=catalog, 

219 band=data_id["band"], 

220 imageForRedistribution=exposure, 

221 removeScarletData=True, 

222 updateFluxColumns=False, 

223 ) 

224 # The config and table are harmless dummies 

225 catexps[band] = CatalogExposurePsfs( 

226 catalog=catalog, 

227 exposure=exposure, 

228 table_psf_fits=astropy.table.Table(), 

229 dataId=data_id, 

230 id_tract_patch=data_id["patch"], 

231 channel=g2f.Channel.get(band), 

232 config_fit=MultiProFitSourceConfig(), 

233 ) 

234 return cls( 

235 catalog_multi=catalog_multi, 

236 catexps=list(catexps.values()), 

237 ) 

238 

239 def load_deblended_object( 

240 self, 

241 idx_row: int, 

242 ) -> list[g2f.ObservationD]: 

243 """Load a deblended object from catexps. 

244 

245 Parameters 

246 ---------- 

247 idx_row 

248 The index of the object to load. 

249 

250 Returns 

251 ------- 

252 observations 

253 The observations of the object (deblended if it is a child). 

254 """ 

255 observations = [] 

256 for catexp in self.catexps: 

257 observations.append(catexp.get_source_observation(catexp.get_catalog()[idx_row])) 

258 return observations 

259 

260 

261class ModelRebuilder(DataLoader): 

262 """A rebuilder of MultiProFit models from their inputs and best-fit 

263 parameter values. 

264 """ 

265 

266 fit_results: astropy.table.Table = pydantic.Field(doc="Multiprofit model fit results") 

267 task_fit: MultiProFitSourceTask = pydantic.Field(doc="The task") 

268 

269 @cached_property 

270 def config_data(self) -> CatalogSourceFitterConfigData: 

271 config_data = self.make_config_data() 

272 return config_data 

273 

274 @classmethod 

275 def from_quantumGraph( 

276 cls, 

277 butler: dafButler.Butler, 

278 quantumgraph: QuantumGraph, 

279 dataId: dict = None, 

280 ): 

281 """Make a rebuilder from a butler and quantumgraph. 

282 

283 Parameters 

284 ---------- 

285 butler 

286 The butler that the quantumgraph was built for. 

287 quantumgraph 

288 The quantum graph file from a CoaddMultibandFitTask using the 

289 MultiProFitSourceTask. 

290 dataId 

291 The dataId for the fit, including skymap, tract and patch. 

292 

293 Returns 

294 ------- 

295 rebuilder 

296 A ModelRebuilder instance initialized with the necessary kwargs. 

297 """ 

298 if dataId is None: 

299 quantum = next(iter(quantumgraph.outputQuanta)).quantum 

300 else: 

301 quantum = None 

302 for node in quantumgraph.outputQuanta: 

303 if node.quantum.dataId.to_simple().dataId == dataId: 

304 quantum = node.quantum 

305 break 

306 if quantum is None: 

307 raise ValueError( 

308 f"{dataId=} not found in {[x.quantum.dataId for x in quantumgraph.outputQuanta]=}" 

309 ) 

310 taskDef = next(iter(quantumgraph.iterTaskGraph())) 

311 butlerQC = QuantumContext(butler, quantum) 

312 config = butler.get(f"{taskDef.label}_config") 

313 # I have no idea what to put for initInputs. 

314 # quantum.initInputs looks wrong - the values can be lists 

315 # quantumgraph.initInputRefs(taskDef) returns a list of DatasetRefs... 

316 # ... but I'm not sure how to map that to connection names? 

317 task: CoaddMultibandFitTask = taskDef.taskClass(config=config, initInputs={}) 

318 if not isinstance(task, CoaddMultibandFitTask): 

319 raise ValueError(f"{task=} type={type(task)} !isinstance of {CoaddMultibandFitTask=}") 

320 task_fit: MultiProFitSourceTask = task.fit_coadd_multiband 

321 if not isinstance(task_fit, MultiProFitSourceTask): 

322 raise ValueError(f"{task_fit=} type={type(task_fit)} !isinstance of {MultiProFitSourceTask=}") 

323 inputRefs, outputRefs = taskDef.connections.buildDatasetRefs(quantum) 

324 inputs = butlerQC.get(inputRefs) 

325 catexps = task.build_catexps(butlerQC, inputRefs, inputs) 

326 catexps = [task_fit.make_CatalogExposurePsfs(catexp) for catexp in catexps] 

327 cat_output: astropy.table.Table = butler.get(outputRefs.cat_output, storageClass="ArrowAstropy") 

328 return cls( 

329 catexps=catexps, 

330 task_fit=task_fit, 

331 catalog_multi=inputs["cat_ref"], 

332 fit_results=cat_output, 

333 ) 

334 

335 def make_config_data(self): 

336 """Make a ConfigData object out of self's channels and fit task 

337 config. 

338 """ 

339 config_data = CatalogSourceFitterConfigData(channels=self.channels, config=self.task_fit.config) 

340 return config_data 

341 

342 def make_model( 

343 self, 

344 idx_row: int, 

345 config_data: CatalogSourceFitterConfigData = None, 

346 init: bool = True, 

347 ) -> g2f.ModelD: 

348 """Make a ModelD for a single row from the originally fitted catalog. 

349 

350 Parameters 

351 ---------- 

352 idx_row 

353 The index of the row to make a model for. 

354 config_data 

355 The model configuration data object. 

356 init 

357 Whether to initialize the model parameters as they would have been 

358 prior to fitting. 

359 

360 Returns 

361 ------- 

362 model 

363 The rebuilt model. 

364 """ 

365 if config_data is None: 

366 config_data = self.config_data 

367 model = self.task_fit.get_model( 

368 idx_row=idx_row, 

369 catalog_multi=self.catalog_multi, 

370 catexps=self.catexps, 

371 config_data=config_data, 

372 results=self.fit_results, 

373 set_flux_limits=False, 

374 ) 

375 if init: 

376 self.set_model(idx_row, config_data) 

377 return model 

378 

379 def set_model(self, idx_row: int, config_data: CatalogSourceFitterConfigData = None) -> None: 

380 """Set model parameters to the best-fit values for a given row. 

381 

382 Parameters 

383 ---------- 

384 idx_row 

385 The index of the row in the fit parameter table to initialize from. 

386 config_data 

387 The model configuration data object. 

388 """ 

389 if config_data is None: 

390 config_data = self.config_data 

391 row = self.fit_results[idx_row] 

392 prefix = config_data.config.prefix_column 

393 offsets = {} 

394 offset_cen = config_data.config.centroid_pixel_offset 

395 if offset_cen != 0: 

396 offsets[g2f.CentroidXParameterD] = -offset_cen 

397 offsets[g2f.CentroidYParameterD] = -offset_cen 

398 for key, param in config_data.parameters.items(): 

399 param.value = row[f"{prefix}{key}"] + offsets.get(type(param), 0.0) 

400 

401 

402class PatchModelMatches(pydantic.BaseModel): 

403 """Storage for MultiProFit tables matched to a reference catalog.""" 

404 

405 model_config = pydantic.ConfigDict(arbitrary_types_allowed=True, frozen=True) 

406 

407 matches: astropy.table.Table | None = pydantic.Field(doc="Catalogs of matches") 

408 quantumgraph: QuantumGraph | None = pydantic.Field(doc="Quantum graph for fit task") 

409 rebuilder: DataLoader | ModelRebuilder | None = pydantic.Field(doc="MultiProFit object model rebuilder") 

410 

411 

412class PatchCoaddRebuilder(pydantic.BaseModel): 

413 """A rebuilder for patch-level coadd catalog/exposure fits.""" 

414 

415 model_config = pydantic.ConfigDict(arbitrary_types_allowed=True, frozen=True) 

416 

417 matches: dict[str, PatchModelMatches] = pydantic.Field("Model matches by algorithm name") 

418 name_model_ref: str = pydantic.Field(doc="The name of the reference model in matches") 

419 objects: astropy.table.Table = pydantic.Field(doc="Object table") 

420 objects_multiprofit: astropy.table.Table | None = pydantic.Field(doc="Object table for MultiProFit fits") 

421 reference: astropy.table.Table = pydantic.Field(doc="Reference object table") 

422 

423 skymap: str = pydantic.Field(doc="The skymap name") 

424 tract: int = pydantic.Field(doc="The tract index") 

425 patch: int = pydantic.Field(doc="The patch index") 

426 

427 @classmethod 

428 def from_butler( 

429 cls, 

430 butler: dafButler.Butler, 

431 skymap: str, 

432 tract: int, 

433 patch: int, 

434 collection_merged: str, 

435 matches: dict[str, QuantumGraph | None], 

436 bands: Iterable[str] = None, 

437 name_model_ref: str = None, 

438 format_collection: str = "{run}", 

439 load_multiprofit: bool = True, 

440 dataset_type_ref: str = "truth_summary", 

441 ): 

442 """Construct a PatchCoaddRebuilder from a single Butler collection. 

443 

444 Parameters 

445 ---------- 

446 butler 

447 The butler to load from. 

448 skymap 

449 The skymap for the collection. 

450 tract 

451 The skymap tract id. 

452 patch 

453 The skymap patch id. 

454 collection_merged 

455 The name of the collection with the merged objectTable(s). 

456 matches 

457 A dictionary of model names with corresponding QuantumGraphs. 

458 These may be None but must be provided for MultiProFit model 

459 reconstruction to be possible. 

460 bands 

461 The list of bands to load data for. 

462 name_model_ref 

463 The name of the model to use as a reference. Must be a key in 

464 `matches`. 

465 format_collection 

466 A format string for the output collection(s) defined in the 

467 `matches` QuantumGraphs. 

468 load_multiprofit 

469 Whether to attempt to load an objectTable_tract_multiprofit. 

470 dataset_type_ref 

471 The dataset type of the reference catalog. 

472 

473 Returns 

474 ------- 

475 rebuilder 

476 The fully-configured PatchCoaddRebuilder. 

477 """ 

478 if name_model_ref is None: 

479 for name, quantumgraph in matches.items(): 

480 if quantumgraph is not None: 

481 name_model_ref = name 

482 break 

483 if name_model_ref is None: 

484 raise ValueError("Must supply name_model_ref or at least one matches with a quantumgraph") 

485 dataId = dict(skymap=skymap, tract=tract, patch=patch) 

486 objects = butler.get( 

487 "objectTable_tract", collections=[collection_merged], storageClass="ArrowAstropy", **dataId 

488 ) 

489 objects = objects[objects["patch"] == patch] 

490 if load_multiprofit: 

491 objects_multiprofit = butler.get( 

492 "objectTable_tract_multiprofit", 

493 collections=[collection_merged], 

494 storageClass="ArrowAstropy", 

495 **dataId, 

496 ) 

497 objects_multiprofit = objects_multiprofit[objects_multiprofit["patch"] == patch] 

498 else: 

499 objects_multiprofit = None 

500 reference = butler.get( 

501 dataset_type_ref, collections=[collection_merged], storageClass="ArrowAstropy", **dataId 

502 ) 

503 skymap_tract = butler.get(BaseSkyMap.SKYMAP_DATASET_TYPE_NAME, skymap=skymap)[tract] 

504 unit_coord_ref = get_radec_unit(reference, "ra", "dec", default=geom.degrees) 

505 if "patch" not in reference.columns: 

506 patches = find_patches(skymap_tract, reference["ra"], reference["dec"], unit=unit_coord_ref) 

507 reference["patch"] = patches 

508 elif reference["patch"].dtype != int: 

509 # the ci_imsim truth_summary still has string patches 

510 index_patch = skymap_tract[patch].index 

511 str_patch = f"{index_patch.y},{index_patch.x}" 

512 reference = reference[ 

513 (reference["patch"] == str_patch) & (reference["is_unique_truth_entry"] == True) # noqa: E712 

514 ] 

515 del reference["patch"] 

516 reference["patch"] = patch 

517 reference = reference[reference["patch"] == patch] 

518 points = skymap_tract.wcs.skyToPixel( 

519 [geom.SpherePoint(row["ra"], row["dec"], units=geom.degrees) for row in reference] 

520 ) 

521 reference["x"] = [point.x for point in points] 

522 reference["y"] = [point.y for point in points] 

523 matches_name = {} 

524 for name, quantumgraph in matches.items(): 

525 is_mpf = quantumgraph is not None 

526 matched = butler.get( 

527 f"matched_{dataset_type_ref}_objectTable_tract{'_multiprofit' if is_mpf else ''}", 

528 collections=[ 

529 ( 

530 format_collection.format(run=quantumgraph.metadata["output"], name=name) 

531 if is_mpf 

532 else collection_merged 

533 ) 

534 ], 

535 storageClass="ArrowAstropy", 

536 **dataId, 

537 ) 

538 # unmatched ref objects don't have a patch set 

539 # should probably be fixed in diff_matched 

540 # but need to decide priority on matched - ref first? or target? 

541 unit_coord_ref = get_radec_unit( 

542 matched, 

543 "refcat_ra", 

544 "refcat_dec", 

545 default=geom.degrees, 

546 ) 

547 unmatched = ( 

548 matched["patch"].mask if np.ma.is_masked(matched["patch"]) else ~(matched["patch"] >= 0) 

549 ) & np.isfinite(matched["refcat_ra"]) 

550 patches_unmatched = find_patches( 

551 skymap_tract, 

552 matched["refcat_ra"][unmatched], 

553 matched["refcat_dec"][unmatched], 

554 unit=unit_coord_ref, 

555 ) 

556 matched["patch"][np.where(unmatched)[0]] = patches_unmatched 

557 matched = matched[matched["patch"] == patch] 

558 rebuilder = ( 

559 ModelRebuilder.from_quantumGraph(butler, quantumgraph, dataId=dataId) 

560 if is_mpf 

561 else DataLoader.from_butler( 

562 butler, data_id=dataId, bands=bands, collections=[collection_merged] 

563 ) 

564 ) 

565 matches_name[name] = PatchModelMatches( 

566 matches=matched, quantumgraph=quantumgraph, rebuilder=rebuilder 

567 ) 

568 return cls( 

569 matches=matches_name, 

570 objects=objects, 

571 objects_multiprofit=objects_multiprofit, 

572 reference=reference, 

573 skymap=skymap, 

574 tract=tract, 

575 patch=patch, 

576 name_model_ref=name_model_ref, 

577 )