Skip to content
Merged
Changes from 2 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
71 changes: 25 additions & 46 deletions src/jimgw/core/single_event/likelihood.py
Original file line number Diff line number Diff line change
Expand Up @@ -552,12 +552,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_bins"]
freq_grid_high: Float[Array, " n_bins"]
bin_widths: Float[Array, " n_bins"]
waveform_low_ref: dict[str, Complex[Array, " n_bins"]]
waveform_high_ref: dict[str, Complex[Array, " n_bins"]]
summary_data: dict[str, Complex[Array, "4 n_bins"]]

def __init__(
self,
Expand Down Expand Up @@ -690,34 +690,15 @@ def __init__(
)
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.n_bins = n_bins
freq_grid = self._make_binning_scheme(
jnp.array(frequency_original), n_bins=self.n_bins
)
self.bin_widths = self.freq_grid_high - self.freq_grid_low
self.freq_grid_low = freq_grid[:-1]
self.freq_grid_high = freq_grid[1:]
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated

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 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]

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

Expand All @@ -738,7 +719,6 @@ def __init__(
detector.sliced_psd,
detector.sliced_frequencies,
freq_grid,
freq_grid_center,
)

def evaluate(self, params: dict[str, Float]) -> FloatScalar:
Expand Down Expand Up @@ -813,7 +793,7 @@ 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)
Expand All @@ -825,7 +805,7 @@ def _make_binning_scheme(
freqs: Float[Array, " n_freq"],
n_bins: int,
chi: float = 1.0,
) -> tuple[Float[Array, " n_bins + 1"], Float[Array, " n_bins"]]:
) -> 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
Expand All @@ -835,42 +815,41 @@ def _make_binning_scheme(
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)
return jnp.array(f_bins)

@staticmethod
def _compute_coefficients(
data: Complex[Array, " n_freq"],
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"]:
f_bins: Float[Array, " n_bins+1"],
) -> Complex[Array, "4 n_bins"]:
df = freqs[1] - freqs[0]
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)
freq_bins_left = f_bins[:-1] # Shape: (n_bins)
freq_bins_right = f_bins[1:] # Shape: (n_bins)
freq_bins_center = (freq_bins_left + freq_bins_right) / 2 # Shape: (n_bins)

# 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)
left_bounds = freq_bins_left[:, None] # Shpae: (n_bins, 1)
right_bounds = freq_bins_right[:, None] # Shape: (n_bins, 1)
freq_bins_center_broadcast = freq_bins_center[:, None] # Shape: (n_bins, 1)

# Shape: (n_bins, n_freq)
mask = (freqs_broadcast >= left_bounds) & (freqs_broadcast < right_bounds)
# 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
freq_shift_matrix = (freqs_broadcast - freq_bins_center_broadcast) * mask

# 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 Down
Loading