diff --git a/src/hflow/statistics.py b/src/hflow/statistics.py index 838c8be0..5ef8be5c 100644 --- a/src/hflow/statistics.py +++ b/src/hflow/statistics.py @@ -6,8 +6,12 @@ from dataclasses import dataclass from itertools import pairwise +import numpy as np + def _finite_number(value: float, name: str) -> float: + if isinstance(value, np.generic): + value = value.item() if isinstance(value, bool) or not isinstance(value, (int, float)): raise ValueError(f"{name} must be a finite number") try: diff --git a/tests/test_statistics.py b/tests/test_statistics.py index 60f9e1d2..a26e8782 100644 --- a/tests/test_statistics.py +++ b/tests/test_statistics.py @@ -3,6 +3,7 @@ import itertools import math +import numpy as np import pytest from hflow import ( @@ -165,6 +166,19 @@ def test_mean_combines_tiny_contributions_before_rounding_to_smallest_float() -> assert distribution.mean == smallest_float +def test_measurements_accept_finite_numpy_numeric_scalars() -> None: + observation = WeightedValue(np.float32(0.5), np.int64(30)) + assert observation == WeightedValue(0.5, 30) + + +@pytest.mark.parametrize("invalid_numpy_bool", [np.bool_(True), np.bool_(False)]) +def test_measurements_reject_numpy_booleans(invalid_numpy_bool: np.bool_) -> None: + with pytest.raises(ValueError, match="value must be a finite number"): + WeightedValue(invalid_numpy_bool, 1) # ty: ignore[invalid-argument-type] + with pytest.raises(ValueError, match="weight must be a finite number"): + WeightedValue(1, invalid_numpy_bool) # ty: ignore[invalid-argument-type] + + @pytest.mark.parametrize("invalid_value", [True, "1", None, math.inf, -math.inf, math.nan]) def test_measurements_reject_nonfinite_or_nonnumeric_values(invalid_value: object) -> None: with pytest.raises(ValueError, match="value must be a finite number"):