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-10 09:59 +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__ = ( 

23 "ConsolidateAstropyTableConfig", 

24 "ConsolidateAstropyTableConfigBase", 

25 "ConsolidateAstropyTableConnections", 

26 "ConsolidateAstropyTableTask", 

27) 

28 

29from collections import defaultdict 

30 

31import astropy.table as apTab 

32import numpy as np 

33 

34import lsst.pex.config as pexConfig 

35import lsst.pipe.base as pipeBase 

36import lsst.pipe.base.connectionTypes as connectionTypes 

37 

38from .input_config import InputConfig 

39 

40 

41class ConsolidateAstropyTableConfigBase(pexConfig.Config): 

42 """Config for ConsolidateAstropyTableTask.""" 

43 

44 inputs = pexConfig.ConfigDictField( 

45 doc="Mapping of input dataset type config by name", 

46 keytype=str, 

47 itemtype=InputConfig, 

48 default={}, 

49 ) 

50 

51 

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

58 

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 ) 

65 

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) 

75 

76 

77class ConsolidateAstropyTableConfig( 

78 pipeBase.PipelineTaskConfig, 

79 ConsolidateAstropyTableConfigBase, 

80 pipelineConnections=ConsolidateAstropyTableConnections, 

81): 

82 """PipelineTaskConfig for ConsolidateAstropyTableTask.""" 

83 

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 ) 

103 

104 

105class ConsolidateAstropyTableTask(pipeBase.PipelineTask): 

106 """Write patch-merged astropy tables to a tract-level astropy table.""" 

107 

108 _DefaultName = "consolidateAstropyTable" 

109 ConfigClass = ConsolidateAstropyTableConfig 

110 

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 

118 

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] 

125 

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

130 

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 

135 

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) 

141 

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} 

146 

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) 

155 

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) 

164 

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

176 

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] 

208 

209 data[name] = data_name 

210 

211 self.log.info("Concatenating %s per-patch astropy Tables", len(patches)) 

212 

213 tables_read = [] 

214 check_columns = self.config.drop_duplicate_columns or self.config.validate_duplicate_columns 

215 n_bands = len(bands_sorted) 

216 

217 for name, data_name in data.items(): 

218 config_input = self.config.inputs[name] 

219 tables = [] 

220 bands_missing = False 

221 

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 

245 

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 ) 

251 

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] 

272 

273 data[name] = table_new 

274 tables_read.append(name) 

275 

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 ) 

290 

291 butlerQC.put(pipeBase.Struct(cat_output=table), outputRefs)