#!/usr/bin/env python """Run Schlafly crowdsource on a single per-exposure destreak_o007_crf file. This matches the photutils iter3 pipeline's per-exposure flow so we can compare apples-to-apples on the same input pixel data. Now with the sky-parameter sparse-matrix shape bug patched (see benchmark/README.md note), nskyx=nskyy=3 is enabled by default — the i2d-style background-subtracted mosaic flow is no longer needed. """ from __future__ import annotations import argparse import sys import time from pathlib import Path import numpy as np from astropy.io import fits from astropy.table import Table from astropy.wcs import WCS sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "shared")) import config # noqa: E402 import crowdsource # noqa: E402 from crowdsource import crowdsource_base # noqa: E402 from astropy.convolution import Gaussian2DKernel, convolve_fft, interpolate_replace_nans # noqa: E402 # JWST DQ bit definitions (stable; from jwst.datamodels.dqflags). DQ_DO_NOT_USE = 1 DQ_SATURATED = 2 DQ_JUMP_DET = 4 DQ_DEAD = 1024 DQ_HOT = 2048 DQ_PERSISTENCE = 16384 # Match production crowdsource_default_kwargs from # crowdsource_catalogs_long.py: maxstars must be huge for crowded fields. CROWDSOURCE_KWARGS = {"maxstars": 500_000} # Bad-DQ bitmask (NIRCam) — drop DO_NOT_USE, dead, hot, jump_det, persistence. # SATURATED is masked separately so its pixels become NaN before interpolation. NIRCAM_BAD_DQ = DQ_DO_NOT_USE | DQ_DEAD | DQ_HOT | DQ_JUMP_DET | DQ_PERSISTENCE def load_gridded_psf(filt: str, detector: str): candidates = list(config.SICKLE_ROOT.glob(f"nircam_{detector}_{filt.lower()}_*.fits")) if not candidates: raise FileNotFoundError( f"no gridded PSF for {filt} {detector} under {config.SICKLE_ROOT}" ) grid_path = candidates[0] from photutils.psf import GriddedPSFModel from astropy.nddata import NDData hdul = fits.open(grid_path) data = hdul[0].data hdr = hdul[0].header grid_xypos = [] n = hdr.get("NPSF", data.shape[0]) for i in range(n): if f"DET_YX{i}" not in hdr: break y_, x_ = eval(hdr[f"DET_YX{i}"]) grid_xypos.append((x_, y_)) if len(grid_xypos) != data.shape[0]: raise ValueError("DET_YX missing") ndd = NDData(data, meta={"oversampling": hdr.get("OVERSAMP", 4), "grid_xypos": grid_xypos}) return GriddedPSFModel(ndd) class WrappedPSFModel(crowdsource.psf.SimplePSF): def __init__(self, psfgridmodel, stampsz=19): self.psfgridmodel = psfgridmodel self.default_stampsz = stampsz def __call__(self, col, row, stampsz=None, deriv=False): if stampsz is None: stampsz = self.default_stampsz parshape = np.broadcast(col, row).shape tparshape = parshape if len(parshape) > 0 else (1,) rows, cols = np.indices((stampsz, stampsz)) - (np.array([stampsz, stampsz]) - 1)[:, None, None] / 2.0 col = np.atleast_1d(col) row = np.atleast_1d(row) stamps = [] for i in range(len(col)): stamps.append(self.psfgridmodel.evaluate(cols + col[i], rows + row[i], 1, col[i], row[i])) stamps = np.array(stamps) stamps /= stamps.sum(axis=(1, 2))[:, None, None] if deriv: dpsfdrow, dpsfdcol = np.gradient(stamps, axis=(1, 2)) ret = stamps if parshape != tparshape: ret = ret.reshape(stampsz, stampsz) if deriv: dpsfdrow = dpsfdrow.reshape(stampsz, stampsz) dpsfdcol = dpsfdcol.reshape(stampsz, stampsz) if deriv: ret = (ret, dpsfdcol, dpsfdrow) return ret def render_model(self, col, row, stampsz=None): if stampsz is not None: self.stampsz = stampsz rows, cols = np.indices(self.stampsz, dtype=float) - (np.array(self.stampsz) - 1)[:, None, None] / 2.0 return self.psfgridmodel.evaluate(cols, rows, 1, col, row).T.squeeze() def detector_token(filename: str) -> str: """Extract NRCB1..4 / NRCBLONG token from JWST filename.""" for d in ("nrcb1", "nrcb2", "nrcb3", "nrcb4", "nrcblong"): if f"_{d}_" in filename: return "nrcb5" if d == "nrcblong" else d raise ValueError(f"can't parse detector from {filename}") def _fwhm_pix_for(filt: str) -> float: """FWHM in pixels — same source as config (filter FWHM in arcsec).""" return config.fwhm_pix(filt) def run_one(crf_path: Path, filt: str, outdir: Path, nskyx=0, nskyy=0): """Per-exposure crowdsource fit, mirroring the preprocessing in `brick2221/analysis/crowdsource_catalogs_long.py`: - NaN/saturated/bad-DQ pixels are flagged - Saturated pixels are set to NaN, then NaNs interpolated with a Gaussian kernel (σ = FWHM/2.355) before passing to fit_im - Mask + DQ flow into fit_im; weight is 1/err with bad pixels zeroed - maxstars = 500_000 to accommodate crowded fields """ print(f"[crowdsource indivexp] {crf_path.name}") hdul = fits.open(crf_path) sci = hdul["SCI"].data.astype(np.float32) err = hdul["ERR"].data.astype(np.float32) if "ERR" in hdul else None dq = hdul["DQ"].data.astype(np.int32) if "DQ" in hdul else np.zeros(sci.shape, dtype=np.int32) wcs = WCS(hdul["SCI"].header) if err is None: sigma = float(np.nanmedian(np.abs(sci - np.nanmedian(sci)))) * 1.4826 err = np.full_like(sci, sigma) # ---- Production-style mask construction ---------------------------- # Identify saturated and other bad-DQ pixels. is_saturated = (dq & DQ_SATURATED) != 0 is_baddq = (dq & NIRCAM_BAD_DQ) != 0 # Inverse-sigma weight; zero out bad pixels. with np.errstate(divide="ignore", invalid="ignore"): weight = np.where(err > 0, 1.0 / err, 0.0) bad_w = ( ~np.isfinite(sci) | (sci == 0) | ~np.isfinite(weight) | (weight == 0) | (err == 0) | is_baddq | is_saturated ) weight[bad_w] = 0.0 # ---- NaN interpolation (production approach) ---------------------- fwhm_pix = _fwhm_pix_for(filt) kernel = Gaussian2DKernel(x_stddev=fwhm_pix / 2.355) data_for_phot = sci.copy() # Saturated pixels become NaN so interpolate_replace_nans fills them # using neighbouring pixels rather than retaining junk values. data_for_phot[is_saturated] = np.nan nan_replaced_data = interpolate_replace_nans(data_for_phot, kernel, convolve=convolve_fft) # ---- Satstar model subtraction (production line ~2664) ----------- # Production reads the iter1 satstar_model.fits and subtracts it before # fit_im so the saturated-star wings don't bias the regular fits. # Path convention: _satstar_model.fits in the same dir. satstar_path = crf_path.parent / (crf_path.stem + "_satstar_model.fits") if satstar_path.exists(): try: sat_img = fits.getdata(str(satstar_path)).astype(float) if sat_img.shape == nan_replaced_data.shape: finite_model = np.where(np.isfinite(sat_img), sat_img, 0.0) nan_replaced_data = nan_replaced_data - finite_model print(f" subtracted satstar_model ({satstar_path.name}) " f"sum={float(np.nansum(finite_model)):.3e}", flush=True) else: print(f" satstar_model shape mismatch; skipping", flush=True) except (OSError, ValueError) as e: print(f" could not read satstar_model: {e}", flush=True) else: print(f" no satstar_model at {satstar_path.name}; skipping", flush=True) # Final mask passed to fit_im as the "bad" multiplicative factor on weight # (production does `weight = weight * (~mask)`); we already applied bad_w # to weight, so this is the union of all bad pixels. mask = bad_w detector = detector_token(crf_path.name) grid = load_gridded_psf(filt, detector) psf = WrappedPSFModel(grid, stampsz=19) t0 = time.time() # Production passes a fresh zeros int array as dq (the JWST DQ ints # overflow crowdsource's int32 indexing in peakfind). Bad-pixel info # is already conveyed via weight=0. crowdsource_dq = np.zeros(sci.shape, dtype="int") result = crowdsource_base.fit_im( nan_replaced_data, psf, weights=weight * (~mask), dq=crowdsource_dq, nskyx=nskyx, nskyy=nskyy, refit_psf=False, verbose=False, threshold=5.0, **CROWDSOURCE_KWARGS, ) dt = time.time() - t0 cat = result["stars"] print(f" done in {dt:.1f}s, {len(cat)} sources") tab = Table(cat) ra, dec = wcs.all_pix2world(tab["x"], tab["y"], 0) tab["ra"] = ra tab["dec"] = dec tab.meta["FILTER"] = filt tab.meta["TOOL"] = "crowdsource" tab.meta["INPUT"] = crf_path.name tab.meta["RUNTIME_S"] = dt tab.meta["NSKYX"] = nskyx outdir.mkdir(parents=True, exist_ok=True) base = crf_path.stem # e.g. jw03958007001_03102_00001_nrcb1_destreak_o007_crf outpath = outdir / f"{base}_crowdsource.fits" tab.write(outpath, overwrite=True) print(f" wrote {outpath}") return outpath def main(): ap = argparse.ArgumentParser() ap.add_argument("--filter", required=True, choices=config.FILTERS) ap.add_argument("--index", type=int, required=True, help="0-based index into the per-filter destreak file list") # crowdsource recommends nsky=0 (no sky-parameter columns); rely on # destreak/JWST-pipeline background already removed at the input level. ap.add_argument("--nskyx", type=int, default=0) ap.add_argument("--nskyy", type=int, default=0) args = ap.parse_args() src_dir = config.SICKLE_ROOT / args.filter / "pipeline" if args.filter in {"F187N", "F210M"}: files = sorted([p for p in src_dir.glob("jw*_destreak_o007_crf.fits") if any(f"_{d}_destreak" in p.name for d in ["nrcb1","nrcb2","nrcb3","nrcb4"])]) else: files = sorted([p for p in src_dir.glob("jw*_nrcblong_destreak_o007_crf.fits")]) if args.index >= len(files): print(f"index {args.index} out of range (have {len(files)} files)", file=sys.stderr) sys.exit(1) crf = files[args.index] outdir = config.BENCHMARK_ROOT / "crowdsource" / "catalogs_indivexp" / args.filter run_one(crf, args.filter, outdir, nskyx=args.nskyx, nskyy=args.nskyy) if __name__ == "__main__": main()