diff --git a/lectures/bayes_nonconj.md b/lectures/bayes_nonconj.md index c70df051d..995408317 100644 --- a/lectures/bayes_nonconj.md +++ b/lectures/bayes_nonconj.md @@ -332,7 +332,11 @@ NumPyro builds this by feeding a `TruncatedNormal` through an `ExpTransform`. def truncated_lognormal(μ, σ): "Log-normal distribution truncated to the unit interval (0, 1]." base = dist.TruncatedNormal(loc=μ, scale=σ, low=-jnp.inf, high=0.0) - return dist.TransformedDistribution(base, dist.transforms.ExpTransform()) + # Declare the (0, 1] support: ExpTransform alone advertises (0, ∞), + # which would let the sampler propose θ > 1. + class _UnitLogNormal(dist.TransformedDistribution): + support = dist.constraints.interval(0.0, 1.0) + return _UnitLogNormal(base, dist.transforms.ExpTransform()) prior_ln = truncated_lognormal(0.0, 1.0) mcmc_ln = run_nuts(binomial_model, prior_ln, k, n)