From 8f5f08e02faa4919f2aedf3abaa3aba4e12007e7 Mon Sep 17 00:00:00 2001 From: Matt McKay Date: Fri, 25 Sep 2026 15:28:13 +1000 Subject: [PATCH] bayes_nonconj: declare the truncated log-normal prior's (0, 1] support (numpyro 0.22) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit numpyro 0.22.0 validates distribution arguments by default. The prior was a TransformedDistribution(TruncatedNormal(high=0), ExpTransform()), whose declared support is (0, ∞), so NUTS initialisation could propose θ > 1 and dist.Binomial(n, θ) raised "BinomialProbs distribution got invalid probs parameter". Declaring interval(0, 1) keeps θ in range. Co-Authored-By: Claude Opus 5.5 (1M context) --- lectures/bayes_nonconj.md | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) 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)