From 34fb48912bd2efff58908544c2e4b718a12f5a48 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Mon, 13 Jul 2026 14:55:18 +0200 Subject: [PATCH 1/2] Reject boolean cardinality parameters --- .../tracking/nonparametric_cardinality.py | 19 ++++++++++++------- 1 file changed, 12 insertions(+), 7 deletions(-) diff --git a/src/pyrecest/tracking/nonparametric_cardinality.py b/src/pyrecest/tracking/nonparametric_cardinality.py index 8ef663f1fa..3f1ab61a10 100644 --- a/src/pyrecest/tracking/nonparametric_cardinality.py +++ b/src/pyrecest/tracking/nonparametric_cardinality.py @@ -10,10 +10,19 @@ from dataclasses import dataclass from math import exp, isfinite, lgamma, log -from numbers import Integral +from numbers import Integral, Real from typing import Sequence +def _as_finite_real_scalar(value, name: str) -> float: + if isinstance(value, bool) or not isinstance(value, Real): + raise TypeError(f"{name} must be a real scalar.") + value = float(value) + if not isfinite(value): + raise ValueError(f"{name} must be finite.") + return value + + @dataclass(frozen=True) class PitmanYorCardinalityPrior: """Pitman--Yor Chinese-restaurant prior over target-generated clusters. @@ -48,12 +57,8 @@ class PitmanYorCardinalityPrior: discount: float = 0.0 def __post_init__(self) -> None: - strength = float(self.strength) - discount = float(self.discount) - if not isfinite(strength): - raise ValueError("strength must be finite.") - if not isfinite(discount): - raise ValueError("discount must be finite.") + strength = _as_finite_real_scalar(self.strength, "strength") + discount = _as_finite_real_scalar(self.discount, "discount") if not 0.0 <= discount < 1.0: raise ValueError("discount must satisfy 0 <= discount < 1.") if strength <= -discount: From f8b4fade9f43330cdae291e38d4618459ea148c6 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Mon, 13 Jul 2026 14:55:45 +0200 Subject: [PATCH 2/2] Test boolean cardinality parameter rejection --- tests/tracking/test_nonparametric_cardinality.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/tests/tracking/test_nonparametric_cardinality.py b/tests/tracking/test_nonparametric_cardinality.py index 397c90e9fb..630948f537 100644 --- a/tests/tracking/test_nonparametric_cardinality.py +++ b/tests/tracking/test_nonparametric_cardinality.py @@ -1,5 +1,6 @@ import math +import numpy as np import pytest from pyrecest.tracking.nonparametric_cardinality import ( @@ -81,6 +82,20 @@ def test_parameter_validation(): DirichletProcessCardinalityPrior(strength=0.0) +@pytest.mark.parametrize( + "kwargs", + [ + pytest.param({"strength": True}, id="python-bool-strength"), + pytest.param({"strength": np.bool_(True)}, id="numpy-bool-strength"), + pytest.param({"discount": False}, id="python-bool-discount"), + pytest.param({"discount": np.bool_(False)}, id="numpy-bool-discount"), + ], +) +def test_parameter_validation_rejects_boolean_scalars(kwargs): + with pytest.raises(TypeError, match="must be a real scalar"): + PitmanYorCardinalityPrior(**kwargs) + + def test_cluster_size_validation(): prior = PitmanYorCardinalityPrior()