Skip to content
Open
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
22 changes: 18 additions & 4 deletions pymc/distributions/dist_math.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,8 @@
from pytensor.tensor import gammaln
from pytensor.tensor.elemwise import Elemwise
from pytensor.utils import lazy_scipy_module
from pytensor.xtensor.basic import tensor_from_xtensor, xtensor_from_tensor
from pytensor.xtensor.type import XTensorVariable

from pymc.distributions.shape_utils import to_tuple
from pymc.logprob.utils import CheckParameterValue
Expand Down Expand Up @@ -65,13 +67,25 @@ def check_parameters(
expression under the normal parameter support as it can be disabled by the user via
check_bounds = False in pm.Model()
"""
expr_dims = None
if isinstance(expr, XTensorVariable):
expr_dims = expr.dims
expr = tensor_from_xtensor(expr)

# pt.all does not accept True/False, but accepts np.array(True)/np.array(False)
conditions_ = [
cond if (cond is not True and cond is not False) else np.array(cond) for cond in conditions
]
conditions_ = []
for cond in conditions:
if cond is True or cond is False:
cond = np.array(cond)
elif isinstance(cond, XTensorVariable):
cond = tensor_from_xtensor(cond)
conditions_.append(cond)
all_true_scalar = pt.all([pt.all(cond) for cond in conditions_])

return CheckParameterValue(msg, can_be_replaced_by_ninf)(expr, all_true_scalar)
checked_expr = CheckParameterValue(msg, can_be_replaced_by_ninf)(expr, all_true_scalar)
if expr_dims is not None:
return xtensor_from_tensor(checked_expr, dims=expr_dims)
return checked_expr


check_icdf_parameters = partial(check_parameters, can_be_replaced_by_ninf=False)
Expand Down
54 changes: 54 additions & 0 deletions tests/distributions/test_dist_math.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,9 @@

from pytensor import config, function
from pytensor.tensor.random.basic import multinomial
from pytensor.tensor.variable import TensorVariable
from pytensor.xtensor import as_xtensor
from pytensor.xtensor.type import XTensorVariable
from scipy import interpolate

import pymc as pm
Expand Down Expand Up @@ -68,6 +71,57 @@ def test_check_parameters_shape():
assert check_parameters(1, *conditions).eval().shape == ()


def test_check_parameters_xtensor_expression_and_conditions():
expr = as_xtensor(np.array([1.0, 2.0]), dims=("batch",))

result = check_parameters(expr, expr > 0, expr < 3)

assert isinstance(result, XTensorVariable)
assert result.dims == expr.dims
assert result.dtype == expr.dtype
assert result.type.shape == expr.type.shape
np.testing.assert_array_equal(result.eval(), expr.eval())


def test_check_parameters_invalid_xtensor_condition():
expr = as_xtensor(np.array([1.0, 2.0]), dims=("batch",))
result = check_parameters(expr, expr < 2, msg="parameter check msg")

with pytest.raises(ParameterValueError, match="^parameter check msg*"):
result.eval()


def test_check_parameters_xtensor_expression_replaced_by_ninf():
expr = as_xtensor(np.array([1.0, 2.0]), dims=("batch",))
result = check_parameters(expr, False)

np.testing.assert_array_equal(pm.compile([], result)(), [-np.inf, -np.inf])


def test_check_parameters_tensor_expression_xtensor_condition():
expr = pt.as_tensor_variable([1.0, 2.0])
condition = as_xtensor(np.array([True, True]), dims=("batch",))

result = check_parameters(expr, condition)

assert isinstance(result, TensorVariable)
np.testing.assert_array_equal(result.eval(), expr.eval())


@pytest.mark.parametrize("python_condition, succeeds", [(True, True), (False, False)])
def test_check_parameters_mixed_conditions(python_condition, succeeds):
expr = as_xtensor(np.array([1.0, 2.0]), dims=("batch",))
tensor_condition = pt.as_tensor_variable([True, True])

result = check_parameters(expr, expr > 0, tensor_condition, python_condition)

if succeeds:
np.testing.assert_array_equal(result.eval(), expr.eval())
else:
with pytest.raises(ParameterValueError):
result.eval()


class MultinomialA(Discrete):
rv_op = multinomial

Expand Down