Skip to content
Merged
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
13 changes: 10 additions & 3 deletions src/pyrecest/filters/replay_grid_likelihood.py
Original file line number Diff line number Diff line change
Expand Up @@ -254,9 +254,10 @@ def adaptive_position_proposal_probability(


def grid_proposal_weights(log_likelihood) -> np.ndarray:
"""Convert finite grid log likelihoods to normalized proposal weights."""
"""Convert valid grid log likelihoods to normalized proposal weights."""

values = np.asarray(log_likelihood, dtype=float)
_validate_grid_log_likelihood_values(values)
finite = np.isfinite(values)
if not np.any(finite):
raise ValueError("all grid log-likelihoods are non-finite")
Expand Down Expand Up @@ -356,9 +357,15 @@ def _coerce_grid_values(values, expected_size: int) -> np.ndarray:
values = np.asarray(values, dtype=float)
if values.shape != (expected_size,):
raise ValueError(f"log_likelihood must have shape ({expected_size},)")
_validate_grid_log_likelihood_values(values)
return values


def _validate_grid_log_likelihood_values(values: np.ndarray) -> None:
if np.any(np.isnan(values) | np.isposinf(values)):
raise ValueError("log_likelihood must contain finite values or -np.inf")


def _linear_rectilinear_grid_values(
positions: np.ndarray,
values: np.ndarray,
Expand Down Expand Up @@ -433,8 +440,8 @@ def _nearest_grid_values(

indices = _nearest_bin_indices(positions[finite_positions], bin_tree)
nearest_values = values[indices]
# A non-finite grid value represents zero or undefined likelihood. Borrowing
# the minimum finite value from another bin would create spurious mass.
# A -inf grid value represents zero likelihood. Borrowing the minimum finite
# value from another bin would create spurious mass.
nearest_values = np.where(
np.isfinite(nearest_values), nearest_values, float(log_zero)
)
Expand Down
47 changes: 47 additions & 0 deletions tests/filters/test_replay_grid_log_likelihood_validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
import unittest

import numpy as np

# pylint: disable=no-name-in-module
from pyrecest.filters import grid_proposal_weights, replay_grid_log_likelihood_values


class TestReplayGridLogLikelihoodValidation(unittest.TestCase):
def test_rejects_nan_and_positive_infinity(self):
bin_centers = np.asarray([[0.0], [1.0]])
positions = np.asarray([[0.0]])

for invalid_value in (np.nan, np.inf):
log_likelihood = np.asarray([0.0, invalid_value])
with self.subTest(invalid_value=invalid_value, api="lookup"):
with self.assertRaisesRegex(
ValueError, "finite values or -np.inf"
):
replay_grid_log_likelihood_values(
positions,
log_likelihood,
bin_centers,
interpolation="nearest",
)
with self.subTest(invalid_value=invalid_value, api="proposal"):
with self.assertRaisesRegex(
ValueError, "finite values or -np.inf"
):
grid_proposal_weights(log_likelihood)

def test_negative_infinity_remains_zero_likelihood(self):
bin_centers = np.asarray([[0.0], [1.0]])
log_likelihood = np.asarray([0.0, -np.inf])

np.testing.assert_allclose(grid_proposal_weights(log_likelihood), [1.0, 0.0])
looked_up = replay_grid_log_likelihood_values(
np.asarray([[1.0]]),
log_likelihood,
bin_centers,
interpolation="nearest",
)
self.assertTrue(np.isneginf(looked_up[0]))


if __name__ == "__main__":
unittest.main()
Loading