Coverage for python/lsst/ap/association/transformDiaSourceCatalog.py: 84%
164 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 12:00 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 12:00 +0000
1# This file is part of ap_association
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__ = ("TransformDiaSourceCatalogConnections",
23 "TransformDiaSourceCatalogConfig",
24 "TransformDiaSourceCatalogTask",
25 "UnpackApdbFlags")
27import os
28import yaml
30import numpy as np
32from lsst.resources import ResourcePath
33from lsst.daf.base import DateTime
34import lsst.pex.config as pexConfig
35import lsst.pipe.base as pipeBase
36import lsst.pipe.base.connectionTypes as connTypes
37from lsst.pipe.tasks.postprocess import TransformCatalogBaseTask, TransformCatalogBaseConfig
38from lsst.pipe.tasks.functors import Column
39from lsst.utils.timer import timeMethod
41from lsst.pipe.tasks.schemaUtils import convertDataFrameToSdmSchema, readSdmSchemaFile
44class TransformDiaSourceCatalogConnections(pipeBase.PipelineTaskConnections,
45 dimensions=("instrument", "visit", "detector"),
46 defaultTemplates={"coaddName": "deep", "fakesType": ""}):
47 diaSourceSchema = connTypes.InitInput(
48 doc="Schema for DIASource catalog output by ImageDifference.",
49 storageClass="SourceCatalog",
50 name="{fakesType}{coaddName}Diff_diaSrc_schema",
51 )
52 diaSourceCat = connTypes.Input(
53 doc="Catalog of DiaSources produced during image differencing.",
54 name="{fakesType}{coaddName}Diff_candidateDiaSrc",
55 storageClass="SourceCatalog",
56 dimensions=("instrument", "visit", "detector"),
57 )
58 diffIm = connTypes.Input(
59 doc="Difference image on which the DiaSources were detected.",
60 name="{fakesType}{coaddName}Diff_differenceExp",
61 storageClass="ExposureF",
62 dimensions=("instrument", "visit", "detector"),
63 )
64 diaSourceTable = connTypes.Output(
65 doc=".",
66 name="{fakesType}{coaddName}Diff_diaSrcTable",
67 storageClass="ArrowAstropy",
68 dimensions=("instrument", "visit", "detector"),
69 )
72class TransformDiaSourceCatalogConfig(TransformCatalogBaseConfig,
73 pipelineConnections=TransformDiaSourceCatalogConnections):
74 flagMap = pexConfig.Field(
75 dtype=str,
76 doc="Yaml file specifying SciencePipelines flag fields to bit packs.",
77 default=os.path.join("${AP_ASSOCIATION_DIR}",
78 "data",
79 "association-flag-map.yaml"),
80 )
81 flagRenameMap = pexConfig.Field(
82 dtype=str,
83 doc="Yaml file specifying specifying rules to rename flag names",
84 default=os.path.join("${AP_ASSOCIATION_DIR}",
85 "data",
86 "flag-rename-rules.yaml"),
87 )
88 doRemoveSkySources = pexConfig.Field(
89 dtype=bool,
90 default=False,
91 doc="Input DiaSource catalog contains SkySources that should be "
92 "removed before storing the output DiaSource catalog.",
93 )
94 # TODO: remove on DM-41532
95 doPackFlags = pexConfig.Field(
96 dtype=bool,
97 default=False,
98 doc="Do pack the flags into one integer column named 'flags'."
99 "If False, instead produce one boolean column per flag.",
100 deprecated="This field is no longer used. Will be removed after v28."
101 )
102 doUseApdbSchema = pexConfig.Field(
103 dtype=bool,
104 default=False,
105 doc="Use the APDB schema to coerce the data types of the output columns.",
106 deprecated="This field has been renamed to doUseSchema, and will be "
107 "removed after v30."
108 )
109 doUseSchema = pexConfig.Field(
110 dtype=bool,
111 default=False,
112 doc="Use an existing schema to coerce the data types of the output columns."
113 )
114 schemaDir = pexConfig.Field(
115 dtype=str,
116 doc="Path to the directory containing schema definitions.",
117 default=os.path.join("${SDM_SCHEMAS_DIR}",
118 "yml"),
119 )
120 schemaFile = pexConfig.Field(
121 dtype=str,
122 doc="Yaml file specifying the schema of the output catalog.",
123 default="apdb.yaml",
124 )
125 schemaName = pexConfig.Field(
126 dtype=str,
127 doc="Name of the table in the schema file to read.",
128 default="ApdbSchema",
129 deprecated="This config is no longer used, and will be removed after v30"
130 )
132 def setDefaults(self):
133 super().setDefaults()
134 self.functorFile = os.path.join("${AP_ASSOCIATION_DIR}",
135 "data",
136 "DiaSource.yaml")
139class TransformDiaSourceCatalogTask(TransformCatalogBaseTask):
140 """Transform a DiaSource catalog by calibrating and renaming columns to
141 produce a table ready to insert into the Apdb.
143 Parameters
144 ----------
145 initInputs : `dict`
146 Must contain ``diaSourceSchema`` as the schema for the input catalog.
147 """
148 ConfigClass = TransformDiaSourceCatalogConfig
149 _DefaultName = "transformDiaSourceCatalog"
150 # Needed to create a valid TransformCatalogBaseTask, but unused
151 inputDataset = "deepDiff_diaSrc"
152 outputDataset = "deepDiff_diaSrcTable"
154 def __init__(self, initInputs, **kwargs):
155 super().__init__(**kwargs)
156 self.funcs = self.getFunctors()
157 self.inputSchema = initInputs['diaSourceSchema'].schema
158 self._create_bit_pack_mappings()
159 if self.config.doUseSchema:
160 schemaFile = os.path.join(self.config.schemaDir, self.config.schemaFile)
161 self.schema = readSdmSchemaFile(schemaFile)
162 else:
163 self.schema = None
165 if not self.config.doPackFlags: 165 ↛ exitline 165 didn't return from function '__init__' because the condition on line 165 was always true
166 # get the flag rename rules
167 with open(os.path.expandvars(self.config.flagRenameMap)) as yaml_stream:
168 self.rename_rules = list(yaml.safe_load_all(yaml_stream))
170 def _create_bit_pack_mappings(self):
171 """Setup all flag bit packings.
172 """
173 self.bit_pack_columns = []
174 flag_map_file = os.path.expandvars(self.config.flagMap)
175 with open(flag_map_file) as yaml_stream:
176 table_list = list(yaml.safe_load_all(yaml_stream))
177 for table in table_list: 177 ↛ 185line 177 didn't jump to line 185
178 if table['tableName'] == 'DiaSource': 178 ↛ 177line 178 didn't jump to line 177 because the condition on line 178 was always true
179 self.bit_pack_columns = table['columns']
180 break
182 # Test that all flags requested are present in the input schemas.
183 # Output schemas are flexible, however if names are not specified in
184 # the Apdb schema, flag columns will not be persisted.
185 for outputFlag in self.bit_pack_columns:
186 bitList = outputFlag['bitList']
187 for bit in bitList:
188 try:
189 self.inputSchema.find(bit['name'])
190 except KeyError:
191 raise KeyError(
192 "Requested column %s not found in input DiaSource "
193 "schema. Please check that the requested input "
194 "column exists." % bit['name'])
196 def runQuantum(self, butlerQC, inputRefs, outputRefs):
197 inputs = butlerQC.get(inputRefs)
198 inputs["band"] = butlerQC.quantum.dataId["band"]
200 outputs = self.run(**inputs)
202 butlerQC.put(outputs, outputRefs)
204 @timeMethod
205 def run(self,
206 diaSourceCat,
207 diffIm,
208 band,
209 reliability=None):
210 """Convert input catalog to ParquetTable/Pandas and run functors.
212 Additionally, add new columns for stripping information from the
213 exposure and into the DiaSource catalog.
215 Parameters
216 ----------
217 diaSourceCat : `lsst.afw.table.SourceCatalog`
218 Catalog of sources measured on the difference image.
219 diffIm : `lsst.afw.image.Exposure`
220 Result of subtracting template and science images.
221 band : `str`
222 Filter band of the science image.
223 reliability : `lsst.afw.table.SourceCatalog`
224 Reliability (e.g. real/bogus) scores, row-matched to
225 ``diaSourceCat``.
227 Returns
228 -------
229 results : `lsst.pipe.base.Struct`
230 Results struct with components.
232 - ``diaSourceTable`` : Catalog of DiaSources with calibrated values
233 and renamed columns.
234 (`lsst.pipe.tasks.ParquetTable` or `pandas.DataFrame`)
235 """
236 self.log.info(
237 "Transforming/standardizing the DiaSource table for visit,detector: %i, %i",
238 diffIm.visitInfo.id, diffIm.detector.getId())
240 diaSourceDf = diaSourceCat.asAstropy().to_pandas()
241 if self.config.doRemoveSkySources: 241 ↛ 242line 241 didn't jump to line 242 because the condition on line 241 was never true
242 diaSourceDf = diaSourceDf[~diaSourceDf["sky_source"]]
243 diaSourceCat = diaSourceCat[~diaSourceCat["sky_source"]]
245 # Need UTC time but without a timezone because pandas requires a
246 # naive datetime.
247 diaSourceDf["timeProcessedMjdTai"] = DateTime.now().get(system=DateTime.MJD, scale=DateTime.TAI)
248 diaSourceDf["snr"] = getSignificance(diaSourceCat)
249 diaSourceDf["bboxSize"] = self.computeBBoxSizes(diaSourceCat)
250 diaSourceDf["visit"] = diffIm.visitInfo.id
251 # int16 instead of uint8 because databases don't like unsigned bytes.
252 diaSourceDf["detector"] = np.int16(diffIm.detector.getId())
253 diaSourceDf["band"] = band
254 diaSourceDf["midpointMjdTai"] = diffIm.visitInfo.date.get(system=DateTime.MJD)
255 diaSourceDf["exposureTime"] = diffIm.visitInfo.exposureTime
256 diaSourceDf["diaObjectId"] = 0
257 diaSourceDf["ssObjectId"] = 0
259 # TODO: this has been formally deprecated and should be removed too
260 if self.config.doPackFlags: 260 ↛ 262line 260 didn't jump to line 262 because the condition on line 260 was never true
261 # either bitpack the flags
262 self.bitPackFlags(diaSourceDf)
263 else:
264 # or add the individual flag functors
265 self.addUnpackedFlagFunctors()
266 # and remove the packed flag functor
267 if 'flags' in self.funcs.funcDict: 267 ↛ 270line 267 didn't jump to line 270 because the condition on line 267 was always true
268 del self.funcs.funcDict['flags']
270 df = self.transform(band,
271 diaSourceDf,
272 self.funcs,
273 dataId=None).df
274 if self.config.doUseSchema:
275 df = convertDataFrameToSdmSchema(self.schema, df, tableName="DiaSource")
277 return pipeBase.Struct(
278 diaSourceTable=df,
279 )
281 def addUnpackedFlagFunctors(self):
282 """Add Column functor for each of the flags to the internal functor
283 dictionary.
284 """
285 for flag in self.bit_pack_columns[0]['bitList']:
286 flagName = flag['name']
287 targetName = self.funcs.renameCol(flagName, self.rename_rules[0]['flag_rename_rules'])
288 self.funcs.update({targetName: Column(flagName)})
290 def computeBBoxSizes(self, inputCatalog):
291 """Compute the size of a square bbox that fully contains the detection
292 footprint.
294 Parameters
295 ----------
296 inputCatalog : `lsst.afw.table.SourceCatalog`
297 Catalog containing detected footprints.
299 Returns
300 -------
301 outputBBoxSizes : `np.ndarray`, (N,)
302 Array of bbox sizes.
303 """
304 # Schema validation requires that this field is int.
305 outputBBoxSizes = np.empty(len(inputCatalog), dtype=int)
306 for i, record in enumerate(inputCatalog):
307 footprintBBox = record.getFootprint().getBBox()
308 # Compute twice the size of the largest dimension of the footprint
309 # bounding box. This is the largest footprint we should need to cover
310 # the complete DiaSource assuming the centroid is within the bounding
311 # box.
312 maxSize = 2 * np.max([footprintBBox.getWidth(),
313 footprintBBox.getHeight()])
314 recX = record.getCentroid().x
315 recY = record.getCentroid().y
316 bboxSize = int(
317 np.ceil(2 * np.max(np.fabs([footprintBBox.maxX - recX,
318 footprintBBox.minX - recX,
319 footprintBBox.maxY - recY,
320 footprintBBox.minY - recY]))))
321 if bboxSize > maxSize: 321 ↛ 322line 321 didn't jump to line 322 because the condition on line 321 was never true
322 bboxSize = maxSize
323 outputBBoxSizes[i] = bboxSize
325 return outputBBoxSizes
327 def bitPackFlags(self, df):
328 """Pack requested flag columns in inputRecord into single columns in
329 outputRecord.
331 Parameters
332 ----------
333 df : `pandas.DataFrame`
334 DataFrame to read bits from and pack them into.
335 """
336 for outputFlag in self.bit_pack_columns:
337 bitList = outputFlag['bitList']
338 value = np.zeros(len(df), dtype=np.uint64)
339 for bit in bitList:
340 # Hard type the bit arrays.
341 value += (df[bit['name']]*2**bit['bit']).to_numpy().astype(np.uint64)
342 df[outputFlag['columnName']] = value
345class UnpackApdbFlags:
346 """Class for unpacking bits from integer flag fields stored in the Apdb.
348 Attributes
349 ----------
350 flag_map_file : `lsst.resources.ResourcePathExpression`
351 Absolute or relative URI to a yaml file specifiying mappings of flags
352 to integer bits.
353 table_name : `str`
354 Name of the Apdb table the integer bit data are coming from.
355 """
357 def __init__(self, flag_map_file, table_name):
358 self.bit_pack_columns = []
359 flag_map_file = os.path.expandvars(flag_map_file)
360 with ResourcePath(flag_map_file, forceDirectory=False).open("r") as yaml_stream:
361 table_list = list(yaml.safe_load_all(yaml_stream))
362 for table in table_list: 362 ↛ 367line 362 didn't jump to line 367
363 if table['tableName'] == table_name: 363 ↛ 362line 363 didn't jump to line 362 because the condition on line 363 was always true
364 self.bit_pack_columns = table['columns']
365 break
367 self.output_flag_columns = {}
369 for column in self.bit_pack_columns:
370 names = {}
371 for bit in column["bitList"]:
372 names[bit["name"]] = bit["bit"]
373 self.output_flag_columns[column["columnName"]] = names
375 def unpack(self, input_flag_values, flag_name):
376 """Determine individual boolean flags from an input array of unsigned
377 ints.
379 Parameters
380 ----------
381 input_flag_values : array-like of type uint
382 Array of integer packed bit flags to unpack.
383 flag_name : `str`
384 Apdb column name from the loaded file, e.g. "flags".
386 Returns
387 -------
388 output_flags : `numpy.ndarray`
389 Numpy structured array of booleans, one column per flag in the
390 loaded file.
391 """
392 output_flags = np.zeros(len(input_flag_values),
393 dtype=[(name, bool) for name in self.output_flag_columns[flag_name]])
395 for name in self.output_flag_columns[flag_name]:
396 masked_bits = np.bitwise_and(input_flag_values,
397 2**self.output_flag_columns[flag_name][name])
398 output_flags[name] = masked_bits
400 return output_flags
402 def flagExists(self, flagName, columnName='flags'):
403 """Check if named flag is in the bitpacked flag set.
405 Parameters:
406 ----------
407 flagName : `str`
408 Flag name to search for.
409 columnName : `str`, optional
410 Name of bitpacked flag column to search in.
412 Returns
413 -------
414 flagExists : `bool`
415 `True` if `flagName` is present in `columnName`.
417 Raises
418 ------
419 ValueError
420 Raised if `columnName` is not defined.
421 """
422 if columnName not in self.output_flag_columns:
423 raise ValueError(f'column {columnName} not in flag map: {self.output_flag_columns}')
425 return flagName in [c for c in self.output_flag_columns[columnName]]
427 def makeFlagBitMask(self, flagNames, columnName='flags'):
428 """Return a bitmask corresponding to the supplied flag names.
430 Parameters:
431 ----------
432 flagNames : `list` [`str`]
433 Flag names to include in the bitmask.
434 columnName : `str`, optional
435 Name of bitpacked flag column.
437 Returns
438 -------
439 bitmask : `np.unit64`
440 Bitmask corresponding to the supplied flag names given the loaded configuration.
442 Raises
443 ------
444 ValueError
445 Raised if a flag in `flagName` is not included in `columnName`.
446 """
447 bitmask = np.uint64(0)
449 for flag in flagNames:
450 if not self.flagExists(flag, columnName=columnName):
451 raise ValueError(f"flag '{flag}' not included in '{columnName}' flag column")
453 for outputFlag in self.bit_pack_columns:
454 if outputFlag['columnName'] == columnName: 454 ↛ 453line 454 didn't jump to line 453 because the condition on line 454 was always true
455 bitList = outputFlag['bitList']
456 for bit in bitList:
457 if bit['name'] in flagNames:
458 bitmask += np.uint64(2**bit['bit'])
460 return bitmask
463def getSignificance(catalog):
464 """Return the significance value of the first peak in each source
465 footprint, or NaN for peaks without a significance field.
467 Parameters
468 ----------
469 catalog : `lsst.afw.table.SourceCatalog`
470 Catalog to process.
472 Returns
473 -------
474 significance : `np.ndarray`, (N,)
475 Signficance of the first peak in each source footprint.
476 """
477 result = np.full(len(catalog), np.nan)
478 for i, record in enumerate(catalog):
479 peaks = record.getFootprint().peaks
480 if "significance" in peaks.schema:
481 result[i] = peaks[0]["significance"]
482 return result