diff --git a/pymc/distributions/dist_math.py b/pymc/distributions/dist_math.py index d9d1b97eee..62443d06fb 100644 --- a/pymc/distributions/dist_math.py +++ b/pymc/distributions/dist_math.py @@ -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 @@ -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) diff --git a/tests/distributions/test_dist_math.py b/tests/distributions/test_dist_math.py index eaea67bd2a..f3c506226e 100644 --- a/tests/distributions/test_dist_math.py +++ b/tests/distributions/test_dist_math.py @@ -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 @@ -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