diff --git a/libs/ledger/ledger/events.py b/libs/ledger/ledger/events.py index f2592a415..0fa615da5 100644 --- a/libs/ledger/ledger/events.py +++ b/libs/ledger/ledger/events.py @@ -5,10 +5,10 @@ import ligo.segments import numpy as np -from ledger.injections import InterferometerResponseSet +from ledger.injections import InterferometerResponseSet, shift_mask from ledger.ledger import Ledger, metadata, parameter -SECONDS_IN_YEAR = 31556952 +SECONDS_PER_YEAR = 31556952 # 60 * 60 * 24 * 365.2425 F = TypeVar("F", np.ndarray, float) @@ -77,10 +77,7 @@ def get_shift(self, shift: np.ndarray) -> "EventSet": EventSet containing only events from the specified shift. """ # downselect to all events from a given shift - mask = self.shift == shift - if self.shift.ndim == 2: - mask = mask.all(axis=-1) - return self[mask] + return self[shift_mask(self.shift, shift)] def nb(self, threshold: F) -> F: """Calculate number of events above detection threshold. @@ -121,7 +118,7 @@ def min_far(self): Returns: Minimum FAR in yr^-1. """ - return (1 / self.Tb) * SECONDS_IN_YEAR + return (1 / self.Tb) * SECONDS_PER_YEAR def far(self, threshold: F) -> F: """Calculate false alarm rate (FAR) for a given detection threshold. @@ -137,7 +134,7 @@ def far(self, threshold: F) -> F: FAR in yr^-1. Returns min_far if threshold exceeds all events. """ nb = self.nb(threshold) - far = SECONDS_IN_YEAR * nb / self.Tb + far = SECONDS_PER_YEAR * nb / self.Tb return np.maximum(far, self.min_far) def significance(self, threshold: F, T: float) -> F: diff --git a/libs/ledger/ledger/injections.py b/libs/ledger/ledger/injections.py index 9aeeebb53..67ed13216 100644 --- a/libs/ledger/ledger/injections.py +++ b/libs/ledger/ledger/injections.py @@ -18,6 +18,14 @@ MSUN = 1.988409902147041637325262574352366540e30 +def shift_mask(shifts: np.ndarray, target) -> np.ndarray: + """Boolean mask of rows in `shifts` equal to `target`.""" + mask = shifts == np.asarray(target) + if shifts.ndim == 2: + mask = mask.all(axis=-1) + return mask + + def chirp_mass( m1: float | np.ndarray, m2: float | np.ndarray ) -> float | np.ndarray: @@ -764,10 +772,7 @@ def get_shift(self, shift): Returns: InjectionParameterSet with injections with the specified shift. """ - mask = self.shift == shift - if self.shift.ndim == 2: - mask = mask.all(axis=-1) - return self[mask] + return self[shift_mask(self.shift, shift)] def get_times(self, start: float | None = None, end: float | None = None): """ diff --git a/libs/ledger/ledger/ledger.py b/libs/ledger/ledger/ledger.py index 4b9f6458e..966920c16 100644 --- a/libs/ledger/ledger/ledger.py +++ b/libs/ledger/ledger/ledger.py @@ -258,6 +258,20 @@ def read(cls, fname: PATH): with h5py.File(fname, "r") as f: return cls._load_with_idx(f, None) + @classmethod + def read_idx(cls, fname: PATH, idx: np.ndarray): + """Read only the rows at `idx` from `fname`. + + Args: + fname: Path to HDF5 file to read. + idx: Row indices to load. + + Returns: + Instance of the ledger class populated with just those rows. + """ + with h5py.File(fname, "r") as f: + return cls._load_with_idx(f, idx) + @classmethod def sample_from_file(cls, fname: PATH, N: int, replace: bool = False): """Sample data from HDF5 file for out-of-memory operations. diff --git a/libs/ledger/tests/test_events.py b/libs/ledger/tests/test_events.py index b0791b061..04d583453 100644 --- a/libs/ledger/tests/test_events.py +++ b/libs/ledger/tests/test_events.py @@ -69,7 +69,7 @@ def test_far(self): det_stats = np.arange(10) times = np.arange(10) shifts = np.array([0] * 4 + [1] * 5 + [2]) - Tb = 2 * events.SECONDS_IN_YEAR + Tb = 2 * events.SECONDS_PER_YEAR obj = events.EventSet(det_stats, times, shifts, Tb) assert obj.far(5) == 2.5 diff --git a/libs/p_astro/p_astro/background.py b/libs/p_astro/p_astro/background.py index 36c1706e4..25ad0da06 100644 --- a/libs/p_astro/p_astro/background.py +++ b/libs/p_astro/p_astro/background.py @@ -1,5 +1,5 @@ import numpy as np -from ledger.events import SECONDS_IN_YEAR, EventSet +from ledger.events import SECONDS_PER_YEAR, EventSet from numpy.polynomial import Polynomial from scipy.stats import gaussian_kde @@ -56,7 +56,7 @@ def __init__(self, *args, split: float | None = None, **kwargs): @property def scale_factor(self): """Scale factor to convert from rate density to count rate""" - return len(self.background) * SECONDS_IN_YEAR / self.background.Tb + return len(self.background) * SECONDS_PER_YEAR / self.background.Tb def fit(self): """ diff --git a/libs/p_astro/uv.lock b/libs/p_astro/uv.lock index 7567c64df..b9890df26 100644 --- a/libs/p_astro/uv.lock +++ b/libs/p_astro/uv.lock @@ -1355,6 +1355,7 @@ version = "0.1.0" source = { editable = "../ledger" } dependencies = [ { name = "h5py" }, + { name = "ligo-segments" }, { name = "numpy" }, { name = "pycbc" }, { name = "utils" }, @@ -1363,6 +1364,7 @@ dependencies = [ [package.metadata] requires-dist = [ { name = "h5py", specifier = ">=3.10.0,<4" }, + { name = "ligo-segments", specifier = ">=1.4.0,<2" }, { name = "numpy", specifier = ">=2" }, { name = "pycbc", specifier = ">=2.5.1,<3" }, { name = "utils", editable = "../utils" }, @@ -1525,7 +1527,7 @@ wheels = [ [[package]] name = "ml4gw" -version = "0.7.11" +version = "0.8.3" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "jaxtyping" }, @@ -1534,9 +1536,9 @@ dependencies = [ { name = "torch" }, { name = "torchaudio" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/6f/0a/722f553635ffc91b32623e69a4c93591c11ce2c24a10e4bda35ab0d8e6ae/ml4gw-0.7.11.tar.gz", hash = "sha256:8df9ebecd97ed6a6e8ba07fab40882f5966e646897f5187a9ccf7913faf6464e", size = 119593, upload-time = "2026-01-29T20:34:30.794Z" } +sdist = { url = "https://files.pythonhosted.org/packages/4a/1b/78a1d86e3253e3e8626b8079152caf92fb502c58ca8942d293426ad71139/ml4gw-0.8.3.tar.gz", hash = "sha256:d34aadd5d977498c3ac8922664a33874b1e5f5a29079033f0345683f7f9d1868", size = 124777, upload-time = "2026-06-24T17:16:52.947Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/89/7d/f8c3e695d52cd9e70fd3f7bb51efd29848a3eb481dc1b94228f481dd05f8/ml4gw-0.7.11-py3-none-any.whl", hash = "sha256:0a6645f27444d266fb94afe988450bc2d00e24bd70328b0a5903194e1900acdb", size = 129588, upload-time = "2026-01-29T20:34:29.357Z" }, + { url = "https://files.pythonhosted.org/packages/dd/14/3e96be1b039d7476e5ce32b701df31dc33651b53db4230e9ac38fbde003d/ml4gw-0.8.3-py3-none-any.whl", hash = "sha256:4601d2034b19b4e485c7f71a97983ea2dae9d407a8e2c126c1164c337c15462b", size = 136438, upload-time = "2026-06-24T17:16:51.813Z" }, ] [[package]] @@ -2920,7 +2922,7 @@ dependencies = [ requires-dist = [ { name = "astropy", specifier = ">=6.0.1" }, { name = "h5py", specifier = "~=3.6" }, - { name = "ml4gw", specifier = ">=0.7.10" }, + { name = "ml4gw", specifier = ">=0.8.0" }, { name = "numpy", specifier = ">=2" }, { name = "s3fs", specifier = ">=2024,<2025" }, ] diff --git a/projects/plots/plots/core/compute.py b/projects/plots/plots/core/compute.py index 621685365..5d1b44bfb 100644 --- a/projects/plots/plots/core/compute.py +++ b/projects/plots/plots/core/compute.py @@ -1,43 +1,38 @@ -from concurrent.futures import ProcessPoolExecutor - import numpy as np -from tqdm import trange - - -def init_fn(det_stats, w): - global detection_statistics, weights - detection_statistics = det_stats - weights = w - - -def compute_sv(threshold): - mask = detection_statistics >= threshold - mus = (weights * mask).sum(-1, keepdims=True) - var_summands = weights * (mask - mus) - stds = (var_summands**2).sum(-1) ** 0.5 - return mus[:, 0], stds def sensitive_volume(detection_statistics, weights, thresholds): - y = np.empty((len(weights), len(thresholds))) - err = np.empty((len(weights), len(thresholds))) - ex = ProcessPoolExecutor( - 8, initializer=init_fn, initargs=(detection_statistics, weights) - ) - with ex: - fs = {ex.submit(compute_sv, t): i for i, t in enumerate(thresholds)} - for _ in trange(len(thresholds)): - while True: - for future in fs: - if future.done(): - future.exception() - break - else: - continue - break - - i = fs.pop(future) - mu, std = future.result() - y[:, i] = mu - err[:, i] = std - return y, err + """Mean and standard error of sensitive volume at each threshold. + + Args: + detection_statistics: `(N,)` array of injection detection stats. + weights: `(C, N)` array of per-injection weights, one row per + mass combo. + thresholds: `(T,)` array of detection statistic thresholds. + + Returns: + `(y, err)`, each `(C, T)`. + """ + order = np.argsort(detection_statistics) + ds_sorted = detection_statistics[order] + weights_sorted = weights[:, order] + + # Compute the reverse cumulative sums of the weights and the + # weights squared. Because the weights are sorted by the + # detection statistics, this avoids needing to compute a mask + # for each threshold. + n = len(ds_sorted) + rev_sum_w = np.zeros((weights.shape[0], n + 1)) + rev_sum_w2 = np.zeros_like(rev_sum_w) + rev_sum_w[:, :-1] = np.cumsum(weights_sorted[:, ::-1], axis=-1)[:, ::-1] + rev_sum_w2[:, :-1] = np.cumsum(weights_sorted[:, ::-1] ** 2, axis=-1)[ + :, ::-1 + ] + + # idxs is the indices of the first detection statistic greater than + # each threshold. + idxs = np.searchsorted(ds_sorted, thresholds) + mu = rev_sum_w[:, idxs] + var = (1 - 2 * mu) * rev_sum_w2[:, idxs] + mu**2 * rev_sum_w2[:, :1] + err = np.sqrt(np.maximum(var, 0)) + return mu, err diff --git a/projects/plots/plots/core/constants.py b/projects/plots/plots/core/constants.py index 635bb1bff..1e9628ddc 100644 --- a/projects/plots/plots/core/constants.py +++ b/projects/plots/plots/core/constants.py @@ -1,2 +1,5 @@ +from ledger.events import SECONDS_PER_YEAR + SECONDS_PER_MONTH = 60 * 60 * 24 * 30 -SECONDS_PER_YEAR = 60 * 60 * 24 * 365.25 +DEFAULT_NUM_FAR_POINTS = 500 +DEFAULT_VIZAPP_MAX_FAR = 100 * SECONDS_PER_YEAR / SECONDS_PER_MONTH diff --git a/projects/plots/plots/core/data.py b/projects/plots/plots/core/data.py index 7c6bb6a0d..28291ba0d 100644 --- a/projects/plots/plots/core/data.py +++ b/projects/plots/plots/core/data.py @@ -44,7 +44,8 @@ def load( """Read the ledgers and drop unphysical foreground events. Args: - background: HDF5 file readable by `EventSet.read` + background: HDF5 file readable by `EventSet.read`. Gets sorted by + detection statistic if it isn't already. foreground: HDF5 file readable by `RecoveredInjectionSet.read` rejected: HDF5 file readable by `InjectionParameterSet.read` """ @@ -54,6 +55,11 @@ def load( foreground=drop_unphysical(RecoveredInjectionSet.read(foreground)), rejected=InjectionParameterSet.read(rejected), ) + + # Background should be sorted by infer already, but just in case + if not data.background.is_sorted_by("detection_statistic"): + data.background = data.background.sort_by("detection_statistic") + logging.info("Read in:") logging.info(f"\t{len(data.background)} background events") logging.info(f"\t{len(data.foreground)} foreground events") diff --git a/projects/plots/plots/core/gwtc3.py b/projects/plots/plots/core/gwtc3.py index 11a3765bf..61d535ed6 100644 --- a/projects/plots/plots/core/gwtc3.py +++ b/projects/plots/plots/core/gwtc3.py @@ -7,7 +7,6 @@ import numpy as np import scipy.stats as stats from astropy.utils.data import download_file -from tqdm import tqdm from utils.cosmology import DEFAULT_COSMOLOGY catalog_results = { @@ -212,31 +211,70 @@ def logdiffexp(x, y): return x + np.log1p(-np.exp(y - x)) -# Changed this function to take log_dN as an argument so that -# all of them can be calculated up front -def get_logVT(log_dN, selection, T_obs, N_draw, p_draw): - """Convienient function that returns log_VT, log_sigma_VT, and N_eff""" +def _cumulative_logsumexp(x: np.ndarray, reverse: bool) -> np.ndarray: + """`log(sum(exp(x)))` accumulated forward or in reverse.""" + n = len(x) + out = np.full(n + 1, -np.inf) + if n == 0: + return out + with np.errstate(divide="ignore"): + if not reverse: + out[1:] = np.log(np.cumsum(np.exp(x))) + else: + out[:n] = np.log(np.cumsum(np.exp(x)[::-1]))[::-1] + return out + + +def _vectorized_pipeline_sv( + det_stat_p: np.ndarray, + log_dNs: list[np.ndarray], + mass_combos: list[tuple], + p_draw: np.ndarray, + T_obs: float, + N_draw: float, + thresholds: np.ndarray, + criterion: str, +) -> tuple[dict, dict]: + """Vectorized SV/err for every mass combo of one pipeline. + + `far` selects a growing forward sum (`det_stat < thresh`, ascending + sort); `pastro` selects a shrinking reverse sum (`det_stat > + thresh`), so it needs `side="right"` and `reverse=True`. + """ + order = np.argsort(det_stat_p) + ds_sorted = det_stat_p[order] + n = len(ds_sorted) + reverse = criterion == "pastro" + + if criterion == "far": + idxs = np.searchsorted(ds_sorted, thresholds, side="left") + k = idxs + else: + idxs = np.searchsorted(ds_sorted, thresholds, side="right") + k = n - idxs - # Calculate VT - log_dN = log_dN[selection] - p_draw = p_draw[selection] - log_VT = ( - np.log(T_obs) - - np.log(N_draw) - + np.logaddexp.reduce(log_dN - np.log(p_draw)) - ) + log_T = np.log(T_obs) + log_N = np.log(N_draw) + log_p_draw_sorted = np.log(p_draw[order]) - # Calculate uncertainty of VT and effective number - log_s2 = ( - 2 * np.log(T_obs) - - 2 * np.log(N_draw) - + np.logaddexp.reduce(2 * (log_dN - np.log(p_draw))) - ) - log_sig2 = logdiffexp(log_s2, 2.0 * log_VT - np.log(N_draw)) - log_sig = log_sig2 / 2 - N_eff = np.exp(2 * log_VT - log_sig2) + sv, err = {}, {} + for (m1, m2), log_dN in zip(mass_combos, log_dNs, strict=True): + key = f"{m1}-{m2}" + x = log_dN[order] - log_p_draw_sorted + csum_x = _cumulative_logsumexp(x, reverse=reverse) + csum_2x = _cumulative_logsumexp(2 * x, reverse=reverse) - return log_VT, log_sig, N_eff + log_VT = log_T - log_N + csum_x[idxs] + log_s2 = 2 * log_T - 2 * log_N + csum_2x[idxs] + + with np.errstate(invalid="ignore"): + log_sig2 = logdiffexp(log_s2, 2 * log_VT - log_N) + log_sig2 = np.where(k == 0, -np.inf, log_sig2) + + sv[key] = np.exp(log_VT) / T_obs + err[key] = np.exp(log_sig2 / 2) / T_obs + + return sv, err def get_logdNs( @@ -270,6 +308,28 @@ def get_logdNs( return log_dNs +def _write_result_file( + output_dir: Path, + detection_criterion: str, + detection_thresholds: np.ndarray, + pipelines: list[str], + mass_combos: list[tuple], + sv: dict, + err: dict, +) -> None: + """Write the `gwtc-3_pipeline_sv.hdf5` data file.""" + output_dir.mkdir(parents=True, exist_ok=True) + outfile = output_dir / "gwtc-3_pipeline_sv.hdf5" + with h5py.File(outfile, "w") as f: + f.create_dataset(f"{detection_criterion}", data=detection_thresholds) + for p in pipelines: + g = f.create_group(p) + for m1, m2 in mass_combos: + h = g.create_group(f"{m1}-{m2}") + h.create_dataset("sv", data=np.array(sv[p][f"{m1}-{m2}"])) + h.create_dataset("err", data=np.array(err[p][f"{m1}-{m2}"])) + + def main( mass_combos: list[float], detection_criterion: str, @@ -315,33 +375,24 @@ def main( sv, err = {}, {} for p in pipelines: logging.info(f"Calculating SV for {p}") - sv[p], err[p] = {}, {} - for (m1, m2), log_dN in zip((mass_combos), log_dNs, strict=True): - sv[p][f"{m1}-{m2}"] = np.zeros_like(detection_thresholds) - err[p][f"{m1}-{m2}"] = np.zeros_like(detection_thresholds) - for i, thresh in enumerate(tqdm(detection_thresholds)): - if detection_criterion == "far": - selection = det_stat[p] < thresh - else: - selection = det_stat[p] > thresh - - log_vt, log_sigma_vt, _ = get_logVT( - log_dN, selection, T_obs, N_draw, p_draw - ) - vt = np.exp(log_vt) - sigma_vt = np.exp(log_sigma_vt) - - sv[p][f"{m1}-{m2}"][i] = vt / T_obs - err[p][f"{m1}-{m2}"][i] = sigma_vt / T_obs - - outfile = output_dir / "gwtc-3_pipeline_sv.hdf5" - with h5py.File(outfile, "w") as f: - f.create_dataset(f"{detection_criterion}", data=detection_thresholds) - for p in pipelines: - g = f.create_group(p) - for m1, m2 in mass_combos: - h = g.create_group(f"{m1}-{m2}") - h.create_dataset("sv", data=np.array(sv[p][f"{m1}-{m2}"])) - h.create_dataset("err", data=np.array(err[p][f"{m1}-{m2}"])) + sv[p], err[p] = _vectorized_pipeline_sv( + det_stat[p], + log_dNs, + mass_combos, + p_draw, + T_obs, + N_draw, + detection_thresholds, + detection_criterion, + ) + _write_result_file( + output_dir, + detection_criterion, + detection_thresholds, + pipelines, + mass_combos, + sv, + err, + ) return sv, err diff --git a/projects/plots/plots/core/sv.py b/projects/plots/plots/core/sv.py index acc051485..9f4bd0ae8 100644 --- a/projects/plots/plots/core/sv.py +++ b/projects/plots/plots/core/sv.py @@ -13,7 +13,7 @@ from utils.cosmology import DEFAULT_COSMOLOGY, get_astrophysical_volume from plots.core import compute, style -from plots.core.constants import SECONDS_PER_YEAR +from plots.core.constants import DEFAULT_NUM_FAR_POINTS, SECONDS_PER_YEAR from plots.core.data import AnalysisData if TYPE_CHECKING: @@ -81,14 +81,17 @@ def read(cls, path: Path) -> "SensitiveVolumeResult": ) -def _far_grid(background, max_far: float): +def _far_grid(background, max_far: float, num_far_points: int): """Build the FAR grid and the thresholds that produce it.""" Tb = background.Tb / SECONDS_PER_YEAR max_events = min(int(max_far * Tb), len(background)) if not max_events: return np.array([]), np.array([]) - fars = np.arange(1, max_events + 1) / Tb - thresholds = np.sort(background.detection_statistic)[::-1][:max_events] + counts = np.unique(np.round(np.geomspace(1, max_events, num_far_points))) + counts = np.clip(counts, 1, max_events) + fars = counts / Tb + stats = background.detection_statistic + thresholds = stats[len(stats) - counts.astype(int)] return fars, thresholds @@ -158,6 +161,7 @@ def compute_sensitive_volume( source_prior: "PriorDict", dt: float | None = None, max_far: float = 365, + num_far_points: int = DEFAULT_NUM_FAR_POINTS, sigma: float = 0.1, ) -> SensitiveVolumeResult: """Compute sensitive volume vs. false alarm rate. @@ -170,10 +174,11 @@ def compute_sensitive_volume( dt: if given, discard injections recovered more than `dt` seconds from their injection time max_far: largest false alarm rate to compute out to, per year + num_far_points: number of points in the FAR grid sigma: width of the log normal mass distributions """ v0 = _astrophysical_volume(source_prior) - fars, thresholds = _far_grid(data.background, max_far) + fars, thresholds = _far_grid(data.background, max_far, num_far_points) weights = _weights(data, mass_combos, source_prior, dt, sigma) logging.info("Computing sensitive volume at thresholds") diff --git a/projects/plots/plots/main.py b/projects/plots/plots/main.py index 1e553e153..982cc9da0 100644 --- a/projects/plots/plots/main.py +++ b/projects/plots/plots/main.py @@ -5,6 +5,7 @@ from utils.cosmology import DEFAULT_COSMOLOGY from utils.logging import configure_logging +from plots.core.constants import DEFAULT_NUM_FAR_POINTS from plots.core.data import AnalysisData from plots.core.gwtc3 import main as gwtc3_pipeline_sv from plots.core.sv import ( @@ -47,6 +48,7 @@ def sensitive_volume( log_file: Path | None = None, dt: float | None = None, max_far: float = 365, + num_far_points: int = DEFAULT_NUM_FAR_POINTS, sigma: float = 0.1, verbose: bool = False, vetos: list[VETO_CATEGORIES] | None = None, @@ -76,6 +78,9 @@ def sensitive_volume( max_far: The maximum FAR to compute the sensitive volume out to in units of years^-1 + num_far_points: + Number of points in the FAR grid to compute the sensitive + volume at sigma: The width of the log normal mass distribution to use verbose: @@ -113,6 +118,7 @@ def sensitive_volume( source_prior=source, dt=dt, max_far=max_far, + num_far_points=num_far_points, sigma=sigma, ) diff --git a/projects/plots/plots/vizapp/app.py b/projects/plots/plots/vizapp/app.py index 4edfe195f..b506ec874 100644 --- a/projects/plots/plots/vizapp/app.py +++ b/projects/plots/plots/vizapp/app.py @@ -9,6 +9,7 @@ from utils.logging import configure_logging from utils.s3 import open_file +from plots.core.constants import DEFAULT_NUM_FAR_POINTS from plots.vetos import VETO_CATEGORIES from plots.vizapp.data import DataManager from plots.vizapp.pages import Analysis, Summary @@ -44,6 +45,7 @@ def __init__( fftlength: float, device: str = "cpu", vetos: list[VETO_CATEGORIES] | None = None, + num_far_points: int = DEFAULT_NUM_FAR_POINTS, verbose: bool = False, ) -> None: configure_logging(verbose=verbose) @@ -68,6 +70,7 @@ def __init__( self.fduration = fduration self.valid_frac = valid_frac self.device = device + self.num_far_points = num_far_points self.verbose = verbose self.weights = weights self.device = device @@ -82,6 +85,10 @@ def __init__( results_dir, waveforms_dir, ifos, vetos ) + self.background, self.foreground = self.data_manager.update_vetos( + None, None, [] + ) + # initialize all our pages and their constituent plots self.pages: list[Page] = [] tabs = [] @@ -96,7 +103,6 @@ def __init__( self.veto_selecter = self.data_manager.get_veto_selecter() self.veto_selecter.on_change("value", self.update) - self.update(None, None, []) # set up a header with a title and the selecter title = Div(text="

aframe Performance Dashboard

", width=500) @@ -122,11 +128,13 @@ def load_model(self): def update(self, attr, old, new): # update the vetos - background, foreground = self.data_manager.update_vetos(attr, old, new) + self.background, self.foreground = self.data_manager.update_vetos( + attr, old, new + ) # update pages with latest background and foreground for page in self.pages: - page.update(background, foreground) + page.update(self.background, self.foreground) def __call__(self, doc): doc.add_root(self.layout) diff --git a/projects/plots/plots/vizapp/data.py b/projects/plots/plots/vizapp/data.py index 39dce9a23..08dcd5a76 100644 --- a/projects/plots/plots/vizapp/data.py +++ b/projects/plots/plots/vizapp/data.py @@ -1,5 +1,4 @@ import logging -from copy import deepcopy from pathlib import Path from bokeh.models import MultiChoice @@ -41,24 +40,19 @@ def __init__( self.rejected_params = data.rejected self.logger.info("Data loaded") - # create copies of the background and foreground - # for applying vetos - self._background = deepcopy(self.background) - self._foreground = deepcopy(self.foreground) - self.background_masks = None self.foreground_masks = None if self.categories: - start = self._background.detection_time.min() - stop = self._background.detection_time.max() + start = self.background.detection_time.min() + stop = self.background.detection_time.max() segments = load_or_fetch_segments( self.categories, self.ifos, start, stop ) self.background_masks = compute_veto_masks( - self._background, self.categories, self.ifos, segments + self.background, self.categories, self.ifos, segments ) self.foreground_masks = compute_veto_masks( - self._foreground, self.categories, self.ifos, segments + self.foreground, self.categories, self.ifos, segments ) def get_veto_selecter(self): @@ -67,10 +61,10 @@ def get_veto_selecter(self): def update_vetos(self, attr, old, new): if not self.categories: - return self._background, self._foreground + return self.background, self.foreground back_mask = combine_masks(self.background_masks, new) fore_mask = combine_masks(self.foreground_masks, new) - background = self._background[~back_mask] - foreground = self._foreground[~fore_mask] + background = self.background[~back_mask] + foreground = self.foreground[~fore_mask] return background, foreground diff --git a/projects/plots/plots/vizapp/infer/analyzer.py b/projects/plots/plots/vizapp/infer/analyzer.py index f389ef781..6da645ca7 100644 --- a/projects/plots/plots/vizapp/infer/analyzer.py +++ b/projects/plots/plots/vizapp/infer/analyzer.py @@ -1,11 +1,16 @@ from collections.abc import Sequence +from functools import cached_property from pathlib import Path import h5py import numpy as np import torch from gwpy.timeseries import TimeSeries -from ledger.injections import InterferometerResponseSet, waveform_class_factory +from ledger.injections import ( + InterferometerResponseSet, + shift_mask, + waveform_class_factory, +) from utils.preprocessing import BackgroundSnapshotter, BatchWhitener from plots.vizapp.infer.utils import get_indices, get_strain_fname @@ -149,15 +154,30 @@ def find_strain(self, time: float, shifts: Sequence[float]): return torch.stack(strain, axis=0), time + self.times[0] + @cached_property + def _injection_index(self): + """`(injection_time, shift, left_pad)` of the waveform file. + + Read once and reused. + """ + with h5py.File(self.response_set, "r") as f: + times = f["parameters"]["injection_time"][:] + shifts = f["parameters"]["shift"][:] + left_pad = f.attrs["duration"] - f.attrs["right_pad"] + return times, shifts, left_pad + def find_waveform(self, time: float, shifts: np.ndarray): """ - find the closest injection that corresponds to event + Find the closest injection that corresponds to event time and shifts from waveform dataset """ - waveform = self.waveform_class.read( - self.response_set, time - 0.1, time + 0.1, shifts - ) - return waveform + times, all_shifts, left_pad = self._injection_index + mask = (times + left_pad) >= (time - 0.1) + mask &= (times - left_pad) <= (time + 0.1) + mask &= shift_mask(all_shifts, shifts) + + idx = np.where(mask)[0] + return self.waveform_class.read_idx(self.response_set, idx) def integrate(self, y): integrated = np.convolve(y, self.window, mode="full") diff --git a/projects/plots/plots/vizapp/pages/analysis/distribution.py b/projects/plots/plots/vizapp/pages/analysis/distribution.py index b2148ae5e..2efa31496 100644 --- a/projects/plots/plots/vizapp/pages/analysis/distribution.py +++ b/projects/plots/plots/vizapp/pages/analysis/distribution.py @@ -31,10 +31,12 @@ "injection_time", "chirp_mass", ] -BACK_ATTRS = ["detection_statistic", "detection_time"] SLIDER_ATTRS = ["mass_1_source", "mass_2_source", "snr"] +MAX_SCATTER_POINTS = 5_000 +MAX_FOREGROUND_POINTS = 20_000 + class DistributionPlot: def __init__(self, page, event_inspector) -> None: @@ -43,15 +45,16 @@ def __init__(self, page, event_inspector) -> None: self.bckgd_color = palette[4] self.frgd_color = palette[2] - def asdict(self, background, foreground): - background = {attr: getattr(background, attr) for attr in BACK_ATTRS} - _foreground = {attr: getattr(foreground, attr) for attr in FORE_ATTRS} - + def asdict(self, foreground, idx=slice(None)): + _foreground = { + attr: getattr(foreground, attr)[idx] for attr in FORE_ATTRS + } + ifo_snrs = foreground.ifo_snrs[idx] for i, ifo in enumerate(foreground.ifos): - _foreground[f"{ifo}_snr"] = foreground.ifo_snrs[:, i] - sorted_snrs = np.sort(foreground.ifo_snrs, axis=-1) + _foreground[f"{ifo}_snr"] = ifo_snrs[:, i] + sorted_snrs = np.sort(ifo_snrs, axis=-1) _foreground["snr_ratio"] = sorted_snrs[:, -1] / sorted_snrs[:, -2] - return background, _foreground + return _foreground def initialize_sources(self): self.bar_source = ColumnDataSource( @@ -295,55 +298,86 @@ def update_background(self, attr, old, new): stats = np.array(self.bar_source.data["center"]) min_ = min([stats[i] for i in new]) max_ = max([stats[i] for i in new]) - mask = self.background.detection_statistic >= min_ - mask &= self.background.detection_statistic <= max_ - self.background_plot.title.text = ( - f"{mask.sum()} events with detection statistic in the range" - f"({min_:0.1f}, {max_:0.1f})" - ) - events = self.background.detection_statistic[mask] - times = self.background.detection_time[mask] - shifts = self.background.shift[mask] + ds = self.background.detection_statistic + low = np.searchsorted(ds, min_, side="left") + high = np.searchsorted(ds, max_, side="right") + n_selected = high - low + + if n_selected == 0: + self.background_plot.title.text = ( + f"0 events with detection statistic in " + f"({min_:0.1f}, {max_:0.1f})" + ) + self.background_source.data = { + "x": [], + "detection_time": [], + "detection_statistic": [], + "shifts": [], + "size": [], + } + self.background_source.selected.indices = [] + return + + if n_selected > MAX_SCATTER_POINTS: + rng = np.random.default_rng() + sample = np.sort( + rng.choice(n_selected, size=MAX_SCATTER_POINTS, replace=False) + ) + idx = low + sample + self.background_plot.title.text = ( + f"showing {MAX_SCATTER_POINTS:,} of {n_selected:,} events " + f"with detection statistic in ({min_:0.1f}, {max_:0.1f})" + ) + else: + idx = np.arange(low, high) + self.background_plot.title.text = ( + f"{n_selected} events with detection statistic in " + f"({min_:0.1f}, {max_:0.1f})" + ) + + events = ds[idx] + times = self.background.detection_time[idx] + shifts = self.background.shift[idx] t0 = times.min() self.background_plot.xaxis.axis_label = f"Time from {t0:0.3f} [hours]" - x = times - t0 - x /= 3600 + x = (times - t0) / 3600 - self.background_source.data.update( - { - "x": x - + shifts.sum( - axis=-1 - ), # give unique time to events at same H1 time - "detection_time": times, - "detection_statistic": events, - "shifts": shifts, - "size": np.ones(len(events)) * 8, - } - ) + self.background_source.data = { + "x": x + shifts.sum(axis=-1) / 3600, + "detection_time": times, + "detection_statistic": events, + "shifts": shifts, + "size": np.full(len(events), 8), + } self.background_source.selected.indices = [] def update(self, background, foreground): self.background = background self.foreground = foreground + n_fg = len(foreground) + if n_fg > MAX_FOREGROUND_POINTS: + rng = np.random.default_rng() + idx = np.sort( + rng.choice(n_fg, size=MAX_FOREGROUND_POINTS, replace=False) + ) + fg_note = f" (showing {MAX_FOREGROUND_POINTS:,} of {n_fg:,})" + else: + idx = slice(None) + fg_note = "" + title = ( f"{len(self.background)} background events from " f"{self.background.Tb / 3600 / 24:0.2f} days worth " - f"of data; {len(self.foreground)} injections overlayed" - ) - background_dict, foreground_dict = self.asdict( - self.background, self.foreground + f"of data; {n_fg} injections overlayed{fg_note}" ) - - self.background_source.data = background_dict - self.foreground_source.data = foreground_dict + self.foreground_source.data = self.asdict(self.foreground, idx) for attr, slider in self.sliders.items(): - values = foreground_dict[attr] + values = getattr(self.foreground, attr) if not len(values) > 0: continue low, high = float(np.min(values)), float(np.max(values)) @@ -354,18 +388,16 @@ def update(self, background, foreground): self.distribution_plot.title.text = title - # update bar plot - hist, bins = np.histogram( - self.background.detection_statistic, bins=100 - ) - hist = np.cumsum(hist[::-1])[::-1] + ds = self.background.detection_statistic + edges = np.histogram_bin_edges(ds, bins=100) + top = len(ds) - np.searchsorted(ds, edges[:-1], side="left") self.distribution_plot.y_range.start = 0.1 - self.distribution_plot.y_range.end = 2 * hist.max() + self.distribution_plot.y_range.end = 2 * top.max() if len(top) else 1 self.bar_source.data.update( - center=(bins[:-1] + bins[1:]) / 2, - top=hist, - width=0.95 * (bins[1:] - bins[:-1]), + center=(edges[:-1] + edges[1:]) / 2, + top=top, + width=0.95 * (edges[1:] - edges[:-1]), ) # update snr axis of plot @@ -379,15 +411,13 @@ def update(self, background, foreground): # clear the background plot until we select another # range of detection characteristics to plot - self.background_source.data.update( - { - "x": [], - "detection_time": [], - "detection_statistic": [], - "shifts": [], - "size": [], - } - ) + self.background_source.data = { + "x": [], + "detection_time": [], + "detection_statistic": [], + "shifts": [], + "size": [], + } self.bar_source.selected.indices = [] self.foreground_source.selected.indices = [] self.background_source.selected.indices = [] diff --git a/projects/plots/plots/vizapp/pages/analysis/page.py b/projects/plots/plots/vizapp/pages/analysis/page.py index 4f08769d9..2f69cc78a 100644 --- a/projects/plots/plots/vizapp/pages/analysis/page.py +++ b/projects/plots/plots/vizapp/pages/analysis/page.py @@ -49,6 +49,9 @@ def get_layout(self): distribution = self.distribution_plot.get_layout( height=400, width=1500 ) + # Need to update the distribution plot here, or the + # tab would be blank until a veto was toggled. + self.distribution_plot.update(self.app.background, self.app.foreground) return column( distribution, event_inspector, sizing_mode="stretch_both" ) diff --git a/projects/plots/plots/vizapp/pages/summary/page.py b/projects/plots/plots/vizapp/pages/summary/page.py index 90eff52d9..a43921998 100644 --- a/projects/plots/plots/vizapp/pages/summary/page.py +++ b/projects/plots/plots/vizapp/pages/summary/page.py @@ -1,6 +1,6 @@ from typing import TYPE_CHECKING -from plots.core.constants import SECONDS_PER_MONTH, SECONDS_PER_YEAR +from plots.core.constants import DEFAULT_VIZAPP_MAX_FAR from plots.core.data import AnalysisData from plots.core.gwtc3 import main as gwtc3_pipeline_sv from plots.core.sv import ( @@ -13,24 +13,22 @@ if TYPE_CHECKING: from ledger.events import EventSet, RecoveredInjectionSet -VIZAPP_MAX_FAR = 100 * SECONDS_PER_YEAR / SECONDS_PER_MONTH - class Summary(Page): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - data_manager = self.app.data_manager data = AnalysisData( - data_manager.background, - data_manager.foreground, - data_manager.rejected_params, + self.app.background, + self.app.foreground, + self.app.data_manager.rejected_params, ) result = compute_sensitive_volume( data, mass_combos=self.app.mass_combos, source_prior=self.app.source_prior, - max_far=VIZAPP_MAX_FAR, + max_far=DEFAULT_VIZAPP_MAX_FAR, + num_far_points=self.app.num_far_points, ) gwtc3_sv, gwtc3_err = gwtc3_pipeline_sv( @@ -59,6 +57,7 @@ def update( data, mass_combos=self.app.mass_combos, source_prior=self.app.source_prior, - max_far=VIZAPP_MAX_FAR, + max_far=DEFAULT_VIZAPP_MAX_FAR, + num_far_points=self.app.num_far_points, ) self.sv.update(result) diff --git a/projects/plots/tests/test_gwtc3.py b/projects/plots/tests/test_gwtc3.py new file mode 100644 index 000000000..04facfb4c --- /dev/null +++ b/projects/plots/tests/test_gwtc3.py @@ -0,0 +1,108 @@ +import h5py +import numpy as np +import pytest +from plots.core import gwtc3 +from utils.cosmology import DEFAULT_COSMOLOGY + +MASS_COMBOS = [(35, 20)] +PIPELINES = ["gstlal"] + + +def _get_logVT(log_dN, selection, T_obs, N_draw, p_draw): + log_dN = log_dN[selection] + p_draw = p_draw[selection] + if log_dN.size == 0: + return -np.inf, -np.inf + + log_VT = ( + np.log(T_obs) + - np.log(N_draw) + + np.logaddexp.reduce(log_dN - np.log(p_draw)) + ) + log_s2 = ( + 2 * np.log(T_obs) + - 2 * np.log(N_draw) + + np.logaddexp.reduce(2 * (log_dN - np.log(p_draw))) + ) + log_sig2 = gwtc3.logdiffexp(log_s2, 2.0 * log_VT - np.log(N_draw)) + return log_VT, log_sig2 / 2 + + +def _compute_sv_loop( + log_dN, det_stat_p, p_draw, T_obs, N_draw, thresholds, criterion +): + n = len(thresholds) + sv = np.zeros(n, dtype=float) + err = np.zeros(n, dtype=float) + for i, thresh in enumerate(thresholds): + if criterion == "far": + selection = det_stat_p < thresh + else: + selection = det_stat_p > thresh + log_vt, log_sigma_vt = _get_logVT( + log_dN, selection, T_obs, N_draw, p_draw + ) + sv[i] = np.exp(log_vt) / T_obs + err[i] = np.exp(log_sigma_vt) / T_obs + return sv, err + + +@pytest.fixture +def injection_file(tmp_path): + rng = np.random.default_rng(0) + n = 200 + m1 = rng.uniform(20, 60, size=n) + m2 = rng.uniform(10, m1) + path = tmp_path / "injections.hdf5" + with h5py.File(path, "w") as f: + f.attrs["analysis_time_s"] = 365.25 * 24 * 3600 + f.attrs["total_generated"] = n * 10 + g = f.create_group("injections") + g.create_dataset("mass1_source", data=m1) + g.create_dataset("mass2_source", data=m2) + for s in ["spin1x", "spin1y", "spin1z", "spin2x", "spin2y", "spin2z"]: + g.create_dataset(s, data=rng.uniform(-0.1, 0.1, size=n)) + g.create_dataset("redshift", data=rng.uniform(0.01, 1.0, size=n)) + g.create_dataset("sampling_pdf", data=rng.uniform(0.5, 1.5, size=n)) + for p in PIPELINES: + g.create_dataset(f"far_{p}", data=rng.exponential(10, size=n)) + g.create_dataset(f"pastro_{p}", data=rng.uniform(0, 1, size=n)) + return path + + +@pytest.mark.parametrize("criterion", ["far", "pastro"]) +def test_gwtc3_vectorized_matches_reference_loop( + injection_file, tmp_path, criterion +): + thresholds = np.geomspace(1e-3, 10, 30) + + sv, err = gwtc3.main( + mass_combos=MASS_COMBOS, + detection_criterion=criterion, + detection_thresholds=thresholds, + output_dir=tmp_path / "run", + injection_file=injection_file, + pipelines=PIPELINES, + ) + + T_obs, N_draw, injection_params, p_draw, det_stat = ( + gwtc3.get_injection_data(PIPELINES, criterion, injection_file) + ) + log_dNs = gwtc3.get_logdNs( + MASS_COMBOS, injection_params, 0.1, 0.4, 0.998, DEFAULT_COSMOLOGY + ) + + for p in PIPELINES: + for (m1, m2), log_dN in zip(MASS_COMBOS, log_dNs, strict=True): + key = f"{m1}-{m2}" + expected_sv, expected_err = _compute_sv_loop( + log_dN, + det_stat[p], + p_draw, + T_obs, + N_draw, + thresholds, + criterion, + ) + np.testing.assert_allclose(sv[p][key], expected_sv, rtol=1e-9) + np.testing.assert_allclose(err[p][key], expected_err, rtol=1e-9) diff --git a/projects/plots/tests/test_sv.py b/projects/plots/tests/test_sv.py index c199b9180..93a89112a 100644 --- a/projects/plots/tests/test_sv.py +++ b/projects/plots/tests/test_sv.py @@ -1,6 +1,8 @@ import numpy as np import pytest from plots.core import compute +from plots.core.constants import SECONDS_PER_YEAR +from plots.core.sv import _far_grid MASS_COMBOS = [(35, 35), (35, 20), (20, 20), (20, 10)] @@ -40,3 +42,17 @@ def test_compute_threshold_extremes(det_stats_and_weights): np.testing.assert_allclose(sv[:, 0], 0) # everything clears a threshold below the quietest np.testing.assert_allclose(sv[:, 1], weights.sum(axis=-1)) + + +def test_far_grid(make_background): + background = make_background(n=2000, Tb=0.25 * SECONDS_PER_YEAR) + background = background.sort_by("detection_statistic") + + fars, thresholds = _far_grid(background, max_far=365, num_far_points=50) + + assert 0 < len(fars) <= 50 + assert len(fars) == len(thresholds) + assert np.all(np.diff(fars) > 0) # ascending + assert np.all(np.diff(thresholds) <= 0) # non-increasing + assert thresholds[0] == background.detection_statistic[-1] + assert fars[0] == pytest.approx(SECONDS_PER_YEAR / background.Tb)