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-14 10:31 +0000
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-14 10:31 +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/>.
22__all__ = ["ModelRebuilder", "PatchCoaddRebuilder", "PatchModelMatches"]
24from collections.abc import Iterable
25from functools import cached_property
27import astropy.table
28import astropy.units as u
29import numpy as np
30import pydantic
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
45from .fit_coadd_multiband import (
46 CatalogExposurePsfs,
47 CatalogSourceFitterConfigData,
48 MultiProFitSourceConfig,
49 MultiProFitSourceTask,
50)
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}
61def astropy_unit_to_geom(unit: u.Unit, default=None) -> geom.AngleUnit:
62 """Convert an astropy unit to an lsst.geom unit.
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.
71 Returns
72 -------
73 unit_geom
74 The equivalent unit, if found.
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
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.
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.
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
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.
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.
127 Returns
128 -------
129 unit
130 The unit of the RA/dec columns or None if none is found.
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
146class DataLoader(pydantic.BaseModel):
147 """A collection of data that can be used to rebuild models."""
149 model_config = pydantic.ConfigDict(arbitrary_types_allowed=True, frozen=True)
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 )
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
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.
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`.
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"]
194 catalog_multi = butler.get(
195 CoaddMultibandFitInputConnections.cat_ref.name.format(name_coadd=name_coadd), **data_id, **kwargs
196 )
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 )
239 def load_deblended_object(
240 self,
241 idx_row: int,
242 ) -> list[g2f.ObservationD]:
243 """Load a deblended object from catexps.
245 Parameters
246 ----------
247 idx_row
248 The index of the object to load.
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
261class ModelRebuilder(DataLoader):
262 """A rebuilder of MultiProFit models from their inputs and best-fit
263 parameter values.
264 """
266 fit_results: astropy.table.Table = pydantic.Field(doc="Multiprofit model fit results")
267 task_fit: MultiProFitSourceTask = pydantic.Field(doc="The task")
269 @cached_property
270 def config_data(self) -> CatalogSourceFitterConfigData:
271 config_data = self.make_config_data()
272 return config_data
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.
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.
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 )
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
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.
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.
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
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.
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)
402class PatchModelMatches(pydantic.BaseModel):
403 """Storage for MultiProFit tables matched to a reference catalog."""
405 model_config = pydantic.ConfigDict(arbitrary_types_allowed=True, frozen=True)
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")
412class PatchCoaddRebuilder(pydantic.BaseModel):
413 """A rebuilder for patch-level coadd catalog/exposure fits."""
415 model_config = pydantic.ConfigDict(arbitrary_types_allowed=True, frozen=True)
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")
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")
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.
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.
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 )