Coverage for python/lsst/analysis/ap/apdbReconstruct.py: 68%

135 statements  

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

1# This file is part of analysis_ap. 

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"""Reconstruct APDB-shaped catalogs from DiaPipelineTask butler outputs. 

23 

24`lsst.ap.association.diaPipe.DiaPipelineTask` writes APDB-bound catalogs as 

25side-effects of its quantum execution: per-(visit, detector) DiaSource, 

26DiaObject, and DiaForcedSource datasets that mirror what gets inserted 

27into the APDB. `ApdbReconstructor` walks those datasets in a butler 

28collection and concatenates them into single DataFrames matching the APDB 

29SDM schema. 

30 

31By default the reconstructor uses the dataset names configured by the 

32production AP pipeline (``ap_pipe/pipelines/_ingredients/ApPipe.yaml``'s 

33``associateApdb`` task: ``dia_source_apdb``, ``dia_object_apdb``, 

34``dia_forced_source_apdb``). For runs that use the raw 

35``DiaPipelineConnections`` defaults instead, pass ``dataset_names`` 

36explicitly. 

37""" 

38 

39from __future__ import annotations 

40 

41__all__ = ["ApdbReconstruction", "ApdbReconstructor"] 

42 

43import dataclasses 

44import logging 

45 

46import pandas as pd 

47 

48from lsst.pipe.tasks.schemaUtils import convertDataFrameToSdmSchema 

49 

50from .apdb import DbQuery, _apdb_schema 

51 

52_log = logging.getLogger(__name__) 

53 

54 

55@dataclasses.dataclass 

56class ApdbReconstruction: 

57 """APDB-shaped DataFrames reconstructed from DiaPipelineTask outputs. 

58 

59 Attributes 

60 ---------- 

61 diaSources : `pandas.DataFrame` 

62 DiaSource rows after association and standardization, deduped on 

63 ``diaSourceId``. 

64 diaObjects : `pandas.DataFrame` 

65 DiaObject rows. By default deduped to the latest snapshot per 

66 ``diaObjectId`` (``history=False`` in `ApdbReconstructor.reconstruct`). 

67 diaForcedSources : `pandas.DataFrame` 

68 DiaForcedSource rows, deduped on the schema primary key 

69 ``(diaObjectId, visit, detector)``. 

70 """ 

71 diaSources: pd.DataFrame 

72 diaObjects: pd.DataFrame 

73 diaForcedSources: pd.DataFrame 

74 

75 

76class ApdbReconstructor: 

77 """Reconstruct APDB-shaped catalogs from `DiaPipelineTask` butler outputs. 

78 

79 Parameters 

80 ---------- 

81 butler : `lsst.daf.butler.Butler` 

82 Butler initialized with the collections that hold the pipeline run. 

83 dataset_names : `dict` [`str`, `list` [`str`]], optional 

84 Override the default dataset types used for each table. Keys are 

85 ``"diaSource"``, ``"diaObject"``, ``"diaForcedSource"``; each 

86 value is a list of dataset-type names that get concatenated and 

87 deduplicated. Defaults to ``DEFAULT_DATASET_NAMES``, which includes the 

88 ``preloaded_*`` datasets so the reconstruction includes both "history" 

89 rows from preload and the new ones. 

90 collections : `list` [`str`], optional 

91 Specify the butler collections to query if not using the default 

92 configured for the butler. 

93 where : `str`, optional 

94 Butler ``where`` clause passed through to ``queryDatasets`` 

95 (e.g. ``"instrument='LSSTComCam' AND visit > 1000"``). 

96 """ 

97 

98 DEFAULT_DATASET_NAMES = { 

99 "diaSource": ["dia_source_apdb", "preloaded_dia_source"], 

100 "diaObject": ["dia_object_apdb", "preloaded_dia_object"], 

101 "diaForcedSource": ["dia_forced_source_apdb", "preloaded_dia_forced_source"], 

102 } 

103 

104 def __init__(self, butler, dataset_names=None, *, 

105 collections=None, where=None): 

106 self.butler = butler 

107 self.dataset_names = (dataset_names if dataset_names is not None 

108 else self.DEFAULT_DATASET_NAMES) 

109 self.collections = collections 

110 self.where = where 

111 self.log = _log 

112 

113 def _query_kwargs(self): 

114 kwargs = {"findFirst": True} 

115 if self.collections is not None: 

116 kwargs["collections"] = self.collections 

117 if self.where is not None: 

118 kwargs["where"] = self.where 

119 return kwargs 

120 

121 def _load_tables(self, dataset_names): 

122 """Load and concatenate one or more dataset types into a single 

123 DataFrame. Returns an empty DataFrame if no dataset is present. 

124 

125 Each dataset is loaded independently via `_load_table` and the 

126 results are concatenated. Dedup happens later in `finalize`, so 

127 overlap between sources (e.g. the same diaSource appearing in 

128 both ``dia_source_apdb`` and ``preloaded_dia_source`` after a 

129 prior pipeline run) is harmless. 

130 """ 

131 frames = [self._load_table(name) for name in dataset_names] 

132 frames = [f for f in frames if len(f)] 

133 if not frames: 

134 return pd.DataFrame() 

135 if len(frames) == 1: 

136 return frames[0] 

137 return pd.concat(frames, ignore_index=True) 

138 

139 def _load_table(self, dataset_name): 

140 """Load every instance of ``dataset_name`` from the butler and 

141 concatenate into a single DataFrame. Returns an empty DataFrame if 

142 the dataset type doesn't exist or no refs match. 

143 """ 

144 try: 

145 refs = list(self.butler.registry.queryDatasets( 

146 dataset_name, **self._query_kwargs())) 

147 except Exception as e: 

148 self.log.info("Skipping %s: query failed (%s)", 

149 dataset_name, e) 

150 return pd.DataFrame() 

151 if not refs: 

152 return pd.DataFrame() 

153 frames = [self.butler.get(ref, storageClass="DataFrame") for ref in refs] 

154 return pd.concat(frames, ignore_index=True) 

155 

156 def reconstruct(self, *, coerce_to_schema=True, history=False): 

157 """Load all per-quantum catalogs and return APDB-shaped DataFrames. 

158 

159 Parameters 

160 ---------- 

161 coerce_to_schema : `bool`, optional 

162 If True (default), coerce each output to the SDM ``apdb.yaml`` 

163 schema. 

164 history : `bool`, optional 

165 If True, keep every diaObject row written across all quanta. 

166 If False (default), dedupe to the atest row per ``diaObjectId``, 

167 mirroring the "current APDB state" view that ``DiaObjectLast`` 

168 would give. 

169 

170 Returns 

171 ------- 

172 result : `ApdbReconstruction` 

173 """ 

174 diaSources = self._load_tables(self.dataset_names["diaSource"]) 

175 diaObjects = self._load_tables(self.dataset_names["diaObject"]) 

176 diaForcedSources = self._load_tables(self.dataset_names["diaForcedSource"]) 

177 return self.finalize(diaSources, diaObjects, diaForcedSources, 

178 coerce_to_schema=coerce_to_schema, 

179 history=history) 

180 

181 @staticmethod 

182 def finalize(diaSources, diaObjects, diaForcedSources, *, 

183 coerce_to_schema=True, history=False): 

184 """Dedup and (optionally) schema-coerce already-loaded catalogs. 

185 

186 Parameters 

187 ---------- 

188 diaSources, diaObjects, diaForcedSources : `pandas.DataFrame` 

189 Concatenated per-quantum catalogs. 

190 coerce_to_schema, history : see `reconstruct`. 

191 

192 Returns 

193 ------- 

194 result : `ApdbReconstruction` 

195 """ 

196 # DiaSource: Primary Key (PK) is diaSourceId. 

197 if len(diaSources) and "diaSourceId" in diaSources.columns: 

198 diaSources = diaSources.drop_duplicates(subset="diaSourceId", 

199 keep="last") 

200 # DiaForcedSource: PK is (diaObjectId, visit, detector). 

201 fkey = ["diaObjectId", "visit", "detector"] 

202 if (len(diaForcedSources) 

203 and set(fkey).issubset(diaForcedSources.columns)): 

204 diaForcedSources = diaForcedSources.drop_duplicates( 

205 subset=fkey, keep="last") 

206 # DiaObject: each quantum that touches a diaObject emits a row for 

207 # it. Many of those rows are "passthrough" snapshots: the diaObject 

208 # was in the quantum's preloaded working set but wasn't actually 

209 # updated, and the writer leaves ``nDiaSources`` (and other 

210 # update-only fields) as NaN/NULL. The validity timestamps on those 

211 # passthrough rows still advance to the quantum's processing time, 

212 # so a naive "sort by validityStart, keep last" dedup picks them 

213 # over the older snapshot that carries the real count. 

214 # 

215 # Fix: include ``nDiaSources`` as the primary sort key with NaN at 

216 # the front, so dedup ``keep="last"`` prefers any informative 

217 # snapshot (highest ``nDiaSources``, ties broken by latest validity) 

218 # over a passthrough one. Falls back to validity-only sort when 

219 # ``nDiaSources`` is absent. 

220 if (len(diaObjects) and "diaObjectId" in diaObjects.columns 

221 and not history): 

222 validity_col = next( 

223 (c for c in ("validityStartMjdTai", "validityStart") 

224 if c in diaObjects.columns), None) 

225 sort_keys = [] 

226 if "nDiaSources" in diaObjects.columns: 226 ↛ 228line 226 didn't jump to line 228 because the condition on line 226 was always true

227 sort_keys.append("nDiaSources") 

228 if validity_col is not None: 228 ↛ 230line 228 didn't jump to line 230 because the condition on line 228 was always true

229 sort_keys.append(validity_col) 

230 if sort_keys: 230 ↛ 233line 230 didn't jump to line 233 because the condition on line 230 was always true

231 diaObjects = diaObjects.sort_values(sort_keys, 

232 na_position="first") 

233 diaObjects = diaObjects.drop_duplicates(subset="diaObjectId", 

234 keep="last") 

235 

236 if coerce_to_schema: 

237 schema = _apdb_schema() 

238 if len(diaSources): 238 ↛ 241line 238 didn't jump to line 241 because the condition on line 238 was always true

239 diaSources = convertDataFrameToSdmSchema( 

240 schema, diaSources, "DiaSource", skipIndex=True) 

241 if len(diaObjects): 241 ↛ 244line 241 didn't jump to line 244 because the condition on line 241 was always true

242 diaObjects = convertDataFrameToSdmSchema( 

243 schema, diaObjects, "DiaObject", skipIndex=True) 

244 if len(diaForcedSources): 244 ↛ 249line 244 didn't jump to line 249 because the condition on line 244 was always true

245 diaForcedSources = convertDataFrameToSdmSchema( 

246 schema, diaForcedSources, "DiaForcedSource", 

247 skipIndex=True) 

248 

249 return ApdbReconstruction( 

250 diaSources=diaSources.reset_index(drop=True), 

251 diaObjects=diaObjects.reset_index(drop=True), 

252 diaForcedSources=diaForcedSources.reset_index(drop=True), 

253 ) 

254 

255 def to_query(self, *, coerce_to_schema=True, history=False): 

256 """Reconstruct and wrap the result as a `DbQuery`-compatible adapter 

257 so the in-memory frames can be passed to `lightcurve`, 

258 `PlotDiaSourceLightcurveTask`, and other tools that expect the 

259 ``DbQuery`` interface. 

260 

261 Returns 

262 ------- 

263 query : `InMemoryDbQuery` 

264 """ 

265 recon = self.reconstruct(coerce_to_schema=coerce_to_schema, 

266 history=history) 

267 return InMemoryDbQuery(recon.diaSources, 

268 recon.diaObjects, 

269 recon.diaForcedSources) 

270 

271 

272class InMemoryDbQuery(DbQuery): 

273 """`DbQuery` backed by in-memory DataFrames (as from `ApdbReconstructor`). 

274 

275 Implements the same load_* methods as `ApdbSqliteQuery`/`ApdbPostgresQuery` 

276 so the reconstructed data can be passed directly to ``lightcurve`` and 

277 ``PlotDiaSourceLightcurveTask``. 

278 """ 

279 

280 def __init__(self, diaSources, diaObjects, diaForcedSources): 

281 self._diaSources = diaSources 

282 self._diaObjects = diaObjects 

283 self._diaForcedSources = diaForcedSources 

284 self.diaSource_flags_exclude = [] 

285 

286 def set_excluded_diaSource_flags(self, flag_list): 

287 # Docstring inherited. 

288 missing = [f for f in flag_list if f not in self._diaSources.columns] 

289 if missing: 

290 raise ValueError( 

291 f"flag(s) {missing} not present in reconstructed DiaSource columns") 

292 self.diaSource_flags_exclude = list(flag_list) 

293 

294 def _apply_flag_exclusion(self, df): 

295 if not self.diaSource_flags_exclude: 295 ↛ 296line 295 didn't jump to line 296 because the condition on line 295 was never true

296 return df 

297 mask = pd.Series(False, index=df.index) 

298 for flag in self.diaSource_flags_exclude: 

299 if flag in df.columns: 299 ↛ 298line 299 didn't jump to line 298 because the condition on line 299 was always true

300 mask |= df[flag].fillna(False).astype(bool) 

301 return df[~mask] 

302 

303 def load_sources_for_object(self, dia_object_id, exclude_flagged=False, 

304 limit=100000): 

305 # Docstring inherited. 

306 df = self._diaSources 

307 result = df[df["diaObjectId"] == dia_object_id] 

308 if exclude_flagged: 308 ↛ 309line 308 didn't jump to line 309 because the condition on line 308 was never true

309 result = self._apply_flag_exclusion(result) 

310 return result.head(limit).reset_index(drop=True) 

311 

312 def load_forced_sources_for_object(self, dia_object_id, 

313 exclude_flagged=False, limit=100000): 

314 # Docstring inherited. 

315 df = self._diaForcedSources 

316 result = df[df["diaObjectId"] == dia_object_id] 

317 return result.head(limit).reset_index(drop=True) 

318 

319 def load_source(self, id): 

320 # Docstring inherited. 

321 match = self._diaSources[self._diaSources["diaSourceId"] == id] 

322 if len(match) == 0: 322 ↛ 324line 322 didn't jump to line 324 because the condition on line 322 was always true

323 raise RuntimeError(f"diaSourceId={id} not found in DiaSource table") 

324 return match.iloc[0] 

325 

326 def load_sources(self, exclude_flagged=False, limit=100000): 

327 # Docstring inherited. 

328 df = self._diaSources 

329 if exclude_flagged: 

330 df = self._apply_flag_exclusion(df) 

331 return df.head(limit).reset_index(drop=True) 

332 

333 def load_object(self, id): 

334 # Docstring inherited. 

335 match = self._diaObjects[self._diaObjects["diaObjectId"] == id] 

336 if len(match) == 0: 336 ↛ 337line 336 didn't jump to line 337 because the condition on line 336 was never true

337 raise RuntimeError(f"diaObjectId={id} not found in DiaObject table") 

338 return match.iloc[0] 

339 

340 def load_objects(self, limit=100000, latest=True): 

341 # Docstring inherited. 

342 return self._diaObjects.head(limit).reset_index(drop=True) 

343 

344 def load_forced_source(self, id): 

345 # Docstring inherited.. 

346 if "diaForcedSourceId" not in self._diaForcedSources.columns: 346 ↛ 347line 346 didn't jump to line 347 because the condition on line 346 was never true

347 raise RuntimeError("Reconstructed DiaForcedSource has no " 

348 "diaForcedSourceId column") 

349 match = self._diaForcedSources[ 

350 self._diaForcedSources["diaForcedSourceId"] == id] 

351 if len(match) == 0: 351 ↛ 352line 351 didn't jump to line 352 because the condition on line 351 was never true

352 raise RuntimeError( 

353 f"diaForcedSourceId={id} not found in DiaForcedSource table") 

354 return match.iloc[0] 

355 

356 def load_forced_sources(self, limit=100000): 

357 # Docstring inherited. 

358 return self._diaForcedSources.head(limit).reset_index(drop=True)