diff --git a/pymc/distributions/multivariate.py b/pymc/distributions/multivariate.py index 3edba7936e..defd92c87d 100644 --- a/pymc/distributions/multivariate.py +++ b/pymc/distributions/multivariate.py @@ -1294,17 +1294,12 @@ def _LKJCholeksyCovRV_logp(op, values, rng, n, eta, sd_dist, **kwargs): det_invjac = pt.log(corr_diag) - idx * pt.log(sd_vals) det_invjac = det_invjac.sum() - # TODO: _lkj_normalizing_constant currently requires `eta` and `n` to be constants + # TODO: _lkj_normalizing_constant currently requires `n` to be a constant try: n = int(get_underlying_scalar_constant_value(n)) except NotScalarConstantError: raise NotImplementedError("logp only implemented for constant `n`") - try: - eta = float(get_underlying_scalar_constant_value(eta)) - except NotScalarConstantError: - raise NotImplementedError("logp only implemented for constant `eta`") - norm = _lkj_normalizing_constant(eta, n) return norm + logp_lkj + logp_sd + det_invjac diff --git a/tests/distributions/test_multivariate.py b/tests/distributions/test_multivariate.py index b1ea31975b..40f345e5c9 100644 --- a/tests/distributions/test_multivariate.py +++ b/tests/distributions/test_multivariate.py @@ -975,6 +975,30 @@ def test_no_warning_logp(self): warnings.simplefilter("error") m.logp() + def test_lkj_cholesky_cov_symbolic_eta(self): + with pm.Model() as model: + eta = pm.HalfNormal("eta", sigma=1.0) + sd_dist = pm.Exponential.dist(1.0) + chol = pm.LKJCholeskyCov("chol", eta=eta, n=3, sd_dist=sd_dist) + + logp_fn = model.compile_logp() + dlogp_fn = model.compile_dlogp() + + ip = model.initial_point() + logp_val = logp_fn(ip) + assert np.isfinite(logp_val) + + dlogp_val = dlogp_fn(ip) + assert dlogp_val.size == 7 + assert np.all(np.isfinite(dlogp_val)) + + with pm.Model() as m_invalid: + n_sym = pt.lscalar("n") + sd_dist = pm.Exponential.dist(1.0) + pm.LKJCholeskyCov("chol_invalid", eta=2.0, n=n_sym, sd_dist=sd_dist, compute_corr=False) + with pytest.raises(NotImplementedError, match="logp only implemented for constant `n`"): + m_invalid.logp() + @pytest.mark.parametrize( "sd_dist", [