#!/usr/bin/env python """Extended cross-tool comparison v2. Adds: - photutils iter1, iter2, iter3 as three separate reference rows - per-tool variant comparison (default vs indivexp for crowdsource/starbug2) - flux normalisation to make tools' native units comparable in the plots Outputs to writeup/figures/v2/* and writeup/tables/extended_summary_v2.ecsv. """ from __future__ import annotations import sys from pathlib import Path import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt from astropy.coordinates import SkyCoord from astropy.table import Table from astropy import units as u sys.path.insert(0, str(Path(__file__).resolve().parent)) import config # noqa: E402 FIGDIR = config.BENCHMARK_ROOT / "writeup" / "figures" / "v2" TABDIR = config.BENCHMARK_ROOT / "writeup" / "tables" FIGDIR.mkdir(parents=True, exist_ok=True) TABDIR.mkdir(parents=True, exist_ok=True) # Reference tools (photutils iterations). # iter1 is the apples-to-apples primary reference: it is per-frame blind # DAOStarFinder + PSFPhotometry, the same algorithmic level as crowdsource, # starbug2, and DOLPHOT (which all do per-frame detection too). iter2 and # iter3 add custom photutils-pipeline machinery (per-filter seeding and # cross-band union seeding) and are reported here as additional context but # not as the primary "alt-tool vs photutils" comparison. REF_TOOLS = ["photutils_iter1", "photutils_iter2", "photutils_iter3"] PRIMARY_REF = "photutils_iter1" # Alt tools to compare against each photutils iter ALT_TOOLS = ["crowdsource", "starbug2", "dolphot"] def load_bm(tool: str) -> Table | None: p = config.tool_bandmerged(tool) if not p.exists(): print(f"WARN: {p} missing") return None return Table.read(p) def safe(t, col): if col not in t.colnames: return np.full(len(t), np.nan) return np.asarray(t[col]).astype(float) def compare_pair(ref_tool: str, alt_tool: str, ref: Table, alt: Table, summary_rows: list): c_ref = SkyCoord(ref["ra"] * u.deg, ref["dec"] * u.deg) c_oth = SkyCoord(alt["ra"] * u.deg, alt["dec"] * u.deg) idx, sep, _ = c_ref.match_to_catalog_sky(c_oth) match = sep.arcsec < config.MATCH_RADIUS_ARCSEC idx_b, sep_b, _ = c_oth.match_to_catalog_sky(c_ref) match_b = sep_b.arcsec < config.MATCH_RADIUS_ARCSEC for FILT in config.FILTERS: rf = safe(ref, f"{FILT}_flux") re_ = safe(ref, f"{FILT}_flux_err") af_all = safe(alt, f"{FILT}_flux") ae_all = safe(alt, f"{FILT}_flux_err") af = af_all[idx] ae = ae_all[idx] both = match & np.isfinite(rf) & np.isfinite(af) & (rf > 0) & (af > 0) n_both = int(both.sum()) if n_both == 0: summary_rows.append({ "ref": ref_tool, "alt": alt_tool, "filter": FILT, "n_ref_total": int(np.sum(np.isfinite(rf))), "n_alt_total": int(np.sum(np.isfinite(af_all))), "n_matched_both": 0, "median_flux_ratio": np.nan, "log10_scatter_dex": np.nan, "rel_err_median_ref": np.nan, "rel_err_median_alt": np.nan, "n_alt_only": int(np.sum(~match_b & np.isfinite(af_all))), }) continue ratio = af[both] / rf[both] med = float(np.median(ratio)) log_ratio = np.log10(ratio / med) sigma_dex = float(1.4826 * np.median(np.abs(log_ratio))) rel_ref = re_[both] / rf[both] rel_alt = ae[both] / af[both] summary_rows.append({ "ref": ref_tool, "alt": alt_tool, "filter": FILT, "n_ref_total": int(np.sum(np.isfinite(rf))), "n_alt_total": int(np.sum(np.isfinite(af_all))), "n_matched_both": n_both, "median_flux_ratio": med, "log10_scatter_dex": sigma_dex, "rel_err_median_ref": float(np.nanmedian(rel_ref)), "rel_err_median_alt": float(np.nanmedian(rel_alt)), "n_alt_only": int(np.sum(~match_b & np.isfinite(af_all))), }) # Per-pair figure: histogram of normalised log-ratio fig, ax = plt.subplots(figsize=(5, 4)) ax.hist(log_ratio, bins=80, range=(-1, 1), histtype="step") ax.axvline(0.0, color="r", lw=0.8) ax.set_xlabel(r"log10(alt/ref) - log10(median)") ax.set_ylabel("N") ax.set_title(f"{alt_tool} vs {ref_tool}, {FILT} N={n_both} σ={sigma_dex:.3f} dex") fig.tight_layout() fig.savefig(FIGDIR / f"flux_logratio_{alt_tool}_vs_{ref_tool}_{FILT}.png", dpi=120) plt.close(fig) def main(): refs = {tool: load_bm(tool) for tool in REF_TOOLS} alts = {tool: load_bm(tool) for tool in ALT_TOOLS} have_refs = [t for t, v in refs.items() if v is not None] have_alts = [t for t, v in alts.items() if v is not None] print(f"References available: {have_refs}") print(f"Alt tools available: {have_alts}") # Detection counts table (one row per tool) det_rows = [] for tool, t in {**refs, **alts}.items(): if t is None: continue row = {"tool": tool, "n_band_merged": len(t)} for FILT in config.FILTERS: col = f"{FILT}_detected" row[f"{FILT}_n_detected"] = int(np.sum(t[col])) if col in t.colnames else \ int(np.sum(np.isfinite(safe(t, f"{FILT}_flux")))) det_rows.append(row) det_tab = Table(det_rows) det_tab.write(TABDIR / "detection_counts_v2.ecsv", overwrite=True) print("=== Detection counts ===") det_tab.pprint_all() # Pairwise comparison summary_rows = [] for ref_tool in have_refs: ref = refs[ref_tool] for alt_tool in have_alts: alt = alts[alt_tool] print(f"\nComparing {alt_tool} vs {ref_tool} ...") compare_pair(ref_tool, alt_tool, ref, alt, summary_rows) summary = Table(summary_rows) summary.write(TABDIR / "extended_summary_v2.ecsv", overwrite=True) # Compact pivot of (ref, alt, filter) -> log scatter print("\n=== Log10 flux-ratio scatter (dex) ===") fmt = "{:>18s} {:>14s}".format("alt vs ref", "filter") for FILT in config.FILTERS: fmt += " {:>9s}".format(FILT) print(fmt) for ref_tool in have_refs: for alt_tool in have_alts: line = f"{alt_tool:>18s} vs {ref_tool[-5:]:<5s} " for FILT in config.FILTERS: row = next((r for r in summary_rows if r["ref"] == ref_tool and r["alt"] == alt_tool and r["filter"] == FILT), None) line += " {:9.4f}".format(row["log10_scatter_dex"] if row and np.isfinite(row["log10_scatter_dex"]) else np.nan) print(line) # ---- Primary headline table: alt vs iter1 only ---- print() print("=== PRIMARY COMPARISON: alt-tool vs photutils iter1 ===") print(" (iter1 is the per-frame blind-detection level; this is the apples-to-apples") print(" comparison. iter3 numbers above are reported only as context.)") print() primary_rows = [r for r in summary_rows if r["ref"] == PRIMARY_REF] headline = Table() for col in ["alt", "filter", "n_ref_total", "n_alt_total", "n_matched_both", "median_flux_ratio", "log10_scatter_dex", "rel_err_median_ref", "rel_err_median_alt", "n_alt_only"]: headline[col] = [r[col] for r in primary_rows] headline.write(TABDIR / "primary_summary_alt_vs_iter1.ecsv", overwrite=True) headline.pprint_all() # Recovery fraction: matched/ref total per (ref, alt, filter) fig, axes = plt.subplots(1, len(have_refs), figsize=(5 * len(have_refs), 4), sharey=True) if len(have_refs) == 1: axes = [axes] width = 0.25 x = np.arange(len(config.FILTERS)) for ax, ref_tool in zip(axes, have_refs): for i, alt_tool in enumerate(have_alts): rec = [] for FILT in config.FILTERS: row = next((r for r in summary_rows if r["ref"] == ref_tool and r["alt"] == alt_tool and r["filter"] == FILT), None) if row is None or row["n_ref_total"] == 0: rec.append(0.0) else: rec.append(row["n_matched_both"] / row["n_ref_total"]) ax.bar(x + (i - 1) * width, rec, width, label=alt_tool) ax.set_xticks(x) ax.set_xticklabels(config.FILTERS) ax.set_title(f"vs {ref_tool}") ax.set_ylim(0, 1) if ax is axes[0]: ax.set_ylabel("matched fraction") ax.legend() fig.suptitle(f"Recovery fraction (within {config.MATCH_RADIUS_ARCSEC:.2f}″)") fig.tight_layout() fig.savefig(FIGDIR / "recovery_fraction_grid.png", dpi=120) plt.close(fig) print(f"\nWrote {FIGDIR}/recovery_fraction_grid.png") if __name__ == "__main__": main()