Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 65 additions & 0 deletions dingo/gw/dataset/_multibanded_domain_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
"""Shared utilities for generating and evaluating MultibandedFrequencyDomain settings."""

from copy import deepcopy
from typing import Dict

import numpy as np

from dingo.gw.prior import build_prior_with_defaults


def build_extreme_prior(settings: dict):
"""Build a BBH prior with extreme parameter values to stress-test multibanding.

Fixes ``chirp_mass`` to its minimum prior value and ``geocent_time`` to 0.12 s
(the typical prior boundary plus the Earth-radius light-crossing time). All other
parameters are sampled from the distributions specified in ``settings``.

Parameters
----------
settings : dict
Dataset settings dict containing an ``'intrinsic_prior'`` key. Not modified.

Returns
-------
BBHPriorDict
Prior with extreme fixed values for ``chirp_mass`` and ``geocent_time``.
"""
nominal_prior = build_prior_with_defaults(settings["intrinsic_prior"])
extreme_settings = deepcopy(settings["intrinsic_prior"])
extreme_settings["geocent_time"] = 0.12
# Pin the chirp mass to (essentially) its minimum -- the longest, hardest-to-decimate
# signal. A bare scalar would become a bilby DeltaFunction, which breaks the
# *constrained* sampling required by the mass_1/mass_2 Constraint priors (it raises
# "non-broadcastable output operand with shape ()" in PriorDict.sample). A
# negligibly narrow Uniform samples cleanly while keeping every draw at the minimum.
mc_min = nominal_prior["chirp_mass"].minimum
extreme_settings["chirp_mass"] = (
f"bilby.core.prior.Uniform(minimum={mc_min}, maximum={mc_min * (1 + 1e-9)})"
)
return build_prior_with_defaults(extreme_settings)


def print_mismatch_stats(mismatches: np.ndarray, num_samples: int) -> None:
"""Print a summary of mismatch statistics to stdout.

Parameters
----------
mismatches : np.ndarray
1D array of mismatch values across all polarisations and samples.
num_samples : int
Number of waveform samples used, reported in the header line.
"""
print("\nMismatches between UFD waveforms and MFD waveforms interpolated to UFD.")
print(
"This is a conservative estimate of the MFD performance when training networks."
)
print(f"num_samples = {num_samples}")
print(f" Mean mismatch = {np.mean(mismatches)}")
print(f" Standard deviation = {np.std(mismatches)}")
print(f" Max mismatch = {np.max(mismatches)}")
print(f" Median mismatch = {np.median(mismatches)}")
print(" Percentiles:")
print(f" 99 -> {np.percentile(mismatches, 99)}")
print(f" 99.9 -> {np.percentile(mismatches, 99.9)}")
print(f" 99.99 -> {np.percentile(mismatches, 99.99)}")
40 changes: 9 additions & 31 deletions dingo/gw/dataset/evaluate_multibanded_domain.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,14 +5,13 @@
from scipy.interpolate import interp1d

from dingo.gw.dataset import generate_parameters_and_polarizations
from dingo.gw.domains import build_domain, MultibandedFrequencyDomain
from dingo.gw.dataset._multibanded_domain_utils import (build_extreme_prior,
print_mismatch_stats)
from dingo.gw.domains import MultibandedFrequencyDomain, build_domain
from dingo.gw.gwutils import get_mismatch
from dingo.gw.prior import build_prior_with_defaults
from dingo.gw.waveform_generator import (
NewInterfaceWaveformGenerator,
WaveformGenerator,
generate_waveforms_parallel,
)
from dingo.gw.waveform_generator import (NewInterfaceWaveformGenerator,
WaveformGenerator,
generate_waveforms_parallel)


def _evaluate_multibanding_main(
Expand All @@ -26,15 +25,7 @@ def _evaluate_multibanding_main(
if "compression" in settings:
del settings["compression"]

# Update prior to challenge the multi-banding:
#
# (a) Set geocent_time = 0.12 s (boundary of usual prior + Earth-radius crossing time)
# (b) Set chirp mass to bottom end of prior.
prior = build_prior_with_defaults(settings["intrinsic_prior"])
settings["intrinsic_prior"]["geocent_time"] = 0.12
settings["intrinsic_prior"]["chirp_mass"] = prior["chirp_mass"].minimum
# Rebuild prior with updated settings.
prior = build_prior_with_defaults(settings["intrinsic_prior"])
prior = build_extreme_prior(settings)
print("Prior")
for k, v in prior.items():
print(f"{k}: {v}")
Expand Down Expand Up @@ -88,21 +79,8 @@ def _evaluate_multibanding_main(
asd_file="aLIGO_ZERO_DET_high_P_asd.txt",
)

print("\nMismatches between UFD waveforms and MFD waveforms interpolated to MFD.")
print(
"This is a conservative estimate of the MFD performance when training "
"networks."
)
mismatches = np.concatenate([v for v in mismatches.values()])
print(f"num_samples = {num_samples}")
print(" Mean mismatch = {}".format(np.mean(mismatches)))
print(" Standard deviation = {}".format(np.std(mismatches)))
print(" Max mismatch = {}".format(np.max(mismatches)))
print(" Median mismatch = {}".format(np.median(mismatches)))
print(" Percentiles:")
print(" 99 -> {}".format(np.percentile(mismatches, 99)))
print(" 99.9 -> {}".format(np.percentile(mismatches, 99.9)))
print(" 99.99 -> {}".format(np.percentile(mismatches, 99.99)))
mismatches = np.concatenate(list(mismatches.values()))
print_mismatch_stats(mismatches, num_samples)


def parse_args():
Expand Down
Loading
Loading