Coverage for python/lsst/analysis/ap/compare.py: 74%
131 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-09-11 04:03 -0700
« prev ^ index » next coverage.py v7.15.4, created at 2026-09-11 04:03 -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"""Catalog cross-matching and pair-wise comparison utilities for
23DiaSource-like tables.
24"""
26__all__ = ["match_catalogs", "flux_residuals", "match_to_truth"]
28import numpy as np
29import pandas as pd
31import astropy.units as u
32from astropy.coordinates import SkyCoord, match_coordinates_sky
35def _as_arcsec(radius):
36 """Coerce a number-or-Quantity radius to a float in arcseconds."""
37 if isinstance(radius, u.Quantity):
38 return float(radius.to_value(u.arcsec))
39 return float(radius)
42def _group_keys(srcs1, srcs2, on):
43 """Yield the union of distinct ``on``-tuples present in either frame."""
44 if not on:
45 yield None
46 return
47 s1 = set(map(tuple, srcs1[list(on)].itertuples(index=False, name=None)))
48 s2 = set(map(tuple, srcs2[list(on)].itertuples(index=False, name=None)))
49 for key in sorted(s1 | s2):
50 yield key
53def _select_group(df, on, key):
54 """Return the slice of ``df`` whose ``on`` columns equal ``key``."""
55 if key is None or not on:
56 return df
57 mask = pd.Series(True, index=df.index)
58 for col, val in zip(on, key):
59 mask &= df[col] == val
60 return df[mask]
63def match_catalogs(srcs1, srcs2, radius=0.5*u.arcsec, on=("visit", "detector"),
64 ra_col="ra", dec_col="dec", id_col="diaSourceId"):
65 """Spatially cross-match two DiaSource-like DataFrames.
67 Sources are first partitioned by the columns in
68 ``on`` (e.g. matched only within the same (visit, detector)), then matched
69 via nearest-neighbor on the sphere.
71 Parameters
72 ----------
73 srcs1, srcs2 : `pandas.DataFrame`
74 Source tables. Each must contain ``ra_col``, ``dec_col``, ``id_col``,
75 and every column listed in ``on``.
76 radius : `astropy.units.Quantity` or `float`
77 Maximum separation for a pair to count as matched. A bare float is
78 interpreted as arcseconds.
79 on : `tuple` [`str`]
80 Columns to group on before matching. Pass an empty tuple to match the
81 full catalog with no grouping.
82 ra_col, dec_col : `str`
83 Column names for sky coordinates, in degrees.
84 id_col : `str`
85 Column name for the source id in both catalogs.
87 Returns
88 -------
89 matched : `pandas.DataFrame`
90 Rows from ``srcs1`` that found a partner in ``srcs2`` within
91 ``radius``. Two columns are added: ``<id_col>_2`` with the partner's
92 id, and ``xmatch_dist_arcsec`` with the on-sky separation in arcsec.
93 unique1 : `pandas.DataFrame`
94 Rows from ``srcs1`` with no partner.
95 unique2 : `pandas.DataFrame`
96 Rows from ``srcs2`` not pointed at by any matched pair.
97 """
98 rad_arcsec = _as_arcsec(radius)
99 on = tuple(on)
100 id2_col = f"{id_col}_2"
102 matched_chunks = []
103 unique1_chunks = []
104 unique2_chunks = []
106 for key in _group_keys(srcs1, srcs2, on):
107 gs1 = _select_group(srcs1, on, key).copy()
108 gs2 = _select_group(srcs2, on, key).copy()
110 if len(gs1) == 0:
111 unique2_chunks.append(gs2)
112 continue
113 if len(gs2) == 0:
114 unique1_chunks.append(gs1)
115 continue
117 coords1 = SkyCoord(ra=gs1[ra_col].values*u.deg,
118 dec=gs1[dec_col].values*u.deg)
119 coords2 = SkyCoord(ra=gs2[ra_col].values*u.deg,
120 dec=gs2[dec_col].values*u.deg)
121 idx, sep, _ = match_coordinates_sky(coords1, coords2)
123 gs1["xmatch_dist_arcsec"] = sep.to_value(u.arcsec)
124 gs1[id2_col] = gs2[id_col].values[idx]
126 has_match = gs1["xmatch_dist_arcsec"] <= rad_arcsec
127 m = gs1[has_match]
128 u1 = gs1[~has_match].drop(columns=["xmatch_dist_arcsec", id2_col])
129 u2 = gs2[~gs2[id_col].isin(set(m[id2_col]))]
131 matched_chunks.append(m)
132 unique1_chunks.append(u1)
133 unique2_chunks.append(u2)
135 matched = (pd.concat(matched_chunks) if matched_chunks
136 else srcs1.iloc[0:0].assign(**{"xmatch_dist_arcsec": np.nan,
137 id2_col: pd.NA}))
138 unique1 = pd.concat(unique1_chunks) if unique1_chunks else srcs1.iloc[0:0].copy()
139 unique2 = pd.concat(unique2_chunks) if unique2_chunks else srcs2.iloc[0:0].copy()
140 return matched, unique1, unique2
143def flux_residuals(matched, srcs2, flux_col="psfFlux", err_col="psfFluxErr",
144 id_col="diaSourceId", plot=False):
145 """Compute per-pair flux residuals from a `match_catalogs` result.
147 Parameters
148 ----------
149 matched : `pandas.DataFrame`
150 Output of `match_catalogs` (rows from catalog 1 with partner ids).
151 srcs2 : `pandas.DataFrame`
152 Catalog 2, indexed implicitly by ``id_col`` for partner lookup.
153 flux_col, err_col : `str`
154 Column names for the flux and its error in both catalogs.
155 id_col : `str`
156 Source-id column in both catalogs. ``matched`` is assumed to have
157 ``f"{id_col}_2"`` populated by `match_catalogs`.
158 plot : `bool`
159 If True, return a histogram + Q-Q plot in addition to the residuals.
161 Returns
162 -------
163 residuals : `pandas.DataFrame`
164 One row per pair with columns ``flux1``, ``flux2``, ``err1``, ``err2``,
165 ``delta_flux``, ``delta_flux_sigma``, and ``xmatch_dist_arcsec``.
166 fig : `matplotlib.figure.Figure`, optional
167 Only returned when ``plot=True``.
168 """
169 id2_col = f"{id_col}_2"
170 s2 = srcs2.set_index(id_col)
171 partner_ids = matched[id2_col].values
173 f1 = matched[flux_col].to_numpy(dtype=float, copy=True)
174 e1 = matched[err_col].to_numpy(dtype=float, copy=True)
175 f2 = s2.loc[partner_ids, flux_col].to_numpy(dtype=float, copy=True)
176 e2 = s2.loc[partner_ids, err_col].to_numpy(dtype=float, copy=True)
178 delta = f1 - f2
179 sigma = np.sqrt(e1**2 + e2**2)
180 with np.errstate(divide="ignore", invalid="ignore"):
181 delta_sigma = np.where(sigma > 0, delta / sigma, np.nan)
183 residuals = pd.DataFrame({
184 id_col: matched[id_col].values,
185 id2_col: partner_ids,
186 "flux1": f1,
187 "flux2": f2,
188 "err1": e1,
189 "err2": e2,
190 "delta_flux": delta,
191 "delta_flux_sigma": delta_sigma,
192 "xmatch_dist_arcsec": matched["xmatch_dist_arcsec"].values,
193 })
195 if not plot: 195 ↛ 197line 195 didn't jump to line 197 because the condition on line 195 was always true
196 return residuals
197 return residuals, _plot_flux_residuals(residuals, flux_col)
200def _plot_flux_residuals(residuals, flux_col):
201 """Histogram + Q-Q plot of normalized flux residuals."""
202 import matplotlib.pyplot as plt
203 sigma = residuals["delta_flux_sigma"].dropna().values
204 if len(sigma) == 0:
205 fig, ax = plt.subplots()
206 ax.text(0.5, 0.5, "no finite residuals", ha="center", va="center")
207 return fig
209 fig, axes = plt.subplots(1, 2, figsize=(10, 4))
210 axes[0].hist(sigma, bins=50, color="C0")
211 axes[0].axvline(0, color="grey", lw=0.5)
212 axes[0].set_xlabel(rf"$\Delta {flux_col} / \sigma$")
213 axes[0].set_ylabel("count")
214 axes[0].set_title(f"N={len(sigma)}, "
215 f"med={np.median(sigma):.3f}, "
216 f"MAD={np.median(np.abs(sigma - np.median(sigma))):.3f}")
218 # Quick Q-Q vs standard normal without depending on scipy.
219 sp = np.sort(sigma)
220 # Inverse-CDF approximation for the standard normal via erfinv.
221 quantiles = (np.arange(len(sp)) + 0.5) / len(sp)
222 expected = np.sqrt(2) * _erfinv(2*quantiles - 1)
223 axes[1].plot(expected, sp, ".", ms=2)
224 lim = max(abs(expected[0]), abs(expected[-1]), 3.0)
225 axes[1].plot([-lim, lim], [-lim, lim], "k--", lw=0.5)
226 axes[1].set_xlabel("expected (N(0,1))")
227 axes[1].set_ylabel("observed")
228 axes[1].set_title("Q-Q vs standard normal")
229 fig.tight_layout()
230 return fig
233def _erfinv(y):
234 """Approximate inverse error function (vectorized).
236 Uses the formula from Winitzki (2008); accurate to ~4e-3 across the
237 domain, which is plenty for plotting Q-Q lines.
238 """
239 a = 0.147
240 sign = np.sign(y)
241 ln1 = np.log(np.clip(1 - y*y, 1e-300, 1.0))
242 term = 2/(np.pi*a) + ln1/2
243 return sign * np.sqrt(np.sqrt(term*term - ln1/a) - term)
246def match_to_truth(srcs, truth, radius=0.5*u.arcsec,
247 src_ra="ra", src_dec="dec", src_id="diaSourceId",
248 truth_ra="ra", truth_dec="dec", truth_id="injection_id"):
249 """Match a detected catalog against a truth/injection catalog.
251 Performs the match in both directions to compute purity (fraction of
252 detections that correspond to a real injected source) and completeness
253 (fraction of injected sources recovered).
255 Parameters
256 ----------
257 srcs : `pandas.DataFrame`
258 Detected sources.
259 truth : `pandas.DataFrame`
260 Truth catalog, e.g. an injection catalog.
261 radius : `astropy.units.Quantity` or `float`
262 Maximum separation for a match.
263 src_ra, src_dec, src_id : `str`
264 Column names in ``srcs``.
265 truth_ra, truth_dec, truth_id : `str`
266 Column names in ``truth``.
268 Returns
269 -------
270 out : `dict`
271 With keys:
273 - ``"srcs"``: a copy of ``srcs`` with three columns appended:
274 ``"is_real"`` (bool), ``"truth_dist_arcsec"`` (float), and
275 ``f"{truth_id}_match"`` (matched truth id, NA if no match).
276 - ``"truth"``: a copy of ``truth`` with three columns appended:
277 ``"detected"``, ``"detection_dist_arcsec"``, and
278 ``f"{src_id}_match"``.
279 - ``"purity"``: float, fraction of ``srcs`` rows with ``is_real``.
280 - ``"completeness"``: float, fraction of ``truth`` rows with
281 ``detected``.
282 """
283 rad_arcsec = _as_arcsec(radius)
284 srcs_out = srcs.copy()
285 truth_out = truth.copy()
287 truth_id_match = f"{truth_id}_match"
288 src_id_match = f"{src_id}_match"
290 if len(srcs_out) == 0 or len(truth_out) == 0: 290 ↛ 291line 290 didn't jump to line 291 because the condition on line 290 was never true
291 srcs_out["is_real"] = False
292 srcs_out["truth_dist_arcsec"] = np.nan
293 srcs_out[truth_id_match] = pd.NA
294 truth_out["detected"] = False
295 truth_out["detection_dist_arcsec"] = np.nan
296 truth_out[src_id_match] = pd.NA
297 return {"srcs": srcs_out, "truth": truth_out,
298 "purity": 0.0, "completeness": 0.0}
300 sc_src = SkyCoord(srcs_out[src_ra].values*u.deg,
301 srcs_out[src_dec].values*u.deg)
302 sc_tru = SkyCoord(truth_out[truth_ra].values*u.deg,
303 truth_out[truth_dec].values*u.deg)
305 idx_to_truth, sep_to_truth, _ = match_coordinates_sky(sc_src, sc_tru)
306 srcs_out["truth_dist_arcsec"] = sep_to_truth.to_value(u.arcsec)
307 srcs_out[truth_id_match] = truth_out[truth_id].values[idx_to_truth]
308 srcs_out["is_real"] = srcs_out["truth_dist_arcsec"] <= rad_arcsec
309 srcs_out.loc[~srcs_out["is_real"], truth_id_match] = pd.NA
311 idx_to_src, sep_to_src, _ = match_coordinates_sky(sc_tru, sc_src)
312 truth_out["detection_dist_arcsec"] = sep_to_src.to_value(u.arcsec)
313 truth_out[src_id_match] = srcs_out[src_id].values[idx_to_src]
314 truth_out["detected"] = truth_out["detection_dist_arcsec"] <= rad_arcsec
315 truth_out.loc[~truth_out["detected"], src_id_match] = pd.NA
317 purity = float(srcs_out["is_real"].sum() / len(srcs_out))
318 completeness = float(truth_out["detected"].sum() / len(truth_out))
319 return {"srcs": srcs_out, "truth": truth_out,
320 "purity": purity, "completeness": completeness}