Skip to content

Commit b0f5e49

Browse files
committed
Support xtensor inputs in check_parameters
1 parent 47bdf54 commit b0f5e49

2 files changed

Lines changed: 65 additions & 4 deletions

File tree

pymc/distributions/dist_math.py

Lines changed: 18 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,8 @@
3131
from pytensor.tensor import gammaln
3232
from pytensor.tensor.elemwise import Elemwise
3333
from pytensor.utils import lazy_scipy_module
34+
from pytensor.xtensor.basic import tensor_from_xtensor, xtensor_from_tensor
35+
from pytensor.xtensor.type import XTensorVariable
3436

3537
from pymc.distributions.shape_utils import to_tuple
3638
from pymc.logprob.utils import CheckParameterValue
@@ -65,13 +67,25 @@ def check_parameters(
6567
expression under the normal parameter support as it can be disabled by the user via
6668
check_bounds = False in pm.Model()
6769
"""
70+
expr_dims = None
71+
if isinstance(expr, XTensorVariable):
72+
expr_dims = expr.dims
73+
expr = tensor_from_xtensor(expr)
74+
6875
# pt.all does not accept True/False, but accepts np.array(True)/np.array(False)
69-
conditions_ = [
70-
cond if (cond is not True and cond is not False) else np.array(cond) for cond in conditions
71-
]
76+
conditions_ = []
77+
for cond in conditions:
78+
if cond is True or cond is False:
79+
cond = np.array(cond)
80+
elif isinstance(cond, XTensorVariable):
81+
cond = tensor_from_xtensor(cond)
82+
conditions_.append(cond)
7283
all_true_scalar = pt.all([pt.all(cond) for cond in conditions_])
7384

74-
return CheckParameterValue(msg, can_be_replaced_by_ninf)(expr, all_true_scalar)
85+
checked_expr = CheckParameterValue(msg, can_be_replaced_by_ninf)(expr, all_true_scalar)
86+
if expr_dims is not None:
87+
return xtensor_from_tensor(checked_expr, dims=expr_dims)
88+
return checked_expr
7589

7690

7791
check_icdf_parameters = partial(check_parameters, can_be_replaced_by_ninf=False)

tests/distributions/test_dist_math.py

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,9 @@
1818

1919
from pytensor import config, function
2020
from pytensor.tensor.random.basic import multinomial
21+
from pytensor.tensor.variable import TensorVariable
22+
from pytensor.xtensor import as_xtensor
23+
from pytensor.xtensor.type import XTensorVariable
2124
from scipy import interpolate
2225

2326
import pymc as pm
@@ -68,6 +71,50 @@ def test_check_parameters_shape():
6871
assert check_parameters(1, *conditions).eval().shape == ()
6972

7073

74+
def test_check_parameters_xtensor_expression_and_conditions():
75+
expr = as_xtensor(np.array([1.0, 2.0]), dims=("batch",))
76+
77+
result = check_parameters(expr, expr > 0, expr < 3)
78+
79+
assert isinstance(result, XTensorVariable)
80+
assert result.dims == expr.dims
81+
assert result.dtype == expr.dtype
82+
assert result.type.shape == expr.type.shape
83+
np.testing.assert_array_equal(result.eval(), expr.eval())
84+
85+
86+
def test_check_parameters_invalid_xtensor_condition():
87+
expr = as_xtensor(np.array([1.0, 2.0]), dims=("batch",))
88+
result = check_parameters(expr, expr < 2, msg="parameter check msg")
89+
90+
with pytest.raises(ParameterValueError, match="^parameter check msg*"):
91+
result.eval()
92+
93+
94+
def test_check_parameters_tensor_expression_xtensor_condition():
95+
expr = pt.as_tensor_variable([1.0, 2.0])
96+
condition = as_xtensor(np.array([True, True]), dims=("batch",))
97+
98+
result = check_parameters(expr, condition)
99+
100+
assert isinstance(result, TensorVariable)
101+
np.testing.assert_array_equal(result.eval(), expr.eval())
102+
103+
104+
@pytest.mark.parametrize("python_condition, succeeds", [(True, True), (False, False)])
105+
def test_check_parameters_mixed_conditions(python_condition, succeeds):
106+
expr = as_xtensor(np.array([1.0, 2.0]), dims=("batch",))
107+
tensor_condition = pt.as_tensor_variable([True, True])
108+
109+
result = check_parameters(expr, expr > 0, tensor_condition, python_condition)
110+
111+
if succeeds:
112+
np.testing.assert_array_equal(result.eval(), expr.eval())
113+
else:
114+
with pytest.raises(ParameterValueError):
115+
result.eval()
116+
117+
71118
class MultinomialA(Discrete):
72119
rv_op = multinomial
73120

0 commit comments

Comments
 (0)