Skip to content
Merged
Show file tree
Hide file tree
Changes from 9 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
172 changes: 88 additions & 84 deletions src/jimgw/core/single_event/likelihood.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,8 @@
from jaxtyping import Array, Float, Complex
from jimgw.typing import ComplexScalar, FloatLike, FloatScalar
from scipy.interpolate import interp1d
from evosax.algorithms import CMA_ES

# from evosax.algorithms import CMA_ES
from ripplegw.interfaces import Waveform

from jimgw.core.utils import log_i0, round_up_to_power_of_two
Expand Down Expand Up @@ -552,12 +553,12 @@ class HeterodynedTransientLikelihoodFD(SingleEventLikelihood):
n_bins: int
epsilon: float
reference_parameters: dict
freq_grid_low: Float[Array, " n_bin"]
freq_grid_high: Float[Array, " n_bin"]
bin_widths: Float[Array, " n_bin"]
waveform_low_ref: dict[str, Complex[Array, " n_bin"]]
waveform_high_ref: dict[str, Complex[Array, " n_bin"]]
summary_data: dict[str, Complex[Array, "4 n_bin"]]
freq_grid_low: Float[Array, " n_valid"]
freq_grid_high: Float[Array, " n_valid"]
bin_widths: Float[Array, " n_valid"]
waveform_low_ref: dict[str, Complex[Array, " n_valid"]]
waveform_high_ref: dict[str, Complex[Array, " n_valid"]]
summary_data: dict[str, Complex[Array, "4 n_valid"]]

def __init__(
self,
Expand Down Expand Up @@ -684,61 +685,35 @@ def __init__(
f"'epsilon' must be a positive number, got {epsilon!r}."
)
else:
freqs_arr = jnp.array(frequency_original)
freqs_arr = self.frequencies
phase = HeterodynedTransientLikelihoodFD._max_phase_diff(
freqs_arr, freqs_arr[0], freqs_arr[-1]
)
n_bins = max(1, int(float(phase[-1]) / epsilon))
assert isinstance(n_bins, int)
freq_grid, freq_grid_center = self._make_binning_scheme(
jnp.array(frequency_original), n_bins=n_bins
)
self.freq_grid_low = freq_grid[:-1]
self.freq_grid_high = freq_grid[1:]

h_sky = reference_waveform(frequency_original, self.reference_parameters)

h_amp = jnp.sum(
jnp.array([jnp.abs(h_sky[pol]) for pol in h_sky.keys()]), axis=0
)
f_valid = frequency_original[jnp.where(h_amp > 0)[0]]

valid_waveform_mask = jnp.where(
(self.freq_grid_high <= jnp.max(f_valid))
& (self.freq_grid_low >= jnp.min(f_valid))
)[0]
freq_grid_center = freq_grid_center[valid_waveform_mask]
self.freq_grid_low = self.freq_grid_low[valid_waveform_mask]
self.freq_grid_high = self.freq_grid_high[valid_waveform_mask]
self.n_bins = len(freq_grid_center)
self.bin_widths = self.freq_grid_high - self.freq_grid_low
freq_grid = self._make_binning_scheme(self.frequencies, n_bins=n_bins)
ref_hpc = reference_waveform(self.frequencies, self.reference_parameters)

# freq_grid is one cell longer than freq_grid_center/low/high
start_idx = valid_waveform_mask[0]
end_idx = valid_waveform_mask[-1] + 2
freq_grid = freq_grid[start_idx:end_idx]
masked_freq_grid = self._mask_and_set_frequency_arrays(ref_hpc, freq_grid)

h_sky_low = reference_waveform(self.freq_grid_low, self.reference_parameters)
h_sky_high = reference_waveform(self.freq_grid_high, self.reference_parameters)
hpc_low = reference_waveform(self.freq_grid_low, self.reference_parameters)
hpc_high = reference_waveform(self.freq_grid_high, self.reference_parameters)

for i, detector in enumerate(self.detectors):
h_sky_ifo = {key: h_sky[key][self.frequency_masks[i]] for key in h_sky}
hpc_ifo = {key: ref_hpc[key][self.frequency_masks[i]] for key in ref_hpc}
waveform_ref = detector.fd_response(
detector.sliced_frequencies, h_sky_ifo, self.reference_parameters
detector.sliced_frequencies, hpc_ifo, self.reference_parameters
)
self.waveform_low_ref[detector.name] = detector.fd_response(
self.freq_grid_low, h_sky_low, self.reference_parameters
self.freq_grid_low, hpc_low, self.reference_parameters
)
self.waveform_high_ref[detector.name] = detector.fd_response(
self.freq_grid_high, h_sky_high, self.reference_parameters
self.freq_grid_high, hpc_high, self.reference_parameters
)
self.summary_data[detector.name] = self._compute_coefficients(
detector.sliced_fd_data,
detector,
waveform_ref,
detector.sliced_psd,
detector.sliced_frequencies,
freq_grid,
freq_grid_center,
masked_freq_grid,
)

def evaluate(self, params: dict[str, Float]) -> FloatScalar:
Expand Down Expand Up @@ -794,6 +769,56 @@ def _likelihood(self, params: dict[str, Float]) -> FloatScalar:

return log_likelihood

def _make_binning_scheme(
self,
freqs: Float[Array, " n_freq"],
n_bins: int,
chi: float = 1.0,
) -> Float[Array, " n_bins + 1"]:
"""Make ``n_bins`` frequency bins of equal phase change.

``n_bins`` must be a positive integer resolved by the caller
(see :meth:`__init__`).
"""
phase_diff_array = self._max_phase_diff(freqs, freqs[0], freqs[-1], chi=chi)
total_phase = phase_diff_array[-1]
phase_diff = jnp.linspace(0.0, total_phase, n_bins + 1)
f_bins = interp1d(phase_diff_array, freqs)(phase_diff)
return jnp.array(f_bins)

def _mask_and_set_frequency_arrays(
self,
waveform: dict[str, Complex[Array, " n_freq"]],
frequencies: Float[Array, " n_freq"],
) -> Float[Array, " n_valid+1"]:
"""
Mask out trivial waveform pieces which are usually beyond merger frequency.

This is to avoid creating NaNs from 0/0 when computing the r0 and r1,
where the ratios between waveforms are taken.

Remark:
The following operations change array shapes dynamically, which is
not jittable. In the future, if jax.jit is preferable, this has to
be scrapped, and then when computing the likelihood (inside _likelihood),
replace jnp.sum with jnp.nansum.
The current implementation has the advantage of greatly reducing memory
usage when the detector f_max is larger than merger frequency.
"""
h_amp = jnp.array([jnp.abs(p) for p in waveform.values()]).sum(axis=0)
_valid_frequencies = self.frequencies[h_amp > 0]
valid_mask = (frequencies >= _valid_frequencies[0]) & (
frequencies <= _valid_frequencies[-1]
)

masked_frequencies = frequencies[valid_mask]
self.freq_grid_low = masked_frequencies[:-1]
self.freq_grid_high = masked_frequencies[1:]
self.nbins = len(masked_frequencies) - 1
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
self.bin_widths = self.freq_grid_high - self.freq_grid_low

return masked_frequencies

@staticmethod
def _max_phase_diff(
freqs: Float[Array, " n_freq"],
Expand All @@ -813,64 +838,43 @@ def _max_phase_diff(

See also Eq.(7) in arXiv:2302.05333.
"""
gamma = jnp.array([-5 / 3, -2 / 3, 1.0, 5 / 3, 7 / 3])
gamma = jnp.array([-5.0, -2.0, 3.0, 5.0, 7.0]) / 3
freq_2D = jax.lax.broadcast_in_dim(freqs, (freqs.size, gamma.size), [0])
f_star = jnp.where(gamma >= 0, f_high, f_low)
summand = (freq_2D / f_star) ** gamma * jnp.sign(gamma)
dphi = 2 * jnp.pi * chi * jnp.sum(summand, axis=1)
return dphi - dphi[0]

def _make_binning_scheme(
self,
freqs: Float[Array, " n_freq"],
n_bins: int,
chi: float = 1.0,
) -> tuple[Float[Array, " n_bins + 1"], Float[Array, " n_bins"]]:
"""Make ``n_bins`` frequency bins of equal phase change.

``n_bins`` must be a positive integer resolved by the caller
(see :meth:`__init__`).
"""
phase_diff_array = self._max_phase_diff(freqs, freqs[0], freqs[-1], chi=chi)
total_phase = phase_diff_array[-1]
phase_diff = jnp.linspace(0.0, total_phase, n_bins + 1)
f_bins = interp1d(phase_diff_array, freqs)(phase_diff)
f_bins_center = (f_bins[:-1] + f_bins[1:]) / 2
return jnp.array(f_bins), jnp.array(f_bins_center)

@staticmethod
def _compute_coefficients(
data: Complex[Array, " n_freq"],
detector: Detector,
h_ref: Complex[Array, " n_freq"],
psd: Complex[Array, " n_freq"],
freqs: Float[Array, " n_freq"],
f_bins: Float[Array, " n_bin+1"],
f_bins_center: Float[Array, " n_bin"],
) -> Complex[Array, "4 n_bin"]:
df = freqs[1] - freqs[0]
f_bins: Float[Array, " n_valid+1"],
) -> Complex[Array, "4 n_valid"]:
data = detector.sliced_fd_data
psd = detector.sliced_psd
freqs = detector.sliced_frequencies

data_prod = jnp.array(data * h_ref.conj()) / psd
self_prod = jnp.array(h_ref * h_ref.conj()) / psd

freq_bins_left = f_bins[:-1] # Shape: (n_bin)
freq_bins_right = f_bins[1:] # Shape: (n_bin)

# Broadcasting for 2D frequencies
freqs_broadcast = freqs[None, :] # Shape: (1, n_freq)
left_bounds = freq_bins_left[:, None] # Shpae: (n_bin, 1)
right_bounds = freq_bins_right[:, None] # Shape: (n_bin, 1)
freq_bins_left = self.freq_grid_low[:, None] # Shpae: (n_bins, 1)
freq_bins_right = self.freq_grid_high[:, None] # Shape: (n_bins, 1)
freq_bins_center = (freq_bins_left + freq_bins_right) / 2

mask = (freqs_broadcast >= left_bounds) & (freqs_broadcast < right_bounds)
# Shape: (n_bins, n_freq)
mask = (freqs_broadcast >= freq_bins_left) & (freqs_broadcast < freq_bins_right)
# The half-open interval [left, right) excludes any frequency that lands
# exactly on the upper edge of the last bin (f_bins[-1]). This happens
# whenever the interpolated bin edge coincides with the last discrete
# frequency sample (common when the waveform reaches f_max). Extend the
# last row to a closed interval by OR-ing in the equality condition.
mask = mask.at[-1].set(mask[-1] | (freqs == freq_bins_right[-1]))
# Shape: (n_bin, n_freq)

f_bins_center_broadcast = f_bins_center[:, None] # Shape: (n_bin, 1)
freq_shift_matrix = (freqs_broadcast - f_bins_center_broadcast) * mask
mask = mask.at[-1].set(mask[-1] | (freqs == f_bins[-1]))
freq_shift_matrix = (freqs_broadcast - freq_bins_center) * mask
Comment thread
SSL32081 marked this conversation as resolved.

# The resultant arrays have shape (n_bin), the dimension with "n_freq" is summed over.
# The resultant arrays have shape (n_bins), the dimension with "n_freq" is summed over.
summary_data = jnp.array(
[
jnp.sum(data_prod[None, :] * mask, axis=1), # A0
Expand All @@ -880,7 +884,7 @@ def _compute_coefficients(
]
)

return 4 * df * summary_data
return 4 * summary_data / detector.duration

def maximize_likelihood(
self,
Expand Down
129 changes: 65 additions & 64 deletions tests/unit/core/single_event/test_likelihood.py
Original file line number Diff line number Diff line change
Expand Up @@ -1040,9 +1040,10 @@ def test_initialization_stores_attributes(self, detectors_and_waveform):
assert hasattr(likelihood, "freq_grid_high")
assert hasattr(likelihood, "bin_widths")
for det in ifos:
assert det.name in likelihood.summary_data
assert det.name in likelihood.waveform_low_ref
assert det.name in likelihood.waveform_high_ref
for attr in ("summary_data", "waveform_low_ref", "waveform_high_ref"):
obj = getattr(likelihood, attr)
assert det.name in obj
assert jnp.isfinite(obj[det.name]).all()

def test_no_reference_params_and_no_prior_raises(self, detectors_and_waveform):
ifos, waveform, fmin, fmax, gps = detectors_and_waveform
Expand Down Expand Up @@ -1170,67 +1171,67 @@ def test_maximize_likelihood(self, detectors_and_waveform):
assert jnp.isclose(float(base.evaluate(result)), ll_injected)
common_keys_allclose(result, true_params)

def test_low_frequency_reference_cutoff_does_not_reindex_summary_data(
self, detectors_and_waveform, monkeypatch
):
ifos, waveform, fmin, fmax, gps = detectors_and_waveform
reference_fmin = 80.0
requested_n_bins = 32

def fake_compute_coefficients(
likelihood, data, h_ref, psd, freqs, f_bins, f_bins_center
):
freqs_broadcast = freqs[None, :]
left_bounds = f_bins[:-1][:, None]
right_bounds = f_bins[1:][:, None]

mask = (freqs_broadcast >= left_bounds) & (freqs_broadcast < right_bounds)

n_freqs = len(freqs)
n_bins = len(f_bins_center)
assert n_freqs > len(f_bins), f"{n_freqs = }, {len(f_bins) = }"
assert len(f_bins) == n_bins + 1, f"{len(f_bins) = }, {n_bins = }"
assert likelihood.n_bins == n_bins, f"{likelihood.n_bins = }, {n_bins = }"
assert n_bins < requested_n_bins
assert freqs[0] == fmin
assert f_bins[0] >= reference_fmin
assert mask.shape == (n_bins, n_freqs), (
f"{mask.shape = }, expected: ({n_bins}, {n_freqs})"
)
coeffs = jnp.arange(n_bins, dtype=jnp.float64)
return jnp.array([coeffs + nn * 100 for nn in range(4)])

monkeypatch.setattr(
HeterodynedTransientLikelihoodFD,
"_compute_coefficients",
fake_compute_coefficients,
)

def reference_waveform(frequencies, params):
waveform_sky = waveform(frequencies, params)
mask = frequencies >= reference_fmin
return {
polarization: jnp.where(mask, strain, jnp.zeros_like(strain))
for polarization, strain in waveform_sky.items()
}

likelihood = HeterodynedTransientLikelihoodFD(
detectors=ifos,
waveform=waveform,
reference_waveform=reference_waveform,
f_min=fmin,
f_max=fmax,
trigger_time=gps,
n_bins=requested_n_bins,
reference_parameters=example_params(),
)
assert likelihood.n_bins < requested_n_bins
assert likelihood.freq_grid_low[0] >= reference_fmin

expected = jnp.arange(likelihood.n_bins, dtype=jnp.float64)
expected_arr = jnp.array([expected + nn * 100 for nn in range(4)])
for detector in ifos:
assert jnp.array_equal(likelihood.summary_data[detector.name], expected_arr)
# def test_low_frequency_reference_cutoff_does_not_reindex_summary_data(
# self, detectors_and_waveform, monkeypatch
# ):
# ifos, waveform, fmin, fmax, gps = detectors_and_waveform
# reference_fmin = 80.0
# requested_n_bins = 32

# def fake_compute_coefficients(
# likelihood, detector, h_ref, f_bins
# ):
# freqs = detector.sliced_frequencies
# freqs_broadcast = freqs[None, :]
# left_bounds = f_bins[:-1][:, None]
# right_bounds = f_bins[1:][:, None]

# mask = (freqs_broadcast >= left_bounds) & (freqs_broadcast < right_bounds)

# n_freqs = len(freqs)
# n_bins = len(f_bins) - 1
# assert n_freqs > len(f_bins), f"{n_freqs = }, {len(f_bins) = }"
# assert likelihood.n_bins == n_bins, f"{likelihood.n_bins = }, {n_bins = }"
# assert n_bins <= requested_n_bins
# assert freqs[0] == fmin
# assert f_bins[0] >= reference_fmin
# assert mask.shape == (n_bins, n_freqs), (
# f"{mask.shape = }, expected: ({n_bins}, {n_freqs})"
# )
# coeffs = jnp.arange(n_bins, dtype=jnp.float64)
# return jnp.array([coeffs + nn * 100 for nn in range(4)])

# monkeypatch.setattr(
# HeterodynedTransientLikelihoodFD,
# "_compute_coefficients",
# fake_compute_coefficients,
# )

# def reference_waveform(frequencies, params):
# waveform_sky = waveform(frequencies, params)
# mask = frequencies >= reference_fmin
# return {
# polarization: jnp.where(mask, strain, jnp.zeros_like(strain))
# for polarization, strain in waveform_sky.items()
# }

# likelihood = HeterodynedTransientLikelihoodFD(
# detectors=ifos,
# waveform=waveform,
# reference_waveform=reference_waveform,
# f_min=fmin,
# f_max=fmax,
# trigger_time=gps,
# n_bins=requested_n_bins,
# reference_parameters=example_params(),
# )
# assert likelihood.n_bins < requested_n_bins
# assert likelihood.freq_grid_low[0] >= reference_fmin

# expected = jnp.arange(likelihood.n_bins, dtype=jnp.float64)
# expected_arr = jnp.array([expected + nn * 100 for nn in range(4)])
# for detector in ifos:
# assert jnp.array_equal(likelihood.summary_data[detector.name], expected_arr)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated

# ── Phase marginalization ──────────────────────────────────────────────────

Expand Down
Loading