Coverage for python/lsst/analysis/ap/compare.py: 74%

131 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-23 03:45 -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/>. 

21 

22"""Catalog cross-matching and pair-wise comparison utilities for 

23DiaSource-like tables. 

24""" 

25 

26__all__ = ["match_catalogs", "flux_residuals", "match_to_truth"] 

27 

28import numpy as np 

29import pandas as pd 

30 

31import astropy.units as u 

32from astropy.coordinates import SkyCoord, match_coordinates_sky 

33 

34 

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) 

40 

41 

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 

51 

52 

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] 

61 

62 

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. 

66 

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. 

70 

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. 

86 

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" 

101 

102 matched_chunks = [] 

103 unique1_chunks = [] 

104 unique2_chunks = [] 

105 

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

109 

110 if len(gs1) == 0: 

111 unique2_chunks.append(gs2) 

112 continue 

113 if len(gs2) == 0: 

114 unique1_chunks.append(gs1) 

115 continue 

116 

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) 

122 

123 gs1["xmatch_dist_arcsec"] = sep.to_value(u.arcsec) 

124 gs1[id2_col] = gs2[id_col].values[idx] 

125 

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

130 

131 matched_chunks.append(m) 

132 unique1_chunks.append(u1) 

133 unique2_chunks.append(u2) 

134 

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 

141 

142 

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. 

146 

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. 

160 

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 

172 

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) 

177 

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) 

182 

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

194 

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) 

198 

199 

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 

208 

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

217 

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 

231 

232 

233def _erfinv(y): 

234 """Approximate inverse error function (vectorized). 

235 

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) 

244 

245 

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. 

250 

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

254 

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

267 

268 Returns 

269 ------- 

270 out : `dict` 

271 With keys: 

272 

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

286 

287 truth_id_match = f"{truth_id}_match" 

288 src_id_match = f"{src_id}_match" 

289 

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} 

299 

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) 

304 

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 

310 

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 

316 

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}