Coverage for python/lsst/analysis/ap/apdbReconstruct.py: 68%
135 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-28 03:20 -0700
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-28 03:20 -0700
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/>.
22"""Reconstruct APDB-shaped catalogs from DiaPipelineTask butler outputs.
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.
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"""
39from __future__ import annotations
41__all__ = ["ApdbReconstruction", "ApdbReconstructor"]
43import dataclasses
44import logging
46import pandas as pd
48from lsst.pipe.tasks.schemaUtils import convertDataFrameToSdmSchema
50from .apdb import DbQuery, _apdb_schema
52_log = logging.getLogger(__name__)
55@dataclasses.dataclass
56class ApdbReconstruction:
57 """APDB-shaped DataFrames reconstructed from DiaPipelineTask outputs.
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
76class ApdbReconstructor:
77 """Reconstruct APDB-shaped catalogs from `DiaPipelineTask` butler outputs.
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 """
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 }
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
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
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.
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)
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)
156 def reconstruct(self, *, coerce_to_schema=True, history=False):
157 """Load all per-quantum catalogs and return APDB-shaped DataFrames.
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.
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)
181 @staticmethod
182 def finalize(diaSources, diaObjects, diaForcedSources, *,
183 coerce_to_schema=True, history=False):
184 """Dedup and (optionally) schema-coerce already-loaded catalogs.
186 Parameters
187 ----------
188 diaSources, diaObjects, diaForcedSources : `pandas.DataFrame`
189 Concatenated per-quantum catalogs.
190 coerce_to_schema, history : see `reconstruct`.
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")
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)
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 )
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.
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)
272class InMemoryDbQuery(DbQuery):
273 """`DbQuery` backed by in-memory DataFrames (as from `ApdbReconstructor`).
275 Implements the same load_* methods as `ApdbSqliteQuery`/`ApdbPostgresQuery`
276 so the reconstructed data can be passed directly to ``lightcurve`` and
277 ``PlotDiaSourceLightcurveTask``.
278 """
280 def __init__(self, diaSources, diaObjects, diaForcedSources):
281 self._diaSources = diaSources
282 self._diaObjects = diaObjects
283 self._diaForcedSources = diaForcedSources
284 self.diaSource_flags_exclude = []
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)
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]
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)
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)
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]
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)
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]
340 def load_objects(self, limit=100000, latest=True):
341 # Docstring inherited.
342 return self._diaObjects.head(limit).reset_index(drop=True)
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]
356 def load_forced_sources(self, limit=100000):
357 # Docstring inherited.
358 return self._diaForcedSources.head(limit).reset_index(drop=True)