Coverage for python/lsst/meas/extensions/multiprofit/consolidate_astropy_table.py: 0%
133 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-16 10:44 +0000
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-16 10:44 +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__ = (
23 "ConsolidateAstropyTableConfig",
24 "ConsolidateAstropyTableConfigBase",
25 "ConsolidateAstropyTableConnections",
26 "ConsolidateAstropyTableTask",
27)
29from collections import defaultdict
31import astropy.table as apTab
32import numpy as np
34import lsst.pex.config as pexConfig
35import lsst.pipe.base as pipeBase
36import lsst.pipe.base.connectionTypes as connectionTypes
38from .input_config import InputConfig
41class ConsolidateAstropyTableConfigBase(pexConfig.Config):
42 """Config for ConsolidateAstropyTableTask."""
44 inputs = pexConfig.ConfigDictField(
45 doc="Mapping of input dataset type config by name",
46 keytype=str,
47 itemtype=InputConfig,
48 default={},
49 )
52class ConsolidateAstropyTableConnections(
53 # Ignore the undocumented inherited config arg in __init__
54 pipeBase.PipelineTaskConnections,
55 dimensions=("tract", "skymap"), # numpydoc ignore=PR01
56):
57 """Connections for ConsolidateAstropyTableTask."""
59 cat_output = connectionTypes.Output(
60 doc="Per-tract horizontal concatenation of the input AstropyTables",
61 name="objectAstropyTable_tract",
62 storageClass="ArrowTable",
63 dimensions=("tract", "skymap"),
64 )
66 def __init__(self, *, config: ConsolidateAstropyTableConfigBase):
67 super().__init__(config=config)
68 for name, config_input in config.inputs.items():
69 if hasattr(self, name):
70 raise ValueError(
71 f"{config_input=} {name=} is invalid, due to being an existing attribute of {self=}"
72 )
73 connection = config_input.get_connection(name)
74 setattr(self, name, connection)
77class ConsolidateAstropyTableConfig(
78 pipeBase.PipelineTaskConfig,
79 ConsolidateAstropyTableConfigBase,
80 pipelineConnections=ConsolidateAstropyTableConnections,
81):
82 """PipelineTaskConfig for ConsolidateAstropyTableTask."""
84 drop_duplicate_columns = pexConfig.Field[bool](
85 doc="Whether to drop columns from a table if they occur in a previous table."
86 " If False, astropy will rename them with its default scheme.",
87 default=True,
88 )
89 join_type = pexConfig.ChoiceField[str](
90 doc="Type of join to perform in the final hstack",
91 allowed={
92 "inner": "Inner join",
93 "outer": "Outer join",
94 "exact": "Exact join",
95 },
96 default="exact",
97 optional=False,
98 )
99 validate_duplicate_columns = pexConfig.Field[bool](
100 doc="Whether to check that duplicate columns are identical in any table they occur in.",
101 default=True,
102 )
105class ConsolidateAstropyTableTask(pipeBase.PipelineTask):
106 """Write patch-merged astropy tables to a tract-level astropy table."""
108 _DefaultName = "consolidateAstropyTable"
109 ConfigClass = ConsolidateAstropyTableConfig
111 def runQuantum(self, butlerQC, inputRefs, outputRefs):
112 inputs = butlerQC.get(inputRefs)
113 bands_ref, patches_ref = None, None
114 band_null, patch_null = "", -1
115 bands_null, patches_null = {band_null}, {patch_null: None}
116 data = dict()
117 bands_sorted = None
119 # inputRefs are usually unsorted lists so they need to be sorted first
120 for name, inputRef_list in inputRefs:
121 inputConfig = self.config.inputs[name]
122 bands, patches = set(), dict()
123 data_name = defaultdict(dict)
124 inputs_name = inputs[name]
126 # if it's not a list, then it's a single object
127 if not hasattr(inputRef_list, "__len__"):
128 inputRef_list = tuple((inputRef_list,))
129 inputs_name = tuple((inputs_name,))
131 # Add every ref by band (if not multiband)
132 for dataRef, data_in in zip(inputRef_list, inputs_name):
133 dataId = dataRef.dataId
134 band = dataId.band.name if not inputConfig.is_multiband else band_null
136 if inputConfig.columns is not None:
137 columns = inputConfig.columns
138 data_in = data_in.get(parameters={"columns": columns})
139 else:
140 columns = tuple(data_in.columns)
142 if inputConfig.storageClass == "DataFrame":
143 data_in = apTab.Table.from_pandas(data_in.reset_index(drop=False))
144 elif inputConfig.storageClass == "ArrowAstropy":
145 data_in.meta = {name: data_in.meta}
147 if not inputConfig.is_multiband:
148 columns_new = [
149 column if column == inputConfig.column_id else f"{band}_{column}"
150 for column in columns
151 ]
152 data_in.rename_columns(columns, columns_new)
153 if inputConfig.action is not None:
154 data_in = inputConfig.action(data_in, datasetType=name)
156 if inputConfig.is_multipatch:
157 patch = patch_null
158 patches[patch] = None
159 else:
160 patch = dataId.patch.id
161 patches[patch] = min(data_in[inputConfig.column_id])
162 data_name[patch][band] = data_in
163 bands.add(band)
165 # Validate the bands
166 if inputConfig.is_multiband:
167 if bands != bands_null:
168 raise RuntimeError(f"multiband {inputConfig=} has non-trivial {bands=}")
169 else:
170 if bands_ref is None:
171 bands_ref = bands
172 bands_sorted = tuple(band for band in sorted(bands_ref))
173 else:
174 if bands != bands_ref:
175 raise RuntimeError(f"{inputConfig=} {bands=} != {bands_ref=}")
177 # Check that every dataset has the same set of patches
178 if inputConfig.is_multipatch:
179 if patches != patches_null:
180 raise RuntimeError(f"{inputConfig=} {patches=} != {patches_null=}")
181 else:
182 column_id = inputConfig.column_id
183 if patches_ref is None:
184 bands = tuple(bands) if inputConfig.is_multiband else bands_sorted
185 for patch in patches:
186 data_patch = data_name[patch]
187 # Make sure any one-time operations are done once
188 # rather than for every band
189 added = False
190 for band in bands:
191 if tab := data_patch.get(band):
192 if not added:
193 # add a patch column to fill in later
194 tab.add_column(np.full(len(tab), patch), name="patch", index=1)
195 # The id column should be objectId
196 tab.rename_column(column_id, "objectId")
197 added = True
198 else:
199 del tab[column_id]
200 patches_objid = {objid: patch for patch, objid in patches.items()}
201 patches_ref = {patch: objid for objid, patch in sorted(patches_objid.items())}
202 elif {patch: patches[patch] for patch in patches_ref.keys()} != patches_ref:
203 raise RuntimeError(f"{inputConfig=} {patches=} != {patches_ref=}")
204 else:
205 for data_patch in data_name.values():
206 for tab in data_patch.values():
207 del tab[column_id]
209 data[name] = data_name
211 self.log.info("Concatenating %s per-patch astropy Tables", len(patches))
213 tables_read = []
214 check_columns = self.config.drop_duplicate_columns or self.config.validate_duplicate_columns
215 n_bands = len(bands_sorted)
217 for name, data_name in data.items():
218 config_input = self.config.inputs[name]
219 tables = []
220 bands_missing = False
222 # If this is a multipatch dataset, loop over patches
223 # Otherwise, loop over the single "null" patch
224 for patch in patches_ref if not config_input.is_multipatch else patches_null:
225 data_name_patch = data_name[patch]
226 # If this is multiband, use the null band, and return an empty
227 # list if there's no corresponding dataset
228 if config_input.is_multiband:
229 tables_patch = data_name_patch.get(band_null, [])
230 else:
231 # Get the tables (or None if it's missing) in sorted order
232 tables_patch = [
233 _tab for band in bands_sorted if (_tab := data_name_patch.get(band)) is not None
234 ]
235 # Check if any bands are missing
236 if not bands_missing and (len(tables_patch) != n_bands):
237 bands_missing = True
238 # Join only if there's something to join
239 if tables_patch:
240 table_patch = apTab.hstack(tables_patch, join_type="exact")
241 tables.append(table_patch)
242 # If there's nothing to join, presumably the task failed
243 # stacking should handle some tasks failing but not others, but
244 # this shouldn't be relied upon
246 table_new = (
247 tables[0]
248 if (len(tables) == 1)
249 else apTab.vstack(tables, join_type="outer" if bands_missing else "exact")
250 )
252 if check_columns:
253 columns_new = set(x for x in table_new.colnames if x != config_input.join_column)
254 for name_previous in tables_read:
255 table_old = data[name_previous]
256 columns_common = columns_new.intersection(
257 x for x in table_old.colnames if x != self.config.inputs[name_previous].join_column
258 )
259 for column_common in columns_common:
260 if self.config.validate_duplicate_columns:
261 if not np.array_equal(
262 table_new[column_common],
263 table_old[column_common],
264 equal_nan=True,
265 ):
266 raise RuntimeError(
267 f"Joined table column={column_common} differs between {name} and"
268 f" {name_previous} tables"
269 )
270 if self.config.drop_duplicate_columns:
271 del table_new[column_common]
273 data[name] = table_new
274 tables_read.append(name)
276 # This will break if all tables have config.join_column
277 # ... but that seems unlikely.
278 table = apTab.hstack(
279 [data[name] for name, config in self.config.inputs.items() if config.join_column is None],
280 join_type=self.config.join_type,
281 )
282 for name, config in self.config.inputs.items():
283 if config.join_column:
284 table = apTab.join(
285 table,
286 data[name],
287 join_type=self.config.join_type,
288 keys=config.join_column,
289 )
291 butlerQC.put(pipeBase.Struct(cat_output=table), outputRefs)