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
7 changes: 1 addition & 6 deletions pymc/distributions/multivariate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
24 changes: 24 additions & 0 deletions tests/distributions/test_multivariate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
[
Expand Down
Loading