Bayesian Modeling & Computation
The engine room of both of your projects. Priors, likelihoods and posteriors; conjugate models; credible intervals and decisions; hierarchical models and shrinkage; model checking; and the computation that makes it all run: MCMC, Hamiltonian Monte Carlo and NUTS, variational inference, the ELBO, SVI, variational guides, your custom training loop, and JAX. Plain English first, then numbers, then plots you can play with.
What is this guide about, in one sentence? Treating every unknown as something you are uncertain about, writing that uncertainty as a probability distribution, and updating it with data, then computing that update when the maths is too hard to do by hand.
Why care? Your A/B framework reports $P(\theta_B \gt \theta_A \mid D)$ from Beta, Dirichlet and hierarchical models. Your forecasting model is a Bayesian model with Laplace priors, fitted with a custom SVI loop and a full-rank or low-rank guide. Every one of those words is a chapter here.
Three ways to say it:
- Picture: start with a cloud of plausible answers (the prior); the data blows away the implausible ones; what is left is the posterior.
- Numbers: before the test you think a conversion rate is "around 10%, give or take 5 points"; after 1 000 users you think "11.8%, give or take 1 point".
- Slogan: posterior ∝ likelihood × prior, and computation is how we get the posterior when we cannot write it down.
Before you start. You need Guide 1 (Probability & Data: Bayes' theorem, the Beta, Dirichlet, Normal, Student-t and Laplace distributions, total variance) and Guide 2 (Inference & Experiments: likelihood, MLE vs MAP, regularization as a prior, confidence intervals). The computation chapters reuse gradients and Adam from the Optimization guide and Cholesky factors from the Linear Algebra guide.
Reading order. Chapters 6.1–6.8 are the modeling half (what to compute). Chapters 6.9–6.17 are the computation half (how to compute it). Your syllabus suggests reading the forecasting-model chapters of Guide 4 before MCMC and VI; both orders work, because Guide 4 links back here whenever it needs SVI.
The four guides
1 · Probability & Data
Probability, distributions, LLN/CLT, descriptive statistics, correlation, KDE, Q-Q plots, transformations. Parts 0–4.
2 · Estimation, Inference & Experiments
Estimators, MLE/MAP, tests, intervals, A/B testing, causal thinking, regression, GLMs, PCA, t-SNE. Parts 5–9, 26, 27.
3 · Bayesian Modeling & Computation (you are here)
Bayesian inference, priors, conjugacy, credible intervals, hierarchical models, checking, MCMC, NUTS, VI, ELBO, SVI, guides, your training loop, JAX, JIT. Parts 10, 11, 19–22, 24, 25.
4 · Time Series & Bayesian Forecasting
Time-series foundations, classical forecasting, the Prophet-style model, evaluation, diagnostics, production, capstone. Parts 12–18, 23, 28–30.
Built around your two projects
The Bayesian A/B testing framework
Beta-Binomial and Dirichlet-Multinomial updating (6.3), $P(\theta_B \gt \theta_A \mid D)$ and practical thresholds (6.4), segments with partial pooling (6.5–6.7), checks (6.8), SVI (6.11–6.14), and the boolean-masking-under-JIT problem (6.17).
The Prophet-style forecasting model
Laplace and other priors with prior predictive checks (6.2), posterior predictive distributions (6.1, 6.8), identifiability between trend, seasonality and regressors (6.8), the ELBO and SVI (6.12), full-rank vs low-rank guides chosen by model size (6.13), and your training loop: relative-ELBO early stopping, patience, best-state checkpointing (6.14).
Try your first interactive
This is what your A/B framework does, in one picture.
The roadmap
- 6.1Bayesian inferencePrior, likelihood, posterior, evidence, predictive · Module 23
- 6.2Choosing priorsPrior types and prior predictive checks · Modules 24, 72
- 6.3Conjugate modelsBeta-Binomial, Dirichlet-Multinomial · Module 25
- 6.4Credible intervals and decisionsHDI, P(B > A), practical thresholds · Module 26
- 6.5Hierarchical modelsGroups, hyperpriors, exchangeability · Module 27
- 6.6Pooling and shrinkageComplete, none, partial · Module 28
- 6.7Centered vs non-centeredThe funnel and how to escape it · Module 29
- 6.8Checking a Bayesian modelPosterior predictive checks, sensitivity, identifiability · Modules 73–75
- 6.9Approximate inference and MCMCWhy, the menu, Metropolis from scratch · Modules 53, 76
- 6.10HMC, NUTS and diagnosticsLeapfrog, U-turns, R-hat, ESS, divergences · Modules 77–78
- 6.11Variational inference and KLInference as optimization · Modules 54–55
- 6.12The ELBO and SVIDerivation, Monte Carlo gradients · Modules 56–57
- 6.13Variational guidesMean-field, full-rank, low-rank · Modules 58–61, 68
- 6.14The custom SVI training loopRelative ELBO stopping, patience, checkpoints · Modules 62–65
- 6.15SVI vs NUTSSpeed vs fidelity · Module 79
- 6.16JAX fundamentalsPure functions, grad, vmap, scan, PRNG keys · Module 66
- 6.17JIT compilation and model sizeTracing, static shapes, masking, complexity · Modules 67–68
What matters most (your P0 list). Bayesian inference and the posterior predictive (6.1), credible intervals and posterior probabilities (6.4), Beta-Binomial and Dirichlet-Multinomial (6.3), hierarchical models and partial pooling (6.5–6.6), MCMC and NUTS (6.9–6.10), VI, KL and the ELBO (6.11–6.12), full-rank vs low-rank guides (6.13), your training loop (6.14), and JAX/JIT (6.16–6.17).
How to read the symbols
| Symbol | Say it as | Meaning |
|---|---|---|
| $p(\theta)$ | "the prior" | What you believe about $\theta$ before seeing the data. |
| $p(D\mid\theta)$ | "the likelihood" | How probable the observed data is if the parameter were $\theta$. |
| $p(\theta\mid D)$ | "the posterior" | What you believe about $\theta$ after seeing the data. |
| $p(D)$ | "the evidence", "marginal likelihood" | The probability of the data averaged over the prior; the normalizing constant. |
| $p(\tilde y\mid D)$ | "posterior predictive" | The distribution of a new observation, averaging over posterior uncertainty. |
| $\mu, \tau$ and $\theta_g$ | "population mean and spread; group g's parameter" | The two levels of a hierarchical model. |
| $q_\phi(\theta)$ | "q with parameters phi" | A variational approximation (in NumPyro: the guide). |
| $KL(q\,\|\,p)$ | "KL from q to p" | How different $q$ is from $p$; zero only if they are equal; not symmetric. |
| ELBO | "evidence lower bound" | The objective SVI maximizes; $\log p(D) = \text{ELBO} + KL$. |
| $\hat R$, ESS | "R-hat", "effective sample size" | MCMC convergence and efficiency diagnostics. |
If a section feels too hard, the fuzzy word is usually "likelihood", "marginal" or "posterior". Each one has a notebook box in 6.1: reread it. Bayesian statistics clicks on the second pass for nearly everyone.
Bayesian inference: prior, likelihood, posterior, evidence, predictive
Five words carry the whole of Bayesian statistics. The prior is what you believed before the data. The likelihood is how well each possible answer explains the data. Multiply them and you get the posterior, once you divide by the evidence. And the posterior predictive turns all of it into the thing a business actually asks for: a forecast of what happens next, with honest uncertainty.
- Name the five objects $p(\theta)$, $p(D\mid\theta)$, $p(\theta\mid D)$, $p(D)$, $p(\tilde y\mid D)$, say each one in plain words, and compute all five in a small example
- Build a posterior by grid approximation: prior × likelihood at every candidate value, then divide by the total area
- Explain the evidence $p(D)$ two ways: the normalizing constant, and the average of the likelihood over the prior
- Update sequentially (today's posterior is tomorrow's prior) and show it gives the same answer as one big update
- Compute and simulate the posterior predictive (draw θ, then draw the future data), and explain why it is wider than a "plug-in" prediction
- Separate the two sources of forecast uncertainty: parameter uncertainty and observation noise
What we need from earlier chapters: Bayes' theorem for events and the base-rate idea (Chapter 4.3); densities and "probability = area" (Chapter 4.4); the Binomial and Normal distributions (Chapters 4.7, 4.9); the Beta distribution and its pseudo-counts (Chapter 4.11); the law of total variance (Chapter 4.6); the likelihood and MAP (Chapter 5.2). Notation for this guide: θ (Greek "theta") is an unknown parameter; $D = \{y_1,\dots,y_n\}$ is the observed data; $\tilde y$ ("y tilde") is a future, not-yet-seen observation. $P(\cdot)$ is the probability of an event; $p(\cdot)$ is a probability mass function (for counts) or a density (for continuous values). $p(a\mid b)$ reads "p of a given b". $N(\mu, \sigma^2)$ is written with the variance; NumPyro and SciPy take the standard deviation $\sigma$. "iid" means independent and identically distributed; $\sim$ reads "is distributed as".
Bayesian inference in one picture core
You launch a new checkout page and want to know its conversion rate θ (the share of visitors who buy). You do not know θ. But you are not clueless either: checkout pages in your product usually convert at around 10%, rarely above 25%. That hunch, written as a spread of plausible values with weights, is your prior.
Then the data arrive: 3 of the first 20 visitors buy. For every possible θ you ask: "if θ were the truth, how probable would exactly this data be?" That score is the likelihood. Values of θ that were plausible and explain the data well keep a lot of weight; the others lose weight. Rescale so the weights add up to 1, and you have the posterior: your updated belief. The rescaling number is the evidence. Finally you use the posterior to answer the real question, "how many of the next 50 visitors will buy?": that is the posterior predictive.
"Inference" just means learning about things you cannot see (θ) from things you can see (the data). Bayesian inference does it with probability distributions all the way through.
Three ways to say it:
- Picture: start with a cloud of plausible answers; the data blow away the implausible ones; what is left, rescaled, is the posterior.
- Numbers: three candidate rates 5%, 10%, 15%, equally likely before the data; after 3 buyers in 20 visitors they have probabilities 0.12, 0.39 and 0.49.
- Slogan: posterior ∝ likelihood × prior, and predictions average over everything you still do not know.
All five objects in one small example. This is the "which rate?" example from Chapter 4.3, now with every Bayesian word attached. Candidates: θ = 0.05, 0.10, 0.15. Data $D$: $k = 3$ buyers among $n = 20$ visitors.
- Prior $p(\theta)$: no reason to favour any candidate, so $1/3$ each.
- Likelihood $p(D\mid\theta) = \binom{20}{3}\theta^3(1-\theta)^{17}$, with $\binom{20}{3} = 1140$: θ = 0.05 gives $0.0596$; θ = 0.10 gives $0.1901$; θ = 0.15 gives $0.2428$.
- Prior × likelihood: $0.0596/3 = 0.0199$, $\;0.1901/3 = 0.0634$, $\;0.2428/3 = 0.0809$.
- Evidence $p(D)$ = the sum: $0.0199 + 0.0634 + 0.0809 = 0.1642$. (It is also the average likelihood: $(0.0596 + 0.1901 + 0.2428)/3 = 0.4925/3 = 0.1642$.)
- Posterior $p(\theta\mid D)$ = each product divided by the evidence: $0.0199/0.1642 \approx 0.121$, $\;0.0634/0.1642 \approx 0.386$, $\;0.0809/0.1642 \approx 0.493$. They add up to 1.
- Posterior predictive, "will the next visitor buy?": average the three candidate rates with their posterior weights: $0.05(0.121) + 0.10(0.386) + 0.15(0.493) = 0.00605 + 0.0386 + 0.07395 \approx 0.119$.
So after the data: the most probable of the three rates is 15%, but the honest prediction for the next visitor is 11.9%, because 5% and 10% are still possible.
A Bayesian model is a story of how the data were made, in two steps: first nature picks the unknown parameter $\theta$ from the prior $p(\theta)$; then the data are generated from the likelihood (the "data model") $p(D\mid\theta)$. Together they give the joint distribution $p(\theta, D) = p(D\mid\theta)\,p(\theta)$. Bayes' theorem turns this around:
$$\underbrace{p(\theta\mid D)}_{\text{posterior}} = \frac{\overbrace{p(D\mid\theta)}^{\text{likelihood}}\;\overbrace{p(\theta)}^{\text{prior}}}{\underbrace{p(D)}_{\text{evidence}}}, \qquad p(D) = \int p(D\mid\theta)\,p(\theta)\,d\theta .$$And the posterior predictive distribution of a new observation $\tilde y$ is
$$p(\tilde y\mid D) = \int p(\tilde y\mid\theta)\,p(\theta\mid D)\,d\theta .$$- Parameter θ: an unknown number inside the model (a conversion rate, an average demand, a trend slope). It can be a whole vector of numbers.
- With a few candidate values the integrals are sums (as in the example). With every value in a range they are integrals: "add up over all θ".
- "$\propto$" reads "is proportional to": equal up to a constant factor that does not depend on θ. Since $p(D)$ does not depend on θ, posterior ∝ likelihood × prior.
- Assumptions: the model (prior and likelihood) is written down before looking at the data, and for prediction, the future data come from the same process with the same θ.
Why do we need it?
We want to say how sure we are about an unknown, combine what we knew with what we saw, and turn that into predictions with honest error bars. Bayes' theorem is the one rule that does all three consistently.
Where is it used?
Bayesian A/B testing ($P(\theta_B \gt \theta_A\mid D)$), Bayesian forecasting models such as Prophet-style models in Stan or NumPyro, spam filters (naive Bayes), Kalman filters for tracking, hierarchical models of many segments, and Bayesian optimization of hyperparameters.
How is it used?
Write the model as "prior, then likelihood" (in NumPyro: one numpyro.sample per unknown, plus one with obs= for the data). Let an algorithm compute the posterior (formula, grid, MCMC or SVI). Then summarize it and simulate predictions from it.
"The posterior is just the likelihood with a different name."
The posterior is likelihood × prior, rescaled. With equal prior weights it has the likelihood's shape, but it is a probability distribution over θ (it adds to 1) and the likelihood is not.
"The best single θ answers every question."
The prediction for the next visitor (11.9%) is not the most probable candidate (15%). Questions about future data need the whole posterior, averaged.
"Bayesian means the answer depends on opinions, so it is not science."
Every model has assumptions; the Bayesian one writes them all down, including the prior, so they can be criticized, checked (Chapter 6.2) and varied (Chapter 6.8).
Both of your projects are this picture. In the A/B framework, θ is a variant's conversion rate (or a vector of category probabilities, a mean revenue, a Poisson rate); the prior is a Beta (or Dirichlet, Normal…), the likelihood is Binomial (or Multinomial, Normal, Student-t, Poisson), and the decision $P(\theta_B \gt \theta_A\mid D)$ is read off the posterior. In the forecasting model, θ is the whole vector of trend, changepoint, seasonality, holiday, regressor and noise parameters, and the forecast for each future day is a posterior predictive distribution.
$p(\theta\mid D) = \dfrac{p(D\mid\theta)\,p(\theta)}{p(D)}$, so posterior ∝ likelihood × prior; $p(D) = \int p(D\mid\theta)p(\theta)d\theta$.
Predictive: $p(\tilde y\mid D) = \int p(\tilde y\mid\theta)\,p(\theta\mid D)\,d\theta$, an average over the posterior.
Three candidates, 3 of 20: posterior 0.12 / 0.39 / 0.49; next-visitor prediction 0.119, not 0.15.
Quick check: in the example, what would the posterior be if the prior had been 0, 0.5, 0.5 (rule out 5%)?
Products: $0$, $0.5 \times 0.1901 = 0.0951$, $0.5 \times 0.2428 = 0.1214$. Evidence $= 0.2165$. Posterior: $0$, $0.0951/0.2165 \approx 0.439$, $0.1214/0.2165 \approx 0.561$. A candidate with prior 0 keeps posterior 0, whatever the data say.
The prior $p(\theta)$: what you believe before the data core
Before a weather forecaster looks at today's satellite images, she already knows it is July, and in July it rarely snows. That background knowledge is a prior. It does not decide the forecast; it sets the starting point that the new evidence moves.
A prior is your belief about an unknown, written as a probability distribution: it says which values are plausible and how plausible, before this dataset. "Prior" simply means "before". It is a distribution over the parameter (the conversion rate), not over the data (the number of buyers).
Three ways to say it:
- Picture: a hill over the possible values of θ: high where you think θ probably is, low where you think it probably is not.
- Numbers: the prior Beta(2, 18) says "about 10%": a 56% chance θ is between 5% and 15%, and only an 8% chance it is above 20%.
- Slogan: a prior is an assumption you can see, check and argue about.
The checkout prior Beta(2, 18) (the Beta distribution is taught in Chapter 4.11).
- Mean: $\alpha/(\alpha+\beta) = 2/(2+18) = 0.10$. "About 10%."
- Pseudo-count reading: it is as if we had already seen $\alpha + \beta = 20$ visitors, 2 of whom bought. The prior is "worth 20 visitors".
- Probabilities are areas: $P(\theta \lt 0.05) \approx 0.245$, $\;P(0.05 \lt \theta \lt 0.15) \approx 0.556$, $\;P(\theta \gt 0.20) \approx 0.083$.
- 90% of the prior lies between its 5% and 95% quantiles: $0.019$ and $0.226$.
- The density at θ = 0.10 is about $5.70$. That is not a probability (it is bigger than 1!); it is probability per unit of θ. The chance that θ falls in a tiny window $0.100$ to $0.101$ is about $5.70 \times 0.001 = 0.0057$.
The prior $p(\theta)$ is a probability distribution over the parameter's possible values, chosen before seeing the data you are about to analyse.
- For a continuous θ it is a density: $p(\theta) \ge 0$ and $\int p(\theta)\,d\theta = 1$ over the support (the set of allowed values; for a rate, $[0, 1]$). Probabilities are areas: $P(a \lt \theta \lt b) = \int_a^b p(\theta)\,d\theta$.
- The numbers inside a prior (here α = 2, β = 18) are called hyperparameters, to tell them apart from the model's parameter θ.
- A value with prior density 0 gets posterior density 0, whatever the data say. So never give zero prior weight to anything that is actually possible.
- The prior is part of the model, like the likelihood. How to choose it (informative, weakly informative, flat, and why "flat" is not "uninformative") is Chapter 6.2.
Why do we need it?
Bayes' theorem needs a starting belief: without $p(\theta)$ there is no posterior. With small data a sensible prior also stops silly answers, such as "100% conversion" after 2 buyers in 2 visitors.
Where is it used?
Every Bayesian model: Beta priors on conversion rates in A/B tests, Laplace priors on changepoint slope changes, Normal priors on regression and Fourier coefficients, half-Normal priors on noise scales, priors on segment-level effects in hierarchical models.
How is it used?
Pick a family that matches the parameter's support (Beta for rates, Normal for any real number, half-Normal or Gamma for positive scales), set its hyperparameters from domain knowledge, then read off what it claims (means, intervals, tail probabilities) and check it with simulations.
"A prior biases the result, so the honest thing is to have no prior."
There is no Bayesian answer without a prior, and "no prior" usually means a hidden one (often flat, which is itself a strong claim on other scales: Chapter 6.2). A visible, sensible prior is more honest than a hidden one, and with lots of data its influence fades.
"The prior density is 5.7 at θ = 0.10, so θ = 0.10 has probability 5.7."
Densities are probability per unit of θ; only areas are probabilities. Any single exact value has probability 0 under a continuous prior.
"Look at the data first, then choose a prior that agrees with it."
That uses the same data twice and makes you over-confident. The prior must come from knowledge you had before this dataset (past experiments, physical limits, domain sense).
In your A/B framework each variant's rate gets a prior such as numpyro.sample("theta_A", dist.Beta(2.0, 18.0)). Watch the parameter names: NumPyro's Beta(concentration1, concentration0) means Beta(α, β) with α counting successes. In your forecasting model every unknown has a prior: the base slope, the changepoint slope changes $\delta_j \sim Laplace(0, b)$, the Fourier, holiday and regressor coefficients, and the noise scale. Each of those priors is an assumption you should be able to state in words ("most slope changes are near zero").
Prior $p(\theta)$: a distribution over the parameter, fixed before the data. Probabilities = areas; densities can exceed 1.
Beta(α, β) prior for a rate: mean α/(α + β), worth α + β pseudo-observations. Beta(2, 18): "about 10%, worth 20 visitors".
Trap: zero prior density can never be undone; never choose the prior from the same data.
Quick check: Beta(2, 18) and Beta(20, 180) have the same mean. What is different, and when would you prefer the second?
Beta(20, 180) is worth 200 pseudo-visitors instead of 20, so it is much narrower (90% between about 6.8% and 13.7% instead of 1.9% to 22.6%). Use it only if you really have that much relevant prior information, for example many past tests of very similar pages.
The likelihood $p(D\mid\theta)$: the data's vote for each θ
The likelihood was built in Chapter 5.2; here is its job inside Bayes' theorem. Think of an election where the candidates are the possible values of θ and the voter is the data. For each candidate the data ask: "if you were the truth, how probable would I be?" Candidates that make the data unsurprising get a big vote; candidates that make the data very surprising get a tiny vote.
The data are fixed (they already happened). What changes along the curve is θ. That is why the likelihood is a function of θ, and why it is not a probability distribution over θ: its votes need not add up to 1.
Three ways to say it:
- Picture: a hill over θ, peaked at the value that explains the data best, narrower when there is more data.
- Numbers: five days averaging 100 orders make "μ = 100" about twice as well supported as "μ = 105" (ratio 0.54) and twelve times as well as "μ = 110" (ratio 0.08).
- Slogan: the prior is your opinion; the likelihood is the data's opinion.
Average daily demand. Five days of orders: 96, 104, 110, 90, 100. Model: each day $y_i \sim N(\mu, \sigma^2)$, independent, with the noise standard deviation $\sigma = 10$ orders assumed known (to keep one unknown). The unknown is the average demand μ.
- Likelihood: $p(D\mid\mu) = \prod_{i=1}^{5} \frac{1}{\sigma\sqrt{2\pi}}\exp\!\Big(-\frac{(y_i-\mu)^2}{2\sigma^2}\Big) \propto \exp\!\Big(-\frac{\sum_i (y_i-\mu)^2}{2\sigma^2}\Big)$.
- Sample mean: $\bar y = (96+104+110+90+100)/5 = 500/5 = 100$.
- Split the sum of squares: $\sum_i (y_i-\mu)^2 = \sum_i (y_i-\bar y)^2 + n(\bar y-\mu)^2 = 232 + 5(100-\mu)^2$. (Distances from the mean are $-4, 4, 10, -10, 0$; their squares add to $16+16+100+100+0 = 232$.)
- The 232 does not depend on μ, so $p(D\mid\mu) \propto \exp\!\big(-5(\mu-100)^2/200\big)$: a bell over μ centred at $\bar y = 100$ with width $\sigma/\sqrt n = 10/\sqrt5 \approx 4.47$.
- Votes relative to the best: $\mu = 105$: $\exp(-5 \cdot 25/200) = e^{-0.625} \approx 0.535$. $\;\mu = 110$: $\exp(-5\cdot100/200) = e^{-2.5} \approx 0.082$.
The likelihood is the data model $p(D\mid\theta)$ read as a function of θ with the observed data $D$ held fixed. For data that are independent given θ:
$$p(D\mid\theta) = \prod_{i=1}^n p(y_i\mid\theta), \qquad \log p(D\mid\theta) = \sum_{i=1}^n \log p(y_i\mid\theta).$$- It is not a distribution over θ; only ratios of likelihoods matter, so factors without θ (like $\binom{20}{3}$ or the 232 above) can be dropped.
- For a Normal mean with known σ it is $\propto \exp\!\big(-n(\mu-\bar y)^2/(2\sigma^2)\big)$: centred at $\bar y$, width $\sigma/\sqrt n$. More data, narrower vote.
- In Bayes' theorem, the likelihood is the only place the data enter. A prior plus a likelihood is a complete model.
Why do we need it?
It is the bridge from data to parameters: it says, for every candidate value, how compatible the observations are with it. Without it the data could not change the prior at all.
Where is it used?
The Binomial likelihood of conversions, the Multinomial likelihood of category counts, the Normal, Student-t and Negative Binomial likelihoods of a forecasting model, the cross-entropy loss of a classifier (a Bernoulli likelihood), and every obs= statement in NumPyro.
How is it used?
Choose a distribution for one observation that matches its support and spread (Chapter 4.11), multiply over observations (in code: add log-probabilities), and hand it to Bayes' theorem together with the prior.
"The likelihood of μ = 105 is 0.535, so there is a 53.5% chance that μ = 105."
0.535 is a ratio of how well 105 and 100 explain the data. A probability statement about μ needs the prior and Bayes' theorem: that is the posterior (next section).
"A wider likelihood means the data are wrong."
It means the data carry less information: fewer points or noisier points. The width $\sigma/\sqrt n$ shrinks with more data.
Likelihood $p(D\mid\theta)$: data fixed, θ varies; iid $\Rightarrow \prod_i p(y_i\mid\theta)$; not a distribution over θ.
Normal mean, σ known: $\propto \exp(-n(\mu-\bar y)^2/2\sigma^2)$, peak $\bar y$, width $\sigma/\sqrt n$.
Trap: a likelihood ratio is not a probability of θ.
Quick check: with the same five days but σ = 20, what is the vote of μ = 110 relative to μ = 100?
$\exp(-5 \cdot 10^2/(2\cdot 20^2)) = \exp(-500/800) = e^{-0.625} \approx 0.535$. Doubling the noise has the same effect as halving every distance: noisier data rule out less.
The posterior: prior × likelihood, rescaled (and the grid approximation) core
Lay the prior curve and the likelihood curve on the same θ axis and multiply their heights at every θ. A value of θ that was plausible beforehand and explains the data well keeps a big product. A value that the prior thought unlikely, or that the data contradict, ends up small. The product has the right shape but the wrong total area, so divide by the area. The result is the posterior ("after" the data): a full probability distribution over θ.
The posterior is a compromise. It lies between the prior and the likelihood, closer to whichever one is sharper (more confident).
Three ways to say it:
- Picture: two hills multiply into a third hill that sits between them and is narrower than both.
- Numbers: prior mean 10%, data 3 of 20 = 15%, posterior mean 12.5%.
- Slogan: posterior ∝ likelihood × prior; divide by the total to make it a distribution.
Prior Beta(2, 18), data 3 buyers in 20 visitors.
- Prior shape: $p(\theta) \propto \theta^{2-1}(1-\theta)^{18-1} = \theta^{1}(1-\theta)^{17}$.
- Likelihood shape: $p(D\mid\theta) \propto \theta^{3}(1-\theta)^{17}$.
- Multiply (add the powers): $\theta^{1+3}(1-\theta)^{17+17} = \theta^{4}(1-\theta)^{34}$. That is the shape of a Beta(5, 35) density, so after dividing by the area the posterior is Beta(5, 35). (Beta prior + Binomial data → Beta posterior: Chapter 6.3 covers this "conjugate" shortcut.)
- Summaries: mean $5/40 = 0.125$; most probable value (mode) $4/38 \approx 0.105$; standard deviation $\approx 0.052$; 90% of the posterior lies between $0.052$ and $0.220$.
- A question only the posterior can answer: $P(\theta \gt 0.10\mid D) \approx 0.650$ (the prior gave it 0.420).
- Check one grid point by hand: at θ = 0.10, prior density $5.704$ × likelihood $0.1901$ = $1.084$. Divide by the area $p(D) = 0.1354$ (next section): $1.084/0.1354 \approx 8.007$, which is exactly the Beta(5, 35) density at 0.10.
The posterior is the distribution of θ given the observed data:
$$p(\theta\mid D) = \frac{p(D\mid\theta)\,p(\theta)}{p(D)} \;\propto\; p(D\mid\theta)\,p(\theta), \qquad \log p(\theta\mid D) = \log p(D\mid\theta) + \log p(\theta) - \log p(D).$$Grid approximation (works for any prior and likelihood with one or two parameters):
- Choose $G$ grid points $\theta_1, \dots, \theta_G$ spaced $\Delta$ apart.
- At each point compute the unnormalized posterior $u_i = p(\theta_i)\,p(D\mid\theta_i)$.
- Approximate the area: $p(D) \approx \sum_i u_i\,\Delta$.
- Divide: $p(\theta_i\mid D) \approx u_i / \big(\sum_j u_j\Delta\big)$. Summaries follow, e.g. $E[\theta\mid D] \approx \sum_i \theta_i u_i / \sum_j u_j$.
- The posterior is a whole distribution: report a centre (mean or median), a spread or interval (Chapter 6.4), and direct probabilities such as $P(\theta \gt 0.10\mid D)$.
- The grid needs $G^d$ points for $d$ parameters: 100 points per parameter and 10 parameters is $10^{20}$ evaluations. That is why real models use MCMC or variational inference (Chapters 6.9–6.12).
Why do we need it?
It is the answer of a Bayesian analysis: everything we believe about θ after seeing the data, including how uncertain we still are. Point estimates, intervals, decision probabilities and predictions are all computed from it.
Where is it used?
$P(\theta_B \gt \theta_A\mid D)$ in Bayesian A/B testing, credible intervals for lifts, the parameter draws behind every forecast fan chart, shrunken segment estimates in hierarchical models, and posterior samples from NUTS or an SVI guide in NumPyro.
How is it used?
Get it exactly (conjugate formulas), on a grid (one or two parameters), or as samples (MCMC) or a fitted approximation (SVI). Then summarize it, compute probabilities of events you care about, and push it through the model to predict.
"The posterior equals the likelihood times the prior."
It is proportional to it. The product has the right shape but its area is $p(D)$, not 1; dividing by $p(D)$ makes it a distribution.
"Strong enough data can revive any value of θ."
Not if the prior is exactly zero there: $0 \times$ anything $= 0$. That is why priors should be zero only where values are truly impossible (a rate below 0 or above 1), never just "unlikely".
"We always need the normalizing constant before we can do anything."
Ratios of posterior densities do not need it: $p(\theta_1\mid D)/p(\theta_2\mid D) = \frac{p(D\mid\theta_1)p(\theta_1)}{p(D\mid\theta_2)p(\theta_2)}$. MCMC uses exactly this trick (Chapter 6.9).
For a Beta prior and Binomial data the posterior is Beta$(\alpha + k, \beta + n - k)$ exactly (a closed form your A/B framework can use to check its SVI fit), so no grid or sampler is needed for that simple case. For the non-conjugate parts (hierarchical segments, Student-t or Poisson likelihoods, and the whole forecasting model) NumPyro builds $\log p(D\mid\theta) + \log p(\theta)$ automatically from your sample statements; SVI and NUTS only ever use this unnormalized log posterior and its gradient, never $p(D)$.
"The posterior is the likelihood multiplied by the prior."
"The posterior is proportional to likelihood times prior; the constant is the evidence $p(D)$."
Model answer: "By Bayes' theorem, $p(\theta\mid D) = p(D\mid\theta)p(\theta)/p(D)$. The evidence does not depend on θ, so it only rescales; the shape comes from likelihood × prior. In practice we work with the unnormalized log posterior, $\log p(D\mid\theta) + \log p(\theta)$, because MCMC and VI only need it up to a constant."
$p(\theta\mid D) \propto p(D\mid\theta)\,p(\theta)$. Grid: $u_i = p(\theta_i)p(D\mid\theta_i)$, $p(D) \approx \sum u_i\Delta$, posterior $= u_i/(\sum u_j\Delta)$.
Beta(2, 18) prior + 3 of 20 → Beta(5, 35): mean 0.125, between prior 0.10 and data 0.15.
Traps: "∝" not "="; zero prior stays zero; grids cost $G^d$.
Quick check: with a flat prior Beta(1, 1), what posterior do 3 buyers in 20 visitors give, and what is its mean?
Prior shape $\theta^0(1-\theta)^0 = 1$, likelihood $\theta^3(1-\theta)^{17}$, product $\theta^3(1-\theta)^{17}$ = shape of Beta(4, 18). Mean $4/22 \approx 0.182$, pulled a little from $3/20 = 0.15$ toward 0.5 by the flat prior's one pseudo-buyer and one pseudo-non-buyer.
The evidence $p(D)$: the normalizing constant, and the average likelihood core
Before the experiment, ask the model: "how probable is it that we will see 3 buyers in 20 visitors?" The model does not know θ, so it must hedge: it averages the probability of that outcome over every θ, weighted by the prior. That average is the evidence $p(D)$. It is also the number we divide by to turn prior × likelihood into the posterior.
Its other name is the marginal likelihood. "Marginal" means "with the other variable averaged out" (as in the margins of a table, where you add across a row). Here θ is averaged out, so what is left depends only on the data and on the model.
A model whose prior bet on θ values that explain the data well earns a high evidence. A model that spread its bets over every possible θ, or bet on the wrong ones, earns a low evidence.
Three ways to say it:
- Picture: the total area under the prior × likelihood curve.
- Numbers: the three-candidate example: $(0.0596 + 0.1901 + 0.2428)/3 = 0.1642$, just the average of the three likelihoods.
- Slogan: the evidence is how well the model predicted the data before it saw them.
3 buyers in 20 visitors, under four different priors.
- General formula for a Beta(α, β) prior: $p(D) = \int_0^1 \binom{n}{k}\theta^k(1-\theta)^{n-k}\,\frac{\theta^{\alpha-1}(1-\theta)^{\beta-1}}{B(\alpha,\beta)}\,d\theta = \binom{n}{k}\frac{B(\alpha+k,\ \beta+n-k)}{B(\alpha,\beta)}$, where $B(\cdot,\cdot)$ is the Beta function (the area under $\theta^{a-1}(1-\theta)^{b-1}$).
- Hunch Beta(2, 18): $p(D) = 1140 \cdot B(5, 35)/B(2, 18) \approx 0.1354$.
- Flat Beta(1, 1): $p(D) = 1/21 \approx 0.0476$. (Under a flat prior every $k$ from 0 to 20 is equally likely before the data: $1/(n+1)$ each.)
- Confident Beta(20, 180), also centred at 10%: $p(D) \approx 0.1821$. It bet more firmly on values near 10–15%, and the data landed there.
- Optimist Beta(18, 2), centred at 90%: $p(D) \approx 3.0 \times 10^{-7}$. It bet on the wrong region.
- Comparing two models by their evidence: hunch vs flat $= 0.1354/0.0476 \approx 2.84$. This ratio is called a Bayes factor: the data are 2.84 times more probable under the hunch model.
If instead the data had been 10 buyers in 20, the flat prior would win: $0.0476$ against $0.0013$ for the hunch, a factor of about 36.
The evidence or marginal likelihood of data $D$ under a model (prior + likelihood) is
$$p(D) = \int p(D\mid\theta)\,p(\theta)\,d\theta = E_{\theta \sim p(\theta)}\big[p(D\mid\theta)\big].$$- It is one number for the observed data, not a function of θ. It makes the posterior integrate to 1.
- It is the average of the likelihood over the prior, so it can be estimated by drawing θ from the prior and averaging $p(D\mid\theta)$ (fine in one dimension; hopeless in many, because almost all prior draws explain the data badly).
- Read as a function of possible datasets, $p(y) = \int p(y\mid\theta)p(\theta)d\theta$ is the prior predictive distribution; the evidence is its value at the data you actually got (Chapter 6.2).
- Bayes factor of model 1 against model 2: $BF_{12} = p(D\mid M_1)/p(D\mid M_2)$. It depends strongly on how wide the priors are.
- For most real models $p(D)$ has no formula and is very hard to compute. MCMC avoids it entirely; variational inference maximizes a lower bound on $\log p(D)$, the ELBO (Chapter 6.12).
Why do we need it?
It turns prior × likelihood into a real probability distribution, and it scores a whole model (prior and likelihood together) by how well it predicted the data, which lets you compare models on the same data.
Where is it used?
Bayes factors for model comparison, the Beta-Binomial normalizing constant, empirical Bayes (choosing prior hyperparameters by maximizing $p(D)$), and variational inference, whose objective (the ELBO) is a lower bound on $\log p(D)$ used in your SVI loop.
How is it used?
For conjugate models use the formula (for the Beta-Binomial, scipy.stats.betabinom.pmf). For small problems use a grid. For model comparison in larger models, prefer predictive checks and cross-validation; treat Bayes factors with care because they swing with prior width.
"A high evidence means the posterior is accurate."
The evidence scores the model's prior predictions of the data. A model can have a sharp, sensible posterior and still a low evidence because its prior was very wide.
"For model comparison, just make the priors very vague so they do not matter."
The opposite: the wider the prior, the thinner it spreads its predictions and the lower its evidence, even when the posterior barely changes. Bayes factors can be pushed around almost at will by prior width (the best-known case is the Jeffreys–Lindley paradox). Posterior estimates are far less sensitive.
"We must compute $p(D)$ to get the posterior's shape."
The shape is fixed by likelihood × prior; $p(D)$ only rescales. That is why samplers and SVI can skip it.
"The marginal likelihood is the likelihood at the best parameter."
"The marginal likelihood is the likelihood averaged over the prior: $p(D) = \int p(D\mid\theta)p(\theta)d\theta$."
Model answer: "The likelihood $p(D\mid\theta)$ is a function of θ. The marginal likelihood integrates θ out against the prior, giving one number per model: the probability the model assigned to the observed data before seeing it. It normalizes the posterior and is the basis of Bayes factors, but it is sensitive to prior width and usually intractable, which is why VI maximizes a lower bound on its log, the ELBO."
$p(D) = \int p(D\mid\theta)p(\theta)\,d\theta = E_{\text{prior}}[p(D\mid\theta)]$: the normalizer and the average likelihood.
Beta-Binomial: $p(D) = \binom nk B(\alpha+k, \beta+n-k)/B(\alpha,\beta)$; 3 of 20 with Beta(2, 18): 0.1354; flat: $1/(n+1)$.
Bayes factor $= p(D\mid M_1)/p(D\mid M_2)$. Trap: very sensitive to prior width.
Quick check: why does the flat prior give every $k$ the same evidence $1/(n+1)$?
$p(k) = \binom nk B(k+1, n-k+1) = \binom nk \frac{k!\,(n-k)!}{(n+1)!} = \frac{n!}{k!(n-k)!}\cdot\frac{k!(n-k)!}{(n+1)!} = \frac{1}{n+1}$. Before any data, a flat prior on the rate makes every count of buyers equally likely: it has no opinion about which $k$ will happen.
Sequential updating: today's posterior is tomorrow's prior
Experiments do not deliver their data all at once; it trickles in day by day. You do not need to start over each day. What you believed on Monday evening (Monday's posterior) is exactly what you believe on Tuesday morning, so it becomes Tuesday's prior. Update it with Tuesday's data and carry on.
And the order does not matter: updating day by day, in any order, or all at once with the totals gives the same final posterior, as long as the days are independent pieces of evidence about the same, unchanging θ.
Three ways to say it:
- Picture: a relay race; each day's posterior hands the baton to the next day as its prior.
- Numbers: Beta(2, 18) → Beta(5, 35) → Beta(6, 54) → Beta(10, 70), the same as one update with 8 buyers in 60 visitors.
- Slogan: the posterior is a running summary of everything learned so far.
Prior Beta(2, 18). Three days, 20 visitors each: 3, then 1, then 4 buyers. Each update adds buyers to α and non-buyers to β.
- Day 1 (3 of 20): Beta(2 + 3, 18 + 17) = Beta(5, 35). Mean $5/40 = 0.125$.
- Day 2 (1 of 20), prior = Beta(5, 35): Beta(5 + 1, 35 + 19) = Beta(6, 54). Mean $6/60 = 0.100$.
- Day 3 (4 of 20), prior = Beta(6, 54): Beta(6 + 4, 54 + 16) = Beta(10, 70). Mean $10/80 = 0.125$.
- All at once: 8 buyers in 60 visitors: Beta(2 + 8, 18 + 52) = Beta(10, 70). The same.
- Other order (day 3, day 1, day 2): Beta(6, 34) → Beta(9, 51) → Beta(10, 70). The same again.
If two batches of data $D_1$ and $D_2$ are conditionally independent given θ (once θ is known, one batch tells you nothing more about the other), then
$$p(\theta\mid D_1, D_2) \propto p(D_1, D_2\mid\theta)\,p(\theta) = p(D_2\mid\theta)\,\underbrace{p(D_1\mid\theta)\,p(\theta)}_{\propto\ p(\theta\mid D_1)} \;\propto\; p(D_2\mid\theta)\;p(\theta\mid D_1).$$- So "posterior after $D_1$" plays the role of the prior for $D_2$. Repeating this gives the same answer as one update with all the data, in any order.
- The original prior is used once; it is already inside every later posterior.
- The model assumes θ stays the same over time. If the true rate drifts (novelty effects, seasonality), old and new data still get equal weight, and you need a model with a time-varying θ instead.
Why do we need it?
Data arrive continuously. Sequential updating lets you keep an always-current posterior without storing or reprocessing all the raw data, and it makes clear that the analysis after day 10 and the analysis of the full dataset are the same thing.
Where is it used?
Daily dashboards of a running A/B test (the Beta parameters are just running counts), Thompson sampling in multi-armed bandits, online learning, and Kalman filters, where yesterday's state estimate is today's prior.
How is it used?
For conjugate models, keep the running sums (buyers, non-buyers) and add each new batch. For other models, refit on all the data so far, or use the previous posterior (or an approximation of it) as the new prior.
"Updating every day counts the prior again every day."
The prior enters once, at the start. Each day's prior is yesterday's posterior, which already contains the original prior and all earlier data.
"Sequential updating automatically tracks a conversion rate that changes over time."
This model assumes one fixed θ, so day 1 and day 30 count equally forever. A drifting rate needs a model in which θ itself changes (a state-space model, Chapter 7.18) or an explicit down-weighting of old data.
"Because the posterior is always valid, I can stop the test the first time $P(\theta_B \gt \theta_A\mid D)$ crosses 95%, with no side effects."
The posterior itself is fine at every look, but a decision rule applied at many looks has different long-run error rates from one applied once. How often such a rule declares a winner when there is none is a design question (Chapter 5.11).
In an A/B framework like yours, the Beta-Binomial posterior after day N is the same whether you update daily or recompute from the cumulative counts, so a dashboard only needs running totals per variant. If your forecasting model is refitted on the whole history each time (check your pipeline), sequential updating still explains why a refit with one more week of data moves the posterior only a little (unless the new week contains a surprise, such as a changepoint).
$p(\theta\mid D_1, D_2) \propto p(D_2\mid\theta)\,p(\theta\mid D_1)$ when the batches are independent given θ.
Beta-Binomial: just add buyers to α and non-buyers to β. Order and batching do not matter.
Traps: the prior counts once; a fixed-θ model cannot track drift; peeking affects decisions, not the posterior.
Quick check: prior Beta(1, 1); Monday 2 of 10, Tuesday 5 of 10. What is the posterior, and would swapping the days change it?
Monday: Beta(1 + 2, 1 + 8) = Beta(3, 9). Tuesday: Beta(3 + 5, 9 + 5) = Beta(8, 14), mean $8/22 \approx 0.364$. Swapping: Beta(6, 6) then Beta(8, 14). Same. All at once: 7 of 20 → Beta(8, 14).
The posterior predictive: what will the next data look like? core
Your manager does not ask "what is θ?". She asks "how many of the next 50 visitors will buy?" or "how many orders do we get tomorrow?". Those are questions about future data, written $\tilde y$. Two separate things make the answer uncertain:
- we do not know θ exactly (parameter uncertainty: the posterior still has a spread), and
- even if we knew θ exactly, the next visitors would still buy or not at random (observation noise).
The posterior predictive handles both by a committee vote. Every plausible θ makes its own prediction for $\tilde y$; the predictions are pooled, each weighted by how plausible its θ is after the data. In practice you simulate it in two steps: draw a θ from the posterior, then draw $\tilde y$ from the model with that θ. Repeat thousands of times; the pile of $\tilde y$ values is the predictive distribution.
Three ways to say it:
- Picture: a committee of plausible θ's, each predicting the future; the committee's pooled prediction is wider than any one member's.
- Numbers: after 3 of 20, the next 50 visitors give 6.25 buyers on average either way, but the spread is 3.46 for the predictive versus 2.34 for "plug in θ = 0.125".
- Slogan: predict with all plausible θ's, not just the best one.
Posterior Beta(5, 35) (prior Beta(2, 18), then 3 buyers in 20 visitors).
- Will the next visitor buy? If θ were known, $P(\text{buy}\mid\theta) = \theta$. Average over the posterior: $P(\tilde y = 1\mid D) = \int \theta\,p(\theta\mid D)\,d\theta = E[\theta\mid D] = 5/40 = 0.125$.
- Will both of the next two visitors buy? Given θ: $\theta^2$. Averaged: $E[\theta^2\mid D] = \frac{5\cdot 6}{40\cdot 41} = \frac{30}{1640} \approx 0.0183$. Plugging in θ = 0.125 would give $0.125^2 \approx 0.0156$: too low, because $E[\theta^2] = Var(\theta) + (E\theta)^2$ and the plug-in ignores $Var(\theta\mid D) \approx 0.00267$.
- How many of the next $m = 50$ buy? Mean $50 \times 0.125 = 6.25$ in both approaches.
- Spread, by the law of total variance (Chapter 4.6): $Var(\tilde y) = E[\,m\theta(1-\theta)\,] + Var(m\theta) = 50(0.125 - 0.0183) + 2500(0.00267) \approx 5.34 + 6.67 = 12.00$. So sd $\approx 3.46$. Observation noise and parameter uncertainty contribute about equally here.
- Plug-in Binomial(50, 0.125): variance $50(0.125)(0.875) = 5.47$, sd $\approx 2.34$. It keeps only the first part.
- Tail question, $P(\tilde y \ge 10)$: predictive $\approx 0.169$; plug-in $\approx 0.088$. The plug-in halves the chance of a good day.
The posterior predictive distribution of new data $\tilde y$, assuming $\tilde y$ and $D$ are independent once θ is known, is
$$p(\tilde y\mid D) = \int p(\tilde y\mid\theta)\,p(\theta\mid D)\,d\theta = E_{\theta\sim p(\theta\mid D)}\big[p(\tilde y\mid\theta)\big].$$Simulation recipe (works for any model): for $s = 1,\dots,S$: draw $\theta^{(s)} \sim p(\theta\mid D)$, then draw $\tilde y^{(s)} \sim p(\tilde y\mid\theta^{(s)})$. The values $\tilde y^{(1)},\dots,\tilde y^{(S)}$ are draws from $p(\tilde y\mid D)$: use their mean, quantiles, and the share above any threshold.
- The plug-in predictive $p(\tilde y\mid\hat\theta)$ uses one estimate $\hat\theta$ and ignores parameter uncertainty; it is always too narrow (variance smaller by $Var(E[\tilde y\mid\theta]\mid D)$).
- Beta prior + Binomial data: the predictive count of buyers among $m$ new visitors is the beta-binomial distribution, $P(\tilde y = j\mid D) = \binom{m}{j}\frac{B(a+j,\ b+m-j)}{B(a, b)}$ for posterior Beta(a, b).
- Replace the posterior by the prior and you get the prior predictive $p(\tilde y) = \int p(\tilde y\mid\theta)p(\theta)d\theta$: what the model expects before any data (Chapter 6.2).
- As data grow, parameter uncertainty shrinks and the predictive approaches $p(\tilde y\mid\theta_{\text{true}})$. It never becomes a single number: observation noise stays.
Why do we need it?
Decisions are about future outcomes: stock levels, capacity, expected revenue, "will the next user convert?". A prediction that ignores parameter uncertainty is overconfident, especially with little data and for tail events.
Where is it used?
Forecast distributions and fan charts of Bayesian forecasting models, P(demand > capacity), expected conversions or revenue in A/B planning, Thompson sampling, and posterior predictive checks of model fit (Chapter 6.8). In NumPyro: numpyro.infer.Predictive.
How is it used?
Take posterior draws (from a formula, NUTS, or an SVI guide), run the model forward with each draw to simulate new data, then read off means, quantiles and probabilities from the simulated values. Report intervals for $\tilde y$, not for θ, when the question is about the future.
"Predicting with the posterior mean of θ is the same as the posterior predictive."
Same centre, too narrow. It drops the parameter-uncertainty part of the variance, so tail probabilities (big days, stock-outs, streaks) come out too small. With 3 of 20: $P(\tilde y \ge 10)$ is 0.169, not 0.088.
"The 90% credible interval for θ tells me where the next value will fall."
That interval is about the rate θ. The next observation also has observation noise, so its interval (the predictive interval) is always wider and in different units (buyers, orders).
"With enough data the prediction becomes certain."
Parameter uncertainty disappears, observation noise does not. With 300 of 2000 the sd of the next-50 count is still about 2.55.
In your forecasting model, the forecast for every future day is a posterior predictive: take a draw of all parameters from the fitted SVI guide, compute trend + seasonality + holidays + regressors for that day, then draw $y$ from the Normal, Student-t or Negative Binomial likelihood; repeat; read quantiles for the intervals and the share above a threshold for $P(\text{demand} \gt \text{capacity})$. In NumPyro this is Predictive(model, guide=guide, params=params, num_samples=1000) (or posterior_samples=… after NUTS). In the A/B framework the probability that the next user converts is the posterior mean, while $P(\theta_B \gt \theta_A\mid D)$ is a statement about parameters, not about future data.
"$P(\theta_B \gt \theta_A\mid D)$ and the forecast distribution are both 'the posterior'."
"The first is a posterior probability about parameters; the second is a posterior predictive distribution about future data."
Model answer: "The posterior $p(\theta\mid D)$ describes what I believe about the parameters, so $P(\theta_B \gt \theta_A\mid D)$ is a posterior probability. The posterior predictive $p(\tilde y\mid D) = \int p(\tilde y\mid\theta)p(\theta\mid D)d\theta$ describes future observations, so it adds the observation noise of the likelihood on top of parameter uncertainty. I simulate it by drawing θ from the posterior and then $\tilde y$ from the likelihood. For forecasting, that is the object I evaluate with coverage and CRPS."
$p(\tilde y\mid D) = \int p(\tilde y\mid\theta)\,p(\theta\mid D)\,d\theta$. Simulate: θ ~ posterior, then ỹ ~ likelihood(θ).
$Var(\tilde y) = E[Var(\tilde y\mid\theta)] + Var(E[\tilde y\mid\theta])$ = noise + parameter uncertainty. Plug-in keeps only the first.
Beta(5, 35), next 50: mean 6.25, sd 3.46 (plug-in 2.34). P(next buys) = E[θ | D] = 0.125.
Quick check: posterior Beta(10, 70). What is the probability that the next visitor buys, and that the next two both buy?
Next one: $E[\theta] = 10/80 = 0.125$. Both: $E[\theta^2] = \frac{10\cdot 11}{80\cdot 81} = 110/6480 \approx 0.0170$, a little above $0.125^2 \approx 0.0156$, because θ is still uncertain (less so than with Beta(5, 35), which gave 0.0183).
Forecasting: parameter uncertainty plus observation noise
Tomorrow's orders are uncertain for two reasons. First, the world is noisy: even if you knew the true average demand μ exactly, tomorrow would land somewhere around it. Second, you do not know μ exactly: with only a few days of history your estimate of μ is itself fuzzy. The predictive distribution adds the two. Their variances add, like stacking two blocks.
With little history, the parameter block is big. With a long history, it shrinks toward zero and only the noise block remains: that is the floor no amount of data can remove. For a trend model there is a twist: a small doubt about the slope becomes a large doubt far in the future, so the parameter block grows with the forecast horizon, and the fan opens up.
Three ways to say it:
- Picture: a wobble on top of a wobble; the forecast fan is the sum of both.
- Numbers: noise sd 20, μ uncertain by sd 10 → predictive sd $\sqrt{20^2 + 10^2} \approx 22.4$, not 20 and not 30.
- Slogan: variances add: predictive = noise + parameter uncertainty.
Daily orders $y \sim N(\mu, 20^2)$ with the noise sd σ = 20 known. Four days of history with mean $\bar y = 110$. With a flat prior on μ (a simplification; Chapter 6.2 discusses flat priors), the posterior is $\mu\mid D \sim N(\bar y, \sigma^2/n)$.
- Parameter uncertainty: $sd(\mu\mid D) = \sigma/\sqrt n = 20/\sqrt4 = 10$.
- Predictive for tomorrow: $\tilde y\mid D \sim N(110,\ 20^2 + 10^2) = N(110, 500)$, sd $\sqrt{500} \approx 22.36$.
- Plug-in (pretend μ = 110 exactly): $N(110, 400)$, sd 20.
- Capacity is 150 orders. Plug-in: $z = (150-110)/20 = 2$, $P(\tilde y \gt 150) \approx 0.023$. Predictive: $z = 40/22.36 \approx 1.79$, $P \approx 0.037$: about 60% more risk than the plug-in says.
- 90% intervals: plug-in 77 to 143; predictive 73 to 147.
- With 100 days of history: $sd(\mu\mid D) = 2$, predictive sd $\sqrt{400 + 4} \approx 20.1$, $P(\tilde y \gt 150) \approx 0.023$. The noise floor remains.
By the law of total variance, for any model,
$$Var(\tilde y\mid D) = \underbrace{E\big[Var(\tilde y\mid\theta)\,\big|\,D\big]}_{\text{observation noise}} + \underbrace{Var\big(E[\tilde y\mid\theta]\,\big|\,D\big)}_{\text{parameter uncertainty}}.$$- Normal mean with known σ and a flat prior: $Var(\tilde y\mid D) = \sigma^2 + \sigma^2/n$.
- Linear trend $y_t = a + b\,t + \varepsilon_t$: the mean at a future time $t$ is $a + b\,t$, so the parameter part is $Var(a + b\,t\mid D) = Var(a) + t^2 Var(b) + 2t\,Cov(a, b)$, which grows roughly like the square of the distance from the middle of the history.
- Other sources a full forecast should consider: model uncertainty (the model form could be wrong; not captured by the posterior), and uncertainty in future regressors that must themselves be forecast (Chapter 7.14).
Why do we need it?
Plans and risk limits depend on the width of the forecast, not just its centre. Knowing which source dominates tells you what helps: more history shrinks parameter uncertainty; nothing shrinks the noise floor except a better model with more explanatory features.
Where is it used?
Forecast fan charts, capacity and inventory planning, prediction intervals in regression (the "+1" in the prediction-interval formula), Bayesian forecasting models such as Prophet-style models, and probabilistic forecast evaluation (coverage, CRPS) in Chapter 7.16.
How is it used?
Simulate: draw parameters from the posterior, then draw future data with noise. Plot quantile bands by horizon. To see each source, simulate once with parameters fixed at their posterior mean (noise only) and once without the noise (parameter only).
"Forecast intervals widen with the horizon because noise piles up."
In a trend model with independent noise, the noise does not pile up: each day gets the same σ. The widening comes from parameter uncertainty (an uncertain slope extrapolated further). In random-walk or autoregressive models, noise does accumulate too (Chapter 7.6).
"With enough history the prediction interval shrinks to zero."
It shrinks to the noise floor $\pm z\sigma$. Only a better model (more explanatory components) lowers σ.
"The posterior predictive covers every kind of uncertainty."
It covers parameter uncertainty and noise within the model. If the model is wrong (a missed changepoint, a wrong likelihood), the intervals can be too narrow; checks in Chapters 6.8 and 7.16 catch this.
In your forecasting model, parameter uncertainty in the trend (base slope and changepoint adjustments $\delta_j$), seasonality and holiday coefficients enters through the posterior draws, and observation noise through the likelihood: a Student-t likelihood adds heavier tails, a Negative Binomial adds count noise that grows with the mean. Note a third trend source: changes that might happen after the end of the history. Prophet's documentation says its trend uncertainty assumes the future sees the same average frequency and size of rate changes as the history, and simulates them; a model that only uses the fitted $\delta_j$ treats the future slope as fixed. Check which your model does.
$Var(\tilde y\mid D) = \underbrace{E[Var(\tilde y\mid\theta)]}_{\text{noise}} + \underbrace{Var(E[\tilde y\mid\theta])}_{\text{parameters}}$.
Normal mean, flat prior: $\sigma^2 + \sigma^2/n$. Example: $\sqrt{400 + 100} = 22.4$; $P(\tilde y \gt 150)$ = 0.037 vs plug-in 0.023.
Trend: the parameter part grows with the horizon; noise is the floor.
Quick check: σ = 20 and 16 days of history (flat prior on μ). What is the predictive sd, and what share of the predictive variance is parameter uncertainty?
$\sigma^2/n = 400/16 = 25$, so $Var = 400 + 25 = 425$, sd $\approx 20.6$. Parameter share $25/425 \approx 5.9\%$. After two weeks of data, noise dominates.
Recap, cheat sheet and practice
- A Bayesian model is a two-step story: θ is drawn from the prior $p(\theta)$, then the data from the likelihood $p(D\mid\theta)$.
- Posterior ∝ likelihood × prior; dividing by the evidence $p(D) = \int p(D\mid\theta)p(\theta)d\theta$ makes it a distribution. The posterior is a compromise, closer to whichever side is sharper.
- Grid approximation: evaluate prior × likelihood on a grid, divide by the total. Exact enough in 1–2 dimensions; impossible in many ($G^d$ points).
- The evidence is the average likelihood over the prior, the prior predictive probability of the data. Ratios of evidences are Bayes factors, sensitive to prior width.
- Sequential updating: today's posterior is tomorrow's prior; same result as one batch, in any order, if θ is fixed and batches are independent given θ.
- The posterior predictive $p(\tilde y\mid D) = \int p(\tilde y\mid\theta)p(\theta\mid D)d\theta$: simulate θ from the posterior, then $\tilde y$ from the likelihood. It is wider than the plug-in prediction.
- Predictive variance = observation noise + parameter uncertainty; the second shrinks with data and grows with the forecast horizon for trends.
Cheat sheet
| Object | Formula | Plain words / remember |
|---|---|---|
| Prior | $p(\theta)$, e.g. Beta(2, 18) | belief before the data; worth α + β pseudo-observations |
| Likelihood | $p(D\mid\theta) = \prod_i p(y_i\mid\theta)$ | data fixed, θ varies; not a distribution over θ |
| Posterior | $p(\theta\mid D) = p(D\mid\theta)p(\theta)/p(D)$ | ∝ likelihood × prior; Beta(α, β) + k of n → Beta(α + k, β + n − k) |
| Evidence | $p(D) = \int p(D\mid\theta)p(\theta)\,d\theta$ | average likelihood over the prior; Beta-Binomial: $\binom nk B(\alpha+k,\beta+n-k)/B(\alpha,\beta)$ |
| Bayes factor | $p(D\mid M_1)/p(D\mid M_2)$ | compares models; swings with prior width |
| Sequential | $p(\theta\mid D_1,D_2) \propto p(D_2\mid\theta)\,p(\theta\mid D_1)$ | prior used once; order irrelevant; assumes fixed θ |
| Posterior predictive | $p(\tilde y\mid D) = \int p(\tilde y\mid\theta)p(\theta\mid D)\,d\theta$ | simulate θ then ỹ; P(next buys) = E[θ | D] |
| Predictive variance | $E[Var(\tilde y\mid\theta)] + Var(E[\tilde y\mid\theta])$ | noise + parameters; Normal: $\sigma^2 + \sigma^2/n$ |
import numpy as np
from scipy import stats
import jax
import numpyro
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS, Predictive
# 1) Grid approximation: prior x likelihood, divide by the area (the evidence)
theta = np.linspace(0, 1, 2001)
prior = stats.beta.pdf(theta, 2, 18) # hunch: about 10%, worth 20 visitors
lik = stats.binom.pmf(3, 20, theta) # 3 buyers among 20 visitors
unnorm = prior * lik
evidence = np.trapezoid(unnorm, theta) # p(D) = area under prior x likelihood (np.trapz in NumPy < 2.0)
post = unnorm / evidence
print(round(evidence, 4), round(np.trapezoid(theta * post, theta), 4)) # 0.1354 0.125
# 2) The exact (conjugate) answers: posterior Beta(2 + 3, 18 + 17), evidence = beta-binomial pmf
print(stats.beta(5, 35).mean(), round(stats.betabinom.pmf(3, 20, 2, 18), 4)) # 0.125 0.1354
# 3) Evidence = the AVERAGE likelihood over draws from the prior
rng = np.random.default_rng(0)
th_prior = rng.beta(2, 18, size=200_000)
print(round(stats.binom.pmf(3, 20, th_prior).mean(), 3)) # 0.136 (Monte Carlo; exact 0.1354)
# 4) Sequential updating gives the same posterior as one batch
a, b = 2, 18
for k, n in [(3, 20), (1, 20), (4, 20)]: # three days
a, b = a + k, b + (n - k) # yesterday's posterior is today's prior
print(a, b, a / (a + b)) # 10 70 0.125 (= Beta(2 + 8, 18 + 52) from 8 of 60)
# 5) Posterior predictive by simulation: draw theta from the posterior, then draw the future data
th_post = rng.beta(5, 35, size=200_000)
y_new = rng.binomial(50, th_post) # buyers among the next 50 visitors
print(round(y_new.mean(), 2), round(y_new.std(), 2), round((y_new >= 10).mean(), 3)) # 6.24 3.45 0.168 (exact 6.25 3.46 0.169)
plug = stats.binom(50, 0.125) # plug-in: pretend theta = 0.125 exactly
print(round(plug.std(), 2), round(plug.sf(9), 3)) # 2.34 0.088 (too narrow, tail too light)
print(round(stats.betabinom(50, 5, 35).std(), 2)) # 3.46 (the exact predictive: beta-binomial)
# 6) The same model in NumPyro: NUTS for the posterior, Predictive for the posterior predictive
def model(n, k=None):
theta = numpyro.sample("theta", dist.Beta(2.0, 18.0)) # Beta(concentration1, concentration0)
numpyro.sample("k", dist.Binomial(total_count=n, probs=theta), obs=k)
mcmc = MCMC(NUTS(model), num_warmup=500, num_samples=2000, progress_bar=False)
mcmc.run(jax.random.PRNGKey(0), n=20, k=3)
draws = mcmc.get_samples()["theta"]
print(round(float(draws.mean()), 3)) # 0.124 (exact 0.125; the gap is Monte Carlo error)
pred = Predictive(model, posterior_samples=mcmc.get_samples())
k_next = pred(jax.random.PRNGKey(1), n=50)["k"] # k=None: the "obs" site is now simulated
print(round(float(k_next.mean()), 1), round(float(k_next.std()), 1)) # 6.1 3.4 (exact 6.25 3.46)
prior_pred = Predictive(model, num_samples=20_000)(jax.random.PRNGKey(2), n=20)["k"]
print(round(float((prior_pred == 3).mean()), 3)) # 0.137: the prior predictive at the observed data IS the evidence (0.1354)
1. Prior Beta(3, 7) for a conversion rate; data: 2 buyers among 10 visitors. The posterior is…
2. What is the evidence (marginal likelihood) $p(D)$?
3. The posterior for a conversion rate is Beta(4, 36). The probability that the next visitor buys is…
4. Compared with the plug-in prediction Binomial(m, θ̂) with θ̂ = the posterior mean, the posterior predictive for the number of buyers among the next m visitors has…
5. You update a Beta(1, 1) prior day by day for four days. A colleague updates once with the four-day totals. Assuming θ did not change, how do the posteriors compare?
6. A grid approximation with 100 points per parameter for a model with 6 parameters needs how many evaluations of prior × likelihood?
Practice problems
A. Three candidate rates 0.05, 0.10, 0.15 with prior weights 0.2, 0.5, 0.3; data 3 buyers in 20 visitors. Find the evidence, the posterior and P(next visitor buys).
- Likelihoods (from the chapter): 0.0596, 0.1901, 0.2428.
- Products: $0.2(0.0596) = 0.0119$, $\;0.5(0.1901) = 0.0951$, $\;0.3(0.2428) = 0.0728$. Evidence $= 0.1798$.
- Posterior: $0.0119/0.1798 \approx 0.066$, $\;0.0951/0.1798 \approx 0.529$, $\;0.0728/0.1798 \approx 0.405$.
- P(next buys) $= 0.05(0.066) + 0.10(0.529) + 0.15(0.405) \approx 0.117$.
B. Average demand μ with prior N(100, 20²); noise σ = 10 known; 4 days with mean 120. For a Normal prior and Normal data, the posterior is Normal with precision (1/variance) = prior precision + n/σ², and mean = precision-weighted average. Find the posterior and the predictive sd for tomorrow.
- Precisions: prior $1/400 = 0.0025$; data $n/\sigma^2 = 4/100 = 0.04$. Posterior precision $0.0425$, variance $1/0.0425 \approx 23.5$, sd $\approx 4.85$.
- Mean: $(0.0025 \times 100 + 0.04 \times 120)/0.0425 = (0.25 + 4.8)/0.0425 \approx 118.8$. The data are 16 times as precise as the prior, so the mean sits close to 120.
- Predictive: variance $\sigma^2 + 23.5 = 123.5$, sd $\approx 11.1$ (noise 10 plus a little parameter uncertainty).
C. A test shows 2 buyers in 2 visitors. Compare the evidence under a flat prior Beta(1, 1) and under Beta(1, 9) ("probably low"). What is the Bayes factor, and what does it mean?
Flat: $1/(n+1) = 1/3 \approx 0.333$. Beta(1, 9): $\binom22 B(3, 9)/B(1, 9) = (2!\,8!/11!)\cdot 9 = 9/495 \approx 0.0182$. Bayes factor flat : low $\approx 0.333/0.0182 \approx 18.3$. The "probably low" model was surprised by two buyers in a row; the flat model was not. With only two visitors, though, both posteriors are still wide, so this is weak practical evidence about the rate itself.
D. (Interview) "What is the difference between the posterior and the posterior predictive, and which one does a forecast use?"
"The posterior $p(\theta\mid D)$ is a distribution over the model's parameters: what I believe about the trend, seasonality and noise after seeing the history. The posterior predictive $p(\tilde y\mid D) = \int p(\tilde y\mid\theta)p(\theta\mid D)d\theta$ is a distribution over future observations: it averages the likelihood over the posterior, so it contains both parameter uncertainty and observation noise. A forecast is a posterior predictive: I draw parameters from the posterior (or the SVI guide), compute the mean for each future day, and draw from the likelihood. Plugging in a single parameter estimate would give intervals that are too narrow."
E. A checkout bug was fixed on day 15 and the conversion rate jumped. You have been updating a Beta posterior every day since day 1. What goes wrong, and what can you do?
The model assumes one fixed θ for all days, so days 1–14 (old rate) and days 15+ (new rate) are pooled with equal weight forever; the posterior concentrates on a blend that is true for neither period, and it gets more confident about that wrong blend as data accumulate. Fixes: restart the analysis at day 15 (a new θ after a known change), model the change explicitly (two rates, or a changepoint), or use a time-varying model that down-weights old data.
F. Posterior Beta(5, 35). What is the probability that none of the next 10 visitors buys? Compare with the plug-in answer.
Given θ: $(1-\theta)^{10}$. Averaged over the posterior: $E[(1-\theta)^{10}] = B(5, 45)/B(5, 35) \approx 0.302$ (the beta-binomial probability of 0 buyers in 10). Plug-in: $(1 - 0.125)^{10} = 0.875^{10} \approx 0.263$. The predictive gives a dry spell more probability, because it allows for the chance that θ is lower than 0.125.
Choosing priors and prior predictive checks
Every Bayesian model needs a prior for every unknown. Some priors carry real knowledge, some only rule out nonsense, some try to say nothing (and fail in a surprising way). This chapter is a tour of the kinds of priors you will meet, how strongly each one pulls, and the single most useful habit for getting them right: simulate fake data from your model before fitting it, and ask whether those worlds could really happen.
- Tell informative, weakly informative, diffuse and flat priors apart, and say what each one claims in plain numbers
- Build an informative prior from past experiments, and decide how much to trust it
- Explain why flat is not uninformative: flat on a rate is a hump on the log-odds, and "vague" on the log-odds piles the rate up at 0% and 100%
- Know what conjugate, hierarchical, shrinkage and regularizing priors are, including Laplace priors on changepoint slope changes
- Weigh prior strength against data size with pseudo-counts and precisions
- Run a prior predictive check for an A/B model and for a trend + changepoints forecasting model, and fix priors that imagine absurd worlds
What we need from earlier chapters: prior, likelihood, posterior and the prior predictive idea (Chapter 6.1); the Beta, Normal, Laplace, Gamma and half-Normal distributions (Chapters 4.9–4.11); log-odds and why MAP depends on the parameter scale (Chapter 5.2); ridge and lasso as Gaussian and Laplace priors (Chapter 5.3); standardization and the global scaler (Chapter 4.18). Words used throughout: the log-odds (or logit) of a probability $p$ is $\text{logit}(p) = \log\frac{p}{1-p}$, any real number; its inverse is the logistic (sigmoid) function $p = 1/(1+e^{-\eta})$. A hyperparameter is a number inside a prior (such as α, β or a prior standard deviation).
What a prior says, and how loudly: informative, weakly informative, diffuse core
Priors are like advisers who speak at different volumes. An informative prior speaks loudly: "I have seen many tests like this; the rate is about 10%, give or take 1.5 points." A weakly informative prior speaks quietly: "I do not know much, but a checkout conversion rate of 80% would be absurd." A diffuse (or vague) prior tries to stay silent: "anything is possible", spread over a huge range. A flat prior gives every value the same density.
A good way to measure the volume is to ask: "this prior is worth how many observations?" A Beta(40, 360) prior on a rate counts like 400 visitors you have already seen; Beta(1, 1) counts like 2.
Three ways to say it:
- Picture: a narrow, tall hill (informative), a broad gentle hill (weakly informative), a nearly flat line (diffuse).
- Numbers: after 3 buyers in 20 visitors, priors worth 400, 20 and 2 visitors give posterior means 10.2%, 12.5% and 18.2%.
- Slogan: every prior says something; know what yours says and how loudly.
Three priors for a checkout conversion rate, each updated with the same small test (3 buyers in 20 visitors; Beta(α, β) + k of n → Beta(α + k, β + n − k), Chapter 6.1).
- Informative Beta(40, 360): mean $40/400 = 0.10$; 90% of it between 7.7% and 12.6%; worth 400 visitors. Posterior Beta(43, 377), mean $43/420 \approx 0.102$. Twenty visitors barely move it.
- Middle Beta(2, 18): mean 0.10; 90% between 1.9% and 22.6%; worth 20 visitors. Posterior Beta(5, 35), mean $5/40 = 0.125$. Prior and data share the say.
- Flat Beta(1, 1): mean 0.5; 90% between 5% and 95%; worth 2 visitors. Posterior Beta(4, 18), mean $4/22 \approx 0.182$. The data dominate (and the flat prior's own pull is toward 50%).
- Same data, three answers. None is "the objective one": each follows from what its prior claimed.
- Informative prior: carries substantial, specific information from outside this dataset (past experiments, published studies, physics). It noticeably affects the posterior unless the data are large.
- Weakly informative prior: deliberately contains less information than you actually have. It puts very little probability on implausible or impossible values and is otherwise wide, so the data decide. It mainly stabilizes estimation.
- Diffuse (vague) prior: very wide, e.g. $N(0, 100^2)$ or $N(0, 1000^2)$ on a coefficient. It tries to say "anything goes".
- Flat (uniform) prior: constant density over the allowed range, e.g. Beta(1, 1) on a rate. A flat density over an unbounded range (the whole real line) does not integrate to 1; it is called improper and is acceptable only if the posterior still integrates to 1.
- Prior strength (effective sample size): for a Beta(α, β) prior it is $\alpha + \beta$ pseudo-observations; for a Normal prior $N(\mu_0, \tau^2)$ on a mean with data noise σ, it is $\sigma^2/\tau^2$ observations.
- Which label applies depends on the parameter's scale and units: $N(0, 1)$ is weakly informative for a standardized coefficient, but on a holiday effect measured in raw orders (on a series of 500 orders a day) it says the holiday changes demand by at most about 2 orders: an extremely strong claim.
Why do we need it?
You must choose a prior for every parameter, and the choice changes results most exactly where data are thin: small tests, small segments, rare holidays, the end of a trend. Knowing the kinds and their strength lets you choose on purpose instead of by default.
Where is it used?
Beta priors on conversion rates in A/B tests, Normal priors on lifts and regression coefficients, half-Normal priors on noise scales, Gamma priors on Poisson rates, Laplace priors on changepoint slope changes, and the default priors of tools such as Prophet, brms and Stan's prior guidance.
How is it used?
For each parameter: pick a family that matches its support, decide how much real knowledge you want to inject, set the hyperparameters, then read off what the prior claims (90% interval, tail probabilities, pseudo-count worth) and check it with a prior predictive simulation.
"The vaguer the prior, the more objective the analysis."
A vague prior still makes claims, often absurd ones (90% checkout conversion as likely as 10%; noise ten times bigger than demand). It also hurts computation and model comparison. "Objective" is not a property you get by widening.
"Weakly informative means it contains no information."
It contains a little, on purpose: enough to exclude the impossible and the absurd. That is exactly what makes estimates stable with small data.
"Gamma(0.001, 0.001) is a safe vague prior for a positive rate."
It has mean 1 and a huge sd, but 97% of its mass is below $10^{-10}$. On a log scale it is nearly flat over an enormous range, which is a strong and strange claim. Check what a prior says before trusting its name.
Informative (real knowledge) → weakly informative (rules out the absurd) → diffuse (very wide) → flat (constant density; improper if unbounded).
Strength = pseudo-observations: Beta(α, β) is worth α + β; $N(\mu_0,\tau^2)$ on a mean is worth $\sigma^2/\tau^2$.
Trap: "vague" is not "neutral"; read the 90% interval and tail probabilities of every prior.
Quick check: a prior Beta(5, 45) on a conversion rate. What is its mean, and how many visitors is it worth?
Mean $5/50 = 0.10$, worth $5 + 45 = 50$ pseudo-visitors (5 pseudo-buyers, 45 pseudo-non-buyers). A test with 50 real visitors would share the say equally with it.
Informative priors: borrowing knowledge from past experiments
Your team has run a dozen tests on the checkout page. In every one, the control converted somewhere between 8% and 12%. It would be strange to analyse test number 13 as if you had never seen a checkout page. An informative prior carries that history into the new analysis.
The danger is trusting history too much. The new page may be different; the season may be different. So it is common to discount the historical prior: keep its centre but shrink its worth, for example from 400 pseudo-visitors to 100.
Three ways to say it:
- Picture: the spread of past results becomes the shape of the new prior, then you loosen it a little.
- Numbers: 12 past rates with mean 10% and sd 1.5 points → Beta(39.9, 359.1), worth 399 visitors; discounted to Beta(10, 90), worth 100.
- Slogan: use what you know, but do not let yesterday outvote today.
Method of moments ("match the mean and the variance"). Past control rates: mean $m = 0.10$, standard deviation $0.015$, so variance $v = 0.000225$.
- A Beta(α, β) has mean $m = \alpha/(\alpha+\beta)$ and variance $v = \frac{m(1-m)}{\alpha+\beta+1}$.
- Solve the variance equation for the worth: $\alpha + \beta = \frac{m(1-m)}{v} - 1 = \frac{0.10 \times 0.90}{0.000225} - 1 = 400 - 1 = 399$.
- Split by the mean: $\alpha = 0.10 \times 399 = 39.9$, $\;\beta = 0.90 \times 399 = 359.1$. Its 90% interval is about 7.7% to 12.6%.
- Discount to a worth of 100: Beta(10, 90). Same mean, 90% interval about 5.6% to 15.3%.
- A new test with 2,000 control visitors: the discounted prior holds $100/(100 + 2000) \approx 4.8\%$ of the weight on the posterior mean; the undiscounted one would hold $399/2399 \approx 16.6\%$.
One subtlety: each past rate was itself measured with sampling noise, so the observed spread is a bit wider than the spread of the true past rates, and this simple fit errs on the cautious (wider) side. A hierarchical model (Chapter 6.5) separates the two properly.
An informative prior encodes substantial knowledge from sources other than the current data. Common sources: past experiments (historical controls), published estimates, physical or business limits, expert elicitation (asking experts for a typical value and a range, then fitting a distribution).
- Method of moments for a Beta prior: $\alpha + \beta = m(1-m)/v - 1$, $\alpha = m(\alpha+\beta)$, $\beta = (1-m)(\alpha+\beta)$ (needs $v \lt m(1-m)$).
- Discounting (the idea behind "power priors"): multiply the pseudo-counts by a factor $a_0$ between 0 and 1, so Beta($a_0\alpha$, $a_0\beta$) keeps the mean and is worth $a_0(\alpha+\beta)$.
- Prior–data conflict: when the data land far in the prior's tail, the posterior is a compromise that matches neither. Check for it (Chapter 6.8) and do not hide it.
- Informative priors on a baseline (the control rate) are common and safe; informative priors on the effect you are testing should be skeptical (centred at no effect) or you are assuming the conclusion.
Why do we need it?
Small tests and small segments give noisy answers. Real knowledge from past tests narrows the posterior honestly, which means faster, more stable decisions, as long as the history is relevant.
Where is it used?
Historical-control priors in A/B testing and clinical trials, Bayesian meta-analysis, informative priors on seasonality strength from similar products, and empirical-Bayes priors fitted to many past segments (a shortcut to the hierarchical model of Chapter 6.5).
How is it used?
Collect relevant past estimates, fit a prior to them (method of moments or maximum likelihood), discount it, and record the choice. After fitting, compare the posterior with the prior to spot conflict, and rerun with a weaker prior to see if the decision changes.
"Use this test's own early data to build the prior, then analyse the whole test with it."
That counts the early data twice and makes the posterior too confident. Priors must come from data you will not include in the likelihood.
"An informative prior is subjective, so it is unscientific."
A prior fitted to documented past experiments is as objective as any modelling choice; the honest move is to state it, discount it sensibly and show how much the result depends on it.
In an A/B framework like yours, the control arm's conversion rate is the natural place for an informative Beta prior built from past tests on the same page (remember: NumPyro Beta(concentration1=α, concentration0=β)). For the treatment effect, keep the prior skeptical or weakly informative, so that a "win" is driven by this test's data. In the forecasting model, past fits of similar series can inform priors on noise scale or seasonality strength, but only if those series really behave alike.
Beta by moments: $\alpha + \beta = m(1-m)/v - 1$; mean 0.10, sd 0.015 → Beta(39.9, 359.1), worth 399.
Discount with $a_0$: Beta($a_0\alpha$, $a_0\beta$), worth $a_0(\alpha+\beta)$.
Traps: never build the prior from the same data; watch for prior–data conflict; be skeptical on effects.
Quick check: past rates have mean 0.20 and sd 0.04. What Beta prior matches them?
$v = 0.0016$. $\alpha + \beta = 0.2 \times 0.8/0.0016 - 1 = 100 - 1 = 99$. $\alpha = 0.2 \times 99 = 19.8$, $\beta = 79.2$: Beta(19.8, 79.2), worth 99 visitors.
Weakly informative priors: rule out the absurd, let the data speak core
Most of the time you do not want a prior to steer the answer. But you also know some things for sure: a new button colour will not change checkout conversion from 10% to 80%; daily demand of about 500 will not have day-to-day noise of 5000. A weakly informative prior encodes only that kind of sanity: it is wide over everything plausible and thin over everything absurd.
The catch is that "wide" only means something relative to the units of the parameter. To judge a prior, translate it into a quantity you have a feel for: a lift in conversion, an amount of orders, a growth over a year.
Three ways to say it:
- Picture: a broad hill that still slopes to nearly zero before reaching the absurd region.
- Numbers: a $N(0, 0.5^2)$ prior on a log-odds lift, with a 10% baseline, keeps the treated rate between 4% and 23% with 95% probability.
- Slogan: wide enough to be surprised, narrow enough to stay sane.
A prior on the lift of variant B. Write B's log-odds as A's log-odds plus a lift Δ, and put $\Delta \sim N(0, s^2)$. Baseline (A) rate 10%, so $\text{logit}(0.10) = \log(0.1/0.9) \approx -2.197$.
- 95% of a Normal lies within $\pm 1.96\,s$. For $s = 0.5$: Δ between $-0.98$ and $+0.98$.
- As an odds ratio: $e^{-0.98} \approx 0.375$ to $e^{0.98} \approx 2.66$.
- As B's conversion rate: $\text{expit}(-2.197 \pm 0.98)$ = $1/(1 + e^{3.177}) \approx 0.040$ to $1/(1 + e^{1.217}) \approx 0.228$. So "B converts between 4% and 23%": big changes allowed, absurd ones not.
- $s = 0.1$ gives 8.4% to 11.9%: an informative claim that effects are small (fine if past tests show that, otherwise it would hide real effects).
- $s = 10$ gives odds ratios from about $3 \times 10^{-9}$ to $3 \times 10^{8}$: B's rate from essentially 0% to essentially 100%. Not weakly informative; absurd.
- Relative lift beyond ±50% (B below 5% or above 15%): prior probability 0.01 for $s = 0.2$, 0.24 for $s = 0.5$, 0.55 for $s = 1$, 0.95 for $s = 10$.
A weakly informative prior is chosen to be wider than your real knowledge, while placing little probability on values that are implausible on the scale of the problem. Typical choices, all rules of thumb that assume the data have been put on a sensible scale (for example standardized):
- Coefficients (effects, regressor weights, log-odds lifts): $N(0, 1)$ to $N(0, 2.5^2)$ on standardized inputs; tighter when effects are known to be small.
- Scales (noise sd, group spread τ): half-Normal or Exponential with a scale near the size of the data's own variation.
- A rate with no other information: Beta(1, 1), or $N(0, 1.5^2)$ on its log-odds (close to flat in the middle, concept 4).
The test of a weakly informative prior is not its name but what it implies for observable quantities, which a prior predictive check shows (concepts 9 and 10).
Why do we need it?
Flat or vague priors let the posterior wander into absurd regions when data are thin (separation in logistic models, huge coefficients, near-zero group spreads) and make samplers and SVI struggle. A gentle prior fixes this without imposing an opinion.
Where is it used?
Default priors in modern Bayesian practice (Stan's prior recommendations, brms and rstanarm defaults, PyMC examples), priors on regression and Fourier coefficients, priors on hierarchical spreads τ, and the scaled priors of Prophet-style models.
How is it used?
Put the data on a known scale (standardize, or divide by a typical value), choose a centred prior with a scale of about 1 on that scale, translate it back into business units to check it, and confirm with a prior predictive simulation.
"$N(0, 1)$ is always a weakly informative prior."
Only on a scale where 1 is a "large but possible" value. On a log-odds lift it is generous. On a holiday effect measured in raw orders, where real effects are in the hundreds, it is absurdly tight. On a standardized regressor coefficient it is about right. Scale first, then choose.
"A prior that looks wide on the parameter is wide on the outcome."
Non-linear links distort widths. $N(0, 10^2)$ looks harmlessly wide on the log-odds, yet on the conversion scale it is a strong claim that effects are gigantic.
Your A/B framework uses one global scaler for all groups. Besides keeping the group differences you want to estimate (per-group scaling would erase them, Chapter 4.18), it means one prior scale such as $N(0, 1)$ means the same thing in every group and segment. Prophet-style models do the same trick for forecasting: Prophet divides $y$ by its largest absolute value and maps time onto $[0, 1]$, so that fixed default priors are on a known scale. Before reading any prior scale in your forecasting model, check what scaling the code applies to $y$, to time and to the regressors.
Weakly informative = wide over the plausible, thin over the absurd. Judge it in business units.
Lift prior $N(0, 0.5^2)$ on log-odds, baseline 10% → B between 4% and 23% (95%). $N(0, 10^2)$ → B anywhere from 0% to 100%.
Trap: a prior's width depends on the scale of the parameter; standardize first.
Quick check: baseline 10%, prior $\Delta \sim N(0, 1)$. Between which conversion rates does B lie with 95% prior probability?
$-2.197 \pm 1.96$ gives $-4.157$ and $-0.237$; expit gives about 1.5% and 44%. Wide, but not absurd: weakly informative for a bold redesign, too wide for a small copy change.
Diffuse and flat priors: why "flat" is not "uninformative" core
A flat prior on a conversion rate, Beta(1, 1), feels neutral: every rate from 0 to 1 gets the same density. But you could equally describe the rate by its log-odds. Ask the same flat prior about the log-odds and it suddenly has an opinion: values near 0 (rates near 50%) are much more likely than values near ±5. Squeezing and stretching the ruler moves probability around (the same reason the MAP moved in Chapter 5.2).
It works the other way too. A "vague" prior that is flat-looking on the log-odds, $N(0, 10^2)$, is anything but flat on the rate: it piles almost two thirds of its probability onto rates below 1% or above 99%. There is no prior that is flat on every scale. So "uninformative" is not something a prior can simply be; every prior says something.
Three ways to say it:
- Picture: flat on one ruler is a hump or a U on another ruler.
- Numbers: $N(0, 10^2)$ on the log-odds gives a 65% prior chance that the rate is below 1% or above 99%.
- Slogan: flat where? There is no view from nowhere.
- Flat on p. $p \sim$ Uniform(0, 1). The log-odds $\eta = \text{logit}(p)$ has $P(|\eta| \lt 1) = P(0.269 \lt p \lt 0.731) = 0.731 - 0.269 = 0.462$ (because expit(1) ≈ 0.731). Nearly half the mass within ±1 on a scale that runs to ±∞: a hump at 0, the standard logistic distribution with sd $\pi/\sqrt3 \approx 1.81$.
- Vague on the log-odds. $\eta \sim N(0, 10^2)$. A rate below 1% means $\eta \lt \text{logit}(0.01) \approx -4.595$, probability $\Phi(-0.4595) \approx 0.323$. Same above 99%. Total $\approx 0.646$.
- Meanwhile the "reasonable" middle $0.2 \lt p \lt 0.8$ ($|\eta| \lt 1.386$) gets only $\approx 0.110$ (flat on p would give 0.6).
- A compromise. $\eta \sim N(0, 1.5^2)$: $P(0.2 \lt p \lt 0.8) \approx 0.645$ and $P(p \lt 0.05 \text{ or } p \gt 0.95) \approx 0.050$ (flat: 0.6 and 0.1). Close to flat in the middle, a little lighter at the extremes.
If $\eta = g(p)$ is a one-to-one change of parameter, densities transform with the stretching factor (the Jacobian):
$$p_\eta(\eta) = p_p\big(g^{-1}(\eta)\big)\,\Big|\frac{d\,g^{-1}(\eta)}{d\eta}\Big|.$$- For $\eta = \text{logit}(p)$, $dp/d\eta = p(1-p)$, so a flat $p_p = 1$ becomes $p_\eta(\eta) = p(1-p) = \frac{e^\eta}{(1+e^\eta)^2}$: the logistic density, a hump at 0.
- A constant density can be flat on at most one scale. "Non-informative" therefore depends on the parameterization. (One classical answer is the Jeffreys prior, built to give the same inferences whatever the parameterization; for a Bernoulli rate it is Beta(½, ½).)
- Diffuse priors such as $N(0, 1000^2)$ are nearly flat over the region where the likelihood lives, so they often give answers close to the MLE for one parameter; but on transformed or combined quantities they can be wildly informative, and in many-parameter models they produce absurd prior predictions (concept 10).
Why do we need it?
"I used a flat prior, so my result is objective" is a common and wrong claim. Knowing how flat priors behave under transformation stops you from adding hidden, strong assumptions to a model, especially logistic and log-link models.
Where is it used?
Priors on logistic-regression coefficients and lifts in conversion models, priors on log-rates in Poisson and Negative Binomial models, priors on log σ, and NumPyro's internal transforms to unconstrained space for SVI and HMC.
How is it used?
Decide on which scale you want the prior to be gentle (usually the scale people reason in: rates, orders, percentages). Write the prior on the model's scale, transform draws to the reasoning scale, and look at them before fitting.
"A flat prior means no prior, so my answer is the same as the frequentist one."
Flat on which parameter? The posterior mode with a flat prior equals the MLE only on the scale where the prior is flat (Chapter 5.2); posterior means and intervals still depend on the scale you called flat.
"$N(0, 1000^2)$ on every coefficient is a safe default."
Through a logit or log link it becomes an extreme claim about rates and counts, and with many coefficients it imagines absurd data. It can also slow down or destabilize sampling and SVI.
"An improper flat prior is always fine as long as the computer runs."
The posterior must integrate to 1. With improper priors and weak data (for example a group spread τ in a hierarchical model) it may not, and samplers can then wander or report nonsense without an error.
"I used an uninformative flat prior."
"I used a prior that is flat on the probability scale, which is not flat on the log-odds scale; I checked that it implies sensible values for the quantities I care about."
Model answer: "Flatness is not invariant to reparameterization: densities pick up a Jacobian. Uniform(0, 1) on a rate is a logistic hump on the log-odds, and a wide Normal on the log-odds is a U-shape on the rate, with about 65% of its mass below 1% or above 99% for sd 10. So I choose weakly informative priors on a sensible scale and confirm them with prior predictive simulation."
$p_\eta(\eta) = p_p(p)\,|dp/d\eta|$; logit: $dp/d\eta = p(1-p)$. Flat on p → logistic hump on η.
$N(0, 10^2)$ on logit → $P(p \lt 1\% \text{ or } \gt 99\%) \approx 0.65$. $N(0, 1.5^2)$ → close to flat in the middle.
Trap: no prior is flat on every scale; "uninformative" is not a property of a density.
Quick check: with $\eta \sim N(0, 5^2)$ on the log-odds, what is the prior probability that the rate is below 1% or above 99%?
$2\Phi(-4.595/5) = 2\Phi(-0.919) \approx 2 \times 0.179 = 0.358$. Still more than a third of the prior on extreme rates.
Conjugate priors: chosen for convenience, not for truth
Some priors and likelihoods fit together like puzzle pieces: multiply them and the posterior is the same kind of distribution as the prior, just with updated numbers. A Beta prior times a Binomial likelihood is a Beta. Updating then means adding counts, no grids, no samplers. Such a prior is called conjugate to that likelihood ("conjugate" = "joined together").
That is a computational gift, not a fact about the world. A conjugate prior is a good choice when its shape can express what you believe; if it cannot, use another prior and let MCMC or SVI do the work. In practice, two priors with the same centre and spread usually give nearly the same posterior.
Three ways to say it:
- Picture: the prior and the posterior come from the same family; only the knobs move.
- Numbers: Beta(2, 18) + 3 of 20 → Beta(5, 35), mean 0.125; a non-conjugate prior with the same mean and sd gives 0.120.
- Slogan: conjugacy is about arithmetic, not about belief.
Conjugate versus non-conjugate, same knowledge. Data: 3 buyers in 20 visitors.
- Conjugate: prior Beta(2, 18) (mean 0.10, sd 0.065). Posterior Beta(5, 35) by adding counts: mean 0.125, 90% interval 0.052 to 0.220.
- Non-conjugate: a "logit-Normal" prior, $\text{logit}(\theta) \sim N(-2.39, 0.715^2)$, chosen to have the same mean (0.10) and sd (0.065). The posterior has no named form; a grid gives mean 0.120, 90% interval 0.050 to 0.217.
- The evidence is almost identical too: 0.1354 versus 0.1349.
- Conclusion: the shape family mattered little; the information (centre and spread) mattered. Conjugacy just made the first one free to compute.
A family of priors $\mathcal F$ is conjugate to a likelihood if a prior in $\mathcal F$ always gives a posterior in $\mathcal F$. The common pairs (derivations and predictive distributions in Chapter 6.3):
| Likelihood (data) | Conjugate prior | Posterior |
|---|---|---|
| Binomial: $k$ successes in $n$ | Beta(α, β) on θ | Beta(α + k, β + n − k) |
| Multinomial: counts $c_1..c_K$ | Dirichlet(α₁..α_K) | Dirichlet(α₁ + c₁, …, α_K + c_K) |
| Poisson: counts $y_1..y_n$ | Gamma(α, rate β) on λ | Gamma(α + Σy, β + n) |
| Normal, σ known: $y_1..y_n$ | $N(\mu_0, \tau^2)$ on μ | Normal, precision $1/\tau^2 + n/\sigma^2$ |
- Conjugate priors have hyperparameters that read as pseudo-data: for the posterior mean, Beta(α, β) acts like α pseudo-successes and β pseudo-failures, so α + β pseudo-trials as a "worth" (for the posterior mode it acts like α − 1 and β − 1).
- Once you add a hierarchy, a Student-t likelihood, a log link or several interacting parameters, conjugacy is usually lost and you need MCMC or variational inference (Chapters 6.9–6.12).
Why do we need it?
Exact, instant posteriors: no sampler, no convergence checks, trivial sequential updating. That makes conjugate models ideal for dashboards that update constantly and for teaching the logic of Bayesian updating.
Where is it used?
Beta-Binomial conversion metrics and Dirichlet-Multinomial category metrics in Bayesian A/B testing, Thompson sampling for bandits, Gamma-Poisson for count rates, Kalman filters (Normal-Normal), and naive Bayes with smoothing (Dirichlet pseudo-counts).
How is it used?
If the conjugate family can express your knowledge, use it and update by formula. If it cannot (you want heavier tails, or the parameter enters through a link), choose the prior you believe and fit with a sampler or SVI; compare both if unsure.
"Beta priors are used because conversion rates really follow a Beta distribution."
They are used because Beta × Binomial = Beta, which makes updating exact and instant. Any prior with a similar centre and spread would give a similar posterior.
"If no conjugate prior exists, the model cannot be fitted."
Grids (1–2 parameters), MCMC and variational inference handle any prior. That is what NumPyro is for.
Your A/B framework uses the two conjugate workhorses: Beta-Binomial for conversion metrics and Dirichlet-Multinomial for categorical metrics, so, for a single variant without pooling, those posteriors are exact and can be updated from running counts. The Normal, Student-t and Poisson likelihoods with hierarchical partial pooling across segments are not conjugate as a whole, which is why the framework also needs NumPyro and SVI.
Conjugate: prior and posterior in the same family. Beta–Binomial, Dirichlet–Multinomial, Gamma–Poisson, Normal–Normal (σ known).
Same centre and spread → similar posterior, conjugate or not (0.125 vs 0.120 in the example).
Trap: conjugacy is a convenience, not evidence that the prior is right.
Quick check: daily order counts 4, 7, 5 with a Gamma(2, rate 0.5) prior on the Poisson rate λ. What is the posterior?
Gamma(α + Σy, β + n) = Gamma(2 + 16, 0.5 + 3) = Gamma(18, 3.5), mean $18/3.5 \approx 5.14$ orders per day (the data mean is 5.33; the prior mean is 4). In SciPy: gamma(a=18, scale=1/3.5).
Prior strength versus data size: the tug of war core
Picture a rope with the prior pulling at one end and the data at the other. The posterior is the knot in the middle. Each side pulls with a strength measured in observations: the prior with its pseudo-observations, the data with their real ones. A prior worth 100 visitors wins against 20 real visitors, ties with 100, and is overwhelmed by 20,000.
For the Beta-Binomial this is exact: the posterior mean is a weighted average of the prior mean and the data rate, with weights proportional to the two strengths.
Three ways to say it:
- Picture: a tug of war; the knot moves toward the side with more people.
- Numbers: prior worth 100 at 10%, data 200 visitors at 15% → posterior mean 13.3%: two thirds of the way to the data.
- Slogan: priors matter most where data are thin.
Prior Beta(10, 90): mean $m_0 = 0.10$, worth $n_0 = 100$.
- Data: 30 buyers in 200 visitors (rate 0.15). Posterior Beta(10 + 30, 90 + 170) = Beta(40, 260), mean $40/300 \approx 0.1333$.
- As a weighted average: data weight $n/(n + n_0) = 200/300 = 2/3$, prior weight $1/3$. $\tfrac23(0.15) + \tfrac13(0.10) = 0.10 + 0.0333 = 0.1333$ ✓.
- Same rate, 100 times the data (3,000 of 20,000): mean $(10 + 3000)/(100 + 20000) = 3010/20100 \approx 0.1498$; data weight $20000/20100 \approx 0.995$.
- Normal version: daily orders with noise σ = 20 and a prior $\mu \sim N(100, 10^2)$. The prior is worth $\sigma^2/\tau^2 = 400/100 = 4$ days of data. After 4 days it holds half the weight; after 36 days a tenth.
Beta-Binomial. Prior Beta(α, β), with $n_0 = \alpha + \beta$ and $m_0 = \alpha/n_0$; data $k$ of $n$:
$$E[\theta\mid D] = \frac{\alpha + k}{n_0 + n} = \underbrace{\frac{n_0}{n_0 + n}}_{\text{prior weight}}\,m_0 + \underbrace{\frac{n}{n_0 + n}}_{\text{data weight}}\,\frac{k}{n}.$$Normal-Normal (σ known, prior $N(\mu_0, \tau^2)$): precisions add, and the mean is precision-weighted:
$$\frac{1}{\tau_n^2} = \frac{1}{\tau^2} + \frac{n}{\sigma^2}, \qquad E[\mu\mid D] = \tau_n^2\Big(\frac{\mu_0}{\tau^2} + \frac{n\bar y}{\sigma^2}\Big).$$- As $n \to \infty$ (with a prior that is positive near the truth), the prior's weight goes to 0, the posterior concentrates near the true value, and it becomes approximately Normal with sd close to the usual standard error. This large-sample result (the Bernstein–von Mises theorem) needs regularity conditions, such as a fixed number of parameters.
- But in real models some parameters always have little data: small segments, a holiday seen twice, changepoints near the end of the history. For those, the prior keeps mattering however big the total dataset is.
Why do we need it?
To know when the prior choice deserves careful thought and when it does not, and to explain to stakeholders, in units they understand ("worth 100 visitors"), how much the prior influenced a result.
Where is it used?
Reporting the effective sample size of priors in A/B tests, choosing discount factors for historical priors, sample-size planning for Bayesian tests, and judging which parameters of a large forecasting model are prior-dominated.
How is it used?
Convert each prior into pseudo-observations, compare with the data each parameter actually sees, and spend your care (and your sensitivity checks, Chapter 6.8) on parameters where the prior holds a large share.
"With a big dataset the prior never matters."
For parameters that see lots of data, yes. But big models contain parameters that see little data (a small segment, a rare holiday, the last changepoint, a tail probability near a decision threshold). There the prior keeps a real say.
"A prior worth 100 visitors has the same effect in every test."
Its share is $n_0/(n_0 + n)$: 83% against 20 visitors, 0.5% against 20,000. Always state the prior's worth next to the data size.
In your A/B framework, a big test overwhelms any reasonable prior on the overall rates, but a segment with 40 users is dominated by its prior or, with partial pooling, by the population of segments (Chapter 6.6). In your forecasting model, the trend slope over a long history is data-dominated, while parameters such as a holiday seen twice or a changepoint near the end of the series stay prior-dominated, which is exactly where prior sensitivity checks matter (Chapter 6.8).
$E[\theta\mid D] = \frac{n_0}{n_0+n}m_0 + \frac{n}{n_0+n}\frac kn$, with $n_0 = \alpha+\beta$.
Normal: precisions add; prior $N(\mu_0,\tau^2)$ is worth $\sigma^2/\tau^2$ observations.
Trap: "lots of data" is per parameter; small segments and rare events stay prior-sensitive.
Quick check: prior Beta(20, 180), data 50 buyers in 200 visitors. What is the data's weight and the posterior mean?
$n_0 = 200$, $n = 200$, weight $1/2$. Posterior Beta(70, 330), mean $70/400 = 0.175 = \tfrac12(0.10) + \tfrac12(0.25)$.
Shrinkage and regularizing priors: Normal and Laplace core
Many unknowns in a model are effects: a holiday bump, a regressor's weight, a change in trend slope, a segment's difference from average. Experience says most effects are small and a few are big. A prior centred at zero says exactly that, and it gently pulls noisy estimates toward zero: this pull is called shrinkage. Because it stops the model from chasing noise, such a prior is also called regularizing, the Bayesian twin of ridge and lasso (Chapter 5.3).
The shape of the prior decides how it pulls. A Normal prior shrinks every estimate by the same fraction, small or big. A Laplace prior, with its sharp peak at 0 and heavier tails, pulls small effects hard toward zero and leaves big effects nearly alone. That is why a Laplace prior suits changepoint slope changes: most are about zero, a few are real.
Three ways to say it:
- Picture: a spring pulling each estimate toward 0; Normal = the same spring for all, Laplace = strong for small effects, weak for big ones.
- Numbers: with equal prior variance, a noisy estimate of 1 becomes 0.50 (Normal) or 0.38 (Laplace); an estimate of 4 becomes 2.0 (Normal) or 2.6 (Laplace).
- Slogan: shrinkage trades a little bias for a lot less noise.
One effect θ, measured once with noise: $x \sim N(\theta, 1)$. Two priors with the same variance 1: Normal $N(0, 1)$ and Laplace(0, b) with $2b^2 = 1$, so $b = 1/\sqrt2 \approx 0.707$.
- Normal prior, Normal data: posterior mean $= x\cdot\frac{\tau^2}{\tau^2 + s^2} = x/2$. So $x = 1 \to 0.50$, $\;x = 4 \to 2.00$, $\;x = 6 \to 3.00$: always half.
- Laplace prior (posterior mean by a grid): $x = 1 \to 0.38$ (more shrinkage), $\;x = 4 \to 2.59$, $\;x = 6 \to 4.59$ (less shrinkage: roughly $x - 1.41$).
- Laplace MAP (the posterior's peak) is soft-thresholding: $\text{sign}(x)\max(0, |x| - s^2/b) = \text{sign}(x)\max(0, |x| - 1.414)$: exactly 0 for $|x| \le 1.414$.
- But the Laplace posterior mean at $x = 1$ is 0.38, not 0: the full posterior is not sparse (Chapter 5.3).
- Shape check: the Laplace prior puts more mass near 0 ($P(|\theta| \lt 0.25) = 0.30$ vs 0.20) and more in the tails ($P(|\theta| \gt 3) = 0.014$ vs 0.0027).
A shrinkage prior concentrates around a reference value (usually 0, or a group mean), pulling estimates toward it. Its negative log is a penalty added to the negative log-likelihood, which is why it regularizes:
$$-\log N(\theta\mid 0, \tau^2) = \frac{\theta^2}{2\tau^2} + c \;\;(\text{L2, ridge}), \qquad -\log \text{Laplace}(\theta\mid 0, b) = \frac{|\theta|}{b} + c \;\;(\text{L1, lasso}).$$- Normal prior, Normal data ($x \sim N(\theta, s^2)$): posterior mean $= x\,\tau^2/(\tau^2 + s^2)$, a constant shrink factor. In linear regression this MAP is ridge with $\lambda = \sigma^2/\tau^2$.
- Laplace prior: MAP = soft-thresholding (exact zeros); posterior mean and median are shrunk but not exactly zero; large effects are shrunk by roughly a constant amount $s^2/b$.
- The scale (τ or b) is the flexibility knob: small scale = strong shrinkage (risk: underfitting real effects); large scale = weak shrinkage (risk: chasing noise).
- Other shrinkage priors with "many near 0, a few large" behaviour exist (Student-t, horseshoe); hierarchical priors (next concept) learn the scale from the data.
Why do we need it?
Models with many effects and limited data overfit: every noisy blip becomes a "holiday effect" or a "trend change". Shrinkage priors keep estimates sensible and forecasts stable, and lower the average error (the bias–variance trade-off of Chapter 5.1).
Where is it used?
Laplace priors on changepoint slope changes in Prophet-style trends, Normal priors on Fourier and holiday coefficients, ridge and lasso regression, weight decay in neural networks, partial pooling of segment effects, and sparse regression with horseshoe priors.
How is it used?
Centre the prior at "no effect", choose its scale relative to the size of effects you consider plausible (on standardized data), and treat that scale as a tuning knob: check it with a prior predictive simulation and by out-of-sample error.
"A Laplace prior sets most changepoint slope changes exactly to zero."
Only the MAP (the posterior's peak) has exact zeros. The full posterior, and therefore the draws your SVI guide or NUTS produce, has no point mass at zero: the changes are shrunk toward zero, "sparse-ish", not sparse.
"Shrinkage biases the estimates, so it makes the model worse."
It adds a little bias and removes a lot of variance. Averaged over many effects, the error goes down (Chapter 5.1). It hurts only when the prior scale is far too small for the real effects.
"Normal and Laplace priors with the same variance shrink the same way."
Same variance, different shapes: Laplace shrinks small estimates more and large estimates less (0.38 vs 0.50 at x = 1; 2.59 vs 2.00 at x = 4).
In your forecasting model the changepoint slope adjustments have Laplace priors, $\delta_j \sim Laplace(0, b)$: the prior says "most candidate changepoints do nothing; a few matter". The scale b is the flexibility knob of the trend (Chapter 7.10): too small and real trend changes are flattened; too large and the trend chases noise. Normal priors on Fourier and holiday coefficients play the same regularizing role. In the A/B framework, partial pooling shrinks each segment toward the population mean, with a strength learned from the data (next concept and Chapter 6.6).
Shrinkage prior: centred at "no effect"; $-\log$ prior = penalty. Normal ↔ L2/ridge, Laplace ↔ L1/lasso (at the MAP).
Normal: $E[\theta\mid x] = x\tau^2/(\tau^2+s^2)$. Laplace: MAP soft-thresholds at $s^2/b$; posterior mean not exactly 0.
Trap: "Laplace prior ⇒ sparse posterior" is false; the scale is a bias–variance knob.
Quick check: Normal prior $N(0, 2^2)$ and an estimate $x = 3$ with noise sd 2. What is the posterior mean, and how much was it shrunk?
Factor $\tau^2/(\tau^2 + s^2) = 4/8 = 0.5$, so $E[\theta\mid x] = 1.5$: shrunk by half, because the prior and the measurement are equally precise.
Hierarchical priors: let the groups set each other's prior
Your A/B test runs in 8 country segments. What prior should segment 7's conversion rate get? You could pick one by hand for each segment, but the best evidence about segment 7's likely rate is the other seven segments. A hierarchical prior says: every segment's rate comes from one shared population distribution, with a centre μ and a spread τ that are themselves unknown and learned from all segments together.
So the prior for each segment is not fixed in advance; it is estimated from the family of segments. If the segments turn out similar (small τ), each one is pulled strongly toward the common centre; if they turn out very different (large τ), each mostly keeps its own data. The full story (pooling and shrinkage) is Chapters 6.5–6.7; here we meet it as a kind of prior.
Three ways to say it:
- Picture: a family tree: hyperpriors at the top, a population (μ, τ) in the middle, the segments below, the data at the bottom.
- Numbers: μ = −2.0 and τ = 0.2 on the log-odds give segment rates like 9.6%, 12.6%, 13.7%, 10.9%: siblings, not strangers.
- Slogan: a prior with learnable knobs; the groups teach each other.
Segment log-odds $\eta_g = \mu + \tau z_g$ with $z_g \sim N(0, 1)$ (the same as $\eta_g \sim N(\mu, \tau^2)$). Suppose the population draw is μ = −2.0, τ = 0.2, and four segments get $z = -1.2, 0.3, 0.8, -0.5$.
- Log-odds: $-2.0 + 0.2(-1.2) = -2.24$; $\;-2.0 + 0.06 = -1.94$; $\;-2.0 + 0.16 = -1.84$; $\;-2.0 - 0.10 = -2.10$.
- Rates, $1/(1 + e^{-\eta})$: 9.6%, 12.6%, 13.7%, 10.9%.
- With τ = 1 instead, the same $z$'s give log-odds $-3.2, -1.7, -1.2, -2.5$, rates 3.9%, 15.4%, 23.2%, 7.6%: segments that barely resemble each other.
- So the prior on τ is a prior on how similar the segments are. A half-Normal(0.3) prior on τ says "segment rates usually differ by a few points, rarely by a factor of two or more".
A hierarchical (multilevel) prior has two or more levels:
$$\theta_g \mid \mu, \tau \sim N(\mu, \tau^2)\ \ (g = 1, \dots, G), \qquad \mu \sim p(\mu), \quad \tau \sim p(\tau).$$- $\theta_g$: group parameters (one per segment). μ and τ: hyperparameters (population centre and spread), shared by all groups. $p(\mu)$, $p(\tau)$: hyperpriors.
- τ controls pooling: τ → 0 makes all groups equal (complete pooling); τ → ∞ makes them independent (no pooling). In between, each group is shrunk toward μ, more when its own data are thin.
- Because τ is learned, the amount of shrinkage is set by the data: a shrinkage prior with an adaptive scale.
- With few groups (say 3–8), the data say little about τ, so its hyperprior matters; half-Normal or half-Student-t priors on τ are common weakly informative choices.
Why do we need it?
Fixed per-group priors either ignore what other groups say (too noisy for small groups) or force all groups to be equal (too rigid). A hierarchical prior learns from the data how much groups should share.
Where is it used?
Segment- and country-level effects in A/B tests, store- or product-level forecasts, the eight-schools example, ratings with few reviews, multi-armed bandits with many arms, and random effects in mixed models.
How is it used?
Write group parameters as draws from a population distribution with learnable μ and τ, put weakly informative hyperpriors on them, and fit everything jointly (NUTS or SVI). Often written in the non-centred form $\theta_g = \mu + \tau z_g$ for better geometry (Chapter 6.7).
"A hierarchical prior uses the data to build the prior, so it double-counts the data."
It is one joint model, $p(\mu, \tau)\prod_g p(\theta_g\mid\mu,\tau)\,p(y_g\mid\theta_g)$, fitted once. Each data point enters the likelihood exactly once; the sharing happens through the shared parameters.
"The hyperprior on τ does not matter; the data will decide."
The data about τ are the number of groups, not the number of users. With 5 segments, τ is weakly identified and its hyperprior visibly changes how much shrinkage you get. Check it (Chapter 6.8).
Partial pooling across groups and segments in your A/B framework is exactly this: segment-level parameters drawn from a shared population, with the spread τ learned. A small segment is pulled toward the population (borrowing strength), a large segment keeps its own estimate. The global scaler makes a single prior on μ and τ meaningful across all segments. Chapters 6.5–6.7 cover the model, the shrinkage formula and the centred vs non-centred parameterization.
$\theta_g \sim N(\mu, \tau^2)$, $\mu \sim p(\mu)$, $\tau \sim p(\tau)$: group parameters, hyperparameters, hyperpriors.
τ = how similar groups are; small τ → strong pooling. The prior on τ is a prior on similarity.
Trap: few groups → τ weakly identified → its hyperprior matters.
Quick check: what happens to the segment estimates as τ → 0 and as τ → ∞?
τ → 0: all $\theta_g$ are forced equal to μ, so every segment gets the overall rate (complete pooling). τ → ∞: the population says nothing, and each segment uses only its own data (no pooling). A learned τ in between gives partial pooling.
Prior predictive checks: does the model imagine plausible worlds? core
It is hard to look at "$\Delta \sim N(0, 10^2)$ on the log-odds" and know whether it is sensible. It is easy to look at a fake A/B test result and know: "a control group converting 99.7% of 1,000 visitors? Impossible." So judge the model where your intuition lives: on data.
A prior predictive check runs the model forwards before seeing any real data: draw parameters from the prior, then draw a fake dataset from the likelihood, many times. Then look at the fake datasets. If the model routinely imagines worlds that could never happen (0% or 100% conversion, negative demand, demand ten times last year's record), something in the priors or the likelihood needs fixing. It catches problems that no single prior shows on its own, because the trouble often comes from priors combining.
Three ways to say it:
- Picture: let the model daydream before it sees the data, and read its daydreams.
- Numbers: with vague $N(0, 10^2)$ priors on the log-odds, 84% of simulated control groups convert below 1% or above 40%; with weakly informative priors, none do.
- Slogan: before fitting, ask "what kind of datasets does my model think are likely?"
An A/B model: control log-odds $\alpha \sim N(\mu_0, s_\alpha^2)$, lift $\Delta \sim N(0, s_\Delta^2)$, then $k_A \sim \text{Binomial}(1000, \text{expit}(\alpha))$ and $k_B \sim \text{Binomial}(1000, \text{expit}(\alpha + \Delta))$. Domain knowledge: checkout rates in this product are between about 2% and 30%; tests rarely move conversion by more than ±20% relative. Call a fake test "absurd" if its control rate is outside 1%–40% or its relative lift is beyond ±50%.
- Vague priors $\mu_0 = 0$, $s_\alpha = 10$, $s_\Delta = 10$: about 84% of simulated control rates fall outside 1%–40% (many are exactly 0 or 1000 buyers), and about 62% of fake tests show a lift beyond ±50% or have no control buyers at all.
- Weakly informative $\mu_0 = -2.2$ (logit of 10%), $s_\alpha = 0.5$, $s_\Delta = 0.2$: control rates between about 4.6% and 20% (90%), no absurd control groups, and about 4% of fake tests show a relative lift beyond ±50% (mostly sampling noise of 1,000 visitors on top of the prior's lift).
- Too tight $s_\Delta = 0.02$: the prior says every lift is within about ±4% relative. Plausible-looking data, but it would hide a real 10% lift. Prior predictive checks also catch priors that are too narrow.
The prior predictive distribution is the distribution of data implied by the model before any observations:
$$p(y) = \int p(y\mid\theta)\,p(\theta)\,d\theta .$$(Its value at the observed data is the evidence of Chapter 6.1.) A prior predictive check:
- For $s = 1,\dots,S$: draw $\theta^{(s)} \sim p(\theta)$, then $y^{(s)} \sim p(y\mid\theta^{(s)})$ (same shapes and sizes as the real data).
- Compute summaries $T(y^{(s)})$ you can judge: rates, lifts, means, maxima, share of zeros, growth over a year, minimum value.
- Compare with domain knowledge (plausible ranges, hard limits), not with the detailed observed data.
- If many worlds are absurd, change priors, scales or the likelihood/link (for example a log link to keep demand positive); if every world is the same, the priors may be too tight. Repeat.
In NumPyro: Predictive(model, num_samples=1000)(rng_key, ...) without posterior samples simulates from the prior.
Why do we need it?
Priors on transformed scales (log-odds, log-rates, scaled time) and priors that combine (many coefficients added together) are impossible to judge one by one. Simulated data are easy to judge, and catching nonsense before fitting avoids slow, unstable fits and silently wrong forecasts.
Where is it used?
The standard Bayesian workflow (Stan, PyMC, NumPyro) before every fit; logistic and Poisson regression with many coefficients; hierarchical models (prior on τ); forecasting models with trend, changepoints and seasonality; checking that a likelihood cannot produce impossible values.
How is it used?
Call Predictive(model, num_samples=…) with the real data's shapes but no observations, compute a handful of interpretable summaries, plot them against plausible ranges, and adjust priors until most simulated worlds are plausible yet still varied.
"A prior predictive check means adjusting the prior until the simulated data look like my data."
That would tune the prior to the data and count them twice. Use only knowledge you would have without this dataset: plausible ranges, physical limits, past experience. Comparing the fit to the data is the job of posterior predictive checks (Chapter 6.8).
"If every prior looks reasonable on its own, the model's predictions will be reasonable."
Priors combine. Ten "reasonable" coefficients added together on a log scale can imply demand multiplied by a million. Only simulating the whole model shows this.
"The simulated data should look exactly like real data."
They should be plausible and usually more varied than your data (the prior covers worlds you have not seen). Absurd worlds are the warning sign, and so are worlds that are all identical.
"Prior predictive checks are for people who are unsure about their priors."
"Prior predictive checks test the whole model's assumptions on the data scale, where priors and likelihood interact; I run them before every fit."
Model answer: "I simulate parameters from the prior and data from the likelihood, with the real data's shapes but no observations, and look at summaries I can judge: conversion rates and lifts for an experiment, levels, growth and minimum values for a forecast. If the model often generates impossible worlds, like negative demand or 100% conversion, I tighten or rescale the priors or change the likelihood. It uses domain knowledge, not the observed data, so it does not overfit the prior."
Prior predictive: $p(y) = \int p(y\mid\theta)p(\theta)d\theta$; simulate θ ~ prior, then y ~ likelihood.
Judge summaries against domain knowledge; fix priors, scales or the likelihood; repeat. NumPyro: Predictive(model, num_samples=…).
Traps: never tune the prior to the observed data; check combinations, not single priors; "too tight" is also a failure.
Quick check: your simulated A/B tests look sensible except that 30% have exactly 0 control buyers among 1,000 visitors. What is the likely problem?
The prior on the control log-odds puts too much mass on very negative values (rates far below 0.1%). Re-centre it on a realistic baseline (for example logit(0.10) ≈ −2.2) and narrow its sd, then rerun the check.
A prior predictive check for a forecasting model: trend, changepoints, noise core
A forecasting model has many priors: the starting level, the slope, a slope change at each candidate changepoint, the noise. Each one alone can look harmless ("slope ~ Normal(0, 1000), that is just vague"). Put them together and simulate a year of daily demand, and the model may imagine shops whose demand goes negative in March and reaches 4,000 orders a day by December. Those are the worlds the model considers plausible before it sees data, and, more worrying, the kind of futures it can extrapolate into after it has seen data, because the end of the trend is weakly pinned down.
The check is the same as before: draw parameters from the priors, draw a year of data, and look. Then adjust the scales until the simulated years look like years your business could have, still varied, never absurd.
Three ways to say it:
- Picture: a spaghetti plot of imagined years; you want a plausible bundle, not an explosion and not a single strand.
- Numbers: vague priors: about 94% of simulated years leave the range 0–2,000 orders a day; weakly informative priors: about 1–2%.
- Slogan: wide trend priors produce absurd demand; check before you fit.
The model (time $t$ in years, $t = 0$ to $364/365$; demand around 500 orders a day; 10 candidate changepoints $s_j = j/11$):
$$g(t) = m + k\,t + \sum_{j=1}^{10}\delta_j\,(t - s_j)_+, \qquad y_t \sim N\big(g(t), \sigma^2\big),$$where $(x)_+ = \max(0, x)$, so each $\delta_j$ changes the slope from $s_j$ on (the trend stays continuous; Chapter 7.8). Plausible world: every day between 0 and 2,000 orders.
- Vague: $m \sim N(500, 1000^2)$, $k \sim N(0, 1000^2)$, $\delta_j \sim Laplace(0, 500)$, $\sigma \sim HalfNormal(500)$. About 94% of simulated years leave 0–2,000 at least once; the end-of-year trend level ranges from about −2,600 to +3,600 (90%). Most of these worlds are impossible.
- Weakly informative: $m \sim N(500, 100^2)$, $k \sim N(0, 150^2)$, $\delta_j \sim Laplace(0, 30)$, $\sigma \sim HalfNormal(30)$. About 1–2% of simulated years leave the range; the end-of-year trend level is between about 180 and 820 (90%). Growth and decline both possible; nothing absurd.
- Too tight: $m \sim N(500, 5^2)$, $k \sim N(0, 5^2)$, $\delta_j \sim Laplace(0, 1)$, $\sigma \sim HalfNormal(5)$. No absurd years, but every year is a flat line near 500 (end of year 488 to 512): this model cannot learn growth, seasonality aside.
- Even the weakly informative model sometimes dips below 0, because a Normal likelihood allows negative demand. If that matters, the fix is in the likelihood (a log link, or a Negative Binomial for counts, Chapter 7.13), not only in the priors.
Prior predictive check for a time-series model: simulate whole series $y^{(s)}_{1:T} \sim p(y_{1:T})$ (parameters from the priors, then data from the likelihood), on the same time grid and with the same regressors as the real data, and examine:
- Range: minimum and maximum (negative counts? values above any physical capacity?).
- Growth: change in level over the history (doubling every month? collapsing to zero?).
- Wiggliness: how often and how sharply the trend bends (set by the changepoint scale $b$); size of seasonal swings (Fourier coefficient priors); size of holiday effects.
- Noise: day-to-day variation relative to the level.
Check components separately as well (trend only, seasonality only), then together. Re-check after rescaling the data, because prior scales are relative to the scaling.
Why do we need it?
A forecast extrapolates the trend beyond the data, where the priors on the last slope changes still carry weight. Priors that imagine absurd trends produce absurd forecast intervals; SVI also converges more slowly and less reliably when priors are needlessly wide.
Where is it used?
Prophet-style additive models (trend + changepoints + Fourier seasonality + holidays + regressors), Bayesian structural time-series models, state-space models with priors on variances, and hierarchical forecasts across many products or stores.
How is it used?
Run Predictive(model, num_samples=…) with the real time index and regressors but y=None, plot 10–20 simulated series against plausible bands, compute summaries such as the share of series that ever leave the plausible range, and adjust prior scales (or the likelihood) until the bundle looks like a believable set of years.
"Wide trend priors are harmless; the data will override them."
Over the history, mostly yes. But the forecast lives beyond the data, where the last slope changes are weakly pinned down by few points, and wide priors there mean wide, sometimes absurd, forecast intervals. They also make SVI and NUTS slower and less stable.
"If the prior predictive shows negative demand, the priors are wrong."
Maybe, but it may be the likelihood: a Normal likelihood can always go negative. Positive counts call for a log or softplus link, or a Negative Binomial likelihood (Chapter 7.13).
"The tighter the priors, the safer the forecast."
Too tight, and the model cannot follow real growth or real trend changes (underfitting). The goal is a varied bundle of plausible worlds, not a single line.
This is the syllabus warning made concrete: in your forecasting model, wide priors on the base slope and on the changepoint adjustments $\delta_j \sim Laplace(0, b)$ can generate absurd demand before any data are seen. Run Predictive(model, num_samples=…) with your real time index, regressors and holiday indicators but without y, and plot the simulated series on the scale your model works on. A reference point: Prophet's published Stan model puts $N(0, 5^2)$ priors on the growth rate and offset and $Laplace(0, \tau)$ on the changepoint deltas with τ = changepoint_prior_scale (default 0.05), all on data scaled by its maximum absolute value with time mapped to [0, 1]; whatever your model's scaling is, its prior scales only mean something relative to it. The choice of b itself is Chapter 7.10.
Simulate whole series from the priors with the real time grid and regressors; check range, growth, wiggliness, noise.
Vague trend priors: ~94% of simulated years absurd; weakly informative: ~1–2%; too tight: every year the same.
Traps: forecasts extrapolate where priors still matter; negative values may be a likelihood problem; prior scales are relative to the data scaling.
Quick check: with $k \sim N(0, 150^2)$ (orders/day per year), how big a change in level over one year does the prior allow with about 95% probability (ignoring changepoints)?
The level changes by $k \times 1$ year, so about $\pm 1.96 \times 150 \approx \pm 294$ orders a day: from roughly 200 to 800 for a shop starting at 500. Generous, but believable.
Recap, cheat sheet and practice
- Priors range from informative (real knowledge, worth many observations) through weakly informative (rule out the absurd) to diffuse and flat. Measure strength in pseudo-observations.
- Informative priors from past experiments: fit by moments, then discount; keep priors on effects skeptical; watch for prior–data conflict.
- Weakly informative priors only make sense on a known scale: standardize (a global scaler), then translate the prior into business units.
- Flat is not uninformative: densities change with the ruler. Flat on p is a hump on the log-odds; $N(0, 10^2)$ on the log-odds is a U on p (65% below 1% or above 99%).
- Conjugate priors make updating exact (Beta–Binomial, Dirichlet–Multinomial, Gamma–Poisson, Normal–Normal); they are a convenience, not a truth.
- Prior strength vs data: posterior mean = weighted average with weights $n_0$ and $n$; priors keep mattering for parameters with little data.
- Shrinkage / regularizing priors pull effects toward 0: Normal ↔ ridge, Laplace ↔ lasso at the MAP; the Laplace posterior is shrunk but not sparse. Hierarchical priors learn the amount of shrinkage.
- Prior predictive checks simulate data from the priors before fitting; fix priors, scales or the likelihood until the imagined worlds are plausible and varied.
Cheat sheet
| Idea | Formula / example | Remember |
|---|---|---|
| Prior strength | Beta(α, β): worth α + β; $N(\mu_0,\tau^2)$ on a mean: worth $\sigma^2/\tau^2$ | compare with the data each parameter sees |
| Informative from history | $\alpha+\beta = m(1-m)/v - 1$; discount by $a_0$ | mean 10%, sd 1.5 pts → worth 399 |
| Weakly informative | lift $N(0, 0.5^2)$ on log-odds → B in 4%–23% (baseline 10%) | judge in business units |
| Change of ruler | $p_\eta(\eta) = p_p(p)\,|dp/d\eta|$, logit: $dp/d\eta = p(1-p)$ | flat on one scale ≠ flat on another |
| Conjugate pairs | Beta–Bin, Dir–Mult, Gamma–Pois, Normal–Normal | exact updates; Chapter 6.3 |
| Tug of war | $E[\theta\mid D] = \frac{n_0}{n_0+n}m_0 + \frac{n}{n_0+n}\frac kn$ | Beta(10, 90) + 30/200 → 0.133 |
| Shrinkage | Normal: $x\tau^2/(\tau^2+s^2)$; Laplace MAP: soft-threshold at $s^2/b$ | Laplace: more pull on small, less on big |
| Hierarchical | $\theta_g \sim N(\mu,\tau^2)$, hyperpriors on μ, τ | τ = similarity of groups |
| Prior predictive | $p(y) = \int p(y\mid\theta)p(\theta)d\theta$; Predictive(model, num_samples=…) | compare with domain limits, not with the data |
import numpy as np
from scipy import stats
from scipy.special import expit, logit
import jax
import jax.numpy as jnp
import numpyro
import numpyro.distributions as dist
from numpyro.infer import Predictive
# 1) Flat is not uninformative: the same priors seen on two rulers
p = stats.uniform.rvs(size=200_000, random_state=0) # flat on the rate p
print(round(np.mean(np.abs(logit(p)) < 1), 3)) # 0.461 (exact 0.462): a hump on the log-odds
eta = stats.norm.rvs(0, 10, size=200_000, random_state=1) # "vague" N(0, 10^2) on the log-odds (scale = sd)
q = expit(eta)
print(round(np.mean((q < 0.01) | (q > 0.99)), 3)) # 0.646: a U on the rate
print(round(2 * stats.norm.cdf(logit(0.01) / 10), 3)) # 0.646 exactly
# 2) An informative Beta prior from past experiments (method of moments), then discounted
m, sd = 0.10, 0.015
worth = m * (1 - m) / sd**2 - 1
print(round(worth, 1), round(m * worth, 1), round((1 - m) * worth, 1)) # 399.0 39.9 359.1
a0 = 100 / worth # keep the mean, shrink the worth to 100 visitors
print(round(a0 * m * worth, 1), round(a0 * (1 - m) * worth, 1)) # 10.0 90.0
# 3) Tug of war: posterior mean = weighted average of prior mean and data rate
a, b, k, n = 10, 90, 30, 200
print(round((a + k) / (a + b + n), 4), round(n / (n + a + b), 4)) # 0.1333 0.6667
# 4) Shrinkage: Normal vs Laplace prior with the same variance, noisy estimate x ~ N(theta, 1)
th = np.linspace(-12, 12, 48001)
def post_mean(x, prior_pdf):
w = prior_pdf * stats.norm.pdf(x, th, 1.0)
return float(np.sum(th * w) / np.sum(w))
normal_pdf, laplace_pdf = stats.norm.pdf(th, 0, 1), stats.laplace.pdf(th, 0, 1 / np.sqrt(2))
print([round(post_mean(x, normal_pdf), 2) for x in (1, 4)]) # [0.5, 2.0]
print([round(post_mean(x, laplace_pdf), 2) for x in (1, 4)]) # [0.38, 2.59] more pull on small, less on big
# 5) Prior predictive check of an A/B model in NumPyro (no data passed: everything is simulated)
def ab_model(n, s_alpha, s_delta, mu0, kA=None, kB=None):
alpha = numpyro.sample("alpha", dist.Normal(mu0, s_alpha)) # control log-odds; Normal(loc, scale=sd)
delta = numpyro.sample("delta", dist.Normal(0.0, s_delta)) # log-odds lift of B
numpyro.sample("kA", dist.Binomial(n, logits=alpha), obs=kA)
numpyro.sample("kB", dist.Binomial(n, logits=alpha + delta), obs=kB)
def check(**priors):
sim = Predictive(ab_model, num_samples=4000)(jax.random.PRNGKey(0), n=1000, **priors)
rA, rB = np.asarray(sim["kA"]) / 1000, np.asarray(sim["kB"]) / 1000
odd_rate = np.mean((rA < 0.01) | (rA > 0.40))
lift = (rB - rA) / np.maximum(rA, 1e-9)
odd_lift = np.mean((rA == 0) | (np.abs(lift) > 0.5))
return round(float(odd_rate), 2), round(float(odd_lift), 2)
print(check(mu0=0.0, s_alpha=10.0, s_delta=10.0)) # (0.84, 0.62): absurd worlds
print(check(mu0=-2.2, s_alpha=0.5, s_delta=0.2)) # (0.0, 0.04): plausible worlds
# 6) Prior predictive check of a trend + changepoints model (one year of daily demand)
T = jnp.arange(365) / 365.0 # time in years
S = jnp.arange(1, 11) / 11.0 # 10 candidate changepoints
def trend_model(s_m, s_k, b, s_sigma, y=None):
m = numpyro.sample("m", dist.Normal(500.0, s_m)) # level (orders/day)
k = numpyro.sample("k", dist.Normal(0.0, s_k)) # slope (orders/day per year)
delta = numpyro.sample("delta", dist.Laplace(0.0, b).expand([10])) # Laplace(loc, scale=b) slope changes
sigma = numpyro.sample("sigma", dist.HalfNormal(s_sigma))
g = m + k * T + jnp.sum(delta * jnp.clip(T[:, None] - S, 0.0), axis=1)
numpyro.sample("y", dist.Normal(g, sigma), obs=y)
for name, pri in [("vague", (1000.0, 1000.0, 500.0, 500.0)), ("weakly informative", (100.0, 150.0, 30.0, 30.0))]:
ysim = np.asarray(Predictive(trend_model, num_samples=2000)(jax.random.PRNGKey(1), *pri)["y"])
absurd = np.mean(((ysim < 0) | (ysim > 2000)).any(axis=1))
print(name, round(float(absurd), 2), np.percentile(ysim[:, -1], [5, 95]).round(-1))
# vague 0.93 [-2650. 3610.] most simulated years leave 0-2000 orders/day
# weakly informative 0.02 [180. 820.] varied but plausible years
1. A conversion rate p has a flat Uniform(0, 1) prior. What does that prior look like on the log-odds scale $\eta = \log(p/(1-p))$?
2. A "vague" prior $\eta \sim N(0, 10^2)$ on the log-odds of a conversion rate. Roughly what prior probability does it give to rates below 1% or above 99%?
3. Prior Beta(30, 270) on a rate; data 60 buyers in 300 visitors. The posterior mean and the data's weight are…
4. Changepoint slope changes have priors $\delta_j \sim Laplace(0, b)$. Which statement is correct?
5. In a prior predictive check, the simulated datasets should be compared with…
6. Daily order counts 2, 4, 6 are Poisson(λ), with a conjugate prior Gamma(3, rate 1). The posterior is…
gamma(a=15, scale=1/4): SciPy uses scale = 1/rate.Practice problems
A. Past tests of a feature had conversion rates with mean 5% and sd 1 point. Fit a Beta prior by moments, then discount it to a worth of 50 visitors.
- $v = 0.0001$. $\alpha + \beta = 0.05 \times 0.95/0.0001 - 1 = 475 - 1 = 474$.
- $\alpha = 0.05 \times 474 = 23.7$, $\beta = 0.95 \times 474 = 450.3$: Beta(23.7, 450.3), worth 474 visitors.
- Discount factor $a_0 = 50/474$: Beta(2.5, 47.5). Same mean 5%, now worth 50 visitors.
B. A model uses $N(0, 2^2)$ on the log-odds of a rate. Compute the prior probability that the rate is below 1% or above 99%, and that it is between 20% and 80%. Is this prior close to flat on the rate?
$2\Phi(-4.595/2) = 2\Phi(-2.298) \approx 0.022$ (flat: 0.02). $P(0.2 \lt p \lt 0.8) = 1 - 2\Phi(-1.386/2) = 1 - 2\Phi(-0.693) \approx 0.512$ (flat: 0.6). It is somewhat heavier at the extremes and lighter in the middle than flat: close-ish, but tilted toward the edges. $N(0, 1.5^2)$ is closer to flat in the middle (0.645).
C. (Interview) "You said you used a non-informative flat prior. Why might an interviewer push back?"
"Because flatness depends on the parameterization. A prior flat on a probability is a logistic hump on the log-odds, and a very wide prior on the log-odds is a U on the probability, with most of its mass near 0 and 1. So 'non-informative' is not well defined; every prior makes claims on some scale. What I actually do is choose weakly informative priors on a sensible, standardized scale and run a prior predictive check to make sure the implied data are plausible."
D. Daily orders have noise sd σ = 30. Prior on the average μ: N(500, 15²). After 12 days with mean 560, what is the posterior, and how many days is the prior worth?
- Worth: $\sigma^2/\tau^2 = 900/225 = 4$ days.
- Posterior mean $= (4 \times 500 + 12 \times 560)/(4 + 12) = (2000 + 6720)/16 = 545$.
- Posterior precision $1/225 + 12/900 = 0.00444 + 0.01333 = 0.01778$, variance 56.25, sd 7.5 ($= \sigma/\sqrt{16}$: like 16 days of data).
E. (Interview, your forecasting model) "Your model has $N(0, 1000)$ priors on most parameters and some forecast paths go negative or explode. What do you do?"
"First a prior predictive check: simulate series from the priors with the real time index and regressors, and look at the minimum, the growth over the history and the size of trend bends. If the trend explodes, I rescale (check how y and time are scaled) and tighten the slope prior and the Laplace scale b on the changepoint deltas. If values go negative even with sensible priors, the likelihood is the issue, so I would consider a log link or a Negative Binomial for counts. Then I refit, compare forecasts, and check sensitivity to b (Chapters 6.8 and 7.10)."
F. A holiday effect (in standardized units) has prior $N(0, 0.5^2)$; the data give an estimate $x = 1.2$ with standard error 0.5. Find the posterior mean and sd.
Shrink factor $\tau^2/(\tau^2 + s^2) = 0.25/0.5 = 0.5$, so the posterior mean is $0.6$. Posterior precision $1/0.25 + 1/0.25 = 8$, sd $= 1/\sqrt8 \approx 0.354$. With a holiday seen only once or twice, the prior and the data carry equal weight: exactly the situation where the prior's scale deserves a sensitivity check.
Conjugate models: Beta-Binomial and Dirichlet-Multinomial
Usually a posterior has to be computed: a grid, a sampler, an optimizer. For a few special pairs of prior and likelihood you can skip all of that. The posterior is the prior's own family again, with new numbers that you get by adding counts. Two of those pairs, Beta-Binomial for conversion rates and Dirichlet-Multinomial for category shares, sit inside your A/B framework. This chapter derives them properly, shows what their numbers mean, predicts future data with them, and shows where the trick stops working.
- Say what conjugate means (the posterior stays in the prior's family) and why it is so convenient
- Derive the Beta-Binomial update Beta($\alpha + k$, $\beta + n - k$), including the evidence $p(D)$, from "multiply powers, add exponents"
- Read the posterior mean as a weighted average of the prior mean and the data rate, with weight $\frac{\alpha+\beta}{\alpha+\beta+n}$ on the prior
- Run the model for two variants side by side, and compute the posterior predictive (the beta-binomial distribution), which is wider than a plug-in Binomial
- Derive and use the Dirichlet-Multinomial update Dirichlet($\boldsymbol\alpha + \mathbf{c}$) for categorical metrics
- Use two more pairs briefly: Gamma-Poisson for count metrics and Normal-Normal with known variance
- Know the table of common conjugate pairs and what to do when there is no conjugate pair (MCMC and VI)
What we need from earlier chapters: prior, likelihood, posterior and evidence (Chapter 6.1) and choosing priors (Chapter 6.2); the Binomial and Multinomial (Chapter 4.7); the Poisson and Negative Binomial (Chapter 4.8); the Gamma (Chapter 4.10); the Beta and Dirichlet with their pseudo-counts and a first look at updating (Chapter 4.11); MAP (Chapter 5.2). Notation: θ is a conversion rate; $k$ conversions among $n$ visitors; $\mathbf{c} = (c_1, \dots, c_K)$ are category counts; $B(a, b)$ is the Beta function $\int_0^1 t^{a-1}(1-t)^{b-1}dt$; $\binom{n}{k}$ is "n choose k". "$\propto$" means "proportional to": equal up to a factor that does not depend on θ.
What "conjugate" means, and why it is so convenient core
Think of a scoreboard at a football match. Its format never changes: "Home – Away". Every goal just changes the numbers. You never need a new scoreboard, and anyone can read the current state at a glance.
Bayesian updating is usually not like that. You multiply the prior curve by the likelihood curve, and out comes a new curve that may have no name and no formula: you need a computer to draw it, find its area, and read anything off it. But for some special pairs of prior and likelihood, the posterior comes out in the same family as the prior: a Beta prior gives a Beta posterior, a Dirichlet prior gives a Dirichlet posterior. Then updating is like the scoreboard: keep the format, change the numbers. Such a prior is called conjugate to that likelihood ("conjugate" here just means "a matching partner").
Three ways to say it:
- Picture: a scoreboard: same format before and after every goal; only the numbers change.
- Numbers: prior Beta(2, 18), then 6 conversions among 40 visitors: posterior Beta(2 + 6, 18 + 34) = Beta(8, 52). No integral, no sampler.
- Slogan: conjugate = the posterior is the prior's family with updated numbers.
Watch the family survive the update. Prior Beta(2, 18) for a conversion rate θ ("about 10%, worth 20 visitors"). Data: $k = 6$ conversions among $n = 40$ visitors.
- Prior density: $p(\theta) = \dfrac{\theta^{1}(1-\theta)^{17}}{B(2, 18)} = 342\,\theta\,(1-\theta)^{17}$, because $B(2,18) = \frac{1!\,17!}{19!} = \frac{1}{18 \times 19} = \frac{1}{342}$.
- Likelihood: $p(D\mid\theta) = \binom{40}{6}\theta^{6}(1-\theta)^{34} = 3\,838\,380\;\theta^{6}(1-\theta)^{34}$.
- Multiply. Numbers that do not contain θ just ride along: $342 \times 3\,838\,380 \times \theta^{1+6}(1-\theta)^{17+34} = (\text{a constant}) \times \theta^{7}(1-\theta)^{51}$.
- Recognize the shape. $\theta^{7}(1-\theta)^{51} = \theta^{8-1}(1-\theta)^{52-1}$ is exactly the θ-part of a Beta(8, 52) density.
- A density must have area 1, so the constant has to be $1/B(8, 52)$. The posterior is exactly Beta(8, 52): mean $8/60 \approx 0.133$, sd $\approx 0.044$.
Nothing was integrated. We only added exponents and read the answer off the shape.
A family of prior distributions $\mathcal{F}$ (for example, "all Beta distributions") is conjugate to a likelihood $p(D\mid\theta)$ if, for every prior in $\mathcal{F}$ and every possible dataset $D$, the posterior $p(\theta\mid D)$ is again in $\mathcal{F}$.
- Then Bayesian updating is a map from old numbers to new numbers: $(\alpha, \beta) \mapsto (\alpha + k,\ \beta + n - k)$. The numbers inside a prior are its hyperparameters.
- The kernel of a density is the part that depends on θ, after dropping constant factors. Conjugacy works when the likelihood's kernel and the prior's kernel are made of the same building blocks (here: powers of θ and powers of $1-\theta$), so their product has the same form.
- A closed form is an exact formula you can write down: no grid, no sampling, no optimization. Conjugate pairs give closed-form posteriors, evidences and predictive distributions.
- Conjugacy is a property of the pair (prior family, likelihood). Beta is conjugate to the Binomial, but not to a Normal likelihood.
Why do we need it?
Most posteriors have no formula and must be computed (MCMC, VI). A conjugate pair gives the exact posterior instantly, by adding counts. That makes results fast, transparent and easy to explain, and it gives an exact answer to test approximate methods against.
Where is it used?
Beta-Binomial conversion tests, Dirichlet-Multinomial tests for categorical metrics, Thompson-sampling bandits, Gamma-Poisson rate models, Kalman filters (Normal-Normal updating at every time step), and Gibbs samplers, which update one block of parameters at a time with conjugate formulas.
How is it used?
Check whether your likelihood has a conjugate prior (table at the end of this chapter). If yes, store the hyperparameters, add the data's counts or sums after each batch, and read summaries with SciPy, e.g. stats.beta(a + k, b + n - k).mean(). If no, use NumPyro with MCMC or SVI.
"Conjugate means the posterior is the same as the prior."
Same family, different numbers. Beta(2, 18) and Beta(8, 52) are both Betas, but the second is centred higher and is much narrower.
"Conjugate priors are the correct priors."
They are the convenient ones. Choose a prior for what it says (Chapter 6.2); if a conjugate one says what you believe, enjoy the shortcut. With plenty of data, reasonable priors from different families give almost the same posterior (the widget at $n = 400$).
"Bayesian A/B testing needs conjugate priors."
Conjugacy is a shortcut, not a requirement. NumPyro fits any prior and likelihood with MCMC or SVI (Chapters 6.9–6.14).
In an A/B framework like yours, the Beta-Binomial (conversion metrics) and the Dirichlet-Multinomial (categorical metrics) are both conjugate pairs. For a single variant on its own, the exact posterior is known, so it is a free unit test: an SVI or NUTS fit of that simple model should reproduce Beta($\alpha + k$, $\beta + n - k$) closely (the code at the end of the chapter does exactly this). Once you add hierarchical partial pooling across segments, the unknown population parameters break conjugacy for the model as a whole, which is one reason to use SVI at all.
"A conjugate prior is a prior that is the same as the likelihood."
A conjugate prior family is one that the likelihood maps back into itself: prior in the family → posterior in the family, for any data.
Model answer: "A prior family is conjugate to a likelihood if the posterior stays in the same family. For a Binomial likelihood the Beta is conjugate: Beta(α, β) plus k successes in n trials gives Beta(α + k, β + n − k). It is convenient because the update is closed-form, but it is a computational convenience, not a statement about which prior is right."
Conjugate: prior family $\mathcal{F}$ + likelihood → posterior in $\mathcal{F}$ for every dataset. Update = change the hyperparameters.
Beta(2, 18) + 6 of 40 → Beta(8, 52): add exponents, recognize the kernel, no integral.
Trap: conjugate = convenient, not "correct"; same family ≠ same distribution.
Quick check: a prior says the log-odds of conversion are Normal. Can you update it by adding counts?
No. A Normal on the log-odds is not conjugate to the Binomial: prior × likelihood is $e^{-(\text{logit}\,\theta - \mu)^2/2s^2} \cdot \theta^{k-1}(1-\theta)^{n-k-1}$, which is not a logit-Normal (nor a Beta). You must compute the posterior numerically (grid, MCMC or VI). That is exactly the situation in logistic regression with Normal priors on the coefficients.
The Beta-Binomial update, derived step by step core
Why does the update turn into simple addition? Because the prior and the likelihood are made of exactly the same two building blocks: a power of θ ("how much do I believe in conversions?") and a power of $1 - \theta$ ("how much do I believe in non-conversions?"). When you multiply two powers of the same thing, the exponents add: $x^2 \cdot x^3 = x^5$. So the prior's exponents plus the data's exponents give the posterior's exponents. That is the whole trick.
The data enter only through two numbers: how many visitors converted ($k$) and how many did not ($n - k$). Who converted first, or which day, does not matter for θ. Numbers that carry all of the data's information about θ are called sufficient statistics.
Three ways to say it:
- Picture: two ledgers, "conversions" and "non-conversions"; the prior opens them with α and β, the data add $k$ and $n - k$.
- Numbers: θ-exponent $1 + 6 = 7$, (1 − θ)-exponent $17 + 34 = 51$, so Beta(2, 18) becomes Beta(8, 52).
- Slogan: multiply densities, add exponents, read off the new Beta.
The full calculation for Beta(2, 18) and 6 of 40, including the evidence.
- Model: $\theta \sim \text{Beta}(2, 18)$, and given θ, $k \sim \text{Binomial}(40, \theta)$. Observed $k = 6$.
- Prior × likelihood $= 342\,\theta(1-\theta)^{17} \times 3\,838\,380\,\theta^6(1-\theta)^{34} = 342 \times 3\,838\,380 \times \theta^{7}(1-\theta)^{51}$.
- Area under $\theta^7(1-\theta)^{51}$ is, by the definition of the Beta function, $B(8, 52) \approx 5.637 \times 10^{-11}$.
- Evidence (the total area under prior × likelihood): $p(D) = 342 \times 3\,838\,380 \times 5.637 \times 10^{-11} \approx 0.0740$.
- Posterior = (prior × likelihood) ÷ evidence $= \dfrac{\theta^{7}(1-\theta)^{51}}{B(8, 52)}$: exactly Beta(8, 52). The big constants cancel.
- Sanity check on the evidence: if θ were exactly 0.10, the chance of 6 of 40 would be $0.107$; at θ = 0.15 it would be $0.174$; at 0.05 only $0.010$; at 0.30 only $0.015$. The prior spreads its belief over all of these, and the weighted average comes out at $0.074$.
The Beta-Binomial model. $\theta \sim \text{Beta}(\alpha, \beta)$ and $k \mid \theta \sim \text{Binomial}(n, \theta)$. Then
$$\begin{aligned} p(\theta \mid k) &\propto \underbrace{\binom{n}{k}\theta^{k}(1-\theta)^{n-k}}_{\text{likelihood}} \cdot \underbrace{\frac{\theta^{\alpha-1}(1-\theta)^{\beta-1}}{B(\alpha,\beta)}}_{\text{prior}} \;\propto\; \theta^{\alpha+k-1}(1-\theta)^{\beta+n-k-1},\\[4pt] \theta \mid k &\sim \text{Beta}(\alpha + k,\ \beta + n - k), \qquad p(k) = \binom{n}{k}\frac{B(\alpha + k,\ \beta + n - k)}{B(\alpha, \beta)} . \end{aligned}$$- $\binom{n}{k}$ and $B(\alpha,\beta)$ do not contain θ, so they cancel from the posterior. They matter only for the evidence $p(k)$, which is the beta-binomial probability of $k$ (you will meet that distribution again for predictions, later in this chapter).
- Sufficient statistics: $(k, n)$. Data entered one visitor at a time (Bernoulli), in batches, or all at once give the same posterior, in any order.
- Assumptions: visitors are independent given θ and share the same θ (the Binomial assumptions of Chapter 4.7); $\alpha, \beta \gt 0$.
- Code: SciPy
stats.beta(a + k, b + n - k); the evidence isstats.betabinom(n, a, b).pmf(k). NumPyro writes the Beta asdist.Beta(concentration1=α, concentration0=β): concentration1 counts successes.
Why do we need it?
Conversion rates are the most common metric in experimentation. This derivation turns "compute a posterior" into "add two numbers", and it tells you exactly which summary of the data you need to keep: the counts $k$ and $n$, nothing else.
Where is it used?
Bayesian A/B tests on conversion, click-through and sign-up rates; Thompson sampling for bandits and ad selection; quality control (defect rates); smoothing click rates for items with few impressions in recommender systems; the evidence $p(k)$ in Bayes-factor model comparison.
How is it used?
Aggregate the data to $(k, n)$ per variant (or per segment). Add them to the prior's $(\alpha, \beta)$. Summarize with stats.beta(a + k, b + n - k): .mean(), .ppf([0.025, 0.975]), .sf(threshold), or draw samples with .rvs(size) for decisions (Chapter 6.4).
"The binomial coefficient $\binom{n}{k}$ changes the posterior."
It does not contain θ, so it cancels. That is also why per-visitor Bernoulli data (no coefficient) and aggregated Binomial data give the same posterior.
"Updating visitor by visitor gives a different answer from one big update."
Adding is order-free: $(\alpha + k_1 + k_2,\ \beta + (n_1 - k_1) + (n_2 - k_2))$ either way. Today's posterior is tomorrow's prior.
"The evidence $p(D) = 0.074$ is the probability that the model is right."
It is the probability of seeing exactly this data under the model, averaged over the prior. It is useful for comparing models (Bayes factors), and it changes a lot with the prior's width even when the posterior hardly moves.
For a conversion metric, the only things the Beta-Binomial needs per variant are $k$ and $n$. If your NumPyro model observes aggregated counts, e.g. numpyro.sample("k", dist.Binomial(total_count=n, probs=theta), obs=k), the posterior is the same as observing every user as a Bernoulli, and much cheaper. Check which form your code uses and make sure $n$ counts exactly the users who were exposed to that variant.
$\theta \sim \text{Beta}(\alpha,\beta)$, $k\mid\theta \sim \text{Bin}(n,\theta)$ ⇒ $\theta\mid k \sim \text{Beta}(\alpha+k,\ \beta+n-k)$.
Evidence $p(k) = \binom{n}{k}B(\alpha+k, \beta+n-k)/B(\alpha,\beta)$ (beta-binomial). Sufficient statistics: $(k, n)$.
Trap: constants cancel from the posterior but not from the evidence; order of the data never matters.
Quick check: prior Beta(1, 1). Day 1: 3 of 20 convert. Day 2: 5 of 30 convert. What is the posterior after both days?
Add everything: $k = 3 + 5 = 8$, $n - k = 17 + 25 = 42$. Posterior Beta(1 + 8, 1 + 42) = Beta(9, 43), mean $9/52 \approx 0.173$. Doing day 1 first (Beta(4, 18)) and then day 2 gives the same Beta(9, 43).
Pseudo-counts: the posterior mean is a weighted average core
A friend has tried a restaurant 20 times and rates it "10% chance of a bad meal". You go 40 times yourself and get 6 bad meals (15%). What should you believe now? Not your friend's 10%, not your own 15%, but something in between, closer to yours because you have more visits. That is exactly what the Beta posterior does. The prior behaves like pretend visitors ("pseudo-counts"): $\alpha$ pretend conversions and $\beta$ pretend non-conversions, $\alpha + \beta$ pretend visitors in total. The posterior mean is the conversion rate of the combined crowd, pretend plus real.
So the prior's influence is not mysterious: it counts as much as $\alpha + \beta$ visitors would. Against 40 real visitors, a prior worth 20 still matters; against 40 000 it is a rounding error.
Three ways to say it:
- Picture: a seesaw with the prior mean on one side and the data rate on the other; the heavier side (more visitors) wins the balance point.
- Numbers: prior 10% worth 20, data 15% from 40: posterior mean $\tfrac13 \times 10\% + \tfrac23 \times 15\% = 13.3\%$.
- Slogan: the prior is worth $\alpha + \beta$ visitors; the posterior mean is a visitor-weighted average.
Prior Beta(2, 18): mean 0.10, worth $\kappa = \alpha + \beta = 20$ visitors.
- Data: 6 of 40 (rate 0.15). Weight on the prior: $w = \frac{20}{20 + 40} = \frac13$; weight on the data: $\frac{40}{60} = \frac23$.
- Posterior mean: $\tfrac13 \times 0.10 + \tfrac23 \times 0.15 = 0.0333 + 0.1000 = 0.1333$. Same as $8/60$. ✓
- Ten times more data at the same rate (60 of 400): $w = \frac{20}{420} = 0.048$; mean $= 0.048 \times 0.10 + 0.952 \times 0.15 = 0.1476$ $(= 62/420)$. The prior barely matters now.
- Posterior sd $= \sqrt{\frac{m(1-m)}{\kappa + n + 1}}$ with $m$ the posterior mean: $\sqrt{\frac{0.1333 \times 0.8667}{61}} = 0.0435$ with 40 visitors, $\sqrt{\frac{0.1476 \times 0.8524}{421}} = 0.0173$ with 400. Ten times the data shrank the spread about $0.0435/0.0173 \approx 2.5$ times, not 10 times: the denominator grew from 61 to 421, and the sd follows its square root.
- Posterior mode (the MAP of Chapter 5.2): $\frac{\alpha + k - 1}{\alpha + \beta + n - 2} = \frac{7}{58} = 0.1207$, a little below the mean because the Beta(8, 52) leans right.
- Zero conversions. Flat prior Beta(1, 1) and 0 of 10 visitors: the observed rate (the MLE) is 0, but the posterior is Beta(1, 11) with mean $\frac{0+1}{10+2} = 0.083$, and a 5% chance that θ is above $1 - 0.05^{1/11} = 0.238$. Zero conversions in 10 visitors does not prove the rate is zero.
Write $\kappa = \alpha + \beta$ (the prior's worth, or "prior sample size", or concentration) and $m_0 = \alpha/\kappa$ (the prior mean). After $k$ conversions in $n$ visitors:
$$E[\theta \mid k] = \frac{\alpha + k}{\kappa + n} = \underbrace{\frac{\kappa}{\kappa + n}}_{w}\, m_0 + \underbrace{\frac{n}{\kappa + n}}_{1 - w}\, \frac{k}{n}, \qquad Var(\theta\mid k) = \frac{m(1-m)}{\kappa + n + 1},$$where $m = E[\theta\mid k]$. (Check the first identity: $\frac{\kappa}{\kappa+n}\cdot\frac{\alpha}{\kappa} + \frac{n}{\kappa+n}\cdot\frac{k}{n} = \frac{\alpha + k}{\kappa + n}$.)
- The posterior mean always lies between the prior mean and the data rate $k/n$ (the MLE). This pull toward the prior mean is called shrinkage.
- $w \to 0$ as $n \to \infty$: the data win. The posterior sd shrinks like $1/\sqrt{\kappa + n}$.
- Mode $= \frac{\alpha + k - 1}{\kappa + n - 2}$ (when both new parameters exceed 1). With the flat prior Beta(1, 1) the mode equals $k/n$, while the mean is $\frac{k+1}{n+2}$ ("Laplace's rule of succession").
Why do we need it?
To state, in plain business terms, how much the prior affects a result ("the prior is worth 20 users; the test had 4 000") and to avoid silly estimates from small data, such as a 0% or 100% conversion rate after a handful of visitors.
Where is it used?
Reporting priors in A/B tests, smoothing click-through rates of new ads or products with few impressions, ranking items by a "Bayesian average" rating, empirical-Bayes priors fitted from past experiments, and the shrinkage of small segments in hierarchical models (Chapter 6.6).
How is it used?
Choose $m_0$ (your best guess) and $\kappa$ (how many visitors your belief is worth), set $\alpha = m_0\kappa$, $\beta = (1-m_0)\kappa$. Report $w = \kappa/(\kappa + n)$ next to the result. If $w$ is large, say so, or run a prior-sensitivity check (Chapter 6.8).
"With a flat prior the posterior mean equals the observed rate $k/n$."
With Beta(1, 1) the posterior mode equals $k/n$; the mean is $(k+1)/(n+2)$. The difference matters only for small $n$, and it is exactly what saves you from "0% after 0 of 10".
"A prior worth 20 visitors means I invented 20 visitors of data."
It is a belief expressed in the unit of data. Say it that way in reports, and never mix pseudo-counts into the real sample size.
"Double the data, halve the posterior sd."
The sd shrinks like $1/\sqrt{\kappa + n}$: you need about four times the visitors to halve it, the same $\sqrt{n}$ law as the standard error (Chapter 5.5).
In an A/B framework like yours, write the prior of each conversion metric in mean–worth form and report the worth next to the traffic: "prior mean 10%, worth 50 users; each variant had 8 000". Then $w \approx 0.006$ and nobody needs to worry about the prior. For small segments the same arithmetic shows when the prior (or, in the hierarchical model, the population distribution) is doing a lot of the work: that is shrinkage, and Chapter 6.6 makes it the centre of partial pooling.
"The Bayesian estimate is biased because it is pulled toward the prior, so it is worse."
It is a compromise between the prior mean and the MLE, weighted by information. The pull adds a little bias but removes variance, which usually lowers the average error for small samples (the bias–variance trade-off of Chapter 5.1), and it vanishes as $n$ grows.
Model answer: "With a Beta(α, β) prior the posterior mean is $\frac{\alpha+\beta}{\alpha+\beta+n}$ times the prior mean plus $\frac{n}{\alpha+\beta+n}$ times $k/n$. The prior acts like α + β extra observations, so its influence is easy to quantify and fades as data accumulate."
$E[\theta\mid k] = w\,m_0 + (1-w)\,\frac{k}{n}$ with $w = \frac{\alpha+\beta}{\alpha+\beta+n}$; $Var = \frac{m(1-m)}{\alpha+\beta+n+1}$.
Prior worth $\kappa = \alpha + \beta$ pretend visitors. Mode $\frac{\alpha+k-1}{\alpha+\beta+n-2}$; flat prior mean $\frac{k+1}{n+2}$.
Trap: flat prior ⇒ mode $= k/n$, not mean. sd shrinks like $1/\sqrt{\kappa+n}$.
Quick check: prior Beta(30, 270), 1 000 visitors with 150 conversions. What is the weight on the prior and the posterior mean?
Prior worth $\kappa = 300$, prior mean $0.10$. $w = 300/1\,300 = 0.231$. Posterior mean $= 0.231 \times 0.10 + 0.769 \times 0.15 = 0.0231 + 0.1154 = 0.1385$, which equals $(30 + 150)/(300 + 1\,000) = 180/1\,300 = 0.1385$.
Two variants side by side: the A/B Beta-Binomial model core
An A/B test is just two copies of the same small model. Variant A has its own unknown conversion rate $\theta_A$; variant B has its own $\theta_B$. Each gets a prior, each is updated by its own visitors. Visitors who saw A tell you nothing directly about $\theta_B$, so the two updates do not touch each other.
After the test you hold two Beta curves, and four numbers describe everything: $(\alpha + k_A,\ \beta + n_A - k_A)$ and $(\alpha + k_B,\ \beta + n_B - k_B)$. The question "is B better?" is then a question about these two curves together, answered in Chapter 6.4.
Three ways to say it:
- Picture: two hills on the same ruler of conversion rates, one blue (A), one orange (B).
- Numbers: 50/500 and 60/500 with flat priors give Beta(51, 451) and Beta(61, 441): means 10.2% and 12.2%.
- Slogan: same prior, separate data, separate posteriors; compare them only at the end.
The checkout test from Chapter 5.6: 50 of 500 converted with the old checkout (A), 60 of 500 with the new one (B).
- Flat prior Beta(1, 1) for both. A: Beta(1 + 50, 1 + 450) = Beta(51, 451), mean $51/502 = 0.1016$, 95% of the belief in $[0.0767, 0.1295]$.
- B: Beta(1 + 60, 1 + 440) = Beta(61, 441), mean $61/502 = 0.1215$, 95% in $[0.0944, 0.1515]$.
- The two 95% ranges overlap from 0.0944 to 0.1295. Overlap alone does not answer "is B better?" (see the careful box).
- A shared informative prior Beta(10, 90) ("about 10%, worth 100 visitors", e.g. from past checkout tests): A becomes Beta(60, 540), mean $0.100$; B becomes Beta(70, 530), mean $0.1167$. Both are pulled toward 10% with weight $100/600 = 1/6$, and the gap shrinks from 2.0 to 1.7 points.
- The gap itself, with flat priors: $\theta_B - \theta_A$ has mean $0.1215 - 0.1016 = 0.0199$ and sd $\sqrt{0.0135^2 + 0.0146^2} = 0.0199$ (variances add because the two posteriors are independent).
The two-variant Beta-Binomial model:
$$\theta_A, \theta_B \overset{\text{iid}}{\sim} \text{Beta}(\alpha, \beta), \qquad k_A \mid \theta_A \sim \text{Bin}(n_A, \theta_A), \qquad k_B \mid \theta_B \sim \text{Bin}(n_B, \theta_B).$$The joint posterior factorizes into two independent Betas:
$$p(\theta_A, \theta_B \mid D) = \text{Beta}(\theta_A;\ \alpha + k_A,\ \beta + n_A - k_A)\ \times\ \text{Beta}(\theta_B;\ \alpha + k_B,\ \beta + n_B - k_B).$$- "Factorizes" means the joint density is a product, so $\theta_A$ and $\theta_B$ are independent given the data. This holds because the priors are independent and each variant's data depend only on its own θ.
- Quantities you care about are functions of both: the difference $\theta_B - \theta_A$, the relative lift $\theta_B/\theta_A - 1$, and probabilities such as $P(\theta_B \gt \theta_A \mid D)$. Their posteriors are not Betas; you get them from random draws (Chapter 6.4).
- More variants (A/B/C/…) just add more independent Betas. Segments with a shared population distribution make the θ's dependent: that is the hierarchical model of Chapter 6.5.
Why do we need it?
To turn raw counts from an experiment into an honest statement of uncertainty about each variant's rate, which is the input to every Bayesian decision (probability B is better, expected loss, practical thresholds).
Where is it used?
Bayesian A/B testing dashboards for conversion, sign-up and click metrics; multi-variant tests; Thompson sampling (draw one θ from each variant's Beta and show the variant whose draw is highest); the conversion part of your A/B framework.
How is it used?
Pick one prior for all variants. Count $k$ and $n$ per variant. Form the posteriors. Report each variant's mean and 95% interval, then draw a few thousand $(\theta_A, \theta_B)$ pairs, e.g. rng.beta(a + kA, b + nA - kA, 4000), to compute differences and decision probabilities.
"The two 95% ranges overlap, so there is no real difference."
Overlapping intervals are a poor test. The gap's own spread is $\sqrt{sd_A^2 + sd_B^2}$, smaller than $sd_A + sd_B$, so intervals can overlap while the gap is clearly positive. Ask about the gap directly: $P(\theta_B \gt \theta_A \mid D)$ (Chapter 6.4).
"Give the new variant an optimistic prior; we expect it to win."
Use the same prior for every variant unless you have specific, pre-registered knowledge about one of them. A favourable prior for B is a thumb on the scale.
"θ_A and θ_B are independent, so the data of A carry no information for B in any model."
That holds in this simple model. In a hierarchical model the variants (or segments) share a population distribution, and data from one inform the others (Chapter 6.5).
This is the conversion-metric core of an A/B framework like yours: one Beta per variant, updated by that variant's counts, compared at the end through $P(\theta_B \gt \theta_A \mid D)$ and practical thresholds (6.4). If your framework fits even this model with SVI, the two closed-form Betas above are the exact answer it should reproduce, and the standard deviation of the gap, $\sqrt{sd_A^2 + sd_B^2}$, is a quick sanity check on the reported uncertainty.
A/B: $\theta_A \mid D \sim \text{Beta}(\alpha + k_A, \beta + n_A - k_A)$, $\theta_B \mid D \sim \text{Beta}(\alpha + k_B, \beta + n_B - k_B)$, independent.
Gap: mean $E[\theta_B] - E[\theta_A]$, sd $\sqrt{sd_A^2 + sd_B^2}$. Same prior for every variant.
Trap: overlapping intervals ≠ no difference; ask about the gap.
Quick check: with a flat prior, A has 30 of 300 and B has 45 of 300. Give both posteriors and their means.
A: Beta(31, 271), mean $31/302 \approx 0.103$. B: Beta(46, 256), mean $46/302 \approx 0.152$.
Predicting the next visitors: the beta-binomial distribution core
Your manager asks: "Out of the next 20 visitors, how many will buy?" If you knew θ exactly, the answer would be a Binomial(20, θ). But you only have a posterior for θ. The honest forecast must carry two kinds of uncertainty: you are not sure about θ (parameter uncertainty), and even with the right θ, buying is random (observation noise).
The posterior predictive (Chapter 6.1) mixes them: imagine drawing a plausible θ from the posterior, then drawing 20 visitors with that θ, and repeating. For the Beta-Binomial model this mixture has a name and a formula: the beta-binomial distribution. It has the same mean as the "plug-in" Binomial that uses the posterior mean, but it is wider.
Three ways to say it:
- Picture: a cloud of Binomials, one for every plausible θ, piled on top of each other; the pile is wider than any single one.
- Numbers: posterior Beta(8, 52), next 20 visitors: mean 2.67 either way, but variance 3.03 (beta-binomial) instead of 2.31 (plug-in Binomial).
- Slogan: predictive = average the model's forecast over everything you still do not know.
Posterior Beta(8, 52) (from Beta(2, 18) and 6 of 40). Predict the number $\tilde y$ of buyers among the next $m = 20$ visitors.
- Next single visitor: $P(\text{buys}\mid D) = E[\theta\mid D] = 8/60 = 0.133$.
- Mean of $\tilde y$: $m \times 8/60 = 20 \times 0.1333 = 2.67$ (same as the plug-in Binomial(20, 0.1333)).
- Variance by the law of total variance (Chapter 4.6): $Var(\tilde y) = E[Var(\tilde y\mid\theta)] + Var(E[\tilde y\mid\theta]) = m\bar p(1-\bar p) + m(m-1)Var(\theta\mid D)$, where $\bar p = 0.1333$ is the posterior mean.
- Numbers: $20 \times 0.1333 \times 0.8667 = 2.311$, plus $20 \times 19 \times 0.001894 = 0.720$, total $3.031$. The plug-in Binomial has only the first part, $2.311$.
- Probability of zero buyers: beta-binomial $\frac{B(8, 72)}{B(8, 52)} = 0.085$; plug-in $0.8667^{20} = 0.057$. Probability of 6 or more: $0.066$ vs $0.041$. The tails are fatter: ignoring parameter uncertainty makes you overconfident.
If $\theta \mid D \sim \text{Beta}(\alpha', \beta')$ and $\tilde y \mid \theta \sim \text{Binomial}(m, \theta)$, the posterior predictive is the beta-binomial distribution $\tilde y \mid D \sim \text{BetaBinomial}(m, \alpha', \beta')$:
$$P(\tilde y = y\mid D) = \int_0^1 \binom{m}{y}\theta^y(1-\theta)^{m-y}\,\frac{\theta^{\alpha'-1}(1-\theta)^{\beta'-1}}{B(\alpha',\beta')}\,d\theta = \binom{m}{y}\frac{B(y+\alpha',\ m-y+\beta')}{B(\alpha',\beta')}, \quad y = 0,\dots,m.$$- Mean $m\bar p$ with $\bar p = \frac{\alpha'}{\alpha'+\beta'}$. Variance $m\bar p(1-\bar p)\,\dfrac{\alpha'+\beta'+m}{\alpha'+\beta'+1}$: the Binomial variance times an inflation factor that is at least 1.
- The integral is done by the same trick as before: the integrand is a Beta kernel with exponents $y + \alpha' - 1$ and $m - y + \beta' - 1$, whose area is a Beta function.
- As the data grow ($\alpha' + \beta' \to \infty$), the factor goes to 1 and the predictive becomes the plug-in Binomial: parameter uncertainty vanishes, observation noise stays.
- Code: SciPy
stats.betabinom(m, a, b); NumPyrodist.BetaBinomial(concentration1=a, concentration0=b, total_count=m).
Why do we need it?
Business questions are about future data ("how many sign-ups next week?"), not about θ. Plugging in a single θ ignores that θ is uncertain and gives intervals that are too narrow, especially with little data.
Where is it used?
Forecasting conversions for capacity or revenue planning, posterior predictive checks of a conversion model (Chapter 6.8), modelling overdispersed proportions (e.g. click rates that vary across pages), and the evidence $p(k)$ in Bayes factors.
How is it used?
Call stats.betabinom(m, a_post, b_post) for probabilities and intervals, or simulate: draw theta = rng.beta(a_post, b_post, S), then y = rng.binomial(m, theta). The simulation recipe works for any model, conjugate or not.
"To forecast, plug the posterior mean into the Binomial."
That keeps the mean right but drops parameter uncertainty, so the forecast interval is too narrow. Average over the posterior instead (beta-binomial, or simulate θ then ỹ).
"The predictive becomes narrow with enough data."
Only the parameter part vanishes. The Binomial noise of 20 new visitors stays forever: even with θ known exactly, "how many of the next 20 buy" is uncertain.
In an A/B framework like yours the posterior predictive answers "if we ship B, how many conversions will the next 10 000 users give?", and posterior predictive checks (Chapter 6.8) compare replicated data with the real counts. In your forecasting model the same idea, with a different likelihood, gives every forecast distribution: draw the parameters from the posterior, then draw $y_t$ from the Normal, Student-t or Negative Binomial likelihood (Chapter 7.14).
"We use a beta-binomial model, so our predictions are beta-binomial." (said without saying which object)
"Beta-binomial" names two different things. The model: Beta prior + Binomial likelihood. The distribution: the marginal or predictive distribution of a count when θ is integrated out. Say which one you mean.
Model answer: "With a Beta(α′, β′) posterior, the number of conversions among the next m users follows a beta-binomial distribution: same mean as Binomial(m, p̄), but variance inflated by $\frac{\alpha'+\beta'+m}{\alpha'+\beta'+1}$ because we are still unsure about the rate."
$\tilde y\mid D \sim \text{BetaBinomial}(m, \alpha', \beta')$: $P(\tilde y = y) = \binom{m}{y}\frac{B(y+\alpha', m-y+\beta')}{B(\alpha',\beta')}$.
Mean $m\bar p$; variance $m\bar p(1-\bar p)\frac{\alpha'+\beta'+m}{\alpha'+\beta'+1}$ ≥ plug-in Binomial variance.
Trap: plug-in forecasts are overconfident; "beta-binomial" = model or distribution, say which.
Quick check: posterior Beta(8, 52). What is the probability that the next two visitors both buy?
Not $(8/60)^2 = 0.0178$. Use the beta-binomial with $m = 2$, $y = 2$: $\frac{B(10, 52)}{B(8, 52)} = \frac{8 \times 9}{60 \times 61} = \frac{72}{3\,660} = 0.0197$. It is higher than the plug-in answer because the two visitors share the same unknown θ: if θ is high, both are likely to buy.
The Dirichlet-Multinomial update: categorical metrics core
Some metrics are not yes/no but "which one?": which plan a new customer picks (Basic, Pro, Enterprise), which star rating a user gives, which device they use. The unknown is now a list of shares that add up to 1, and the data are counts per category.
Everything from the Beta-Binomial carries over, one category at a time. The Dirichlet prior holds one pseudo-count per category (Chapter 4.11). The data add their real counts to the matching pseudo-counts. The posterior is again a Dirichlet. On the triangle picture, the cloud of plausible share-vectors moves from the prior's centre toward the observed shares and gets tighter as customers arrive.
Three ways to say it:
- Picture: a cloud of points on a triangle slides toward the observed mix and shrinks.
- Numbers: Dirichlet(2, 2, 2) plus counts (40, 25, 15) gives Dirichlet(42, 27, 17).
- Slogan: one pseudo-count per category, plus one real count per category.
Prior Dirichlet(2, 2, 2) for the shares of (Basic, Pro, Enterprise): centred at equal shares, worth $\alpha_0 = 6$ customers. Data from one variant: 40 Basic, 25 Pro, 15 Enterprise ($n = 80$).
- Add counts to pseudo-counts: $(2 + 40,\ 2 + 25,\ 2 + 15) = (42, 27, 17)$, total $\alpha_0' = 86$.
- Posterior mean shares: $(42/86,\ 27/86,\ 17/86) = (0.488,\ 0.314,\ 0.198)$; the raw shares were $(0.500, 0.3125, 0.1875)$.
- Weighted-average check for Enterprise: prior weight $6/86 = 0.070$. $0.070 \times \tfrac13 + 0.930 \times 0.1875 = 0.0233 + 0.1744 = 0.198$. ✓
- One share alone is a Beta (Chapter 4.11): Enterprise ~ Beta(17, 86 − 17) = Beta(17, 69), sd $0.043$, 95% of the belief in $[0.121, 0.288]$.
- Prediction: the next customer picks Enterprise with probability $17/86 = 0.198$. Among the next 10 customers, the number of Enterprise picks is beta-binomial(10, 17, 69): variance $1.75$ instead of the plug-in $1.59$.
The Dirichlet-Multinomial model. $\mathbf{p} = (p_1,\dots,p_K) \sim \text{Dirichlet}(\boldsymbol\alpha)$ and the counts $\mathbf{c} \mid \mathbf{p} \sim \text{Multinomial}(n, \mathbf{p})$. Then
$$p(\mathbf{p}\mid\mathbf{c}) \propto \underbrace{\frac{n!}{\prod_k c_k!}\prod_{k} p_k^{c_k}}_{\text{likelihood}} \cdot \underbrace{\frac{1}{B(\boldsymbol\alpha)}\prod_k p_k^{\alpha_k - 1}}_{\text{prior}} \propto \prod_k p_k^{\alpha_k + c_k - 1} \quad\Longrightarrow\quad \mathbf{p}\mid\mathbf{c} \sim \text{Dirichlet}(\boldsymbol\alpha + \mathbf{c}).$$- $B(\boldsymbol\alpha) = \frac{\prod_k \Gamma(\alpha_k)}{\Gamma(\alpha_0)}$ is the multivariate Beta function (it makes the Dirichlet's volume 1).
- Posterior means $\frac{\alpha_k + c_k}{\alpha_0 + n}$: a weighted average of the prior mean share $\alpha_k/\alpha_0$ and the data share $c_k/n$, with weight $\frac{\alpha_0}{\alpha_0 + n}$ on the prior.
- Marginals (one share on its own, ignoring the others): $p_k \mid \mathbf{c} \sim \text{Beta}(\alpha_k + c_k,\ \alpha_0 + n - \alpha_k - c_k)$. Merging categories adds their numbers.
- Predictive for the next $m$ customers: the Dirichlet-multinomial distribution, $P(\tilde{\mathbf{c}}\mid\mathbf{c}) = \frac{m!}{\prod_k \tilde c_k!}\frac{B(\boldsymbol\alpha + \mathbf{c} + \tilde{\mathbf{c}})}{B(\boldsymbol\alpha + \mathbf{c})}$ (SciPy
stats.dirichlet_multinomial). With $K = 2$ everything reduces to the Beta-Binomial. - Assumptions: customers choose independently, all with the same share vector $\mathbf{p}$. NumPyro:
dist.Dirichlet(concentration=alpha)anddist.Multinomial(total_count=n, probs=p).
Why do we need it?
Category metrics have several probabilities tied together by "they add to 1". Testing each category as a separate yes/no metric ignores that link. The Dirichlet-Multinomial updates all shares at once, exactly, by adding counts.
Where is it used?
Categorical metrics in Bayesian A/B tests (plan mix, rating distribution, device mix), topic models such as LDA (Dirichlet priors on topic shares), language-model smoothing of word counts ("add-α smoothing" is a Dirichlet posterior mean), and multi-class bandits.
How is it used?
Count customers per category for each variant, add the prior's pseudo-counts, and draw share vectors with rng.dirichlet(alpha + counts, size=4000). Answer one-category questions with the Beta marginal, and whole-mix questions (e.g. "did revenue per customer change?") from the joint draws.
"Analyse each category as its own Beta-Binomial, independently of the others."
Each marginal is a correct Beta, but the shares are tied: if Enterprise goes up, something else must go down. For any question that involves two or more shares together (the whole mix, revenue per customer from plan prices), use joint draws from the Dirichlet.
"The Dirichlet-Multinomial allows customers to differ."
The model assumes every customer in a variant picks with the same share vector, independently. If segments have different mixes, pooling them can distort the shares (Simpson's paradox, Chapter 4.6); model the segments explicitly (Chapter 6.5).
"A flat prior on the shares means Dirichlet(0, 0, 0)."
Every α must be positive. Dirichlet(1, 1, 1) is flat over the triangle; values below 1 push belief into the corners (a strong "sparse" claim).
This is the categorical-metric model of an A/B framework like yours: per variant, a Dirichlet prior on the category probabilities, Multinomial counts, and a Dirichlet posterior obtained by adding counts. One-category questions ("did the Enterprise share rise?") use the Beta marginals; whole-mix questions use joint draws, as in the right panel above. If the framework fits it with SVI, Dirichlet($\boldsymbol\alpha + \mathbf{c}$) is the exact answer to compare against.
$\mathbf{p}\sim\text{Dir}(\boldsymbol\alpha)$, $\mathbf{c}\mid\mathbf{p}\sim\text{Mult}(n,\mathbf{p})$ ⇒ $\mathbf{p}\mid\mathbf{c}\sim\text{Dir}(\boldsymbol\alpha + \mathbf{c})$; mean $\frac{\alpha_k + c_k}{\alpha_0 + n}$.
One share: Beta($\alpha_k + c_k$, $\alpha_0 + n - \alpha_k - c_k$). Predictive: Dirichlet-multinomial.
Trap: shares are tied (sum to 1): whole-mix questions need joint draws; all α must be positive.
Quick check: prior Dirichlet(1, 1, 1, 1) for four star ratings (1–4). Counts: 5, 10, 30, 55. What is the posterior mean share of 4-star ratings, and its marginal?
Posterior Dirichlet(6, 11, 31, 56), total 104. Mean share of 4 stars $= 56/104 \approx 0.538$ (raw share 0.55). Marginal: Beta(56, 104 − 56) = Beta(56, 48).
Gamma-Poisson: the conjugate pair for count metrics
Many metrics are counts with no upper limit: orders per user in a week, support tickets per day, page views per session. A common model says each unit's count is Poisson with an unknown rate λ (events per unit of exposure, Chapter 4.8). The conjugate prior for λ is a Gamma (Chapter 4.10).
The pseudo-count reading is the same as before, with a twist: a Gamma(a, b) prior acts like a pretend history of a events seen in b units of exposure (b users, or b days). The data add their real events to a and their real exposure to b. The posterior mean is "total events ÷ total exposure", pretend plus real.
Three ways to say it:
- Picture: a running tally with two columns: events and exposure; both the prior and the data write into them.
- Numbers: Gamma(6, 3) ("6 orders in 3 users") plus 26 orders from 10 users gives Gamma(32, 13): about 2.46 orders per user.
- Slogan: add events to the shape, add exposure to the rate.
Orders per user per week. Prior Gamma(shape 6, rate 3): mean $6/3 = 2$ orders per user, worth 3 users. Data: 10 users placed 3, 1, 4, 2, 0, 5, 3, 2, 4, 2 orders (total 26).
- Posterior: Gamma(6 + 26, 3 + 10) = Gamma(32, 13).
- Mean $32/13 = 2.46$. Weighted average: prior weight $3/13 = 0.231$, data mean $26/10 = 2.6$: $0.231 \times 2 + 0.769 \times 2.6 = 0.462 + 2.000 = 2.46$. ✓
- sd $= \sqrt{32}/13 = 0.435$ (a Gamma(a, b) has variance $a/b^2$); 95% of the belief between $1.68$ and $3.38$ (the prior's 95% range was $0.73$ to $3.89$).
- Next user's count (posterior predictive): Negative Binomial with $r = 32$ and $p = 13/14$, i.e. NB2 with mean $2.46$ and concentration $32$. Variance $2.46 + 2.46^2/32 = 2.65$, a bit more than a Poisson(2.46) (variance 2.46).
- $P(\text{next user orders nothing}) = (13/14)^{32} = 0.093$, against $e^{-2.46} = 0.085$ for the plug-in Poisson.
Prior $\lambda \sim \text{Gamma}(a, b)$ with shape $a$ and rate $b$ (density $\propto \lambda^{a-1}e^{-b\lambda}$, mean $a/b$). Data $y_i \mid \lambda \sim \text{Poisson}(\lambda t_i)$ with known exposures $t_i$ (all $t_i = 1$ for "per user"). Then
$$p(\lambda\mid y) \propto \lambda^{\sum y_i} e^{-\lambda\sum t_i} \cdot \lambda^{a-1}e^{-b\lambda} = \lambda^{a + \sum y_i - 1}e^{-(b + \sum t_i)\lambda} \;\Longrightarrow\; \lambda\mid y \sim \text{Gamma}\Big(a + \sum_i y_i,\ \ b + \sum_i t_i\Big).$$- Sufficient statistics: total events $\sum y_i$ and total exposure $\sum t_i$. Posterior mean $= w\cdot\frac{a}{b} + (1-w)\cdot\frac{\sum y_i}{\sum t_i}$ with $w = \frac{b}{b + \sum t_i}$.
- Predictive for a new unit with exposure 1: Negative Binomial, $\tilde y\mid y \sim \text{NB}(r = a',\ p = \frac{b'}{b'+1})$ in SciPy's
nbinom(n, p), which is NB2 with mean $a'/b'$ and concentration $a'$ (NumPyroNegativeBinomial2(mean=a'/b', concentration=a')). - Parameterization trap: here $b$ is a rate. SciPy uses
stats.gamma(a, scale=1/b); NumPyrng.gamma(a, 1/b); NumPyrodist.Gamma(concentration=a, rate=b). - Assumption: given λ, every unit's count is Poisson with the same rate. If users truly differ, the counts are overdispersed (more spread than a Poisson allows, Chapter 4.8), and this posterior for λ is too narrow.
Why do we need it?
Count metrics (orders, sessions, tickets) need a rate model, and the Gamma-Poisson gives the exact posterior of the rate and a predictive that already includes rate uncertainty, by adding two totals.
Where is it used?
Poisson count metrics in Bayesian A/B tests, event-rate monitoring (incidents per day), insurance claim rates, empirical-Bayes rate smoothing for small regions or stores, and as the origin of the Negative Binomial (a Gamma-mixed Poisson, Chapter 4.8).
How is it used?
Sum the events and the exposure per variant, add them to (a, b), then use stats.gamma(a + S, scale=1/(b + T)) for the rate and stats.nbinom(a + S, (b + T)/(b + T + 1)) for the next unit. Check overdispersion first (variance vs mean of the per-user counts).
"stats.gamma(32, 13) is my Gamma(32, 13) posterior."
SciPy's second positional argument is loc (a shift), not the rate! Write stats.gamma(32, scale=1/13). NumPyro's Gamma(32., 13.) does use the rate. Always check which convention a library uses.
"The Poisson model is fine for any count metric."
It forces variance = mean for every unit. Real per-user counts are usually overdispersed (heavy buyers and non-buyers). Then the Gamma-Poisson posterior for λ is overconfident; use a Negative Binomial likelihood (no simple conjugate prior for its dispersion, so MCMC or SVI).
Your A/B framework has Poisson likelihoods for count metrics; for a single variant with a Gamma prior on the rate the posterior is Gamma(a + total events, b + total exposure), a handy check for an SVI fit. Your forecasting model uses a Negative Binomial likelihood for counts: that is the Gamma-mixed Poisson again, but its dispersion parameter is usually learned (check your code), which is one reason that model needs SVI rather than formulas (Chapter 7.13).
$\lambda\sim\text{Gamma}(a, b)$ (rate $b$), $y_i\mid\lambda\sim\text{Pois}(\lambda t_i)$ ⇒ $\lambda\mid y \sim \text{Gamma}(a + \sum y_i,\ b + \sum t_i)$.
Prior = "a events in b units". Predictive: NB2(mean $a'/b'$, concentration $a'$) = SciPy nbinom(a', b'/(b'+1)).
Trap: SciPy gamma uses scale = 1/rate; Poisson assumes no overdispersion.
Quick check: prior Gamma(2, 1) for tickets per day. Over 7 days you see 35 tickets. What is the posterior mean?
Posterior Gamma(2 + 35, 1 + 7) = Gamma(37, 8), mean $37/8 = 4.625$ tickets per day. The data mean is 5; the prior mean 2 has weight $1/8$: $0.125 \times 2 + 0.875 \times 5 = 4.625$.
Normal-Normal with known variance: precisions add
For a continuous metric, such as average order value, a simple model says each order value is Normal around an unknown mean μ with a known spread σ. If the prior on μ is also Normal, the posterior is Normal again. This time the "counts" that add are precisions: precision is 1 / variance, "how sharp" a piece of information is. The prior brings precision $1/s_0^2$; each observation brings $1/\sigma^2$, so $n$ observations bring $n/\sigma^2$. The posterior precision is the sum, and the posterior mean is the precision-weighted average of the prior mean and the sample mean.
Three ways to say it:
- Picture: two bells, the prior and the data; the posterior bell sits between them, nearer the sharper one, and is sharper than both.
- Numbers: prior N(50, 10²), 25 orders with mean 56 and σ = 20: posterior mean 55.2, sd 3.7.
- Slogan: precisions add; means average, weighted by precision.
Average order value in euros. Prior $\mu \sim N(50, 10^2)$. Known order-to-order sd $\sigma = 20$. Data: $n = 25$ orders with mean $\bar y = 56$.
- Prior precision $1/10^2 = 0.01$. Data precision $n/\sigma^2 = 25/400 = 0.0625$ (the data alone have standard error $20/\sqrt{25} = 4$).
- Posterior precision $0.01 + 0.0625 = 0.0725$, so posterior variance $1/0.0725 = 13.79$ and sd $3.71$ (sharper than both 10 and 4).
- Weights: prior $0.01/0.0725 = 0.138$, data $0.0625/0.0725 = 0.862$.
- Posterior mean $0.138 \times 50 + 0.862 \times 56 = 6.90 + 48.28 = 55.17$.
- Next order's value (predictive): $N(55.17,\ 20^2 + 3.71^2)$, sd $\sqrt{400 + 13.79} = 20.34$: mostly order-to-order noise, plus a little uncertainty about μ.
Prior $\mu \sim N(m_0, s_0^2)$, data $y_1, \dots, y_n \mid \mu \overset{\text{iid}}{\sim} N(\mu, \sigma^2)$ with $\sigma$ known. Then $\mu \mid y \sim N(m_n, s_n^2)$ with
$$\frac{1}{s_n^2} = \frac{1}{s_0^2} + \frac{n}{\sigma^2}, \qquad m_n = s_n^2\left(\frac{m_0}{s_0^2} + \frac{n\bar y}{\sigma^2}\right) = w\,m_0 + (1-w)\,\bar y, \quad w = \frac{1/s_0^2}{1/s_0^2 + n/\sigma^2}.$$- Why: the log of each Normal is a quadratic in μ; the sum of two quadratics is a quadratic, i.e. a Normal again ("complete the square").
- Sufficient statistic: $\bar y$ (with $n$). Predictive: $\tilde y\mid y \sim N(m_n,\ \sigma^2 + s_n^2)$.
- Weight formula: the prior counts like $\sigma^2/s_0^2$ observations (here $400/100 = 4$ orders). This precision-weighted average is exactly the partial-pooling formula of Chapter 6.6, with the population distribution playing the prior.
- If σ is unknown too, the conjugate prior is Normal-Inverse-Gamma and the predictive becomes a Student-t. In code, NumPyro's
dist.Normal(loc, scale)takes the sd, not the variance.
Why do we need it?
It shows, in one line, how a Bayesian combines two noisy sources of information on a continuous scale: add precisions, average means by precision. The same rule explains shrinkage, Kalman filters and meta-analysis.
Where is it used?
Continuous A/B metrics with large samples (revenue per user, order value, latency) using a Normal approximation, Kalman filters for tracking and state-space forecasting, meta-analysis of several experiments, and the group means in hierarchical models (Chapter 6.6).
How is it used?
Convert each source to (mean, variance). Add the precisions, take the precision-weighted mean, and report $m_n \pm 1.96\,s_n$. For a large-sample A/B metric, use $\bar y$ with its standard error as the "data" bell.
"Average the prior mean and the sample mean 50/50."
Weight them by precision. A vague prior (large $s_0$) gets almost no weight; a sharp one can dominate a small sample.
"The posterior sd is the average of the prior sd and the standard error."
Precisions add, so the posterior is sharper than both: $3.71 \lt 4 \lt 10$ in the example.
"σ known is a harmless assumption."
It is a simplification. With large samples, plugging in the sample sd is fine; with small samples, treat σ as unknown (Normal-Inverse-Gamma prior, Student-t predictive) or fit the model with MCMC/SVI.
In an A/B framework like yours, a revenue-type metric with a Normal likelihood behaves like this formula when σ is well estimated: each variant's posterior mean is a precision-weighted blend of prior and data. If metrics are standardized with a single global scaler, $m_0$, $s_0$, σ and the effect are all in scaled units; convert back by multiplying by the global sd. The same "precisions add" rule, applied to segment means and a population mean, is the partial pooling of your hierarchical model (Chapters 6.5–6.6).
$\frac{1}{s_n^2} = \frac{1}{s_0^2} + \frac{n}{\sigma^2}$; $m_n = w\,m_0 + (1-w)\,\bar y$, $w = \frac{1/s_0^2}{1/s_0^2 + n/\sigma^2}$.
Predictive $N(m_n, \sigma^2 + s_n^2)$. Prior worth $\sigma^2/s_0^2$ observations.
Trap: weight by precision, not 50/50; the posterior is sharper than both sources.
Quick check: prior N(0, 1) on a standardized effect, data mean 0.5 with standard error 0.5. What is the posterior?
Precisions: prior 1, data $1/0.25 = 4$; total 5, so posterior sd $= \sqrt{1/5} = 0.447$. Mean $= (1 \times 0 + 4 \times 0.5)/5 = 0.4$. Posterior N(0.4, 0.447²).
The table of conjugate pairs, and what to do when there is none core
Conjugate pairs are like keys cut for specific locks. There are a handful of them, one for each common likelihood, and when your model is one of those simple shapes the door opens instantly. But real models quickly stop being one simple shape: a prior that matches your belief better, a heavy-tailed likelihood for outliers, a regression with many coefficients, groups that share a population distribution, a trend with Laplace priors on its slope changes. None of those has a key. Then the posterior must be computed: by a grid for one or two parameters, and by MCMC or variational inference for anything bigger.
Three ways to say it:
- Picture: a small key ring of exact formulas; for every other door you need a general tool.
- Numbers: a grid of 1 500 points is fine for 1 parameter; for 5 parameters it needs $1\,500^5 \approx 7.6 \times 10^{15}$ evaluations.
- Slogan: conjugacy for simple building blocks; MCMC and VI for whole models.
Which of these have a closed-form posterior?
- Conversion rate, Beta prior, Binomial data: yes, Beta($\alpha + k$, $\beta + n - k$).
- Orders per user, Gamma prior, Poisson data: yes, Gamma($a + \sum y$, $b + \sum t$).
- Order value, Normal prior on the mean, Normal data with known σ: yes, Normal with precisions added.
- Order value, Normal prior, Student-t data (robust to outliers): no. The t likelihood is not built from the same pieces as the Normal prior.
- A slope change δ with a Laplace prior and Normal data: no closed-form posterior (its MAP has a formula, soft-thresholding, from Chapter 5.3, but the whole posterior does not have a named family).
- Segment rates θ_g ~ Beta(μκ, (1 − μ)κ) with unknown μ and κ (a hierarchical model): no closed form for the whole posterior, although each θ_g given μ and κ is still a Beta. This "conditionally conjugate" structure is what Gibbs sampling exploits.
Common conjugate pairs (prior hyperparameters → posterior hyperparameters, and the posterior predictive):
| Likelihood (data) | Conjugate prior | Posterior | Posterior predictive |
|---|---|---|---|
| Bernoulli / Binomial($n$, θ): $k$ successes | Beta($\alpha$, $\beta$) | Beta($\alpha + k$, $\beta + n - k$) | beta-binomial |
| Categorical / Multinomial($n$, $\mathbf{p}$): counts $\mathbf{c}$ | Dirichlet($\boldsymbol\alpha$) | Dirichlet($\boldsymbol\alpha + \mathbf{c}$) | Dirichlet-multinomial |
| Poisson(λ$t_i$): counts $y_i$, exposures $t_i$ | Gamma($a$, rate $b$) | Gamma($a + \sum y_i$, $b + \sum t_i$) | Negative Binomial |
| Exponential(λ): waiting times $y_i$ | Gamma($a$, rate $b$) | Gamma($a + n$, $b + \sum y_i$) | Lomax (Pareto II) |
| Normal(μ, σ²), σ known | Normal($m_0$, $s_0^2$) | Normal: precisions add | Normal($m_n$, $\sigma^2 + s_n^2$) |
| Normal(μ, σ²), μ known | Inverse-Gamma($a$, $b$) on σ² | Inv-Gamma($a + n/2$, $b + \tfrac12\sum(y_i-\mu)^2$) | Student-t |
| Normal(μ, σ²), both unknown | Normal-Inverse-Gamma | Normal-Inverse-Gamma | Student-t |
- Not conjugate (posterior must be computed): logistic and Poisson regression with Normal priors on coefficients, Student-t likelihoods, Negative Binomial with unknown dispersion, Laplace priors, hierarchical models with unknown population parameters, mixtures, and almost every model with many interacting parameters.
- The general tools: grid approximation (1–2 parameters), the Laplace approximation (a Normal at the posterior mode), MCMC such as NUTS (Chapters 6.9–6.10), and variational inference such as SVI (Chapters 6.11–6.14).
Why do we need it?
To know instantly whether a model has an exact answer (use the formula and skip the machinery) or needs computation (choose MCMC or VI and check convergence). It also tells you which pieces of a big model can be tested against exact results.
Where is it used?
Conjugate pieces: A/B conversion and categorical metrics, rate monitoring, Kalman filters, Gibbs samplers in topic models. Non-conjugate: logistic regression, robust (Student-t) regression, Prophet-style forecasting models with Laplace priors, hierarchical A/B models fitted with NUTS or SVI in NumPyro.
How is it used?
Look up your likelihood in the table. If the whole model is a conjugate pair, use the update rule. If not, write the model in NumPyro and fit it with NUTS or SVI; use any conjugate sub-model (one variant, no pooling) as a unit test of the fit.
"If my model is not conjugate, I cannot do Bayesian inference on it."
You can always do it; you just cannot do it with a formula. NUTS and SVI handle non-conjugate models routinely. Conjugacy only removes the computation.
"Change the likelihood so that it becomes conjugate."
Choose the likelihood for the data (support, variance, tails), never for algebraic convenience. A Normal likelihood forced onto outlier-heavy data is wrong, conjugate or not.
"A grid works for any model if I am patient."
The cost grows like (points per axis)$^{d}$. Beyond 2 or 3 parameters a grid is hopeless, which is why MCMC and VI exist (Chapter 6.9).
Both of your projects, taken as a whole, are non-conjugate: the A/B framework adds hierarchical pooling and Student-t likelihoods; the forecasting model has Laplace priors on changepoint slopes, Fourier and regressor coefficients, and Normal, Student-t or Negative Binomial likelihoods. That is why both run SVI in NumPyro. The conjugate pieces (one variant's Beta-Binomial or Dirichlet-Multinomial without pooling) are still valuable as exact references to test the inference code against.
Pairs: Beta–Binomial, Dirichlet–Multinomial, Gamma–Poisson, Gamma–Exponential, Normal–Normal (σ known), Inv-Gamma–Normal (μ known), NIG–Normal.
Not conjugate: t likelihoods, Laplace priors, logistic/NB regression, hierarchical models → grid (tiny), NUTS or SVI.
Trap: choose likelihoods for the data, not for conjugacy; grids cost (points)$^d$.
Quick check: waiting times between orders are Exponential(λ). Prior Gamma(2, 10) on λ (rate: about 0.2 orders per minute). You observe 5 waits totalling 20 minutes. Posterior?
Gamma(2 + 5, 10 + 20) = Gamma(7, 30): mean $7/30 \approx 0.233$ orders per minute. For Exponential data, the count of observations goes onto the shape and the total waiting time onto the rate.
Recap, cheat sheet and practice
- A prior family is conjugate to a likelihood when the posterior stays in the family. Updating is then a change of hyperparameters: no integrals, no sampling.
- Beta-Binomial: Beta($\alpha$, $\beta$) + $k$ of $n$ → Beta($\alpha + k$, $\beta + n - k$), by adding exponents. Sufficient statistics $(k, n)$. Evidence = beta-binomial pmf.
- The prior is worth $\alpha + \beta$ pretend observations; the posterior mean is a weighted average of prior mean and data rate with weight $\frac{\alpha+\beta}{\alpha+\beta+n}$ on the prior.
- Two variants: same prior, independent posteriors; decisions use the gap (Chapter 6.4).
- Posterior predictive: beta-binomial, same mean as the plug-in Binomial but wider (parameter + observation uncertainty).
- Dirichlet-Multinomial: Dirichlet($\boldsymbol\alpha$) + counts → Dirichlet($\boldsymbol\alpha + \mathbf{c}$); single shares are Betas; whole-mix questions need joint draws.
- Gamma-Poisson: add events to the shape and exposure to the rate; predictive Negative Binomial. Normal-Normal: precisions add, means average by precision.
- Most real models (both your projects as a whole) are not conjugate: use NUTS or SVI, and use conjugate sub-models as exact tests.
Cheat sheet
| Idea | Formula | In words |
|---|---|---|
| Conjugacy | prior ∈ $\mathcal{F}$ ⇒ posterior ∈ $\mathcal{F}$ | same family, new numbers |
| Beta-Binomial update | Beta($\alpha + k$, $\beta + n - k$) | add successes and failures |
| Evidence | $\binom{n}{k}\frac{B(\alpha+k, \beta+n-k)}{B(\alpha,\beta)}$ | average likelihood over the prior |
| Weighted average | $w\,m_0 + (1-w)\frac{k}{n}$, $w = \frac{\alpha+\beta}{\alpha+\beta+n}$ | prior worth α + β visitors |
| Posterior sd | $\sqrt{\frac{m(1-m)}{\alpha+\beta+n+1}}$ | shrinks like $1/\sqrt{n}$ |
| Beta-binomial predictive | var $= m\bar p(1-\bar p)\frac{\alpha'+\beta'+m}{\alpha'+\beta'+1}$ | wider than plug-in |
| Dirichlet-Multinomial | Dirichlet($\boldsymbol\alpha + \mathbf{c}$); share $k$ ~ Beta($\alpha_k + c_k$, rest) | one pseudo-count per category |
| Gamma-Poisson | Gamma($a + \sum y$, $b + \sum t$); predictive NB2($a'/b'$, $a'$) | events → shape, exposure → rate |
| Normal-Normal | $\frac{1}{s_n^2} = \frac{1}{s_0^2} + \frac{n}{\sigma^2}$, precision-weighted mean | precisions add |
| No conjugacy | grid (1–2 params), NUTS, SVI | compute the posterior |
import numpy as np
from scipy import stats, integrate
# --- Beta-Binomial: prior Beta(2, 18), data 6 conversions of 40 ---
a, b, k, n = 2, 18, 6, 40
post = stats.beta(a + k, b + n - k) # Beta(8, 52)
print(post.mean(), post.std()) # 0.1333 0.0435
print(post.ppf([0.025, 0.975])) # [0.0604 0.2293]
# posterior mean as a weighted average of prior mean and data rate
w = (a + b) / (a + b + n) # 20 / 60
print(w, w * a / (a + b) + (1 - w) * k / n) # 0.3333 0.1333
# the evidence p(D): closed form (beta-binomial pmf) vs brute-force integral
print(stats.betabinom(n, a, b).pmf(k)) # 0.0740
print(integrate.quad(lambda t: stats.binom.pmf(k, n, t) * stats.beta.pdf(t, a, b), 0, 1)[0]) # 0.0740
# --- posterior predictive for the next m = 20 visitors ---
m = 20
bb = stats.betabinom(m, a + k, b + n - k) # beta-binomial predictive
plug = stats.binom(m, post.mean()) # plug-in Binomial (too narrow)
print(bb.mean(), bb.var(), plug.var()) # 2.667 3.031 2.311
print(bb.pmf(0), plug.pmf(0)) # 0.085 0.057
rng = np.random.default_rng(0)
theta = rng.beta(a + k, b + n - k, size=200_000) # 1. draw theta from the posterior
y = rng.binomial(m, theta) # 2. draw new data with that theta
print(y.var(), (y == 0).mean()) # about 3.03 and 0.085
# --- two variants: the checkout test, flat priors ---
for name, kk in [("A", 50), ("B", 60)]:
d = stats.beta(1 + kk, 1 + 500 - kk)
print(name, d.mean(), d.ppf([0.025, 0.975])) # A 0.1016 [0.0767 0.1295] / B 0.1215 [0.0944 0.1515]
# --- Dirichlet-Multinomial: prior (2, 2, 2), counts (40, 25, 15) ---
alpha = np.array([2.0, 2.0, 2.0])
counts = np.array([40, 25, 15])
post_d = alpha + counts # Dirichlet(42, 27, 17)
print(post_d / post_d.sum()) # [0.4884 0.314 0.1977]
print(stats.beta(17, 86 - 17).ppf([0.025, 0.975])) # Enterprise share alone: [0.121 0.2876]
print(stats.dirichlet_multinomial(post_d, 10).pmf([5, 3, 2]), # 0.0758 predictive, next 10 customers
stats.multinomial(10, post_d / post_d.sum()).pmf([5, 3, 2])) # 0.0847 plug-in
# --- Gamma-Poisson: prior Gamma(6, rate 3), 26 orders from 10 users ---
A, B = 6 + 26, 3 + 10
lam = stats.gamma(A, scale=1 / B) # SciPy wants scale = 1/rate !
print(lam.mean(), lam.std(), lam.ppf([0.025, 0.975])) # 2.4615 0.4351 [1.6837 3.3848]
nb = stats.nbinom(A, B / (B + 1)) # predictive for the next user
print(nb.mean(), nb.var(), nb.pmf(0)) # 2.4615 2.6509 0.0933
# --- Normal-Normal (sigma known): prior N(50, 10^2), 25 orders, mean 56, sigma 20 ---
m0, s0, sigma, nn, ybar = 50, 10, 20, 25, 56
prec = 1 / s0**2 + nn / sigma**2 # precisions add
mn = (m0 / s0**2 + nn * ybar / sigma**2) / prec # precision-weighted mean
print(mn, np.sqrt(1 / prec), np.sqrt(sigma**2 + 1 / prec)) # 55.17 3.714 20.34
# --- NumPyro: SVI with a Beta guide should recover the exact Beta(8, 52) ---
import jax
import numpyro
import numpyro.distributions as dist
from numpyro.infer import SVI, Trace_ELBO
def model(k, n):
theta = numpyro.sample("theta", dist.Beta(2.0, 18.0)) # Beta(concentration1, concentration0)
numpyro.sample("k", dist.Binomial(total_count=n, probs=theta), obs=k)
def guide(k, n):
qa = numpyro.param("qa", 1.0, constraint=dist.constraints.positive)
qb = numpyro.param("qb", 1.0, constraint=dist.constraints.positive)
numpyro.sample("theta", dist.Beta(qa, qb))
lr = lambda step: 0.05 * 0.5 ** (step / 1000) # learning rate halves every 1000 steps
svi = SVI(model, guide, numpyro.optim.Adam(lr), Trace_ELBO(num_particles=16))
res = svi.run(jax.random.PRNGKey(0), 6000, k=6, n=40, progress_bar=False)
qa, qb = float(res.params["qa"]), float(res.params["qb"])
print(qa, qb, qa / (qa + qb)) # about 7.97 51.97 0.133 (exact: 8, 52, 0.1333)
1. Prior Beta(3, 7). Then 12 of 40 visitors convert. What is the posterior?
2. A Beta prior is worth $\alpha + \beta = 50$ visitors and the test has $n = 450$ visitors. How much weight does the prior mean get in the posterior mean?
3. Compared with the plug-in Binomial($m$, posterior mean), the beta-binomial posterior predictive has…
4. Prior Dirichlet(1, 1, 1); counts (10, 20, 30). What is the posterior mean share of the third category?
5. Prior Gamma(shape 4, rate 2) on tickets per day. Over 10 days you see 30 tickets. Posterior?
6. Which of these models has no closed-form (conjugate) posterior?
Practice problems
A. Flat prior Beta(1, 1). Visitors arrive one by one: buy, no, no, buy, no. Update after each visitor, then show the final answer equals one batch update.
Start (1, 1). Buy → (2, 1). No → (2, 2). No → (2, 3). Buy → (3, 3). No → (3, 4). Batch: $k = 2$, $n - k = 3$, so Beta(1 + 2, 1 + 3) = Beta(3, 4). Same answer, because adding is order-free. Posterior mean $3/7 \approx 0.43$; the raw rate is $2/5 = 0.4$.
B. With prior Beta(2, 18), how many visitors do you need before the prior's weight in the posterior mean falls to 10% or less?
$w = \frac{20}{20 + n} \le 0.1 \iff 20 \le 2 + 0.1n \iff n \ge 180$. After 180 visitors the data carry at least 90% of the weight.
C. Posterior Beta(8, 52). What is the probability that the next 3 visitors all buy? Compare with the plug-in answer.
Beta-binomial with $m = 3$, $y = 3$: $\frac{B(11, 52)}{B(8, 52)} = \frac{8 \times 9 \times 10}{60 \times 61 \times 62} = \frac{720}{226\,920} = 0.00317$. Plug-in: $(8/60)^3 = 0.00237$. The predictive is about a third higher, because the three visitors share the same unknown θ: if θ is high, all three are likely to buy together.
D. Prior Dirichlet(2, 2, 2) for (Basic, Pro, Enterprise). Variant A: counts (40, 25, 15). Variant B: counts (32, 22, 30). Give both posteriors, the Enterprise posterior mean of each, and the posterior mean of the difference.
A: Dirichlet(42, 27, 17), total 86, Enterprise mean $17/86 = 0.198$. B: Dirichlet(34, 24, 32), total 90, Enterprise mean $32/90 = 0.356$. Because the posteriors are independent, the mean of the difference is the difference of the means: $0.356 - 0.198 = 0.158$. Marginals: A's Enterprise share ~ Beta(17, 69), B's ~ Beta(32, 58). For the probability that B's share is higher, draw from both (Chapter 6.4).
E. Interview: "Your A/B framework fits a Beta-Binomial model with SVI. How would you convince me the fit is right?"
"For a single variant without pooling the model is conjugate, so the exact posterior is Beta(α + k, β + n − k). I compare the SVI posterior's mean, sd and 95% interval with that closed form (and, for the Dirichlet-Multinomial, with Dirichlet(α + c)). If the guide is a Beta, SVI should recover nearly the same two numbers; with an AutoNormal guide on the logit scale, it should be close but not identical. Only once the simple case matches do I trust the fit of the full hierarchical model, where no formula exists, and there I also compare against NUTS on a subset."
F. A standardized lift has prior N(0, 0.5²). The experiment estimates the lift as 0.8 with standard error 0.4 (treat this as a Normal likelihood with known σ). Find the posterior.
Prior precision $1/0.25 = 4$; data precision $1/0.16 = 6.25$; total $10.25$, so posterior sd $= \sqrt{1/10.25} = 0.312$. Posterior mean $= (4 \times 0 + 6.25 \times 0.8)/10.25 = 5/10.25 = 0.488$. The sceptical prior pulls the estimate from 0.8 to about 0.49 and makes it more precise.
Credible intervals and posterior decisions
A posterior is a whole curve, but people ask for a number, a range, or a yes/no. This chapter turns posteriors into those answers honestly: one-number summaries (mean, median, mode) and what each one is best at; credible intervals, equal-tailed or highest-density; the probability that B beats A; the far more useful probability that B beats A by enough to matter; relative lift; expected loss; and decision rules you could defend to a product team. It ends with the precise difference between a credible interval and a confidence interval.
- Summarize a posterior by its mean, median or mode, see when they differ (skewed posteriors), and match each to the "cost of being wrong" it minimizes
- Build and read a credible interval as a direct probability statement about θ given the data
- Compare the equal-tailed interval (ETI) and the highest-density interval (HDI), including what happens under a change of scale
- Compute $P(\theta_B \gt \theta_A \mid D)$ by Monte Carlo (a cloud of draws and the diagonal) and know its Monte Carlo error
- Compute $P(\theta_B - \theta_A \gt \delta \mid D)$ for a practically meaningful δ, and explain why it is more useful than $P(\theta_B \gt \theta_A \mid D)$
- Work with the posterior of the relative lift, the expected loss, and simple decision rules, including how often a rule ships a variant that is not really better
- Explain the difference between a credible interval and a confidence interval precisely, the way an interviewer wants to hear it
What we need from earlier chapters: the posterior and posterior draws (Chapter 6.1); the A/B Beta-Binomial model, two independent Beta posteriors (Chapter 6.3); quantiles and the CDF (Chapter 4.4); MAP and its dependence on the scale (Chapter 5.2); absolute vs relative lift (Chapter 5.7); confidence intervals (Chapter 5.8); peeking and multiple comparisons (Chapter 5.11). Running example: the checkout test, 50 of 500 conversions in A and 60 of 500 in B, with flat Beta(1, 1) priors, so $\theta_A \mid D \sim$ Beta(51, 451) and $\theta_B \mid D \sim$ Beta(61, 441). A "point" (percentage point) is 0.01 on the rate scale: 10% → 11% is +1 point. The syllabus writes $P(\theta_A \gt \theta_B \mid D)$; here B is the new variant, so we ask about $\theta_B \gt \theta_A$. It is the same idea with the labels swapped.
One number from a posterior: mean, median or mode? core
A posterior is a hill of belief. If someone insists on one number, there are three natural ways to pick it. The mode is the top of the hill (the single most plausible value). The median splits the hill's area in half (50% belief below, 50% above). The mean is the balance point (where a cardboard cut-out of the hill would balance on a finger).
For a symmetric hill the three coincide. For a skewed hill, typical for a conversion rate estimated from few visitors or near 0%, they spread out: the long tail pulls the mean furthest, the median a little, the mode not at all. Which one is "right"? None and all: each is the best answer to a different question about how bad different mistakes are.
Three ways to say it:
- Picture: top of the hill (mode), half the area (median), balance point (mean).
- Numbers: 2 conversions out of 41 visitors, flat prior: mode 4.9%, median 6.3%, mean 7.0%.
- Slogan: squared error → mean, absolute error → median, "exactly right or nothing" → mode.
A new landing page: 2 conversions among the first 41 visitors, flat prior Beta(1, 1). Posterior Beta(1 + 2, 1 + 39) = Beta(3, 40).
- Mode (the peak): $\frac{\alpha - 1}{\alpha + \beta - 2} = \frac{2}{41} = 0.0488$. With a flat prior this equals the observed rate $k/n$, the MLE.
- Mean (the balance point): $\frac{\alpha}{\alpha + \beta} = \frac{3}{43} = 0.0698$.
- Median (half the area on each side): no simple formula; SciPy's
stats.beta(3, 40).median()gives $0.0632$. (Rule of thumb for Betas with α, β > 1: $\frac{\alpha - 1/3}{\alpha + \beta - 2/3} = \frac{2.667}{42.33} = 0.0630$.) - Order: mode $\lt$ median $\lt$ mean. That is the signature of a right-skewed posterior (a long tail toward higher rates).
- Expected squared error of each guess: $E[(\theta - d)^2] = Var(\theta) + (E\theta - d)^2$. At the mean: $0.001475$. At the mode: $0.001475 + (0.0698 - 0.0488)^2 = 0.001475 + 0.000441 = 0.001916$, about 30% worse.
A loss function $L(\theta, d)$ says how costly it is to report $d$ when the truth is θ. The posterior expected loss of a report $d$ is $\rho(d) = E[L(\theta, d) \mid D]$, the cost averaged over everything the posterior still allows. The best report minimizes it:
| Loss | $L(\theta, d)$ | Best one-number summary |
|---|---|---|
| squared error | $(\theta - d)^2$ | posterior mean $E[\theta\mid D]$ |
| absolute error | $|\theta - d|$ | posterior median |
| all-or-nothing | 0 if $|\theta - d| \lt \varepsilon$, else 1 | the mode (MAP), as $\varepsilon \to 0$ |
- Why the mean wins for squared error: $E[(\theta - d)^2] = Var(\theta\mid D) + (E[\theta\mid D] - d)^2$; the first term does not depend on $d$, the second is zero exactly at $d = E[\theta\mid D]$.
- Change of scale: the median commutes with any increasing transformation $g$ (median of $g(\theta)$ = $g$(median of θ)). The mean and the mode do not: $E[g(\theta)] \ne g(E[\theta])$ in general, and the mode depends on the scale (Chapter 5.2).
- With lots of data the posterior becomes nearly symmetric and the three summaries agree.
Why do we need it?
Dashboards, reports and downstream systems need one number per metric. Picking the summary that matches how errors cost you avoids systematic bias, especially for rare events and early in an experiment when posteriors are skewed.
Where is it used?
Point estimates in Bayesian A/B reports (usually the mean or median), MAP estimates in regularized ML models (ridge and lasso are MAPs, Chapter 5.3), median forecasts in forecasting (they minimize absolute error), and NumPyro/ArviZ summary tables that show mean, sd and quantiles.
How is it used?
From draws: x.mean(), np.median(x); the mode needs a density estimate or the closed form. Report the mean or median together with an interval, and say which one you report. Prefer the median for skewed quantities such as ratios and lifts.
"The posterior mode (MAP) is the Bayesian answer."
It is one summary, best only under all-or-nothing loss, and it depends on the scale. For Beta(3, 40) the mode of θ is 4.9%, but if you compute the mode of the log-odds $\log\frac{\theta}{1-\theta}$ and convert back, you get 7.0%. The median gives the same answer on every scale.
"Mean, median and mode are interchangeable."
Only for symmetric posteriors. For skewed ones (small data, rates near 0 or 1, ratios and lifts) say which one you report.
Your A/B framework gets posterior draws (for example from SVI). Means and medians of draws are easy and stable; modes need a density estimate and depend on the parameterization (a logit-scale guide and a rate-scale report have different modes). For a conversion rate with thousands of users the three agree; for small segments, and especially for relative lifts, report the median (or mean) and say which.
Mode = peak (MAP, all-or-nothing loss); median = half the area (absolute loss); mean = balance point (squared loss).
$E[(\theta-d)^2] = Var + (E\theta - d)^2$ ⇒ the mean minimizes squared error. Right skew: mode < median < mean.
Trap: the mode (and the mean) change under a change of scale; the median does not.
Quick check: posterior Beta(21, 181), from 20 of 200 with a flat prior. Are mode, median and mean far apart?
Mode $20/200 = 0.1000$, mean $21/202 = 0.1040$, median $\approx \frac{21 - 1/3}{202 - 2/3} = 0.1026$. They differ by less than half a point: with 200 visitors the posterior is nearly symmetric.
Credible intervals: a range with a probability attached core
One number hides the uncertainty. A range shows it. A credible interval is a range of parameter values that holds a chosen share, say 95%, of the posterior's area. Because the posterior is a probability distribution over θ given the data, you can read the interval literally: "given my model, my prior and these data, there is a 95% probability that the conversion rate lies between 6.0% and 22.9%."
That plain sentence is exactly what most people wish a confidence interval meant. For a credible interval it is true (inside the model). The price is that the statement depends on the prior and the likelihood you chose.
Credible intervals are not the only probability statements you can read off a posterior. Any area works: $P(\theta \gt 10\% \mid D)$, $P(\theta \lt 20\% \mid D)$. Intervals are just the most common summary.
Three ways to say it:
- Picture: shade 95% of the area under the posterior; the shaded stretch of the axis is the interval.
- Numbers: Beta(8, 52): 95% credible interval 6.0% to 22.9%; $P(\theta \gt 10\% \mid D) = 0.77$.
- Slogan: a credible interval is a probability statement about θ, given the data.
The posterior from Chapter 6.3: Beta(8, 52) (prior Beta(2, 18), then 6 of 40 visitors).
- Lower end: the 2.5% quantile,
stats.beta(8, 52).ppf(0.025)$= 0.0604$. - Upper end: the 97.5% quantile, $0.2293$.
- Check: $F(0.2293) - F(0.0604) = 0.975 - 0.025 = 0.95$. So $P(0.0604 \le \theta \le 0.2293 \mid D) = 0.95$.
- Other areas: $P(\theta \gt 0.10 \mid D) = 1 - F(0.10) = 0.766$; $P(\theta \lt 0.20 \mid D) = 0.925$.
- From 4 000 posterior draws (what MCMC or SVI give you), the 2.5% and 97.5% sample quantiles typically land within 0.001 to 0.002 of these exact values.
A $100(1-\gamma)\%$ credible interval for θ is any interval $[L, U]$ with
$$P(L \le \theta \le U \mid D) = \int_L^U p(\theta\mid D)\,d\theta = 1 - \gamma .$$- $L$ and $U$ are numbers computed from this dataset; the probability is over θ, which the posterior treats as uncertain.
- Many intervals hold 95%; the two standard choices are the equal-tailed and the highest-density interval (next section). One-sided statements ("θ is above 6.8% with probability 95%") are fine too.
- It is conditional on the whole model: change the prior or the likelihood and the interval changes. With little data the prior matters; with lots of data it hardly does.
- From draws $\theta^{(1)},\dots,\theta^{(S)}$: the equal-tailed interval is the pair of sample quantiles
np.quantile(draws, [0.025, 0.975]).
Why do we need it?
A point estimate without a range invites over-confidence. A credible interval says how precisely the data pin down θ, in a sentence a stakeholder can take literally.
Where is it used?
Bayesian A/B reports (interval for each variant and for the lift), forecast intervals from posterior predictive draws, NumPyro and ArviZ summary tables (quantiles or HDIs of every parameter), and uncertainty bands on fitted curves such as a forecasting model's trend.
How is it used?
Closed form: stats.beta(a, b).ppf([0.025, 0.975]). From draws: np.quantile(x, [0.025, 0.975]). Report it with the model and prior ("flat prior, Binomial likelihood"), and choose the level before looking (95% or 90% are common).
"The 95% credible interval will contain the true θ in 95% of experiments, whatever θ is."
That is a frequentist coverage promise, which credible intervals do not make for every fixed θ. Averaged over θ values drawn from the prior, they cover exactly 95%; for a particular θ far from what the prior expects, coverage can be much worse (the last section shows this).
"The credible interval is objective; it comes from the data."
It comes from the data and the model, including the prior. Always report which prior and likelihood were used.
"Values just outside the 95% interval are impossible."
They are less plausible, not impossible: there is still 5% of belief outside, split between the tails.
When your A/B framework reports intervals from posterior draws, say them literally: "given our model and prior, there is a 95% probability that B's conversion rate is between x and y." If the metric was standardized with the global scaler, transform the draws back to the original units first and take the quantiles afterwards; quantiles commute with that increasing linear map, so both orders agree, but means of nonlinear transforms do not.
$[L, U]$ is a 95% credible interval if $P(L \le \theta \le U \mid D) = 0.95$. From draws: np.quantile(x, [0.025, 0.975]).
Beta(8, 52): 95% interval $[0.060, 0.229]$; $P(\theta \gt 0.10 \mid D) = 0.77$.
Trap: it is a probability about θ given the data and the model; it is not a coverage guarantee for every θ.
Quick check: a posterior has $P(\theta \lt 0.05 \mid D) = 0.02$ and $P(\theta \gt 0.12 \mid D) = 0.03$. Is [0.05, 0.12] a 95% credible interval?
Yes: $P(0.05 \le \theta \le 0.12 \mid D) = 1 - 0.02 - 0.03 = 0.95$. It is not equal-tailed (2% and 3% in the tails), but it is a valid 95% credible interval.
Equal-tailed or highest-density interval? core
There are many ways to pick a stretch of the axis that holds 95% of the posterior. Two are standard.
- The equal-tailed interval (ETI) cuts off 2.5% of the area on the left and 2.5% on the right. It is the simplest: two quantiles.
- The highest-density interval (HDI) keeps the most plausible values. Picture the posterior as an island and flood it: lower the water level until exactly 95% of the island's area sticks out of the water. The dry land is the HDI. Every value inside is more plausible than every value outside, and for a one-humped posterior it is the shortest 95% interval.
For a symmetric hill the two are the same. For a skewed hill the HDI slides toward the peak and gets shorter. For a two-humped hill the HDI can be two separate pieces, which an ETI can never show.
Three ways to say it:
- Picture: ETI = trim equal tails; HDI = flood the island to a level and keep the dry land.
- Numbers: Beta(3, 40): ETI 1.5% to 16.2% (width 14.7 pts), HDI 0.8% to 14.5% (width 13.7 pts).
- Slogan: ETI = equal tails, same on every scale; HDI = most plausible values, shortest, scale-dependent.
Posterior Beta(3, 40) (2 of 41 visitors, flat prior) and 95% mass.
- ETI: the 2.5% and 97.5% quantiles: $[0.0150,\ 0.1616]$, width $0.1467$.
- HDI: search for the shortest interval with 95% inside: $[0.0080,\ 0.1452]$, width $0.1372$ (about 6% shorter).
- Check the "water level": the density is the same at both HDI ends, $p(0.0080) = p(0.1452) = 1.60$. At the ETI ends it is not ($p(0.0150) = 4.3$ but $p(0.1616) = 0.93$): the ETI keeps some low-density values on the right and drops higher-density values on the left.
- Change of scale: on the log-odds scale $\log\frac{\theta}{1-\theta}$ the ETI is $[-4.19, -1.65]$, which maps back to exactly the same $[0.0150, 0.1616]$. But the HDI computed on the log-odds scale maps back to $[0.0170, 0.1744]$, a different interval from the HDI on the θ scale.
- An extreme case: Beta(1, 20) (0 of 19 visitors). Its density is highest at θ = 0. The 95% HDI is $[0,\ 0.139]$ and includes the peak; the 95% ETI $[0.0013,\ 0.168]$ cuts the most plausible values off.
For mass $1 - \gamma$ (e.g. 0.95), with posterior CDF $F$ and density $p(\theta\mid D)$:
$$\text{ETI} = \big[F^{-1}(\gamma/2),\ F^{-1}(1 - \gamma/2)\big], \qquad \text{HDI} = \{\theta : p(\theta\mid D) \ge c\}\ \text{ with } c \text{ chosen so that } P(\theta \in \text{HDI}\mid D) = 1 - \gamma .$$- For a one-humped (unimodal) posterior the HDI is an interval, the shortest one with that mass, and the density is equal at its two ends. For a multi-humped posterior it can be a union of intervals.
- Invariance: for any increasing transformation $g$, the ETI of $g(\theta)$ is $g$(ETI of θ). The HDI has no such property, because densities change shape when you change variables.
- From draws: ETI = sample quantiles; HDI ≈ the shortest window that contains 95% of the sorted draws (
az.hdiin ArviZ,numpyro.diagnostics.hpdiin NumPyro,LA.stats.hdihere). The HDI of draws is noisier than the ETI, especially in the tails, so use many draws.
Why do we need it?
For skewed posteriors, such as rates near 0, scale parameters, and lifts, the choice changes the reported range. You need to know which one a tool reports and when they disagree in a way that matters.
Where is it used?
ArviZ and many Bayesian papers report HDIs, and so does NumPyro's print_summary (a 90% HPDI by default); most A/B dashboards report quantile-based intervals; HDIs appear in "ROPE" decision rules; ETIs are preferred for transformed quantities like lifts and odds ratios.
How is it used?
Default to the ETI when the result will be transformed or compared across scales, or when the posterior is close to symmetric. Use the HDI to show the most plausible values of a strongly skewed or multi-humped posterior. Always label which one, and the mass.
"The HDI is always better because it is shorter."
It is the shortest on this scale. Report the same posterior on another scale (log-odds, relative lift, a standardized metric) and the HDI changes, while the ETI just transforms along. For quantities that will be transformed, the ETI is the safer default.
"An equal-tailed interval always contains the most plausible value."
Not for posteriors piled up at a boundary: for Beta(1, 20) the peak is at 0, and the 95% ETI starts at 0.13%. The HDI keeps the peak.
"The columns 5.0% and 95.0% in NumPyro's print_summary are the 5% and 95% quantiles."
They are the two ends of a 90% highest-density interval (NumPyro calls it HPDI), even though the labels look like quantiles. For a skewed posterior the two differ: for Beta(1, 20) the 90% HDI is [0, 0.109], while the 5% and 95% quantiles are [0.003, 0.139].
"An HDI from 500 draws is precise."
HDI endpoints from draws depend on the sparse tails and wobble from run to run. Use thousands of draws, or the closed form when you have it.
"Equal-tailed and highest-density intervals are the same thing."
They agree for symmetric unimodal posteriors and differ otherwise. ETI: quantiles $\gamma/2$ and $1-\gamma/2$, invariant to monotone transformations. HDI: the region of highest density with the given mass, the shortest such interval for unimodal posteriors, possibly disjoint for multimodal ones, and not invariant to reparameterization.
Model answer: "For a skewed posterior such as a small-sample conversion rate, I report the equal-tailed interval when the number will be transformed or compared, and the HDI when I want the most plausible values; I always say which, and the mass."
ETI $= [F^{-1}(\gamma/2), F^{-1}(1-\gamma/2)]$: equal tails, transforms with the parameter.
HDI $= \{\theta: p(\theta\mid D) \ge c\}$: shortest (unimodal), equal density at the ends, can be several pieces, scale-dependent.
Beta(3, 40), 95%: ETI [0.015, 0.162], HDI [0.008, 0.145].
Quick check: a posterior is exactly Normal(0.02, 0.01²). What are its 95% ETI and HDI?
Both are $0.02 \pm 1.96 \times 0.01 = [0.0004, 0.0396]$. A symmetric unimodal posterior has equal tails at equal density, so the two intervals coincide.
The probability that B beats A, by Monte Carlo core
"Is B better than A?" becomes, in Bayesian language, "how much of my belief says $\theta_B \gt \theta_A$?" The two posteriors from Chapter 6.3 describe everything we believe about the two rates. To get the probability, use the most useful trick in applied Bayesian work: simulate. Estimating a number by random simulation is called Monte Carlo (after the casino). Draw a plausible θ_A from A's posterior and a plausible θ_B from B's posterior. That pair is one plausible world. Repeat 4 000 times and count the worlds in which B wins.
On a picture, each world is a dot with coordinates (θ_A, θ_B). The diagonal line θ_B = θ_A splits the plane: dots above it are worlds where B is better. The answer is simply the share of the cloud above the diagonal. Because it is estimated from random draws it has a small Monte Carlo error, which shrinks as you draw more.
Three ways to say it:
- Picture: a cloud of dots and a diagonal; count the dots above the line.
- Numbers: checkout test: about 3 370 of 4 000 draws have θ_B > θ_A, so $P(\theta_B \gt \theta_A \mid D) \approx 0.84$ (exact 0.843).
- Slogan: any question about the posterior = a share of the posterior draws.
Checkout test with flat priors: $\theta_A \mid D \sim$ Beta(51, 451), $\theta_B \mid D \sim$ Beta(61, 441).
- Draw $\theta_A^{(s)}$ from Beta(51, 451) and $\theta_B^{(s)}$ from Beta(61, 441) for $s = 1, \dots, 4\,000$ (independently, since the two posteriors are independent).
- Count the draws with $\theta_B^{(s)} \gt \theta_A^{(s)}$: about 3 370. Estimate $\hat P = 3\,370/4\,000 \approx 0.843$.
- Monte Carlo standard error: $\sqrt{\hat P(1-\hat P)/S} = \sqrt{0.843 \times 0.157 / 4\,000} = 0.0058$. So report "0.84", not "0.8427".
- Exact value by a one-dimensional integral (next box): $0.8429$.
- Quick approximation: $\theta_B - \theta_A$ has mean $0.0199$ and sd $0.0198$, so $P \approx \Phi(0.0199/0.0198) = \Phi(1.005) \approx 0.842$.
- For comparison, the two-sided p-value of this test was $0.31$ (Chapter 5.6), and the one-sided p-value $0.156$. With flat priors and this much data, $P(\theta_B \gt \theta_A \mid D) \approx 1 - 0.156$: similar numbers, different meanings (careful box).
For two parameters with joint posterior $p(\theta_A, \theta_B \mid D)$:
$$P(\theta_B \gt \theta_A \mid D) = \iint_{\theta_B \gt \theta_A} p(\theta_A, \theta_B\mid D)\,d\theta_A\,d\theta_B \;\approx\; \frac{1}{S}\sum_{s=1}^{S} \mathbf{1}\big[\theta_B^{(s)} \gt \theta_A^{(s)}\big],$$where $\mathbf{1}[\cdot]$ is 1 when the condition holds and 0 otherwise, and $(\theta_A^{(s)}, \theta_B^{(s)})$ are draws from the joint posterior.
- When the posteriors are independent (separate priors, separate data) the integral collapses to one dimension: $\int p(\theta_A \mid D)\,P(\theta_B \gt \theta_A \mid D)\,d\theta_A$, which is how the exact values here are computed.
- Monte Carlo error of a probability estimated from $S$ independent draws: $\sqrt{P(1-P)/S}$; at most $0.5/\sqrt{S}$ (0.008 for $S = 4\,000$). MCMC draws are correlated, so use the effective sample size instead of $S$ (Chapter 6.10).
- In a hierarchical model the θ's are correlated in the posterior. Then you must compare draws with the same index s (the same joint draw), never shuffle or draw them separately.
Why do we need it?
It answers the business question directly ("how likely is it that B is better?") from any posterior, conjugate or not, without formulas: you only need draws. The same recipe gives every other decision quantity.
Where is it used?
Bayesian A/B testing tools (often called "chance to beat control"), Thompson-sampling bandits (choose B with probability P(B best)), multi-variant tests (P(each variant is the best)), and model comparisons from posterior draws in NumPyro.
How is it used?
Get draws: tA = rng.beta(aA, bA, S), tB = rng.beta(aB, bB, S) (or posterior samples from NUTS/SVI). Compute (tB > tA).mean(). Report it with its Monte Carlo error, and alongside the size of the gap (next section).
"P(θ_B > θ_A | D) = 0.84 means B converts 84% better."
It is a probability about the direction of the difference, not its size. B could be better by 0.01 points with probability 0.99. Always report the gap (or lift) and its interval too.
"0.84 is 1 minus the p-value."
The two-sided p-value was 0.31, so 1 − p = 0.69, a different number. With flat priors and large samples, $P(\theta_B \gt \theta_A\mid D)$ is close to 1 minus the one-sided p-value (0.844 here), but the meanings differ: one is the probability of B being better given the data; the other is the probability of data this extreme if there were no difference.
"Draw A's and B's samples separately from my NUTS output and pair them up at random."
Only valid when the posteriors are independent. In a hierarchical model, keep draw $s$ of θ_A with draw $s$ of θ_B: they come from the same joint draw and carry the correlation.
This is the posterior decision quantity of an A/B framework like yours; the syllabus writes it as $P(\theta_A \gt \theta_B \mid D)$ (one minus the number here, with the labels swapped). With SVI, draw from the fitted guide (for example with NumPyro's Predictive(model, guide=guide, params=params, num_samples=4000)), and compute the share of draws with B ahead, using the same draw index for every variant. Report the Monte Carlo error, and remember that with an approximate (SVI) posterior the number is only as good as the guide (Chapters 6.13 and 6.15).
$P(\theta_B \gt \theta_A \mid D) \approx \frac1S\sum_s \mathbf{1}[\theta_B^{(s)} \gt \theta_A^{(s)}]$ = share of the cloud above the diagonal.
MC error $\sqrt{P(1-P)/S} \le 0.5/\sqrt S$. Checkout: 0.843 (two-sided p was 0.31).
Trap: direction, not size; not 1 − p; pair draws by index when the posterior is correlated.
Quick check: how many independent draws do you need so that the Monte Carlo error of a probability near 0.9 is below 0.005?
$\sqrt{0.9 \times 0.1 / S} \lt 0.005 \iff S \gt 0.09/0.000025 = 3\,600$. About 4 000 draws are enough; MCMC draws need an effective sample size of 3 600.
$P(\theta_B - \theta_A \gt \delta \mid D)$: better by enough to matter core
"B is better" is not what a product team really needs to know. Shipping a change costs engineering time, adds risk, and has to be maintained. If B is better by 0.01 percentage points, nobody cares. The real question is: is B better by at least an amount δ that is worth it? Here δ (delta) is a practically meaningful improvement, chosen from business sense before the test, for example "at least half a point of conversion".
Why this matters: with enough traffic, any tiny real difference makes $P(\theta_B \gt \theta_A \mid D)$ approach 1. A test with 200 000 users per variant can be 98% sure that B is better and, at the same time, 99.9% sure that it is better by less than half a point. $P(\theta_B \gt \theta_A)$ says "ship"; $P(\theta_B - \theta_A \gt \delta)$ says "not worth it". On the cloud picture, you simply move the diagonal up by δ and count the dots above the new line.
Three ways to say it:
- Picture: shift the diagonal up by δ; only dots above the shifted line count as "worth shipping".
- Numbers: huge test, 10.0% vs 10.2%: $P(B \gt A) = 0.98$ but $P(B - A \gt 0.5\text{ pt}) = 0.001$.
- Slogan: not "is B better?" but "is B better by enough to matter?"
Choose δ = 0.5 point (an absolute gain of 0.005; that is 5% relative at a 10% baseline). Flat priors throughout.
- Checkout test (50/500 vs 60/500): $P(\theta_B \gt \theta_A) = 0.843$, $P(\theta_B - \theta_A \gt 0.005) = 0.774$, $P(\gt 0.01) = 0.692$, $P(\gt 0.02) = 0.498$. Promising, but far from certain.
- Huge test, tiny gain (20 000/200 000 vs 20 400/200 000, i.e. 10.0% vs 10.2%): $P(\theta_B \gt \theta_A) = 0.982$, but $P(\theta_B - \theta_A \gt 0.005) = 0.0008$.
- In the same huge test, $P(|\theta_B - \theta_A| \lt 0.005) = 0.999$: we are almost sure the two are practically equivalent. The 95% interval of the gap is $[+0.01, +0.39]$ points: real but small.
- Small test, big gain (50/500 vs 75/500): $P(\theta_B \gt \theta_A) = 0.992$ and $P(\theta_B - \theta_A \gt 0.005) = 0.984$. Here both say the same thing.
- Lesson: $P(\theta_B \gt \theta_A)$ mixes "how big is the gain" with "how much data do we have". $P(\theta_B - \theta_A \gt \delta)$ asks the business question directly.
For a practically meaningful threshold $\delta \ge 0$ fixed in advance:
$$P(\theta_B - \theta_A \gt \delta \mid D) \;\approx\; \frac1S\sum_{s=1}^{S}\mathbf{1}\big[\theta_B^{(s)} - \theta_A^{(s)} \gt \delta\big].$$- With $\delta = 0$ it is $P(\theta_B \gt \theta_A \mid D)$. It can only decrease as δ grows.
- The interval $[-\delta, +\delta]$ is often called the region of practical equivalence (ROPE): differences inside it are too small to matter. $P(|\theta_B - \theta_A| \lt \delta \mid D)$ is the probability that the variants are practically the same.
- Three zones for the gap: worse ($\lt 0$, or $\lt -\delta$ if you want a symmetric ROPE), negligible, meaningfully better ($\gt \delta$). Their posterior probabilities add to 1.
- δ may be absolute (points) or relative (a lift, next section). It must be in the same units as the parameter in your model.
Why do we need it?
To stop shipping "statistically real but practically useless" changes from big tests, and to be able to conclude "no meaningful difference" (stop the test) instead of waiting forever for a significance that does not matter.
Where is it used?
Bayesian A/B decision rules with a minimum effect of interest, ROPE-based analyses, non-inferiority checks for guardrail metrics ("B is not worse than A by more than δ"), and clinical trials with a minimal clinically important difference.
How is it used?
Agree on δ before the test, from costs and benefits or the minimum detectable effect you planned for (Chapter 5.10). After the test compute ((tB - tA) > delta).mean() and (abs(tB - tA) < delta).mean() from paired draws and feed them into a decision rule (Section 7).
"Choose δ after looking at the results."
Then δ can be tuned to make any result look good or bad. Fix δ (and the decision threshold) before the test, from costs and benefits, and write it down.
"P(θ_B > θ_A) = 0.98, so ship."
With huge traffic a negligible gain gives 0.98. Check $P(\theta_B - \theta_A \gt \delta)$ and the size of the gap before deciding.
"δ = 0.5 points works for every metric in my model."
δ lives in the units of the parameter. If a continuous metric is standardized with a global scaler (one mean and one sd for all groups), a raw δ of 2 euros becomes $\delta / \text{sd}_{\text{global}}$ in the model's units. (A per-group scaler would make the groups' units different, and the comparison meaningless.)
The syllabus asks for exactly this in your A/B framework: $P(\theta_A - \theta_B \gt \delta \mid D)$ with δ a practically meaningful improvement. Compute it from the same paired posterior draws as $P(\theta_A \gt \theta_B)$; report both, plus the probability of practical equivalence $P(|\theta_A - \theta_B| \lt \delta)$. For standardized metrics, convert δ with the global scaler before comparing.
"Our Bayesian test says B is 98% likely to be better, so the result is important."
98% is about the direction of the effect. Importance is about its size relative to what matters. A huge sample makes a tiny effect almost certainly positive.
Model answer: "I report $P(\theta_B - \theta_A \gt \delta \mid D)$ with δ set in advance as the smallest gain worth shipping, plus the posterior of the gap with a credible interval. That separates 'is it real?' from 'is it worth it?', which $P(\theta_B \gt \theta_A)$ alone mixes up."
$P(\theta_B - \theta_A \gt \delta \mid D)$ = share of paired draws with gap $\gt \delta$; δ = smallest gain worth shipping, fixed in advance.
ROPE $[-\delta, \delta]$: $P(|\theta_B - \theta_A| \lt \delta)$ = probability of practical equivalence.
Huge test 10.0% vs 10.2%: $P(B \gt A) = 0.98$ but $P(B - A \gt 0.5\text{ pt}) = 0.001$.
Quick check: from 4 000 draws, 3 100 have θ_B − θ_A > 0, 2 300 have θ_B − θ_A > 0.01, and 150 have θ_B − θ_A < −0.01. With δ = 1 point, give the three-zone probabilities.
Better by at least δ: $2\,300/4\,000 = 0.575$. Practically equivalent ($-0.01 \lt$ gap $\lt 0.01$): $(4\,000 - 2\,300 - 150)/4\,000 = 1\,550/4\,000 = 0.388$. Worse by more than δ: $150/4\,000 = 0.0375$. (Plain $P(\theta_B \gt \theta_A) = 0.775$.)
The posterior of the relative lift
Stakeholders rarely say "plus 2 points". They say "B lifts conversion by 20%". That relative lift is $\theta_B/\theta_A - 1$: the gain as a fraction of A's rate. It is just another function of the two rates, so its posterior comes from the same draws: for each draw, divide.
Dividing by an uncertain number has a side effect: the posterior of a ratio is skewed to the right. When a draw of θ_A happens to be small, the ratio shoots up; when it is large, the ratio cannot fall nearly as far. So the mean lift sits above the median lift, and the interval is lopsided. Report the median and an equal-tailed interval, and never compute "the lift" only from the two posterior means.
Three ways to say it:
- Picture: a histogram of lifts with a long right tail.
- Numbers: checkout test: median lift 19.8%, mean 21.8%, 95% interval −15.7% to +70.5%.
- Slogan: compute the lift per draw; summarize with the median and quantiles.
Checkout test, flat priors, 200 000 paired draws.
- For each draw: $\text{lift}^{(s)} = \theta_B^{(s)}/\theta_A^{(s)} - 1$.
- Median $0.198$ (≈ 20%); compare the ratio of the posterior means, $\frac{61/502}{51/502} - 1 = \frac{61}{51} - 1 = 0.196$.
- Mean $0.218$: the long right tail pulls it up by 2 points of lift.
- 95% equal-tailed interval $[-0.157,\ +0.705]$: B could be 16% worse or 70% better. Lopsided around the median, as ratios are.
- $P(\text{lift} \gt 0) = 0.843$, the same event as $\theta_B \gt \theta_A$. $P(\text{lift} \gt 5\%) = 0.769$; $P(\text{lift} \gt 10\%) = 0.683$.
- Same absolute gap, different baseline: +2 points on a 40% baseline is only a 5% lift; +2 points on a 2% baseline is a 100% lift. Relative numbers depend heavily on the baseline.
The relative lift of B over A is $L = \frac{\theta_B}{\theta_A} - 1 = \frac{\theta_B - \theta_A}{\theta_A}$. Its posterior is obtained by transforming draws: $L^{(s)} = \theta_B^{(s)}/\theta_A^{(s)} - 1$.
- $P(L \gt 0 \mid D) = P(\theta_B \gt \theta_A \mid D)$, because $\theta_A \gt 0$. A relative threshold $r$ gives $P(L \gt r \mid D) = P(\theta_B \gt (1 + r)\theta_A \mid D)$.
- The posterior of $L$ is not a Beta and is right-skewed; its mean is not the ratio of means, $E[\theta_B/\theta_A] \ne E[\theta_B]/E[\theta_A]$ (Jensen's inequality, Chapter 4.5).
- Quantiles transform cleanly: the median and ETI of $L$ are what you get by applying the same per-draw formula, so they are the natural summaries. The log-ratio $\log(\theta_B/\theta_A)$ is often more symmetric.
Why do we need it?
Business goals and forecasts of impact are usually stated in relative terms ("+5% conversion"). A credible interval for the lift says how sure you are about that number, which the absolute gap alone does not show.
Where is it used?
A/B testing dashboards ("expected lift with 95% interval"), revenue impact projections (lift × baseline revenue), meta-analyses of many experiments, and relative thresholds δ in decision rules ("ship if B is at least 3% better").
How is it used?
From paired draws: lift = tB / tA - 1; report np.median(lift) and np.quantile(lift, [0.025, 0.975]); compute (lift > r).mean() for a relative threshold $r$. Always state the baseline rate next to a relative lift.
"The lift is (mean of B)/(mean of A) − 1, and its interval is the ratio of the two intervals' ends."
Dividing interval ends ignores how the two rates combine. Compute the lift for every paired draw and take quantiles of those lifts.
"Average the lifts of my segments to get the overall lift."
Relative lifts with different baselines and traffic do not average simply. Compute the overall rates per draw (traffic-weighted), then the overall lift per draw.
"A 100% lift is a big win."
On a 0.2% baseline it is +0.2 points. Always give the baseline and the absolute gap next to a relative lift.
When your framework reports a lift for a conversion metric, compute it per posterior draw from the variant rates (same draw index for both), and report the median with an equal-tailed interval. For metrics standardized with the global scaler, transform the draws back to raw units first: a relative lift of standardized values (which can be near 0 or negative) is meaningless.
Lift $L^{(s)} = \theta_B^{(s)}/\theta_A^{(s)} - 1$ per draw; report median + ETI; $P(L \gt 0) = P(\theta_B \gt \theta_A)$.
Ratios are right-skewed: mean > median. Checkout: median 19.8%, 95% [−15.7%, +70.5%].
Trap: never divide interval ends or average segment lifts; always state the baseline.
Quick check: why is $P(\text{lift} \gt 0 \mid D)$ exactly equal to $P(\theta_B \gt \theta_A \mid D)$, but $P(\text{lift} \gt 5\%)$ not equal to $P(\theta_B - \theta_A \gt 0.005)$?
Lift $\gt 0$ ⇔ $\theta_B/\theta_A \gt 1$ ⇔ $\theta_B \gt \theta_A$ (since $\theta_A \gt 0$): the same event. Lift $\gt 5\%$ ⇔ $\theta_B - \theta_A \gt 0.05\,\theta_A$, a threshold that moves with θ_A. Only when θ_A is exactly 10% does it equal 0.005. In the checkout test the two are close (0.769 vs 0.774) because θ_A is near 10%.
From probabilities to decisions: rules and expected loss core
A probability does not make a decision. "P = 0.77" still leaves the question: ship, stop, or keep testing? A decision rule turns posterior quantities into actions, and it should be written down before the test, like the rules of a game. Three families are common:
- Threshold on P(B better): ship B if $P(\theta_B \gt \theta_A \mid D) \ge 0.95$. Simple, but blind to the size of the gain.
- Threshold with a practical δ (three outcomes): ship B if $P(\theta_B - \theta_A \gt \delta) \ge 0.95$; stop and keep A if $P(\theta_B - \theta_A \lt \delta) \ge 0.95$ ("not worth it"); otherwise keep collecting data (up to a planned maximum).
- Expected loss: "if I ship B and I am wrong, how much conversion do I lose, on average over my uncertainty?" Ship B when that average loss is below a small "threshold of caring" ε.
Expected loss weighs how likely a mistake is by how costly it would be. Being wrong by 0.01 points barely matters; being wrong by 3 points matters a lot. A rule based on it stops early when the stakes are small, and waits when a costly mistake is still plausible.
Three ways to say it:
- Picture: the part of the gap's posterior where B is worse, weighted by how much worse, is the risk of choosing B.
- Numbers: checkout test: choosing B risks 0.16 points on average; choosing A risks 2.16 points.
- Slogan: decide with a rule fixed in advance; weigh mistakes by their cost.
Rules fixed in advance: δ = 0.5 point, probability threshold $c = 0.95$, threshold of caring ε = 0.05 point.
- Checkout test (50/500 vs 60/500). $P(\theta_B \gt \theta_A) = 0.843 \lt 0.95$: rule 1 does not ship.
- $P(\theta_B - \theta_A \gt 0.005) = 0.774$ and $P(\theta_B - \theta_A \lt 0.005) = 0.226$; neither reaches 0.95: rule 2 says keep collecting.
- Expected loss of shipping B: $E[\max(\theta_A - \theta_B, 0)\mid D] = 0.0016$ (0.16 points); of keeping A: $E[\max(\theta_B - \theta_A, 0)\mid D] = 0.0216$ (2.16 points). B is the better bet, but its risk 0.16 > ε = 0.05: rule 3 also says keep collecting.
- Check: $L(A) - L(B) = 0.0216 - 0.0016 = 0.0200 = E[\theta_B - \theta_A \mid D]$, as the definition below guarantees.
- Huge test (10.0% vs 10.2%, 200 000 each). Rule 1: $0.982 \ge 0.95$, ship B. Rule 2: $P(\theta_B - \theta_A \lt 0.005) = 0.999 \ge 0.95$: stop, not worth shipping. Rule 3: risk of B is 0.0006 points, below ε: choosing B is safe. Rules 1 and 3 ask "is B a safe choice?"; rule 2 asks "is B worth the change?". Use the one that matches your real costs.
Let $g = \theta_B - \theta_A$. The posterior expected loss of each action is
$$L(\text{choose B}) = E\big[\max(\theta_A - \theta_B,\ 0) \mid D\big] = \int_{g \lt 0} |g|\,p(g\mid D)\,dg, \qquad L(\text{choose A}) = E\big[\max(\theta_B - \theta_A,\ 0)\mid D\big],$$and $L(A) - L(B) = E[\theta_B - \theta_A \mid D]$. Bayesian decision theory chooses the action with the smaller expected loss; a stopping rule ends the test when that smaller loss drops below ε.
- From draws:
np.maximum(tA - tB, 0).mean(). It is in the units of θ (points of conversion), so it can be turned into money. - A decision rule is a function (posterior → action) fixed in advance, together with δ, $c$, ε and a maximum sample size.
- The posterior is a valid summary of belief at any moment. But the frequency with which a rule makes mistakes depends on how it is used: checking every day and stopping at the first crossing ("peeking", Chapter 5.11) raises the rate of shipping variants that are not really better, and many metrics or segments multiply the chances of a lucky "winner". Simulate the rule as you will really run it.
Why do we need it?
Without a rule fixed in advance, decisions drift toward whatever the team hoped for. A rule makes decisions consistent, and expected loss ties them to what a wrong call would actually cost.
Where is it used?
Bayesian experimentation platforms (expected-loss stopping, "chance to beat control" thresholds), bandits that stop exploring when the expected loss of the current best arm is small, ROPE-based decisions in research, and launch reviews with guardrail metrics (non-inferiority: "not worse by more than δ").
How is it used?
Before the test: choose δ, $c$, ε and a maximum sample size. After (or at planned looks): compute the probabilities and expected losses from paired draws, apply the rule, and record the decision. Simulate the rule on fake data (next widget) to know how often it ships a variant that is not really better.
"A Bayesian test cannot be hurt by peeking, because the posterior is always valid."
The posterior is a correct summary of belief at every moment (given the model). But a rule like "stop the first day P(θ_B > θ_A) passes 0.95" ships more no-better variants than the same rule applied once at a planned end. If you look often, choose stricter thresholds or an expected-loss rule, and check the rule's error rates by simulation.
"Expected loss tells me whether shipping is worth it."
Only if the loss includes all the costs. The plain $E[\max(\theta_A - \theta_B, 0)]$ ignores the cost of the change itself; add it (or use δ) when switching is not free.
"Use the same 0.95 threshold for 20 metrics and 10 segments."
With many comparisons some will cross any threshold by luck. Pick one primary metric, treat the others as guardrails, and use hierarchical models for segments, which shrink extreme segment results (Chapter 6.6).
In an A/B framework like yours, write the decision rule into the analysis config before launch: the primary metric, δ in that metric's units (converted with the global scaler where needed), the probability threshold, a maximum sample size, and the guardrail checks ("B not worse than A by more than δ_g"). The framework can then report the same three or four numbers every time: $P(\theta_B \gt \theta_A)$, $P(\theta_B - \theta_A \gt \delta)$, the probability of practical equivalence, and the expected loss, all from the same paired posterior draws.
Expected loss of choosing B: $E[\max(\theta_A - \theta_B, 0)\mid D]$; $L(A) - L(B) = E[\theta_B - \theta_A\mid D]$.
Three-outcome rule: ship if $P(\text{gap} \gt \delta) \ge c$; stop if $P(\text{gap} \lt \delta) \ge c$; else continue (to a max $n$).
Trap: fix the rule in advance; peeking and many metrics raise the rate of bad ships; plain expected loss ignores switching cost.
Quick check: from 4 000 paired draws, the average of max(θ_A − θ_B, 0) is 0.0004 and the average of max(θ_B − θ_A, 0) is 0.0090. What is the posterior mean gap, and which variant does expected loss favour?
$E[\theta_B - \theta_A] = L(A) - L(B) = 0.0090 - 0.0004 = 0.0086$ (0.86 points). Choosing B risks only 0.04 points on average, choosing A risks 0.90, so expected loss favours B; with ε = 0.05 points, B can be chosen now.
Credible interval vs confidence interval: the precise difference core
The two intervals answer two different questions, and they treat different things as random.
- A confidence interval (Chapter 5.8) answers: "If the true rate is some fixed value and I repeated the whole experiment many times, how often would my interval-making recipe catch it?" The rate is fixed; the data, and so the interval, are random. 95% is a property of the recipe, not of today's interval.
- A credible interval answers: "Given the data I actually have, and my model and prior, where is the rate probably?" The data are fixed; the rate is uncertain. 95% is a probability about θ for this dataset.
With a flat prior and plenty of data the two often give almost the same numbers, which is why people mix them up. They part ways with small samples, rates near 0 or 1, and strong priors. And neither is "the true one": each keeps a different promise.
Three ways to say it:
- Picture: confidence = many intervals around one fixed θ (some miss); credible = one dataset, one posterior, 95% of its area shaded.
- Numbers: 60 of 500: Wilson CI [9.44%, 15.14%], flat-prior credible interval [9.44%, 15.15%]: same numbers, different meanings.
- Slogan: confidence is about the recipe over repeated data; credibility is about θ given this data.
- Lots of data, flat prior. 60 of 500: Wald CI $[0.0915, 0.1485]$, Wilson CI $[0.0944, 0.1514]$, flat-prior 95% credible interval $[0.0944, 0.1515]$. Practically identical.
- Small data near 0. 2 of 41: the Wald CI is $0.0488 \pm 1.96 \times 0.0336 = [-0.017, 0.115]$, with an impossible negative end (statsmodels clips it to 0). Wilson: $[0.0135, 0.1614]$. Flat-prior credible: $[0.0150, 0.1616]$. Jeffreys prior Beta(½, ½), a standard default prior for a rate: $[0.0103, 0.1474]$. Now the prior choice visibly matters.
- Strong prior. Prior Beta(20, 180) ("10%, worth 200 visitors") and 10 of 50: posterior Beta(30, 220), credible interval $[0.083, 0.163]$; the data alone (Wilson) say $[0.112, 0.330]$. The credible interval is a correct summary of that prior combined with that data.
- Coverage at a fixed θ (n = 50, exact calculation over all possible datasets): if the true rate is 0.20, the flat-prior credible interval contains it in 95.1% of datasets; the Beta(20, 180)-prior interval in only 0.25%. That prior is confidently wrong about this θ.
- Coverage averaged over the prior. If instead θ itself is drawn from Beta(20, 180) for each experiment, the Beta(20, 180) credible intervals contain it in exactly 95.0% of experiments. Credible intervals are calibrated on average over the prior, not at every fixed θ.
| Confidence interval (frequentist) | Credible interval (Bayesian) | |
|---|---|---|
| What is random? | the data, hence the interval; θ is fixed | θ (as a state of belief); the data are fixed |
| The 95% promise | $P_D\big(L(D) \le \theta \le U(D) \mid \theta\big) \approx 0.95$ for every θ | $P(L \le \theta \le U \mid D) = 0.95$ for this $D$ |
| Correct sentence | "This interval came from a procedure that captures the true value in 95% of repeated samples." | "Given the model, prior and data, θ is in this interval with probability 95%." |
| Needs a prior? | no | yes, and the result depends on it |
| Can it include impossible values? | some recipes can (Wald below 0) | no: the posterior lives on the parameter's support |
- When they agree: for regular models with lots of data and a prior that does not rule out the truth, the posterior is close to Normal(MLE, SE²) (the Bernstein–von Mises theorem), so a 95% credible interval is close to a 95% Wald interval.
- Calibration on average ("calibrated" means the stated probability matches how often the statement turns out true): if θ is truly drawn from the prior and the model is right, $P(\theta \in \text{CrI}(D)) = E_D\big[P(\theta \in \text{CrI}(D) \mid D)\big] = 0.95$ exactly, averaged over θ and D. At a single fixed θ there is no such guarantee.
- This guide owns the Bayesian side; the confidence interval, its coverage simulation and the Wald-vs-Wilson comparison are in Chapter 5.8.
Why do we need it?
Interviewers and stakeholders probe this distinction, and getting it wrong leads to wrong claims ("95% chance the true lift is in the CI"). Knowing when the two agree also tells you when a Bayesian and a frequentist analysis of the same test should give similar numbers.
Where is it used?
Every report of an interval: frequentist A/B tools report confidence intervals, Bayesian ones credible intervals; regulated settings (clinical trials) often require frequentist coverage even for Bayesian designs; calibration checks of Bayesian models compare credible-interval coverage on simulated data.
How is it used?
Say the right sentence for the interval you have. If a Bayesian interval must also behave well in repeated use, check its coverage by simulation at a few plausible θ values (simulation-based calibration checks the average-over-the-prior version of this systematically). With small data, report the prior and try an alternative (Chapter 6.8).
"A 95% confidence interval means there is a 95% probability that θ is inside it."
That sentence describes a credible interval. For a confidence interval, the 95% belongs to the procedure over repeated samples; once computed, the interval either contains the fixed θ or not.
"A 95% credible interval contains the true θ 95% of the time."
Only on average over θ values drawn from the prior (and if the model is right). For a particular θ that the prior finds implausible, coverage can be terrible, as the widget shows.
"With a flat prior, the credible interval is the confidence interval."
The numbers can be close (60 of 500), but the meanings never merge; and with small data the numbers differ too (2 of 41). "Flat" is also scale-dependent (Chapter 6.2).
Intervals taken from your A/B framework's posterior draws are credible intervals, so describe them as probability statements about the metric given the model and prior. If a stakeholder asks "how often would this be wrong?", that is a frequentist question: answer it by simulating the framework on synthetic experiments with known effects (as in the widgets above), which also tests the SVI approximation and the decision rule together.
"The confidence interval and the credible interval are basically the same; Bayesians just use different words."
They answer different questions. Confidence: θ fixed, data random, 95% is the long-run hit rate of the procedure. Credible: data fixed, θ uncertain, 95% is the posterior probability that θ is in this interval, given the prior and the model.
Model answer: "A 95% confidence interval comes from a procedure that would capture the fixed true value in 95% of repeated experiments; it does not give a probability for this particular interval. A 95% credible interval is a direct probability statement: given my model, prior and the observed data, θ lies in it with probability 0.95. With weak priors and lots of data the two are numerically similar, by the Bernstein–von Mises theorem, but with small samples or strong priors they can differ, and a credible interval is only calibrated on average over the prior."
CI: θ fixed, data random; $P_D(\theta \in [L(D), U(D)] \mid \theta) \approx 0.95$, a property of the recipe.
CrI: data fixed, θ uncertain; $P(\theta \in [L, U] \mid D) = 0.95$; depends on the prior; calibrated on average over the prior.
60/500: Wilson [0.0944, 0.1514] ≈ flat CrI [0.0944, 0.1515]; 2/41: Wald goes negative, CrI cannot.
Quick check: a colleague says "our Bayesian 95% interval for the lift is [0.5%, 4%], so if we reran the test 100 times, about 95 intervals would contain the true lift." What is right and wrong?
The correct reading is "given our model, prior and data, the lift is in [0.5%, 4%] with probability 0.95." The repeated-sampling claim is a confidence-interval property; a credible interval has it only on average over the prior (and approximately with weak priors and large samples). To know the repeated-sampling behaviour at plausible true lifts, simulate it.
Recap, cheat sheet and practice
- One-number summaries answer different losses: mean (squared error), median (absolute error, same on every scale), mode/MAP (all-or-nothing, scale-dependent). For skewed posteriors: mode < median < mean.
- A credible interval holds a chosen share of the posterior: $P(L \le \theta \le U \mid D) = 0.95$, a direct probability statement given model, prior and data.
- ETI = equal tails, transforms with the parameter; HDI = highest density, shortest for unimodal posteriors, can split into pieces, depends on the scale. NumPyro's summary columns are HPDI ends.
- $P(\theta_B \gt \theta_A \mid D)$ = share of paired posterior draws above the diagonal; Monte Carlo error $\sqrt{P(1-P)/S}$. It measures direction, not size, and is not 1 − p.
- $P(\theta_B - \theta_A \gt \delta \mid D)$ with a practical δ fixed in advance asks "better by enough to matter?". With big samples it can be near 0 while $P(\theta_B \gt \theta_A)$ is near 1. The ROPE probability $P(|\theta_B - \theta_A| \lt \delta)$ supports "no meaningful difference".
- The relative lift is computed per draw, is right-skewed, and is reported with its median and ETI, always next to the baseline.
- Decision rules are fixed in advance (δ, threshold, ε, max n). Expected loss $E[\max(\theta_A - \theta_B, 0)]$ weighs mistakes by their size. Peeking and many comparisons change how often a rule makes bad calls.
- Credible vs confidence: θ uncertain given fixed data vs θ fixed with data random. Similar numbers with weak priors and big n; different meanings always; credible intervals are calibrated on average over the prior.
Cheat sheet
| Idea | Formula / recipe | In words |
|---|---|---|
| Mean / median / mode | squared / absolute / 0–1 loss | balance point / half area / peak |
| Credible interval | $P(L \le \theta \le U\mid D) = 1 - \gamma$ | a probability about θ given D |
| ETI | $[F^{-1}(\gamma/2), F^{-1}(1-\gamma/2)]$ | equal tails; scale-free |
| HDI | $\{\theta: p(\theta\mid D) \ge c\}$ with mass $1-\gamma$ | water level; shortest; scale-dependent |
| P(B better) | $\frac1S\sum \mathbf{1}[\theta_B^{(s)} \gt \theta_A^{(s)}]$, error $\sqrt{P(1-P)/S}$ | share of the cloud above the diagonal |
| Better by δ | $\frac1S\sum \mathbf{1}[\theta_B^{(s)} - \theta_A^{(s)} \gt \delta]$ | shift the diagonal up by δ |
| ROPE | $P(|\theta_B - \theta_A| \lt \delta\mid D)$ | probability of "practically the same" |
| Relative lift | $L^{(s)} = \theta_B^{(s)}/\theta_A^{(s)} - 1$ | per draw; median + ETI |
| Expected loss of B | $E[\max(\theta_A - \theta_B, 0)\mid D]$ | average loss if choosing B was wrong |
| CI vs CrI | $P_D(\theta \in CI \mid \theta)$ vs $P(\theta \in CrI \mid D)$ | recipe over repeats vs belief given data |
import numpy as np
from scipy import stats, integrate, optimize
from statsmodels.stats.proportion import proportion_confint
rng = np.random.default_rng(0)
# --- one-number summaries of a skewed posterior: 2 of 41, flat prior -> Beta(3, 40) ---
post = stats.beta(3, 40)
print(post.mean(), post.median(), (3 - 1) / (3 + 40 - 2)) # 0.0698 0.0632 0.0488 (mean, median, mode)
# --- 95% equal-tailed and highest-density intervals ---
print(post.ppf([0.025, 0.975])) # ETI [0.0150 0.1616]
def hdi_closed_form(d, mass=0.95):
"""Shortest interval [lo, F^-1(F(lo) + mass)] for a unimodal distribution."""
width = lambda lo: d.ppf(d.cdf(lo) + mass) - lo
lo = optimize.minimize_scalar(width, bounds=(0, d.ppf(1 - mass)), method="bounded").x
return lo, d.ppf(d.cdf(lo) + mass)
def hdi_draws(x, mass=0.95):
"""Shortest window that contains `mass` of the sorted draws."""
s = np.sort(x)
k = int(np.floor(mass * len(s)))
i = np.argmin(s[k:] - s[:len(s) - k])
return s[i], s[i + k]
print(hdi_closed_form(post)) # (0.0080, 0.1452)
x = post.rvs(size=200_000, random_state=rng)
print(np.quantile(x, [0.025, 0.975]), hdi_draws(x)) # about [0.0150 0.1616] and (0.0080, 0.1452)
# --- credible interval and tail probabilities for Beta(8, 52) ---
d = stats.beta(8, 52)
print(d.ppf([0.025, 0.975]), d.sf(0.10)) # [0.0604 0.2293] 0.766
# --- checkout test: A 50/500, B 60/500, flat priors ---
A, B = stats.beta(51, 451), stats.beta(61, 441)
S = 20_000
tA = A.rvs(size=S, random_state=rng)
tB = B.rvs(size=S, random_state=rng) # independent posteriors -> independent draws
p = (tB > tA).mean()
print(p, np.sqrt(p * (1 - p) / S)) # about 0.84, Monte Carlo error about 0.003
exact = integrate.quad(lambda t: A.pdf(t) * B.sf(t), 0, 1, points=[0.1, 0.12])[0]
print(exact) # 0.8429
for delta in [0.005, 0.01, 0.02]: # better by at least delta?
print(delta, ((tB - tA) > delta).mean()) # about 0.77, 0.69, 0.50
print((np.abs(tB - tA) < 0.005).mean()) # about 0.12 practically equivalent (ROPE)
lift = tB / tA - 1 # relative lift, per draw
print(np.median(lift), lift.mean(), np.quantile(lift, [0.025, 0.975])) # about 0.20 0.22 [-0.16 0.71]
print((lift > 0.05).mean()) # about 0.77
print(np.maximum(tA - tB, 0).mean(), np.maximum(tB - tA, 0).mean()) # about 0.0017 0.0215 expected loss of B, of A (exact 0.0016, 0.0216)
# --- huge test, tiny gain: 10.0% vs 10.2% with 200 000 users each ---
hA, hB = stats.beta(20_001, 180_001), stats.beta(20_401, 179_601)
m = hA.mean()
for delta in [0.0, 0.005]:
print(delta, integrate.quad(lambda t: hA.pdf(t) * hB.sf(t + delta), m - 0.01, m + 0.01)[0]) # 0.982, then 0.0008
# --- credible vs confidence intervals ---
print(proportion_confint(60, 500, method="wilson"), stats.beta(61, 441).ppf([0.025, 0.975])) # nearly equal
print(proportion_confint(2, 41, method="wilson"), stats.beta(3, 40).ppf([0.025, 0.975])) # (0.0135, 0.1614) [0.0150 0.1616]
def coverage(theta, n, a, b):
"""Exact frequentist coverage of the 95% equal-tailed credible interval at a fixed theta."""
k = np.arange(n + 1)
lo = stats.beta(a + k, b + n - k).ppf(0.025)
hi = stats.beta(a + k, b + n - k).ppf(0.975)
return (stats.binom(n, theta).pmf(k) * ((lo <= theta) & (theta <= hi))).sum()
print(coverage(0.2, 50, 1, 1), coverage(0.2, 50, 20, 180)) # 0.951 0.0025
# average coverage when theta is drawn from the Beta(20, 180) prior: 0.95 (simulation)
th = rng.beta(20, 180, size=200_000)
k = rng.binomial(50, th)
post_k = stats.beta(20 + k, 180 + 50 - k)
inside = (post_k.ppf(0.025) <= th) & (th <= post_k.ppf(0.975))
print(inside.mean()) # about 0.95
1. A posterior for a conversion rate has a long right tail. How are its summaries ordered?
2. You report a 95% interval for θ, and your colleague converts it to the log-odds scale by transforming its two ends. For which interval is that conversion exactly right?
3. In 4 000 paired posterior draws, θ_B > θ_A in 3 200. What are the estimate and its Monte Carlo error?
4. A huge test gives $P(\theta_B \gt \theta_A \mid D) = 0.98$ and $P(\theta_B - \theta_A \gt \delta \mid D) = 0.01$ for a δ agreed in advance. The best summary is…
5. Why is the posterior median usually reported for a relative lift θ_B/θ_A − 1?
6. Which statement about a 95% credible interval [L, U] is correct?
Practice problems
A. Posterior Beta(4, 6). Give the mean, mode and (approximate) median. Which one would you report under absolute-error loss?
Mean $4/10 = 0.400$. Mode $\frac{4-1}{10-2} = 3/8 = 0.375$. Median $\approx \frac{4 - 1/3}{10 - 2/3} = \frac{3.667}{9.333} = 0.393$ (SciPy: 0.3931). Under absolute-error loss, report the median, 0.393. Order mode < median < mean: a slight right skew.
B. For B's posterior in the checkout test, Beta(61, 441), compute a 90% equal-tailed credible interval and say it in one sentence.
stats.beta(61, 441).ppf([0.05, 0.95]) $= [0.0984, 0.1463]$. "Given a flat prior, the Binomial model and 60 conversions out of 500, B's conversion rate is between 9.8% and 14.6% with probability 90%."
C. A: 200 of 2 000, B: 240 of 2 000, flat priors. Approximate $P(\theta_B \gt \theta_A \mid D)$ with a Normal approximation of the gap.
Posteriors Beta(201, 1 801) and Beta(241, 1 761): means 0.1004 and 0.1204, sds 0.00672 and 0.00727. Gap mean $0.0200$, sd $\sqrt{0.00672^2 + 0.00727^2} = 0.00990$. $P \approx \Phi(0.0200/0.00990) = \Phi(2.02) = 0.978$. The exact value is also 0.978: with 2 000 users per variant the Normal approximation is excellent.
D. Choosing δ. A new checkout costs 60 000 euros to build and run for a year. About 100 000 visitors per month will see it, and each extra conversion is worth 10 euros. What absolute gain δ makes it pay for itself within a year?
Visitors in a year: $12 \times 100\,000 = 1\,200\,000$. A gain δ (as a rate) gives $1\,200\,000\,\delta$ extra conversions, worth $12\,000\,000\,\delta$ euros. Break-even: $12\,000\,000\,\delta = 60\,000 \Rightarrow \delta = 0.005$, i.e. 0.5 percentage point. Fix this δ before the test and decide with $P(\theta_B - \theta_A \gt 0.005 \mid D)$.
E. Five paired posterior draws: θ_A = 0.10, 0.12, 0.09, 0.11, 0.10 and θ_B = 0.13, 0.11, 0.12, 0.14, 0.09. Estimate P(θ_B > θ_A), the expected loss of choosing B and of choosing A, and check the identity that links them.
Gaps θ_B − θ_A: +0.03, −0.01, +0.03, +0.03, −0.01. $P(\theta_B \gt \theta_A) \approx 3/5 = 0.6$. Loss of B: mean of $\max(\theta_A - \theta_B, 0) = (0 + 0.01 + 0 + 0 + 0.01)/5 = 0.004$. Loss of A: $(0.03 + 0 + 0.03 + 0.03 + 0)/5 = 0.018$. Identity: $L(A) - L(B) = 0.014$ = mean gap $(0.03 - 0.01 + 0.03 + 0.03 - 0.01)/5 = 0.014$. ✓ (Five draws are far too few in practice; this only shows the arithmetic.)
F. Interview: "Your Bayesian A/B test reports a 95% credible interval of [9.4%, 15.1%] for B, and a frequentist colleague's Wilson interval is [9.4%, 15.1%]. Are they the same thing?"
"The numbers agree because the prior is flat and 500 users is enough for the likelihood to dominate (Bernstein–von Mises). The meanings differ. Mine says: given the model, flat prior and data, B's rate is in [9.4%, 15.1%] with probability 0.95. Hers says: the Wilson procedure captures the fixed true rate in about 95% of repeated experiments; it makes no probability statement about this particular interval. With small samples or an informative prior the numbers would also differ, for example 2 of 41: Wald goes below zero, the credible interval cannot."
Hierarchical models: groups that are related, not identical
A framework like yours does not stop at one overall number. It also looks at segments: per country, per device, per kind of user. Some segments have thousands of users, some have twenty. A hierarchical model is the honest way to describe such data: every segment gets its own value, and the segments' values are themselves drawn from one shared population. This chapter builds that two-level story, names every part of it, shows what the data can and cannot tell you about how different the segments really are, and tells you when the story is the wrong one.
- Write and read the two-level model $\theta_g \sim N(\mu, \tau^2)$, $y_{gi} \sim p(y \mid \theta_g)$, and simulate data from it level by level
- Name every unknown: group parameters, global parameters, hyperparameters, hyperpriors, and the fixed constants you choose; draw the model as a plate diagram
- Explain how the data teach us the between-group spread $\tau$, why that needs several groups, and how to choose and check a hyperprior for it
- Say what exchangeability means in plain words, why it is not "identical" and not "independent", and what to do when it fails
- Separate between-group variation $\tau^2$ from within-group variation $\sigma^2$, and see why simple estimates of $\tau^2$ can come out negative
- Build hierarchies for conversion rates and counts (logit and log scales, or a Beta population), and run a fake-data check before trusting a fit
What we need from earlier chapters: prior, likelihood, posterior, and the word "hyperparameter" for the numbers inside a prior (Chapter 6.1); prior predictive checks (Chapter 6.2); Beta-Binomial updating and pseudo-counts (Chapter 6.3); the law of total variance and the first look at $\tau^2$ vs $\sigma^2$ (Chapter 4.6); the Normal distribution (Chapter 4.9); the global scaler (Chapter 4.18); bias, variance and shrinkage (Chapter 5.1). Notation: $G$ = number of groups (segments); $g = 1, \dots, G$ indexes them; $n_g$ = number of observations in group $g$; $y_{gi}$ = observation $i$ in group $g$; $\bar y_g$ = the average of group $g$'s observations; $\theta_g$ = group $g$'s true (unknown) value; $\mu$ and $\tau$ = the centre and the spread of the group values; $\sigma$ = the spread of single observations inside a group. In maths, $N(\mu, \tau^2)$ is written with the variance; NumPyro's dist.Normal(mu, tau) takes the standard deviation $\tau$.
The two-level story: a population of groups, then data inside each group core
Think of the children in one family. Each child has their own height. But they share parents, so if you know the heights of four brothers, you already have a decent guess for their sister, before you ever measure her. The children are different, yet they are not strangers to each other.
Segments in an experiment are like that. Mobile users in Spain, desktop users in India, new users, returning users: each segment has its own true average order value or conversion rate. But they are all users of the same product, so their values sit in the same general area. A hierarchical model writes this down as a story with two levels. "Hierarchical" just means "arranged in levels", like an organisation chart; the same thing is also called a multilevel model.
- Level 1 (the population of segments): nature picks each segment's true value $\theta_g$ from one shared distribution, with a centre $\mu$ and a spread $\tau$.
- Level 2 (the data): inside segment $g$, each individual observation scatters around that segment's own $\theta_g$, with spread $\sigma$.
Three ways to say it:
- Picture: a two-stage lottery. First draw a segment's true value out of the "segment hat"; then draw that segment's orders out of its own small hat.
- Numbers: segment averages sit around 40 dollars give or take 4 ($\tau$); single orders sit around their segment's value give or take 12 ($\sigma$).
- Slogan: the groups are different, but they are not strangers.
Simulating the story by hand. The metric is the average order value (AOV) in dollars. Suppose the population of segments has centre $\mu = 40$ and spread $\tau = 4$, and single orders inside a segment have spread $\sigma = 12$. A "standard Normal draw" $z$ is a random number from $N(0, 1)$; we turn it into a draw with mean $m$ and standard deviation $s$ by $m + s \cdot z$.
- What level 1 allows: about 95% of segments have a true AOV within $\mu \pm 1.96\,\tau = 40 \pm 7.84$, so between about 32.2 and 47.8 dollars.
- Draw two segments: segment A gets $z = 0.5$, so $\theta_A = 40 + 4(0.5) = 42$. Segment B gets $z = -1.25$, so $\theta_B = 40 + 4(-1.25) = 35$.
- Draw three orders in segment A: noise draws $0.5, -1, 1$ give $42 + 12(0.5) = 48$, $\;42 + 12(-1) = 30$, $\;42 + 12(1) = 54$. Their average is $(48 + 30 + 54)/3 = 132/3 = 44$, not $42$: an average of 3 noisy orders misses the segment's true value.
- Spread of one order (law of total variance): $\sigma^2 + \tau^2 = 144 + 16 = 160$, an SD of $\sqrt{160} \approx 12.6$. Only $16/160 = 10\%$ of it comes from segment differences.
- Spread of a segment's average of $n = 9$ orders: $\tau^2 + \sigma^2/n = 16 + 144/9 = 16 + 16 = 32$. Half of it is real segment difference, half is noise. That half-and-half split will decide how much a segment should trust its own data (Chapter 6.6).
A two-level hierarchical model for $G$ groups:
$$\theta_g \mid \mu, \tau \;\sim\; N(\mu, \tau^2) \quad (g = 1, \dots, G), \qquad y_{gi} \mid \theta_g \;\sim\; p(y \mid \theta_g) \quad (i = 1, \dots, n_g),$$together with priors for the unknowns at the top, $p(\mu)$ and $p(\tau)$ (and for $\sigma$ if the likelihood has one). The most common data model is Normal: $y_{gi} \mid \theta_g \sim N(\theta_g, \sigma^2)$.
- $\theta_g$: group $g$'s true value (its long-run average). Unknown; one per group.
- $\mu$: the population mean, the centre of all group values.
- $\tau$: the between-group standard deviation: how far a typical group's true value is from $\mu$. $\tau = 0$ means "all groups are identical".
- $\sigma$: the within-group standard deviation: how far a single observation typically is from its own group's $\theta_g$.
- "Given" ($\mid$) matters: the $\theta_g$ are independent once $\mu$ and $\tau$ are known, and the observations are independent once their group's $\theta_g$ is known.
The whole model is one joint distribution, a product of the pieces in the story's order:
$$p(\mu, \tau, \sigma, \boldsymbol\theta, D) = p(\mu)\,p(\tau)\,p(\sigma)\prod_{g=1}^{G}\Big[N(\theta_g \mid \mu, \tau^2)\prod_{i=1}^{n_g} p(y_{gi} \mid \theta_g, \sigma)\Big].$$The posterior is this product divided by the evidence, exactly as in Chapter 6.1; there are just more unknowns.
Why do we need it?
Segments with little data give very noisy estimates. Analysing every segment alone overreacts to that noise; forcing all segments to share one value ignores real differences. The two-level story lets each group have its own value while all groups learn from each other through $\mu$ and $\tau$.
Where is it used?
Per-segment metrics in Bayesian A/B testing frameworks (yours), meta-analysis of many studies (the "eight schools" example), batting averages in sports, polls broken down by state, ranking hospitals or stores, per-store demand forecasts, and user or item effects in recommender systems.
How is it used?
Write the model as the story, top to bottom: priors for $\mu$ and $\tau$, then one $\theta_g$ per group, then the observations. In NumPyro that is numpyro.sample for $\mu$ and $\tau$, a numpyro.plate("groups", G) block for $\theta_g$, and theta[g] indexing for each observation. Fit it, then read off each $\theta_g$ and $\tau$.
"A hierarchical model is any model with many groups."
Fitting a separate model to each of 50 groups is not hierarchical. The defining feature is that the group parameters $\theta_g$ themselves have a distribution, $N(\mu, \tau^2)$, whose parameters are unknown and learned from all groups together.
"$\theta_g \sim N(\mu, \tau^2)$ means the data in each group are Normal."
It is a statement about the groups' true values, not about single observations. The data model $p(y \mid \theta_g)$ can be Binomial, Poisson or Student-t; then $\theta_g$ usually lives on a transformed scale (log-odds, log), as in the last section of this chapter.
"The orange averages are the segments' true values."
They are estimates of $\theta_g$ that still contain noise of variance $\sigma^2/n_g$. With few observations they can sit far from the green ticks, even when all segments are identical ($\tau = 0$).
In an A/B framework like yours, the groups are segments (country, device, new vs returning users), $\theta_g$ is a segment's metric (a conversion rate on the log-odds scale, a mean revenue, a lift), $\mu$ is the overall level and $\tau$ says how much segments truly differ. "Hierarchical partial pooling across groups/segments" is exactly this two-level story fitted in NumPyro; Chapter 6.6 shows what it does to each segment's estimate.
$\theta_g \sim N(\mu, \tau^2)$ (level 1: groups), $\;y_{gi} \sim p(y \mid \theta_g)$ (level 2: data), plus priors on $\mu, \tau$.
$\tau$ = real spread between groups; $\sigma$ = noise inside a group. Normal case: $Var(y) = \sigma^2 + \tau^2$, $Var(\bar y_g) = \tau^2 + \sigma^2/n_g$.
Trap: an observed group average is not the group's true value; it carries noise $\sigma^2/n_g$.
Quick check: what does the two-level model say if $\tau = 0$? And if $\tau$ is huge?
$\tau = 0$: every $\theta_g$ equals $\mu$, so all groups are identical and one shared number describes them (one parameter for everybody). Huge $\tau$: level 1 allows almost any value for each group, so the groups have essentially nothing to do with each other and each one is on its own. Real data sit in between, and the model learns where.
Global, group and hyper: naming every unknown, and the plate diagram core
Think of a school system. Some numbers describe the whole system: the average score over all schools, and how much schools usually differ from each other. Some numbers describe one school: its own average. And some numbers are the students' actual scores. A hierarchical model has the same three kinds of things, and each kind has a name that you will meet in papers, in NumPyro code and in interviews.
- Group parameters (also called local parameters): one per group, here $\theta_1, \dots, \theta_G$.
- Global parameters: one number shared by every group, such as $\mu$, $\tau$ and a common noise level $\sigma$.
- Hyperparameters: the parameters of the distribution the group parameters come from: $\mu$ and $\tau$ in $\theta_g \sim N(\mu, \tau^2)$. "Hyper" is Greek for "over": they sit one level over the group parameters.
- Hyperpriors: the priors we put on the hyperparameters, such as $\mu \sim N(0, 1)$ and $\tau \sim \text{HalfNormal}(1)$.
- Constants: the fixed numbers inside the hyperpriors (the 0, 1 and 1). You choose them; the model does not learn them.
Three ways to say it:
- Picture: an organisation chart: $\mu, \tau$ at the top, one $\theta_g$ per group in the middle, the data at the bottom.
- Numbers: 8 segments give 8 group parameters plus $\mu$, $\tau$, $\sigma$: 11 unknowns.
- Slogan: hyperparameters are learned from the data; the constants in the hyperpriors are chosen by you.
Counting the unknowns for three ways of modelling $G$ segments with a Normal likelihood and one shared noise level $\sigma$.
- One value for everybody ($\theta_g = \theta$ for all $g$): unknowns $\theta, \sigma$. That is $2$, whatever $G$ is.
- A separate value per segment, each with its own fixed prior such as $N(0, 10^2)$: unknowns $\theta_1, \dots, \theta_G, \sigma$. That is $G + 1$; for $G = 8$, $9$.
- Hierarchical: unknowns $\theta_1, \dots, \theta_G$, $\mu$, $\tau$, $\sigma$. That is $G + 3$; for $G = 8$, $11$; for $G = 200$ segments, $203$.
- The hierarchical model has the most unknowns, yet it is the least likely of the last two to chase noise: the shared prior $N(\mu, \tau^2)$ ties the $\theta_g$ together, so they behave like fewer free numbers (Chapter 6.6 makes this precise).
A model graph (also called a directed graphical model or plate diagram) draws the generative story:
- a circle is a random quantity; a shaded circle is observed data; a small square is a fixed constant;
- an arrow $a \to b$ means "$b$ is drawn from a distribution that uses $a$"; the things pointing into $b$ are its parents;
- a plate (a rectangle labelled with an index, such as "$g = 1, \dots, G$") means "repeat everything inside, once per index value". Plates can be nested.
The joint density is one factor per circle, conditional on its parents: $p(\mu)\,p(\tau)\,p(\sigma)\prod_g p(\theta_g \mid \mu, \tau)\prod_{g,i} p(y_{gi} \mid \theta_g, \sigma)$. Unshaded circles are the latent (unobserved) unknowns that inference must handle; their total count, with plates multiplied out, is the model's latent dimension.
Why do we need it?
The words tell you what is learned (hyperparameters, group parameters) and what you chose (constants). The picture makes the story, the joint density and the code line up, and the count of unknowns tells you how big the inference problem is.
Where is it used?
numpyro.plate in NumPyro and Pyro, array-shaped parameters in Stan and PyMC, the figures of every Bayesian textbook, and the latent dimension $d$ that sets the cost of NUTS and of full-rank versus low-rank SVI guides (Chapters 6.13 and 6.17).
How is it used?
Draw the graph before coding. Then write one numpyro.sample per circle, one with numpyro.plate(...) per rectangle, and obs=data for each shaded circle. Count the unshaded circles times their plate sizes to get the number of latent dimensions.
"The hierarchical model has more parameters than separate models, so it must overfit more."
Counting circles is not the whole story. The prior $N(\mu, \tau^2)$ links the $\theta_g$: when the data say the groups are similar ($\tau$ small), the $\theta_g$ are held close together and act like far fewer free numbers.
"$\sigma$ is always a global parameter."
It is a modelling choice. One shared $\sigma$ is global. A noise level per segment, $\sigma_g$, is a group parameter, and it can get its own hierarchy (for example $\log \sigma_g \sim N(a, b^2)$).
"A plate is a loop in Python."
A plate is a statement of conditional independence: the copies inside are independent given their parents. numpyro.plate uses that to vectorize and to scale minibatches; it does not run a Python loop.
"Hyperparameters are the settings I tune, like the learning rate."
In machine learning "hyperparameter" means a training setting. In a hierarchical model, $\mu$ and $\tau$ are unknown random quantities with their own priors (hyperpriors), and the posterior tells you what the data say about them. The fixed constants inside the hyperpriors are what you choose.
Model answer: "In my hierarchical model, the segment effects $\theta_g$ are drawn from $N(\mu, \tau^2)$. $\mu$ and $\tau$ are hyperparameters: they are learned from all segments together, through hyperpriors such as $\tau \sim \text{HalfNormal}(1)$. The only things I fix by hand are the constants in those hyperpriors, and I check them with a prior predictive simulation."
In NumPyro, segment-level parameters normally live inside a numpyro.plate over segments, and $\mu$, $\tau$ (and a shared noise scale) are sampled outside it; check how your own model is laid out. Each unshaded circle times its plate size is one latent dimension for SVI: with segments and variants both inside plates, the latent dimension grows like (number of segments) × (number of variants). That number is what decides how expensive a full-rank guide becomes (Chapter 6.13).
Group (local) parameters $\theta_g$; global parameters $\mu, \tau, \sigma$; hyperparameters $\mu, \tau$ = parameters of the group population; hyperpriors = their priors; constants = numbers inside hyperpriors (chosen).
Plate diagram: circle = random, shaded = observed, square = constant, arrow = "drawn using", plate = "repeat per index". Joint = one factor per circle given its parents.
Counts with $G$ groups: complete 2, separate $G+1$, hierarchical $G+3$.
Quick check: 50 segments, a hierarchy on the segment means, and a separate noise SD $\sigma_g \sim \text{HalfNormal}(1)$ for every segment. How many latent dimensions?
$\theta_1, \dots, \theta_{50}$ (50), $\sigma_1, \dots, \sigma_{50}$ (50), plus $\mu$ and $\tau$ (2): $102$. The $\sigma_g$ sit inside the segment plate; the constant 1 in HalfNormal(1) is not an unknown.
Hyperpriors: letting the data decide how different the groups are core
How different are your segments really? You do not know $\tau$ in advance, and guessing it is risky: guess too small and you treat different segments as the same; guess too big and every segment is left alone with its noise. So we give $\tau$ a prior of its own (a hyperprior) and let the data decide.
How can data say anything about $\tau$? Look at the segment averages. Even if all segments were identical, their averages would differ by noise, and we know how big that noise is ($\sigma/\sqrt{n}$). If the averages are spread out no more than noise alone would produce, the data say "$\tau$ is small". If they are spread out much more, the extra spread is real segment difference, and the data say "$\tau$ is large". Like estimating any spread, this needs several groups: two segments tell you almost nothing about how segments vary.
Three ways to say it:
- Picture: compare how wide the row of segment averages is with how wide the noise bars are; the extra width is $\tau$.
- Numbers: averages of 9 orders have noise SD 4; if the 8 averages have SD 4.9, the left-over spread is about $\sqrt{4.9^2 - 4^2} \approx 2.8$.
- Slogan: $\tau$ is the spread that is left after the noise is taken out.
Eight segments, $n = 9$ orders each, order-level noise $\sigma = 12$. Observed segment averages: 33, 36, 38, 40, 41, 43, 45, 48.
- Noise SD of one average: $\sigma/\sqrt{n} = 12/\sqrt9 = 12/3 = 4$, so noise variance $16$.
- Mean of the averages: $(33 + 36 + 38 + 40 + 41 + 43 + 45 + 48)/8 = 324/8 = 40.5$.
- Squared distances from 40.5: $56.25, 20.25, 6.25, 0.25, 0.25, 6.25, 20.25, 56.25$; sum $= 166$. Sample variance $= 166/7 \approx 23.71$ (SD $\approx 4.87$).
- If $\tau$ were 0, we would expect a variance near 16. The excess is $23.71 - 16 = 7.71$, a rough estimate $\hat\tau \approx \sqrt{7.71} \approx 2.78$.
- The Bayesian answer with hyperprior $\tau \sim \text{HalfNormal}(10)$ (and a flat prior on $\mu$) is a whole distribution: median $\approx 3.05$, 90% credible interval $\approx 0.36$ to $7.6$. Eight segments pin $\tau$ down only loosely; the next widget computes this curve live.
A hyperprior is a prior on a hyperparameter. For the Normal hierarchy the usual choices are
$$\mu \sim N(m_0, s_0^2), \qquad \tau \sim \text{HalfNormal}(s) \;\text{ or }\; \text{Half-Cauchy}(s) \;\text{ or }\; \text{Half-}t \;\text{ or }\; \text{Exponential}(\lambda).$$- "Half" means the distribution folded onto $\tau \ge 0$ (a standard deviation cannot be negative). $\text{HalfNormal}(s)$ puts about 95% of its mass below $2s$.
- With a global scaler (the metric standardized once for all data), $\mu \sim N(0, 1)$ and $\tau \sim \text{HalfNormal}(1)$ or $\text{HalfNormal}(0.5)$ are common weakly informative choices; always translate them back into business units and check them.
How the data inform $\tau$. Average out $\theta_g$: a group average is $\bar y_g = \theta_g + (\text{noise average})$, a sum of two independent Normals, so
$$\bar y_g \mid \mu, \tau \;\sim\; N\!\big(\mu,\; \tau^2 + s_g^2\big), \qquad s_g = \sigma/\sqrt{n_g}.$$The data "see" $\tau$ only through how much the $\bar y_g$ spread beyond their noise $s_g$. With a flat prior on $\mu$, the posterior of $\tau$ is $p(\tau \mid D) \propto p(\tau)\, V_\mu^{1/2}\prod_g (\tau^2 + s_g^2)^{-1/2}\exp\!\Big(-\frac{(\bar y_g - \hat\mu_\tau)^2}{2(\tau^2 + s_g^2)}\Big)$, where $\hat\mu_\tau$ is the average of the $\bar y_g$ weighted by $1/(\tau^2 + s_g^2)$ and $V_\mu = 1/\sum_g (\tau^2 + s_g^2)^{-1}$. You never need to type this: NUTS or SVI handles it. It is shown so you can see that only the spread of the averages matters.
Why do we need it?
$\tau$ controls how much the groups share (Chapter 6.6). Fixing it by hand is a guess; learning it from all groups makes the amount of sharing data-driven, and the hyperprior keeps it away from silly values when there are few groups.
Where is it used?
The eight-schools model (half-Cauchy on $\tau$), the hierarchical examples in the NumPyro, Stan and PyMC documentation, random-effects meta-analysis, segment pooling in Bayesian A/B testing, and hierarchical priors on Fourier or holiday coefficients in forecasting models.
How is it used?
Standardize the metric once (global scaler); choose $\mu \sim N(0, 1)$ and a half-Normal or half-$t$ for $\tau$; simulate from the hyperpriors to see what segment differences they allow; fit; look at the posterior of $\tau$; refit with one or two other reasonable scales and report whether conclusions change (Chapter 6.8).
"$\tau$ is the standard deviation of the observed segment averages."
The averages spread by $\sqrt{\tau^2 + \sigma^2/n}$; noise is part of it. $\tau$ is what is left after subtracting the noise, and it can be near zero even when the averages look quite different.
"InverseGamma(0.001, 0.001) on $\tau^2$ is a safe, non-informative choice."
It looks vague but behaves strongly near $\tau = 0$, and when the data are consistent with small $\tau$ the answer can depend on the arbitrary 0.001 (Gelman, 2006). Put a half-Normal, half-$t$ or half-Cauchy prior on the standard deviation $\tau$ instead.
"With 3 segments, the data will tell me $\tau$."
Estimating a spread from 3 numbers is very uncertain. With few groups the posterior of $\tau$ is wide and the hyperprior matters a lot; check how sensitive your conclusions are to it.
In an A/B framework like yours, the hyperprior on the between-segment spread answers "how different can segments plausibly be?". If the metric is put through the global scaler, a $\text{HalfNormal}(1)$ on $\tau$ allows the spread between segments to be as large as about two standard deviations of the metric (its 95% limit is $1.96$), which is generous. Run the prior predictive simulation above in your own units (dollars, or conversion points after the logit transform of the last section) before trusting that number.
Hyperprior = prior on $\mu$, $\tau$; for $\tau$ use HalfNormal / half-$t$ / half-Cauchy on the SD.
$\bar y_g \mid \mu, \tau \sim N(\mu, \tau^2 + \sigma^2/n_g)$: the data see $\tau$ only through spread beyond noise. Rough $\hat\tau^2 = s^2_{\bar y} - \sigma^2/n$.
Trap: few groups ⇒ $\tau$ poorly learned ⇒ hyperprior matters; check it with a prior predictive simulation.
Quick check: 5 segments, 25 orders each, $\sigma = 10$. The 5 averages have SD 2.2. What is the rough estimate of $\tau$, and what will the posterior look like?
Noise SD of an average $= 10/\sqrt{25} = 2$, noise variance 4. Excess $= 2.2^2 - 4 = 4.84 - 4 = 0.84$, so $\hat\tau \approx 0.92$: the segments might be nearly identical. With only 5 groups, the posterior of $\tau$ will have a lot of mass near 0 but a long tail to larger values, and its exact shape will depend on the hyperprior.
Exchangeability: the assumption that lets groups share core
Write each segment's name on a card and shuffle the cards face down. Before seeing any data, would you say anything different about the third card than about the seventh? If not, the segments are exchangeable for you: you believe the same thing about each one ("its value is around 40, give or take 4"), and swapping the labels changes nothing in your beliefs.
Two things exchangeable does not mean. It does not mean identical: the segments' values really differ. And it does not mean independent: because they come from the same family, learning that one segment is high makes you expect the others to be a bit higher too. That dependence is exactly what lets groups "borrow strength" from each other.
Three ways to say it:
- Picture: shuffle the labels; nothing in your beliefs moves.
- Numbers: before the data every segment is "40 ± 4", the same sentence for every label; if one turns out at 50, your guess for the others rises toward it.
- Slogan: exchangeable means "interchangeable before the data", not "the same".
Exchangeable, yet dependent. Suppose even the centre is uncertain: $\mu \sim N(40, 10^2)$, and $\theta_g = \mu + \tau z_g$ with $\tau = 4$ and independent $z_g \sim N(0,1)$.
- Each segment: $Var(\theta_g) = Var(\mu) + \tau^2 = 100 + 16 = 116$. The same for every label, so the labels are interchangeable.
- Two segments share $\mu$: $Cov(\theta_1, \theta_2) = Var(\mu) = 100$. Correlation $= 100/116 \approx 0.86$.
- Learn that $\theta_1 = 50$. For two jointly Normal values with the same mean $m$ and the same variance, the best guess for one given the other is $m + \rho\,(\text{other} - m)$: here $40 + 0.86 \times (50 - 40) \approx 48.6$. One segment's value moved our guess for another by 8.6 dollars.
Not exchangeable. The groups are weeks 1 to 8 of a growing product, with conversion rising about 0.5 points a week. Before the data you already expect week 8 to be higher than week 1, so swapping the labels "week 1" and "week 8" would change your beliefs. Treating these weeks as exchangeable would pull late weeks down and early weeks up.
Quantities $\theta_1, \dots, \theta_G$ are exchangeable if their joint distribution does not change when you reorder them: $p(\theta_1, \dots, \theta_G) = p(\theta_{\pi(1)}, \dots, \theta_{\pi(G)})$ for every reordering (permutation) $\pi$.
- iid $\Rightarrow$ exchangeable, but exchangeable $\not\Rightarrow$ independent (the example above has correlation 0.86).
- De Finetti's theorem (informally): if you would treat any number of such groups as exchangeable, your beliefs can always be written as "iid given some unknown parameters $\phi$, with a prior on $\phi$": $p(\theta_1, \dots, \theta_G) = \int \prod_g p(\theta_g \mid \phi)\,p(\phi)\,d\phi$. With $\phi = (\mu, \tau)$ this is the hierarchical model. Exchangeability is the assumption that justifies writing $\theta_g \sim N(\mu, \tau^2)$.
- Conditional exchangeability: if groups differ in a known way $x_g$ (device, market size, week), use $\theta_g \sim N(\mu + \beta x_g, \tau^2)$. The groups are then exchangeable given $x_g$: only the left-over differences are treated as interchangeable.
- Typical failures: known systematic differences (add a covariate), an ordering in time (add a trend or a time-series model), nested structure such as cities within countries (add another level), and groups selected because they looked extreme.
Why do we need it?
It is the assumption that makes sharing information fair. If it is false, pooling pulls every group toward the wrong centre and adds systematic error instead of removing noise.
Where is it used?
Every hierarchical model; meta-analysis ("are these studies comparable?"); segment pooling in A/B tests ("are these segments comparable?"); multilevel regression and post-stratification in election polling, where state-level predictors make states exchangeable given those predictors.
How is it used?
Ask: "Before the data, do I know anything that makes one group different from another?" If yes, put that knowledge into the model as a group-level covariate or as another level. Then check the fitted group effects against that knowledge (plot them against the covariate, against time).
"Exchangeable means the groups are the same."
Their values differ (that is what $\tau$ measures). Exchangeable only means you have no information, before the data, that tells the groups apart.
"Exchangeable means independent, so each group's estimate should use only its own data."
Exchangeable groups are dependent through the shared $\mu$ and $\tau$. That dependence is why one group's data change your estimate for another: it is the whole point of partial pooling.
"Mobile and desktop convert differently, so a hierarchical model cannot be used for them."
Put the known difference into the population model, $\theta_g \sim N(\mu + \beta\,x_g, \tau^2)$ with $x_g$ = "is mobile". The groups are then exchangeable given device, and the pooling pulls each segment toward the right centre.
"Partial pooling assumes all the groups are identical."
That is complete pooling. Partial pooling assumes exchangeability: the groups differ, but we have no prior information that distinguishes them (possibly after adjusting for known covariates).
Model answer: "The key assumption is exchangeability: before seeing the data, I would say the same thing about every segment, so their effects can be modelled as draws from one population. If I know something that separates segments, such as device or market size, I add it as a group-level predictor so that only the left-over differences are treated as exchangeable."
In an A/B framework like yours, ask the exchangeability question for every pooling decision. Segments such as countries you know little about are a good fit. Segments with a known systematic difference (mobile vs desktop, paid vs organic traffic) need that difference as a predictor, or the model pulls them toward a shared centre that suits neither. And the variants are not exchangeable groups: you pool segment-level results across segments, you do not pool control and treatment into one population, because their difference is exactly what the experiment measures.
Exchangeable: $p(\theta_1, \dots, \theta_G)$ unchanged by reordering = "interchangeable before the data".
iid ⇒ exchangeable; exchangeable ⇏ independent (shared $\mu$ gives correlation $Var(\mu)/(Var(\mu) + \tau^2)$). De Finetti: exchangeable ≈ iid given unknown $(\mu, \tau)$ = the hierarchical model.
Fails with known differences, time order, nesting, or selection on extremes. Fix: covariates ($\theta_g \sim N(\mu + \beta x_g, \tau^2)$) or more levels.
Quick check: eight stores, four in big cities (known to sell about twice as much) and four in small towns. Are the eight stores exchangeable?
Not as one group: before the data you already expect the city stores to be higher, so swapping a city label with a town label changes your beliefs. They are exchangeable given the store type: model $\theta_g \sim N(\mu + \beta\,\text{city}_g, \tau^2)$, or give the two types their own population means.
Between-group variation $\tau^2$ and within-group variation $\sigma^2$ core
Two orders can differ for two reasons. They may come from different segments (one market spends more than another): that is between-group variation, measured by $\tau$. Or they may come from the same segment and differ because single orders always vary: that is within-group variation, measured by $\sigma$. You met this split as the law of total variance in Chapter 4.6. Here we use it to answer a practical question: how much of a segment's observed average is real signal, and how much is noise?
The answer depends on the segment's size. One order is mostly noise ($\sigma$ is big). The average of many orders has little noise left ($\sigma^2/n$ is small), so it is mostly the segment's real value.
Three ways to say it:
- Picture: a bar for "how much a segment average varies", split into a green part (real segment differences) and a blue part (noise that shrinks as $n$ grows).
- Numbers: with $\tau = 4$, $\sigma = 12$: one order is 10% signal; an average of 9 orders is 50% signal; an average of 144 orders is 94% signal.
- Slogan: $\tau^2$ is the signal between segments; $\sigma^2/n$ is the noise sitting on each segment's average.
$\tau = 4$ and $\sigma = 12$, so $\tau^2 = 16$ and $\sigma^2 = 144$.
- One order: $Var(y_{gi}) = \sigma^2 + \tau^2 = 144 + 16 = 160$. Signal share $16/160 = 0.10$. This share is also the correlation between two orders from the same segment, the intraclass correlation (ICC).
- Average of $n = 9$ orders: $Var(\bar y_g) = \tau^2 + \sigma^2/n = 16 + 144/9 = 16 + 16 = 32$. Signal share $16/32 = 0.5$.
- Average of $n = 144$ orders: $16 + 144/144 = 16 + 1 = 17$. Signal share $16/17 \approx 0.94$.
- The size at which signal and noise are equal: $\sigma^2/n = \tau^2$, so $n = \sigma^2/\tau^2 = 144/16 = 9$ orders. Remember this number: in Chapter 6.6 it becomes "the population is worth 9 orders".
- A rough estimate of $\tau^2$ from 8 segment averages of 9 orders: their sample variance minus the noise, $s^2_{\bar y} - 16$. If the averages happen to have variance 12, this gives $12 - 16 = -4$: a negative variance, which is impossible. Noise alone made the segments look more similar than usual.
For the Normal two-level model ($\theta_g \sim N(\mu, \tau^2)$, $y_{gi} \mid \theta_g \sim N(\theta_g, \sigma^2)$):
$$Var(y_{gi}) = \underbrace{\sigma^2}_{\text{within}} + \underbrace{\tau^2}_{\text{between}}, \qquad \text{ICC} = \frac{\tau^2}{\tau^2 + \sigma^2}, \qquad Var(\bar y_g) = \tau^2 + \frac{\sigma^2}{n_g}.$$- Noise share of a group average: $\dfrac{\sigma^2/n_g}{\tau^2 + \sigma^2/n_g}$. This exact fraction becomes the amount of shrinkage in Chapter 6.6.
- Method-of-moments (ANOVA) estimate with equal sizes $n$: $\hat\tau^2 = s^2_{\bar y} - \hat\sigma^2/n$, where $s^2_{\bar y}$ is the sample variance of the $G$ averages. On average it equals $\tau^2$ (unbiased), but in a single dataset it can be negative. Setting negatives to 0 makes it biased upward.
- The Bayesian posterior of $\tau$ never goes negative: when the data look "too similar", it puts more weight near $\tau = 0$.
Why do we need it?
The split tells you whether segment differences are worth chasing and how much each segment's own average can be trusted. Without it, noise in small segments looks like real segment differences.
Where is it used?
The intraclass correlation and design effect of clustered experiments (users within cities), random-effects ANOVA, variance components in mixed models (statsmodels MixedLM, lme4), meta-analysis heterogeneity ($\tau^2$, $I^2$), and the shrinkage of segment estimates in your A/B framework.
How is it used?
Estimate $\sigma$ from the spread inside segments and $\tau$ from the extra spread between segment averages (or fit the hierarchical model and read both posteriors). Compare $\tau^2$ with $\sigma^2/n_g$ for each segment: if the noise term dominates, that segment's raw average should not drive decisions.
"The ICC is only 0.10, so segments do not matter."
The ICC compares $\tau^2$ with the noise of one observation. A segment's average of 1 000 orders has noise $\sigma^2/1000$, far below $\tau^2$, so segment differences can be large and very clear even with a small ICC.
"A negative estimate of $\tau^2$ means there is a bug."
It is ordinary sampling noise: the segment averages happened to be closer together than noise alone would usually make them. Report $\tau$ as "close to zero, uncertain", or use the posterior of $\tau$, which is always non-negative.
"Within-group variation is the variance of all the data."
The variance of all the data is within + between. The within part is the average variance inside segments, around each segment's own mean.
For continuous metrics in your A/B framework, the global scaler (one mean and one SD for all segments, Chapter 4.18) rescales $\tau^2$ and $\sigma^2$ by the same factor, so their ratio and every signal share above stay the same. Scaling each segment by its own mean and SD would make every segment's average exactly 0: the between part $\tau^2$ in the data becomes 0 and the model can only conclude "no segment differences". That is the syllabus warning in one line of algebra.
$Var(y) = \sigma^2 + \tau^2$; ICC $= \tau^2/(\tau^2 + \sigma^2)$; $Var(\bar y_g) = \tau^2 + \sigma^2/n_g$.
Noise share of a segment average $= (\sigma^2/n_g)/(\tau^2 + \sigma^2/n_g)$; equal signal and noise at $n = \sigma^2/\tau^2$.
Trap: $\hat\tau^2 = s^2_{\bar y} - \sigma^2/n$ can be negative (noise); the Bayesian posterior of $\tau$ stays $\ge 0$.
Quick check: $\tau = 2$ and $\sigma = 20$. How many orders does a segment need before its average is at least half signal?
Signal $\tau^2 = 4$ equals noise $\sigma^2/n = 400/n$ at $n = 100$. Below 100 orders the average is mostly noise; above it, mostly signal.
Beyond Normal data: hierarchies for conversion rates and counts
The bell curve $N(\mu, \tau^2)$ spreads over every real number, negative ones included. A conversion rate must stay between 0 and 1, and an order rate must stay above 0. So we put the bell curve on a scale where the parameter can roam freely, then map back.
- For a rate $p$: the log-odds or logit, $\log\frac{p}{1-p}$. It stretches the ruler from 0 to 1 into a ruler from $-\infty$ to $+\infty$. The way back is the logistic function $\frac{1}{1+e^{-\theta}}$.
- For a positive rate $\lambda$ (counts): the logarithm $\log\lambda$; the way back is $e^{\theta}$.
- Or, for rates, use a population that already lives on $[0, 1]$: a Beta distribution, whose pseudo-counts you met in Chapter 6.3.
Three ways to say it:
- Picture: stretch the 0-to-1 ruler into an endless ruler, put the bell curve there, squeeze it back.
- Numbers: overall rate 10% → $\mu = \text{logit}(0.10) = -2.20$; with $\tau = 0.3$, 95% of segments convert between 5.8% and 16.7%.
- Slogan: put the bell curve where the parameter is free to roam.
Segments of a checkout test, overall conversion about 10%, between-segment spread $\tau = 0.3$ on the logit scale.
- Centre: $\mu = \text{logit}(0.10) = \ln(0.10/0.90) = \ln(0.111) \approx -2.197$.
- 95% range on the logit scale: $\mu \pm 1.96\tau = -2.197 \pm 0.588$, from $-2.785$ to $-1.609$.
- Back to rates: $1/(1 + e^{2.785}) = 1/(1 + 16.2) \approx 0.058$ and $1/(1 + e^{1.609}) = 1/(1 + 5.0) \approx 0.167$. So 95% of segments convert between about 5.8% and 16.7%: 4.2 points below 10% but 6.7 points above it. The range is lopsided, as rates near 0 should be.
- A rule of thumb for small $\tau$: near a rate $p$, one logit SD $\tau$ is about a relative change of $\tau(1-p)$. Here $0.3 \times 0.9 = 0.27$: segments typically differ by about 27% of 10%, roughly 2.7 points (the exact SD is 2.8 points).
- The Beta alternative: $p_g \sim \text{Beta}(10, 90)$ has mean $10/100 = 0.10$ and SD $\sqrt{0.1 \times 0.9/101} \approx 0.030$. The population is "worth 100 visitors" ($\kappa = \alpha + \beta = 100$).
- Logit-Normal hierarchy for conversions: $\theta_g \sim N(\mu, \tau^2)$, $p_g = \text{logistic}(\theta_g)$, $k_g \mid p_g \sim \text{Binomial}(n_g, p_g)$. Note that $\text{logistic}(\mu)$ is the median segment rate; the mean rate is slightly higher (10.3% in the example).
- Log-Normal hierarchy for counts: $\theta_g \sim N(\mu, \tau^2)$, $\lambda_g = e^{\theta_g}$, $y_{gi} \mid \lambda_g \sim \text{Poisson}(\lambda_g)$ (or Negative Binomial). A logit or log scale used this way is called a link, as in GLMs (Chapter 5.14).
- Beta hierarchy: $p_g \sim \text{Beta}(\kappa\phi, \kappa(1-\phi))$ with population mean rate $\phi$ and concentration $\kappa \gt 0$ ("worth $\kappa$ visitors"); $Var(p_g) = \phi(1-\phi)/(\kappa + 1)$. Given $\phi$ and $\kappa$, each segment updates exactly like the Beta-Binomial of Chapter 6.3. Hyperpriors go on $\phi$ (a Beta) and $\kappa$ (a positive distribution).
- Segment lifts in an A/B test: $\text{logit}\,p_{g,B} = \text{logit}\,p_{g,A} + \delta_g$ with $\delta_g \sim N(\Delta, \tau_\delta^2)$: the overall lift $\Delta$ and how much the lift varies by segment, $\tau_\delta$.
Why do we need it?
A Normal population placed directly on a rate can produce impossible rates below 0, and it treats one point at 2% like one point at 50%. On the logit or log scale parameters stay legal and differences are relative, which matches how rates actually vary.
Where is it used?
Hierarchical logistic regression for conversion by segment, Poisson and Negative Binomial models with store or segment effects, Beta-Binomial hierarchies for click-through rates and batting averages, and segment-level lifts in Bayesian A/B testing.
How is it used?
Choose the scale from the likelihood (Binomial → logit, Poisson → log, Normal → the metric itself after the global scaler). Put $\theta_g \sim N(\mu, \tau^2)$ there. Translate $\tau$'s hyperprior into percentage points with a simulation. In NumPyro: dist.Binomial(total_count=n, logits=theta[g]) or dist.Poisson(jnp.exp(theta[g])).
"$\tau = 0.3$ means segments differ by about 0.3 percentage points."
$\tau$ lives on the logit scale. Near a 10% rate, $\tau = 0.3$ means segments typically differ by about 27% relative, roughly 2.7 points. Always translate it with a simulation.
"Simplest is best: $p_g \sim N(0.02, 0.01^2)$ for a 2% metric."
That population gives about 2.3% of segments a negative conversion rate. Use the logit scale or a Beta population, which can never leave $[0, 1]$.
"$\text{logistic}(\mu)$ is the average conversion rate across segments."
It is the median segment rate. Because the logistic curve bends, the mean rate is a little higher (10.3% instead of 10% at $\tau = 0.3$; 13.4% at $\tau = 1$).
Your A/B framework has Beta-Binomial, Dirichlet-Multinomial, Normal, Student-t and Poisson metrics. The natural scale for a hierarchy on each is: logit (or a Beta population with mean $\phi$ and concentration $\kappa$) for conversions, log for Poisson rates, and the mean itself for Normal or Student-t metrics on the globally scaled outcome. For category shares, the same idea uses a shared Dirichlet population or a softmax scale. Both "logit-Normal" and "Beta population" are common for conversions; check which one your code builds, because the meaning of $\tau$ (or $\kappa$) differs.
Rates: $\text{logit}\,p_g = \theta_g \sim N(\mu, \tau^2)$, $k_g \sim \text{Binomial}(n_g, p_g)$. Counts: $\log\lambda_g \sim N(\mu, \tau^2)$. Or $p_g \sim \text{Beta}(\kappa\phi, \kappa(1-\phi))$, $Var = \phi(1-\phi)/(\kappa+1)$.
logit(0.10) = −2.20; τ = 0.3 → 95% of segments in 5.8%–16.7%; near rate $p$, τ ≈ relative spread $\tau(1-p)$.
Trap: τ is not in percentage points; logistic(μ) is the median, not the mean rate.
Quick check: the overall rate is 50% and $\tau = 0.4$ on the logit scale. Roughly how many percentage points does a typical segment differ from 50%?
At $p = 0.5$, the slope of the logistic curve is $p(1-p) = 0.25$, so one logit SD is about $0.4 \times 0.25 = 0.10$: about 10 points. (Relative version: $\tau(1-p) = 0.4 \times 0.5 = 20\%$ of 50% = 10 points.)
Simulate first: fake-data checks for a hierarchical model
Before you weigh flour on a new kitchen scale, you put a known 1 kg weight on it. A fake-data check does the same for a model: choose the true $\mu$, $\tau$, $\sigma$ yourself, simulate a dataset from the two-level story, fit the model to it, and see whether the fit finds the values you chose. Because a hierarchical model is a story of how data are made, it can always generate its own test data.
For hierarchical models the check is especially revealing: it shows you, before any real decision, how poorly $\tau$ is learned from a handful of groups, and it catches coding mistakes and sampler trouble.
Three ways to say it:
- Picture: put a known weight on the scale before you trust it.
- Numbers: with 8 segments and true $\tau = 4$, the 90% interval for $\tau$ is about 7.5 dollars wide on average; with 30 segments, about 3.8.
- Slogan: if the model cannot find the truth in fake data, it will not find it in real data.
The code at the end of this chapter runs one full check.
- Choose the truth: $\mu = 40$, $\tau = 4$, $\sigma = 12$, eight segments with 5 to 200 orders.
- Simulate: draw 8 segment means, then the orders. The 5-order segment's average comes out at 28.8 although its true value is 40.5.
- Fit the hierarchical model (with a global scaler) by NUTS.
- Compare: posterior means $\mu \approx 41.5$ (truth 40), $\sigma \approx 12.2$ (truth 12); the 90% interval for $\tau$ is 1.1 to 6.4 and contains the true 4. The 5-order segment's estimate is 38.3, much closer to its truth than its raw average.
- Repeat many times (the widget below): about 9 in 10 of the 90% intervals should contain the truth, and their width shows how much the data can say.
A fake-data check (also called a parameter-recovery check) is: (1) fix parameter values; (2) simulate a dataset from the model; (3) fit the same model to it; (4) compare the posterior with the values you fixed; (5) repeat.
- If you instead draw the "true" parameters from the prior every time, the check is called simulation-based calibration: when inference is correct, a 90% interval then contains the truth in exactly 90% of runs. For one fixed truth, coverage is only approximately 90%.
- It tests the inference (code, sampler, how much the data can tell you), not whether real data follow the model. That second question is answered by posterior predictive checks (Chapter 6.8).
Why do we need it?
Hierarchical models fail quietly: a bug in indexing, a scale mix-up, a sampler stuck in a hard region, or too few groups to learn $\tau$. On fake data you know the answer, so every one of these problems becomes visible.
Where is it used?
The Bayesian workflow recommended by the Stan and NumPyro developers, simulation-based calibration (SBC) packages, validating A/B decision rules before launch, and testing forecasting pipelines on simulated series with known changepoints.
How is it used?
Simulate with NumPy (or with NumPyro's Predictive after fixing parameters), run your real pipeline end to end, and record whether the intervals contain the truth and how wide they are. Do it for a few scenarios: small $\tau$, large $\tau$, few groups, tiny groups.
"The fit recovered the truth on one fake dataset, so the model works."
One dataset can be lucky. Repeat the check many times and in several scenarios (small $\tau$, few groups, tiny groups); look at how often the intervals contain the truth and how wide they are.
"Some 90% intervals missed the truth, so there is a bug."
About 1 in 10 should miss. Worry when far more miss, or when they all miss in the same direction.
"A good fake-data check proves the model is right for my real data."
It proves the machinery can recover parameters when the model is true. Whether the real data look like the model is a separate question, answered by posterior predictive checks (Chapter 6.8).
For an A/B framework like yours, a fake-data check is the cheapest insurance you can buy: simulate segments with known lifts (most zero, a few positive), push them through the whole pipeline (global scaler, hierarchical model, SVI, decision rule) and count how often the rule declares a winner when the true lift is zero, and how often it finds the real ones. The same habit applies to the forecasting model: simulate a series with known changepoints and seasonality, then check that the fit finds them.
Fake-data check: fix parameters → simulate → fit → compare → repeat. Truth drawn from the prior each time = simulation-based calibration.
With few groups, $\tau$'s interval is wide (8 groups, $\tau = 4$: about 7.5 wide; 30 groups: about 3.8).
Trap: it checks the inference, not whether the model fits real data.
Quick check: you run 100 fake datasets and the 90% intervals for $\tau$ contain the truth only 55 times, all missing on the low side. What could be wrong?
Far too few intervals contain the truth, and the misses are one-sided, so something systematic is wrong: for example a bug that uses the wrong noise level (too large a $\sigma$ makes the model explain the spread of averages as noise and underestimate $\tau$), a scale mix-up such as passing a variance where NumPyro expects a standard deviation, or a sampler that never explores large $\tau$. The check has done its job: fix the cause before using real data.
Recap, cheat sheet and practice
- A hierarchical model is a two-level story: group values $\theta_g \sim N(\mu, \tau^2)$, then data $y_{gi} \sim p(y \mid \theta_g)$, with hyperpriors on $\mu$ and $\tau$.
- Names: group (local) parameters $\theta_g$; global parameters $\mu, \tau, \sigma$; hyperparameters $\mu, \tau$ (learned); hyperpriors (their priors); constants (chosen). A plate diagram draws it; the joint density is one factor per circle given its parents.
- The data see $\tau$ only through how much the group averages spread beyond their noise: $\bar y_g \sim N(\mu, \tau^2 + \sigma^2/n_g)$. Few groups ⇒ $\tau$ is poorly learned and the hyperprior matters; check it with a prior predictive simulation.
- Exchangeability = "interchangeable before the data". It is not "identical" and not "independent". Known differences go into the model as covariates or extra levels.
- Between vs within: $Var(y) = \sigma^2 + \tau^2$, ICC $= \tau^2/(\tau^2 + \sigma^2)$, $Var(\bar y_g) = \tau^2 + \sigma^2/n_g$; simple estimates of $\tau^2$ can be negative.
- Rates and counts: put the Normal on the logit or log scale, or use a Beta population; translate $\tau$ into business units.
- Fake-data checks test the inference before you trust it on real data.
Cheat sheet
| Idea | Formula | Plain words / remember |
|---|---|---|
| Two-level model | $\theta_g \sim N(\mu, \tau^2)$, $y_{gi} \sim p(y \mid \theta_g)$ | groups from a population, data from each group |
| Joint density | $p(\mu)p(\tau)\prod_g N(\theta_g \mid \mu, \tau^2)\prod_i p(y_{gi} \mid \theta_g)$ | one factor per circle of the plate diagram |
| Unknown counts | complete 2 · separate $G+1$ · hierarchical $G+3$ | more circles, yet less overfitting |
| Hyperpriors | $\mu \sim N(0,1)$, $\tau \sim \text{HalfNormal}(1)$ (scaled data) | half-Normal / half-$t$ on the SD, not InvGamma(ε, ε) |
| What the data see | $\bar y_g \mid \mu, \tau \sim N(\mu, \tau^2 + \sigma^2/n_g)$ | $\tau$ = spread beyond noise |
| Rough $\tau^2$ | $s^2_{\bar y} - \sigma^2/n$ | unbiased but can be negative |
| Exchangeable | $p(\theta_1..\theta_G)$ invariant to reordering | ≈ iid given $(\mu, \tau)$ (de Finetti); not independent |
| Shared-centre correlation | $Var(\mu)/(Var(\mu) + \tau^2)$ | borrowing strength |
| Variance split | $\sigma^2 + \tau^2$; ICC $= \tau^2/(\tau^2+\sigma^2)$ | within + between |
| Group average | $\tau^2 + \sigma^2/n_g$; equal at $n = \sigma^2/\tau^2$ | small groups are mostly noise |
| Rates | $\text{logit}\,p_g \sim N(\mu, \tau^2)$ or $\text{Beta}(\kappa\phi, \kappa(1-\phi))$ | τ ≈ relative spread $\tau(1-p)$; κ = "worth κ visitors" |
| Fake-data check | fix → simulate → fit → compare → repeat | tests the inference, not the fit to reality |
import numpy as np
import jax
import numpyro
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS, Predictive
from numpyro.handlers import reparam
from numpyro.infer.reparam import LocScaleReparam
rng = np.random.default_rng(0)
# 1) Generative simulation, level by level (a fake dataset with KNOWN answers)
mu_true, tau_true, sigma_true = 40.0, 4.0, 12.0 # order value in dollars
n_g = np.array([5, 10, 20, 40, 60, 80, 120, 200]) # orders per segment: very unequal
G = len(n_g)
theta_true = rng.normal(mu_true, tau_true, G) # level 1: one true mean per segment
g = np.repeat(np.arange(G), n_g) # segment index of every order
y = rng.normal(theta_true[g], sigma_true) # level 2: orders around their segment's mean
ybar = np.array([y[g == j].mean() for j in range(G)])
print(np.round(theta_true, 1)) # [40.5 39.5 42.6 40.4 37.9 41.4 45.2 43.8]
print(np.round(ybar, 1)) # [28.8 38.2 43.9 43.3 38.3 40.2 43.5 44.2] the 5-order segment is far off
# 2) Law of total variance: Var(y) = sigma^2 + tau^2 = 160; Var(ybar) = tau^2 + sigma^2/n
th = rng.normal(mu_true, tau_true, 4000)
yy = rng.normal(th[:, None], sigma_true, (4000, 9))
print(round(yy.var(), 1), round(yy.mean(axis=1).var(), 1)) # 160.6 31.7 (theory 160 and 16 + 144/9 = 32)
# 3) The rough estimate tau^2_hat = var(averages) - sigma^2/n is often negative when tau is small
yb = rng.normal(40, 1.0, (20000, 8)) + rng.normal(0, 4.0, (20000, 8)) # true tau = 1, noise sd 4
print(round(((yb.var(axis=1, ddof=1) - 16) < 0).mean(), 2)) # 0.52
# 4) The model, on a GLOBALLY scaled outcome (one mean and one sd for every order)
m, s = y.mean(), y.std()
z = (y - m) / s
def model(g, G, z=None):
mu = numpyro.sample("mu", dist.Normal(0.0, 1.0)) # hyperparameter (global)
tau = numpyro.sample("tau", dist.HalfNormal(1.0)) # hyperparameter (global): between-segment sd
sigma = numpyro.sample("sigma", dist.HalfNormal(1.0)) # global: within-segment noise sd
with numpyro.plate("segments", G):
theta = numpyro.sample("theta", dist.Normal(mu, tau)) # group parameters
with numpyro.plate("orders", len(g)):
numpyro.sample("z", dist.Normal(theta[g], sigma), obs=z)
# 5) Prior predictive check of the hyperpriors, in dollars: gap between best and worst segment
prior = Predictive(model, num_samples=2000)(jax.random.PRNGKey(1), g=g, G=G)
gap = s * (prior["theta"].max(axis=1) - prior["theta"].min(axis=1))
print(np.round(np.quantile(gap, [0.1, 0.5, 0.9]), 1)) # [ 4.1 21.8 57.3] generous but not absurd
# 6) Fit with NUTS. The same model in two coordinate systems (Chapter 6.7 explains the difference)
for name, mdl in [("centered", model),
("non-centered", reparam(model, config={"theta": LocScaleReparam(centered=0)}))]:
mcmc = MCMC(NUTS(mdl), num_warmup=500, num_samples=1000, progress_bar=False)
mcmc.run(jax.random.PRNGKey(0), g=g, G=G, z=z, extra_fields=("diverging",))
print(name, "divergences:", int(mcmc.get_extra_fields()["diverging"].sum()))
# centered divergences: 186 <- a warning sign: do not trust that run (Chapter 6.7)
# non-centered divergences: 0
post = mcmc.get_samples() # samples from the non-centered run
# 7) Fake-data check: did we recover the known truth? (undo the global scaler first)
mu_p, tau_p, sig_p = m + s * post["mu"], s * post["tau"], s * post["sigma"]
print(round(float(mu_p.mean()), 1), round(float(sig_p.mean()), 1)) # 41.5 12.2 (truth 40 and 12)
print(np.round(np.quantile(tau_p, [0.05, 0.5, 0.95]), 1)) # [1.1 2.9 6.4] (truth 4: inside, and wide)
theta_p = np.asarray(m + s * post["theta"], dtype=float)
print(np.round(theta_p.mean(axis=0), 1))
# [38.3 40.4 42.9 42.7 39.2 40.5 43.1 43.9] the 5-order segment moved from 28.8 to 38.3 (truth 40.5)
1. In the model $\theta_g \sim N(\mu, \tau^2)$, $y_{gi} \sim N(\theta_g, \sigma^2)$, with $\mu \sim N(0, 1)$ and $\tau \sim \text{HalfNormal}(1)$, which quantities are the hyperparameters?
2. $\tau = 3$, $\sigma = 12$, and each segment average uses 16 orders. What is the variance of a segment's observed average?
3. Segments are exchangeable when…
4. Eight segment averages have SD 3; the noise SD of each average is 4. What is the best summary of the between-segment spread?
5. Overall conversion 10%, segment spread $\tau = 0.3$ on the logit scale. About 95% of segments convert between…
6. What does a fake-data check of a hierarchical model tell you?
Practice problems
A. Simulate by hand: $\mu = 50$, $\tau = 5$, $\sigma = 10$. A segment draws $z = -0.4$; its three orders draw noise $1.2, -0.5, 0.2$. Find the segment's true value, the orders and their average. Then find $Var(y)$, the ICC, and the variance of an average of 4 orders.
- $\theta = 50 + 5(-0.4) = 48$.
- Orders: $48 + 10(1.2) = 60$, $\;48 + 10(-0.5) = 43$, $\;48 + 10(0.2) = 50$. Average $= 153/3 = 51$.
- $Var(y) = \sigma^2 + \tau^2 = 100 + 25 = 125$; ICC $= 25/125 = 0.2$.
- Average of 4: $\tau^2 + \sigma^2/4 = 25 + 25 = 50$. Half signal, half noise ($n = \sigma^2/\tau^2 = 4$).
B. 30 countries × 2 devices. Each country-device cell has a mean $\theta_{c,d} \sim N(\mu_d, \tau^2)$ with a separate population mean for each device, one shared $\tau$, and one shared noise $\sigma$. How many latent dimensions? Draw the plate diagram in words.
Cells: $30 \times 2 = 60$ values $\theta_{c,d}$; two device means $\mu_{\text{mobile}}, \mu_{\text{desktop}}$; $\tau$; $\sigma$. Total $60 + 2 + 1 + 1 = 64$. Diagram: a plate "d = 1..2" holding $\mu_d$; inside it, a plate "c = 1..30" holding $\theta_{c,d}$ (arrows from $\mu_d$ and from $\tau$, which sits outside both plates); inside that, a plate over observations holding the shaded $y$, with an arrow from $\sigma$. Because device is in the model, the 30 countries are exchangeable within each device.
C. (Interview) "Why not just analyse every segment separately?"
"Separate analyses use only each segment's own data, so small segments get very noisy estimates and the most extreme results are mostly noise. Looking at many segments also means some will look significant by chance. A hierarchical model treats the segments as draws from one population: it learns how much segments really differ ($\tau$) from all of them together, and each segment's estimate borrows strength from the others in proportion to how little data it has. Large segments keep their own estimate; tiny ones are pulled toward the overall level."
D. Six segments, 25 orders each, $\sigma = 10$. Averages: 18, 21, 22, 24, 25, 28. Estimate $\tau$ roughly.
- Noise variance of an average: $\sigma^2/n = 100/25 = 4$.
- Mean of averages: $138/6 = 23$. Squared distances: $25, 4, 1, 1, 4, 25$; sum $60$; sample variance $60/5 = 12$.
- $\hat\tau^2 = 12 - 4 = 8$, $\hat\tau \approx 2.83$. With only 6 segments this is rough; a Bayesian fit would give a wide interval around it.
E. Suppose your segments are "new users" and "returning users" in 12 countries (24 segments). Returning users are known to convert about twice as often. Is a single population $\theta_g \sim N(\mu, \tau^2)$ for all 24 segments a good idea?
No: the segments are not exchangeable, because the user type is known to matter. A single population would pull every returning-user segment down and every new-user segment up toward one shared centre, and $\tau$ would be inflated by the type difference. Better: $\theta_g \sim N(\mu + \beta\,\text{returning}_g, \tau^2)$ on the logit scale, so segments are exchangeable given user type (possibly with a country level as well).
F. Overall conversion 4%, segment spread $\tau = 0.5$ on the logit scale. Find the range holding 95% of segments, and the rule-of-thumb typical difference in points.
- $\mu = \text{logit}(0.04) = \ln(0.04/0.96) \approx -3.178$; $\pm 1.96 \times 0.5 = \pm 0.98$: from $-4.158$ to $-2.198$.
- Back: $1/(1 + e^{4.158}) = 1/(1 + 63.9) \approx 0.015$; $\;1/(1 + e^{2.198}) = 1/(1 + 9.0) \approx 0.100$. So about 1.5% to 10.0%.
- Rule of thumb: relative spread $\tau(1 - p) = 0.5 \times 0.96 = 48\%$ of 4%, about 1.9 points. A logit SD of 0.5 is a lot of segment variation for a 4% metric; check that it matches what you believe.
Pooling and shrinkage: how groups borrow strength
A segment with 5 users shows a 40% conversion rate. A segment with 10 000 users shows 13%. Which one is really better? Your syllabus marks this chapter P0 ("must master") for exactly this reason: an experimentation framework reports results for segments with very little data. This chapter derives, step by step, how a hierarchical model combines each segment's own data with what all the other segments say, why small segments are pulled harder than large ones, and why this pulling ("shrinkage") makes the estimates more accurate on average, not less.
- Compare complete pooling, no pooling and partial pooling: what each one assumes and what it does to a small segment
- Derive the partial-pooling estimate as a precision-weighted average, with weight $w_g = \dfrac{n_g/\sigma^2}{n_g/\sigma^2 + 1/\tau^2}$ on the segment's own average
- Read a shrinkage plot: small groups shrink more, large groups less; the population is "worth $\sigma^2/\tau^2$ observations"
- Show that no pooling ($\tau \to \infty$) and complete pooling ($\tau \to 0$) are the two ends of one dial, and count the effective number of parameters
- See how $\mu$ and $\tau$ are learned from all groups, and why empirical Bayes can be overconfident, using the classic eight-schools example
- Prove by formula and by simulation that shrinkage lowers total error, and see who pays for it
- Apply all of it to conversion rates (Beta pseudo-counts) and avoid the winner's curse when ranking segments
What we need from earlier chapters: the two-level model, $\tau$ vs $\sigma$, exchangeability and how $\tau$ is learned (Chapter 6.5); posterior ∝ likelihood × prior (Chapter 6.1); pseudo-counts and the posterior mean as a weighted average (Chapter 6.3); bias, variance, MSE and the first look at shrinkage (Chapter 5.1); the law of total variance (Chapter 4.6). Notation: $\bar y_g$ = segment $g$'s observed average of $n_g$ observations; $\sigma$ = the spread of single observations, so $s_g = \sigma/\sqrt{n_g}$ is the noise (standard error) of $\bar y_g$; $\theta_g \sim N(\mu, \tau^2)$ = the population of true segment values. The precision of a quantity is 1 / its variance: a sharp, reliable quantity has high precision. $\hat\theta_g$ ("theta hat") = our estimate of $\theta_g$.
Complete pooling, no pooling, partial pooling core
You are choosing a restaurant. Restaurant A has one review: 5 stars. Restaurant B has 400 reviews averaging 4.6 stars. Nobody really believes A is a perfect 5.0; one happy customer proves little. But nobody believes A is exactly average either. Your gut says: "A is probably a bit above average, and B is almost surely about 4.6." That gut feeling is partial pooling.
"Pooling" means combining the data of different groups to estimate something. There are three attitudes:
- No pooling: "every group stands alone". Each segment's estimate is its own raw average. A gets 5.0.
- Complete pooling: "all groups are the same". Throw all data together; every segment gets the overall average. A and B both get the city average.
- Partial pooling: "a group is like the others until its own data prove otherwise". Each estimate sits between its own average and the overall centre, closer to its own average when it has a lot of data.
Three ways to say it:
- Picture: a dial from "each alone" to "all the same"; partial pooling sets the dial for each group from its amount of data.
- Numbers: a segment with one order of 58 dollars: no pooling says 58, complete pooling says 40, partial pooling says 41.8.
- Slogan: trust a group's own data in proportion to how much of it there is.
Four segments of an order-value metric, with single orders spreading $\sigma = 12$ dollars around their segment's value. Suppose the population of segments is known from hundreds of earlier segments: centre $\mu = 40$, spread $\tau = 4$ (Section 5 learns these from the data instead).
| Segment | orders $n_g$ | raw average $\bar y_g$ | no pooling | complete pooling | weight $w_g = \frac{n_g}{n_g + 9}$ | partial pooling |
|---|---|---|---|---|---|---|
| A | 1 | 58 | 58 | 40 | 0.1 | $40 + 0.1 \times 18 = 41.8$ |
| B | 9 | 46 | 46 | 40 | 0.5 | $40 + 0.5 \times 6 = 43.0$ |
| C | 36 | 38 | 38 | 40 | 0.8 | $40 + 0.8 \times (-2) = 38.4$ |
| D | 81 | 40 | 40 | 40 | 0.9 | $40 + 0.9 \times 0 = 40.0$ |
- Complete pooling: all $1 + 9 + 36 + 81 = 127$ orders together: $(1 \times 58 + 9 \times 46 + 36 \times 38 + 81 \times 40)/127 = (58 + 414 + 1368 + 3240)/127 = 5080/127 = 40$. Every segment gets 40.
- The weight on a segment's own average is $w_g = n_g/(n_g + 9)$, where $9 = \sigma^2/\tau^2 = 144/16$ (derived in the next section). A: $1/10 = 0.1$; B: $9/18 = 0.5$; C: $36/45 = 0.8$; D: $81/90 = 0.9$.
- Partial pooling: $\hat\theta_g = \mu + w_g(\bar y_g - \mu)$: start at the centre and move a fraction $w_g$ of the way toward the segment's own average.
- The ranking changes. No pooling ranks A first (58). Partial pooling ranks B first (43.0) and A second (41.8): one order of 58 dollars is weak evidence; nine orders averaging 46 are stronger.
With groups $g = 1, \dots, G$, observed averages $\bar y_g$ of $n_g$ observations each:
- Complete pooling: $\hat\theta_g = \bar y = \dfrac{\sum_h n_h \bar y_h}{\sum_h n_h}$ for every $g$. Model: one shared parameter, $\theta_g = \theta$ (the hierarchical model with $\tau = 0$).
- No pooling: $\hat\theta_g = \bar y_g$. Model: separate parameters with no link between them (the hierarchical model with $\tau \to \infty$).
- Partial pooling: $\hat\theta_g = w_g\,\bar y_g + (1 - w_g)\,\mu$ with $0 \lt w_g \lt 1$: the posterior mean of $\theta_g$ in the hierarchical model $\theta_g \sim N(\mu, \tau^2)$, $\bar y_g \mid \theta_g \sim N(\theta_g, \sigma^2/n_g)$. Groups share information through the population distribution.
Shrinkage is the movement from $\bar y_g$ toward $\mu$: $\bar y_g - \hat\theta_g = (1 - w_g)(\bar y_g - \mu)$. "The estimate is shrunk toward the centre."
Why do we need it?
Every per-segment report secretly picks one of the three. No pooling crowns tiny lucky segments as winners; complete pooling hides real segment differences. Partial pooling is the principled middle, and the data decide where in the middle each segment lands.
Where is it used?
Segment-level results in Bayesian A/B testing (yours), cold-start items and new stores in demand forecasting, small-area estimates such as state-level polls, "Bayesian average" ratings (IMDb's Top-250 weighted rating has exactly the form $w R + (1 - w) C$ with $w = v/(v + m)$), and credibility weighting of insurance premiums.
How is it used?
Fit the hierarchical model (in NumPyro: $\theta_g$ inside a plate with Normal(mu, tau), priors on $\mu$ and $\tau$). Report each segment's posterior mean and interval instead of its raw average. Read off $\tau$ to see how much the segments really differ.
"Partial pooling is a compromise number the analyst picks by taste."
The weight $w_g$ is computed from $n_g$, $\sigma$ and $\tau$, and $\tau$ is learned from the data. The amount of pooling is an output of the model, not a knob you set.
"No pooling is unbiased, so it is the honest, neutral choice."
It is unbiased for each segment, but its variance is huge for small segments, so its total error is far above partial pooling's (Section 6 proves it). Unbiased is not the goal; small error is (Chapter 5.1).
"Complete pooling is always wrong because segments differ."
If the segments truly differ very little ($\tau \approx 0$), complete pooling is nearly right. The hierarchical model finds that out by itself: with $\tau$ near 0, partial pooling becomes complete pooling.
"Partial pooling just averages each segment with the global mean."
It is a weighted average, and the weight differs by segment: segments with more data (or less noise) keep more of their own average. The weights come from the model, through $n_g$, $\sigma$ and the learned $\tau$.
Model answer: "Complete pooling uses one shared parameter, so all segments get the same estimate. No pooling fits each segment independently, so small segments get noisy estimates. Partial pooling puts the segment parameters in a hierarchical model, $\theta_g \sim N(\mu, \tau^2)$; each segment's posterior mean is a precision-weighted average of its own data and the population mean, so small segments borrow more strength from the population and large segments mostly keep their own estimate."
This is the P0 idea behind your A/B framework's hierarchical partial pooling. A segment with a handful of users and an extreme lift is, under no pooling, the "winner" of your report; under partial pooling it is pulled toward the overall lift until its own data are strong enough to resist. Large segments keep their own result. So the per-segment numbers a hierarchical model reports are pooled posterior summaries (usually means or medians), not raw segment averages, and they should be read that way.
Complete: $\hat\theta_g = \bar y$ (all the same, $\tau = 0$). None: $\hat\theta_g = \bar y_g$ ($\tau = \infty$). Partial: $\hat\theta_g = w_g \bar y_g + (1 - w_g)\mu$.
Example ($\sigma = 12$, $\tau = 4$, $\mu = 40$): $n = 1, 9, 36, 81$ → $w = 0.1, 0.5, 0.8, 0.9$.
Trap: the weight is not a taste choice and not the same for every segment.
Quick check: in the table, a fifth segment E has 81 orders averaging 58 dollars (like A's single order). What is its partial-pooling estimate?
$w = 81/90 = 0.9$, so $\hat\theta_E = 40 + 0.9 \times 18 = 56.2$. Same raw average as A (58), but A's estimate is 41.8: the amount of data, not the raw value, decides how far the estimate moves.
Deriving partial pooling: the precision-weighted average core
Two witnesses tell you a segment's value. The population says: "segments are usually around 40, give or take 4." The segment's own data say: "my average is 46, give or take 4" (the noise of an average of 9 orders with $\sigma = 12$ is $12/3 = 4$). Both are equally sure, so you go halfway: 43. If the segment had 81 orders, its "give or take" would be only 1.3, and you would believe it much more.
The rule is to weigh each witness by its precision, 1 / variance: a witness who is twice as sharp (half the standard deviation) gets four times the weight. It is how you would combine a good thermometer with a cheap one.
Three ways to say it:
- Picture: a tug-of-war: the population pulls toward $\mu$ with strength $1/\tau^2$, the data pull toward $\bar y_g$ with strength $n_g/\sigma^2$; the estimate lands where the pulls balance.
- Numbers: data precision $9/144 = 0.0625$, population precision $1/16 = 0.0625$: equal, so the estimate is halfway, 43.
- Slogan: precision is weight; the sharper witness wins.
Segment B: $n = 9$ orders, average $\bar y = 46$, order noise $\sigma = 12$; population $\mu = 40$, $\tau = 4$.
- Noise of the average: variance $\sigma^2/n = 144/9 = 16$, so data precision $= 1/16 = n/\sigma^2 = 9/144 = 0.0625$.
- Population precision $= 1/\tau^2 = 1/16 = 0.0625$.
- Posterior precision = sum $= 0.125$, so posterior variance $= 1/0.125 = 8$ and posterior SD $= \sqrt 8 \approx 2.83$ (smaller than both 4's).
- Weight on the data: $w = 0.0625/0.125 = 0.5$.
- Posterior mean: $0.5 \times 46 + 0.5 \times 40 = 43$.
Setting. Treat $\mu$, $\tau$, $\sigma$ as known for now. Two facts: the population says $\theta_g \sim N(\mu, \tau^2)$; and given $\theta_g$, the average of $n_g$ Normal observations is $\bar y_g \mid \theta_g \sim N(\theta_g, s_g^2)$ with $s_g^2 = \sigma^2/n_g$. (For Normal data the average carries everything the segment's data say about $\theta_g$, so we can work with $\bar y_g$ alone.)
Derivation (posterior ∝ likelihood × prior, then "complete the square"):
- $p(\theta_g \mid \bar y_g) \propto \exp\!\Big(-\dfrac{(\bar y_g - \theta_g)^2}{2 s_g^2}\Big)\exp\!\Big(-\dfrac{(\theta_g - \mu)^2}{2\tau^2}\Big)$.
- Expand the exponent and keep only terms with $\theta_g$ (the rest is a constant absorbed by "∝"): $-\tfrac12\Big[\theta_g^2\big(\tfrac{1}{s_g^2} + \tfrac{1}{\tau^2}\big) - 2\theta_g\big(\tfrac{\bar y_g}{s_g^2} + \tfrac{\mu}{\tau^2}\big)\Big]$.
- Name the pieces: $P = \dfrac{1}{s_g^2} + \dfrac{1}{\tau^2}$ and $m = \dfrac{\bar y_g/s_g^2 + \mu/\tau^2}{P}$. The bracket is $P\theta_g^2 - 2Pm\,\theta_g$.
- Complete the square: $P\theta_g^2 - 2Pm\,\theta_g = P(\theta_g - m)^2 - Pm^2$, and $Pm^2$ does not involve $\theta_g$. So the posterior is $\propto \exp\!\big(-\tfrac{P}{2}(\theta_g - m)^2\big)$: a Normal with mean $m$ and variance $1/P$.
Result.
$$\theta_g \mid \bar y_g \sim N\!\Big(w_g\,\bar y_g + (1 - w_g)\,\mu,\;\; \frac{1}{n_g/\sigma^2 + 1/\tau^2}\Big), \qquad w_g = \frac{n_g/\sigma^2}{n_g/\sigma^2 + 1/\tau^2}.$$- Precisions add: posterior precision = data precision $n_g/\sigma^2$ + population precision $1/\tau^2$.
- Equivalent forms of the weight: $w_g = \dfrac{\tau^2}{\tau^2 + \sigma^2/n_g} = \dfrac{n_g}{n_g + \sigma^2/\tau^2}$.
- Shrinkage factor $B_g = 1 - w_g = \dfrac{\sigma^2/n_g}{\sigma^2/n_g + \tau^2}$: exactly the noise share of the segment's average from Chapter 6.5. A segment is shrunk by the fraction of its average that is noise.
- The last form says the population acts like $k = \sigma^2/\tau^2$ extra observations sitting at $\mu$: $\hat\theta_g = \dfrac{n_g\bar y_g + k\mu}{n_g + k}$, the same pseudo-count idea as $\alpha + \beta$ in the Beta-Binomial (Chapter 6.3).
Why do we need it?
It answers "how much should a segment trust its own data?" exactly, and shows the answer depends on only three things: how much data the segment has, how noisy single observations are, and how different segments really are.
Where is it used?
Posterior means in every Normal hierarchical model, the Normal-Normal conjugate update, the update step of the Kalman filter, Bühlmann credibility in insurance ($Z = n/(n + k)$), the James–Stein estimator, and ridge regression, where an L2 penalty is a Normal prior (Chapter 5.3).
How is it used?
Compute $s_g^2 = \sigma^2/n_g$, then $w_g = \tau^2/(\tau^2 + s_g^2)$, then $w_g\bar y_g + (1-w_g)\mu$. In a fitted NumPyro model you do not code this: the posterior means of theta are this formula averaged over the uncertainty in $\mu$, $\tau$ and $\sigma$. Use it to sanity-check the fit by hand.
"The weight depends only on the sample size: a segment with 100 users always gets the same weight."
It depends on $n_g/\sigma^2$ compared with $1/\tau^2$. The same 100 users count for more when single observations are less noisy (small $\sigma$) or when segments are known to differ a lot (large $\tau$).
"Shrinkage only moves the estimate; the uncertainty stays the raw standard error."
Precisions add, so the posterior SD $1/\sqrt{n_g/\sigma^2 + 1/\tau^2}$ is smaller than the raw standard error $\sigma/\sqrt{n_g}$: borrowing strength also narrows the interval (in Section 5 we add back the uncertainty about $\mu$ and $\tau$ themselves).
"$N(\mu, \tau)$ in the code means variance $\tau$."
NumPyro's dist.Normal(mu, tau) takes the standard deviation. The formulas here use variances $\tau^2$ and $\sigma^2$; precision is $1/\tau^2$, not $1/\tau$.
When your A/B framework reports a segment's posterior mean, it is (for Normal-type metrics) this precision-weighted average, averaged over the posterior of the global parameters. A quick hand check you can do in an interview or a code review: take a segment's raw average, its standard error and the fitted $\tau$; compute $w = \tau^2/(\tau^2 + SE^2)$; the reported estimate should sit about a fraction $w$ of the way from the population mean to the raw average.
$\theta_g \mid \bar y_g \sim N\big(w_g\bar y_g + (1-w_g)\mu,\; 1/(n_g/\sigma^2 + 1/\tau^2)\big)$.
$w_g = \dfrac{n_g/\sigma^2}{n_g/\sigma^2 + 1/\tau^2} = \dfrac{\tau^2}{\tau^2 + \sigma^2/n_g} = \dfrac{n_g}{n_g + \sigma^2/\tau^2}$; precisions add; the population is worth $\sigma^2/\tau^2$ observations.
Trap: shrinkage factor $1 - w_g$ = noise share; NumPyro's Normal takes the SD, precision uses the variance.
Quick check: $\sigma = 10$, $\tau = 5$, a segment with $n = 4$ and $\bar y = 70$, $\mu = 50$. Estimate and posterior SD?
$\sigma^2/\tau^2 = 100/25 = 4$, so $w = 4/(4 + 4) = 0.5$ and $\hat\theta = 50 + 0.5 \times 20 = 60$. Posterior precision $= 4/100 + 1/25 = 0.04 + 0.04 = 0.08$, SD $= 1/\sqrt{0.08} \approx 3.54$ (raw SE $10/2 = 5$).
The shrinkage plot: small groups shrink more, large groups less core
Put every segment on one picture. On a top line, mark each segment's raw average. On a bottom line, mark its partial-pooling estimate. Join each pair with an arrow. This is a shrinkage plot, and it shows two rules at a glance:
- Every arrow points toward the centre $\mu$ and stops a fraction $1 - w_g$ of the way there.
- Small segments have small $w_g$, so they get long arrows; large segments get short arrows.
Three ways to say it:
- Picture: arrows from raw averages toward the centre; thin segments travel far, thick segments barely move.
- Numbers: two segments both 18 dollars above the centre: with 1 order the arrow is $0.9 \times 18 = 16.2$ long; with 81 orders it is $0.1 \times 18 = 1.8$ long.
- Slogan: shrinkage = noise share × distance from the centre.
The segments of Section 1 ($\mu = 40$, $\sigma = 12$, $\tau = 4$, so $w_g = n_g/(n_g + 9)$), plus segment E with 81 orders averaging 58.
- A ($n = 1$, 58): arrow $= (1 - 0.1) \times (58 - 40) = 0.9 \times 18 = 16.2$, lands at $58 - 16.2 = 41.8$.
- B ($n = 9$, 46): arrow $= 0.5 \times 6 = 3$, lands at $43$.
- C ($n = 36$, 38): arrow $= 0.2 \times (-2) = -0.4$, so it moves up by 0.4 to $38.4$, toward the centre from below.
- E ($n = 81$, 58): arrow $= 0.1 \times 18 = 1.8$, lands at $56.2$. Same raw value as A, very different estimate.
- No arrow crosses the centre, and two segments with the same $n$ keep their order. Segments with different $n$ can swap places (A and B did).
The movement of segment $g$ is
$$\bar y_g - \hat\theta_g = B_g\,(\bar y_g - \mu), \qquad B_g = 1 - w_g = \frac{\sigma^2/n_g}{\sigma^2/n_g + \tau^2} = \frac{\sigma^2}{\sigma^2 + n_g\tau^2}.$$- $B_g$ (the shrinkage factor) is between 0 and 1 and falls as $n_g$ grows: more data, less shrinkage.
- $\hat\theta_g$ always lies between $\bar y_g$ and $\mu$.
- The estimates are less spread out than the raw averages, and also less spread out than the true values: with equal $n$ and known $\mu$, $Var(\hat\theta_g) = w\,\tau^2 \lt \tau^2$. A histogram of shrunk estimates therefore understates how different segments really are; for that, look at $\tau$ or at posterior draws.
Why do we need it?
It is the quickest way to see what pooling did to every segment at once, and to spot trouble: a large segment that moved a lot, or a cluster of segments all pulled the same way, suggests the population model or exchangeability is wrong.
Where is it used?
Multilevel model reports (Gelman and Hill's books, brms, lme4, NumPyro notebooks), the classic Efron–Morris baseball batting-average study, per-segment dashboards in experimentation platforms, and per-store forecast reconciliation.
How is it used?
For each segment plot the raw average (top) and the posterior mean (bottom) with an arrow, and size the points by $n_g$. Check that small segments move most. Investigate any big segment that moves far, and any pattern that lines up with a known covariate.
"Shrinkage moves every segment by the same amount."
Each segment moves a fraction $B_g = \sigma^2/(\sigma^2 + n_g\tau^2)$ of its distance to the centre. Small segments move a large fraction, large segments a tiny one.
"The segment furthest from the centre is always shrunk the most."
In dollars it moves $B_g \times$ distance, so distance matters, but the fraction is set by its size. A large segment far from the centre (E above) hardly moves; a tiny segment near the centre still moves most of the way.
"The spread of the shrunk estimates shows how different the segments really are."
Shrunk estimates are less spread out than the true values (variance $w\tau^2$ instead of $\tau^2$). Read the segment-to-segment variation from $\tau$, or from posterior draws, not from a histogram of posterior means.
A shrinkage plot of segment lifts is a strong artefact for an A/B framework like yours: raw segment lifts on top, pooled lifts below, points sized by users. It explains to stakeholders in one picture why the "+30% in a 40-user segment" did not survive, and it lets you check that large segments were left alone. A large segment that moved a lot is a signal to look for a missing covariate (exchangeability, Chapter 6.5).
Movement $= B_g(\bar y_g - \mu)$ with $B_g = \sigma^2/(\sigma^2 + n_g\tau^2) = 1 - w_g$: falls as $n_g$ grows.
Estimates stay between $\bar y_g$ and $\mu$; same $n$ keeps order; different $n$ can reorder.
Trap: shrunk estimates understate the true spread; use $\tau$ or posterior draws.
Quick check: $\sigma = 12$, $\tau = 4$, $\mu = 40$. Segment F has 3 orders averaging 22. Where does it land, and how far did it move?
$B = 144/(144 + 3 \times 16) = 144/192 = 0.75$. Movement $= 0.75 \times (22 - 40) = -13.5$, so it moves up by 13.5 to $22 + 13.5 = 35.5$. (Check: $w = 0.25$, $40 + 0.25 \times (-18) = 35.5$.)
The two ends of the dial: $\tau \to 0$ and $\tau \to \infty$ core
Complete pooling and no pooling are not rivals of the hierarchical model. They are two of its settings. The dial is $\tau$, "how different are segments really?"
- Turn $\tau$ down to 0: segments are identical, so every segment's estimate becomes the shared centre: complete pooling.
- Turn $\tau$ up toward infinity: segments have nothing in common, so each keeps its own average: no pooling.
The amount of data works the same way: a segment with endless data needs no help ($w \to 1$), and a segment with no data gets the centre ($w \to 0$). So the question "should I pool?" turns into "what is $\tau$?", and the data answer it.
Three ways to say it:
- Picture: one dial from "all the same" to "each alone"; the data set it.
- Numbers: a segment of 9 orders with $\sigma = 12$: $w = 0$ at $\tau = 0$, $0.2$ at $\tau = 2$, $0.5$ at $\tau = 4$, $0.9$ at $\tau = 12$, $0.99$ at $\tau = 40$.
- Slogan: no pooling and complete pooling are partial pooling with the dial turned all the way.
The four segments of Section 1 ($n = 1, 9, 36, 81$; averages $58, 46, 38, 40$; $\sigma = 12$), now with the centre $\hat\mu$ estimated from them using weights $1/(\tau^2 + \sigma^2/n_g)$.
- $w = \dfrac{9}{9 + 144/\tau^2}$ for a 9-order segment: $\tau = 2$: $9/(9 + 36) = 0.2$; $\tau = 4$: $9/18 = 0.5$; $\tau = 12$: $9/10 = 0.9$; $\tau = 40$: $9/9.09 \approx 0.99$.
- $\tau \to 0$: all weights $1/(\sigma^2/n_g) \propto n_g$, so $\hat\mu$ is the grand mean of all orders, $5080/127 = 40$, and every segment gets 40.
- $\tau \to \infty$: all weights become equal, so $\hat\mu$ is the plain average of the four segment averages, $(58 + 46 + 38 + 40)/4 = 45.5$; and every segment keeps its own average.
- $\tau = 4$: $\hat\mu \approx 41.4$, and the estimates sit in between.
- Effective number of parameters at $\tau = 4$: about $2.6$ for these 4 segments; it is $1$ at $\tau = 0$ and $4$ at $\tau = \infty$.
With $w_g = \tau^2/(\tau^2 + \sigma^2/n_g)$:
$$\tau \to 0:\; w_g \to 0,\; \hat\theta_g \to \hat\mu = \frac{\sum n_g\bar y_g}{\sum n_g}\;\text{(complete pooling)}; \qquad \tau \to \infty:\; w_g \to 1,\; \hat\theta_g \to \bar y_g\;\text{(no pooling)}.$$At a fixed $\tau \gt 0$: $n_g \to \infty$ gives $w_g \to 1$, and $n_g \to 0$ gives $w_g \to 0$.
Effective number of parameters (how many independent numbers the estimates really use): $p_{\text{eff}} = \sum_g \dfrac{\partial\hat\theta_g}{\partial\bar y_g} = \sum_g w_g + \sum_g (1 - w_g)\dfrac{v_g}{\sum_h v_h}$, with $v_g = 1/(\tau^2 + \sigma^2/n_g)$. It runs from 1 ($\tau = 0$: one shared number) to $G$ ($\tau \to \infty$: one number per group). This is the same "trace of the smoother" idea as the degrees of freedom of ridge regression.
Why do we need it?
It removes the false choice between "analyse segments separately" and "ignore segments". Both are special cases with an extreme, usually unjustified assumption about $\tau$; the hierarchical model lets the data place the dial.
Where is it used?
Model comparison of pooled vs unpooled vs multilevel regressions, the effective-parameter counts behind information criteria such as DIC and WAIC, ridge degrees of freedom, and the heterogeneity question in A/B testing ("do treatment effects differ by segment?").
How is it used?
After fitting, look at the posterior of $\tau$. Mass near 0 means "segments barely differ: report the pooled result". Large $\tau$ means "segments really differ: the segment estimates carry their own information". Report $p_{\text{eff}}$ or the weights to show how much each segment relied on itself.
"When unsure, analysing segments separately is the safe, assumption-free choice."
No pooling is the hierarchical model with $\tau = \infty$: the strong claim that segments have nothing in common. It is the most extreme assumption on the dial, and it gives the noisiest small-segment estimates.
"As $\tau \to \infty$ the centre $\hat\mu$ becomes the grand mean of all observations."
At $\tau = 0$ each observation counts equally (grand mean). As $\tau \to \infty$ each segment counts equally (plain average of the segment averages). In the example: 40 vs 45.5.
"The hierarchical model has $G + 3$ parameters, so it is the most complex option."
Its effective number of parameters sits between 1 and $G$ (2.6 of 4 in the example), and the data choose where.
In your A/B framework, "do treatment effects differ by segment?" is a question about the between-segment spread of the lifts, $\tau_\delta$. A posterior for $\tau_\delta$ piled near 0 says the segments share one lift (report the overall result; segment "differences" are noise). A posterior for $\tau_\delta$ well away from 0 says the lift really varies, and the large segments' estimates carry their own information. Look at that posterior before presenting any segment breakdown.
$\tau \to 0$: $w_g \to 0$, complete pooling (centre = grand mean). $\tau \to \infty$: $w_g \to 1$, no pooling (centre = plain average of segment averages).
$n_g \to \infty$: $w_g \to 1$; $n_g \to 0$: $w_g \to 0$.
$p_{\text{eff}} = \sum_g \partial\hat\theta_g/\partial\bar y_g \in [1, G]$. Trap: "no pooling" is not assumption-free; it assumes $\tau = \infty$.
Quick check: $\sigma = 20$. How large must $\tau$ be before a 25-order segment puts at least half its weight on its own average?
$w = \tau^2/(\tau^2 + \sigma^2/n) \ge 0.5$ when $\tau^2 \ge \sigma^2/n = 400/25 = 16$, so $\tau \ge 4$.
Learning $\mu$ and $\tau$ from all groups: the eight-schools example core
So far someone handed us $\mu$ and $\tau$. In practice the model learns them from all the groups together: $\mu$ from where the group averages sit (precise groups count more), and $\tau$ from how much the averages spread beyond their noise (Chapter 6.5). A surprising consequence: every group's estimate depends on every other group's data. Add one wildly different group and $\hat\tau$ grows, so every group is shrunk less.
There are two ways to use what the data say about $\tau$:
- Empirical Bayes: find the single best value $\hat\tau$, then plug it into the formula as if it were known exactly.
- Full Bayes: keep the whole posterior of $\tau$. For every plausible $\tau$ compute the shrunk estimates, then average them, weighting each $\tau$ by how plausible it is. This is what NUTS or SVI do for you in NumPyro.
The classic eight-schools example shows why the difference matters: the single best $\tau$ is exactly 0, so empirical Bayes declares all eight schools identical. But $\tau = 10$ is also quite plausible, and full Bayes keeps that possibility alive.
Three ways to say it:
- Picture: a fan of lines, one per school, showing its estimate for every $\tau$; full Bayes averages along the fan, weighted by the posterior of $\tau$.
- Numbers: school A: empirical Bayes 7.7 ± 4.1; full Bayes 11.4 ± 8.3.
- Slogan: do not pretend you know $\tau$.
Eight schools (Rubin, 1981; a famous example in Gelman et al., Bayesian Data Analysis). Eight high schools each ran their own experiment on an SAT coaching programme. Each reported an estimated effect $y_j$ (in SAT points) and its standard error $\sigma_j$, treated as known.
| School | A | B | C | D | E | F | G | H |
|---|---|---|---|---|---|---|---|---|
| effect $y_j$ | 28 | 8 | −3 | 7 | −1 | 1 | 18 | 12 |
| std. error $\sigma_j$ | 15 | 10 | 16 | 11 | 9 | 11 | 10 | 18 |
- No pooling: school A's effect is $28 \pm 15$. But the standard errors are large: all eight results are noisy.
- Complete pooling: weight each school by its precision $1/\sigma_j^2$: $\hat\mu = 7.7$ with standard error $4.1$.
- Are the schools different? A chi-square test of "all effects equal" gives $\chi^2 = 4.7$ on 7 degrees of freedom ($p \approx 0.70$): the data are consistent with equal effects. Consistent with is not the same as proven.
- Posterior of $\tau$ (flat priors on $\mu$ and $\tau$, as in the book): most probable value 0, median 5.2, 90% interval about 0.5 to 17.
- Empirical Bayes plugs in $\hat\tau = 0$: every school gets $7.7 \pm 4.1$. Full Bayes averages over $\tau$: school A gets $11.4$ with SD $8.3$; the eight posterior means range from 5.1 (E) to 11.4 (A). Pulled strongly toward the centre, but not all the way, and with honest uncertainty.
With noise SDs $s_j$ ($= \sigma/\sqrt{n_j}$ for segment averages, or a reported standard error) and $v_j = 1/(\tau^2 + s_j^2)$:
- Centre for a given $\tau$: $\hat\mu(\tau) = \dfrac{\sum_j v_j\,y_j}{\sum_j v_j}$, with uncertainty $V_\mu(\tau) = 1/\sum_j v_j$.
- Empirical Bayes (also "type-II maximum likelihood"): choose $\hat\tau$ to maximize the marginal likelihood $p(y \mid \tau)$ (the probability of the data with every $\theta_j$, and $\mu$, averaged out), then report $E[\theta_j \mid \hat\tau, y]$.
- Full Bayes: $E[\theta_j \mid y] = \int E[\theta_j \mid \tau, y]\,p(\tau \mid y)\,d\tau$ with $E[\theta_j \mid \tau, y] = w_j(\tau)\,y_j + (1 - w_j(\tau))\,\hat\mu(\tau)$; and by the law of total variance, $Var(\theta_j \mid y) = E\big[Var(\theta_j \mid \tau, y)\big] + Var\big(E[\theta_j \mid \tau, y]\big)$, where $Var(\theta_j \mid \tau, y) = w_j s_j^2 + (1 - w_j)^2 V_\mu$.
- The second term is the extra uncertainty from not knowing $\tau$; empirical Bayes drops it, so its intervals are too narrow, badly so when $\hat\tau$ sits at 0 and there are few groups.
Why do we need it?
The amount of pooling must come from the data, and its uncertainty must reach every segment's interval. Otherwise, with few segments, you report confident "all segments are the same" answers that the data do not support.
Where is it used?
Eight schools is the standard test case in Stan, NumPyro and PyMC. Empirical Bayes shines with thousands of groups (ad click-through rates, gene-expression tools such as limma, sports statistics), where $\tau$ is pinned down well. Full Bayes is the default for tens of segments.
How is it used?
With many groups, empirical Bayes and full Bayes agree and empirical Bayes is cheaper. With few groups (a rough rule of thumb: under about 20), fit $\mu$ and $\tau$ with priors (NUTS or SVI), report the posterior of $\tau$, and check the answer with one or two other hyperpriors.
"The chi-square test did not reject equal effects, so the schools are identical and complete pooling is right."
Failing to reject is not proof (absence of evidence is not evidence of absence). The posterior of $\tau$ is wide: 0 is the single most likely value, but values above 10 are well supported too. Full Bayes keeps both.
"Empirical Bayes is Bayesian, so its intervals are fine."
It treats $\hat\tau$ as known and drops the uncertainty about it. With few groups that makes intervals too narrow; at $\hat\tau = 0$ it reports every group with the precision of the pooled mean.
"Each segment's estimate depends only on its own data and the global mean."
The global mean and $\tau$ are learned from all segments, so every segment's estimate depends on all the others. Move one segment and all the estimates move (try it in the shrinkage-arrows widget with $\tau$ learned).
An A/B framework like yours fits $\mu$ and $\tau$ inside the model (with SVI), so it is full Bayes over the global parameters, up to the accuracy of the variational approximation. Two practical consequences: with only a handful of segments, the hyperprior on $\tau$ matters and deserves a sensitivity check (Chapter 6.8); and a mean-field guide can understate the uncertainty in $\tau$ and its link to the segment effects, which narrows segment intervals much like empirical Bayes does (Chapters 6.11 and 6.13). Eight schools is also the standard example of the funnel-shaped posterior that makes these models hard to fit (Chapter 6.7).
$\hat\mu(\tau) = \sum v_j y_j/\sum v_j$, $v_j = 1/(\tau^2 + s_j^2)$; $\tau$ from the spread of the $y_j$ beyond noise.
Empirical Bayes: plug in $\hat\tau$ (too narrow with few groups). Full Bayes: average $E[\theta_j \mid \tau, y]$ over $p(\tau \mid y)$; variance gains $Var(E[\theta_j \mid \tau, y])$.
Eight schools: pooled 7.7 ± 4.1; $\tau$ mode 0, median 5.2; school A full Bayes 11.4 ± 8.3.
Quick check: why does adding a ninth school with $y = 60$, $\sigma = 10$ change school A's estimate, even though A's own data did not change?
The new school is far from the others, so the spread of the results beyond their noise grows and the posterior of $\tau$ moves to larger values. A larger $\tau$ means larger weights $w_j$ on every school's own result, so A is shrunk less (its estimate moves toward 28). The centre $\hat\mu$ also moves up a little.
Why shrinkage lowers the total error, and who pays for it core
Shrinking looks like cheating: we deliberately move every estimate away from its own data. Why would that make estimates better? Recall from Chapter 5.1: the average squared error of an estimate (its MSE) is bias² + variance. A small segment's raw average has no bias but a huge variance: it jumps around with every new sample. Pulling it toward the centre adds a little bias but removes a lot of that jumping. For small segments the trade is excellent; for large segments it barely changes anything.
There is a price. A segment that truly is far from the centre is pulled too far in, on average. So partial pooling does not make every segment better; it makes the total error smaller. A few unusual segments lose a little; most segments gain a lot.
Three ways to say it:
- Picture: darts. Raw estimates of small segments land all over the board; shrunk estimates land in a tighter group, slightly off-centre for the unusual segments.
- Numbers: a 9-order segment ($\sigma = 12$, $\tau = 4$): raw error 16, partial pooling error 8. Half.
- Slogan: a little bias buys a big cut in noise.
A segment with $n = 9$ orders, $\sigma = 12$, so its raw average has noise variance $s^2 = 144/9 = 16$. Population $\mu = 40$, $\tau = 4$, so $w = 0.5$.
- No pooling: $\bar y - \theta$ is pure noise: MSE $= s^2 = 16$, bias 0.
- Partial pooling: $\hat\theta - \theta = w(\bar y - \theta) + (1 - w)(\mu - \theta)$. The noise part has variance $w^2 s^2 = 0.25 \times 16 = 4$. The bias part $(1 - w)(\mu - \theta)$ depends on how far this segment truly is from the centre.
- Averaged over segments ($\theta$ drawn from $N(40, 16)$): the bias part averages $(1 - w)^2\tau^2 = 0.25 \times 16 = 4$. Total $4 + 4 = 8 = w\,s^2$. Half of the raw error.
- An unusual segment 12 dollars from the centre (3τ): $4 + 0.25 \times 144 = 40$, worse than 16. This segment pays.
- Break-even: partial pooling loses only when $(1 - w)^2 d^2 \gt (1 - w^2)s^2$, which simplifies to $d^2 \gt 2\tau^2 + s^2 = 32 + 16 = 48$, i.e. $|d| \gt 6.93 = 1.73\tau$. Only 8.3% of a $N(\mu, \tau^2)$ population is that far out; the other 91.7% gain.
- Eight segments of 2 to 250 orders: total MSE $144.6$ with no pooling, $46.1$ with partial pooling: a 68% cut, almost all of it from the small segments.
For a segment whose true value is $\theta$, at distance $d = \theta - \mu$ from the centre, with $s^2 = \sigma^2/n$ ($\mu$, $\tau$, $\sigma$ known):
$$\text{MSE}_{\text{none}} = s^2, \qquad \text{MSE}_{\text{partial}}(d) = \underbrace{w^2 s^2}_{\text{variance}} + \underbrace{(1 - w)^2 d^2}_{\text{bias}^2}.$$- Averaged over the population ($d \sim N(0, \tau^2)$): $\text{MSE}_{\text{partial}} = w\,s^2 = \dfrac{1}{n/\sigma^2 + 1/\tau^2} \lt s^2$. The gain is $(1 - w)\,s^2$: biggest when $s^2$ (noise) is large compared with $\tau^2$.
- Who loses: segments with $d^2 \gt 2\tau^2 + s^2$, a minority (at most about 16% of the population, fewer when $n$ is small).
- Stein's paradox (Chapter 5.1): when estimating 3 or more Normal means with known noise, a James–Stein estimator that shrinks by a data-chosen amount has lower total expected squared error than the raw averages for every set of true means, with no population assumption at all. The hierarchical model is the Bayesian version of the same effect.
- When $\mu$ and $\tau$ are learned from the data rather than known, most of the gain remains (the next widget shows it).
Why do we need it?
"You moved my segment's number away from its data" is the objection every stakeholder raises. This is the answer: across segments, shrunk estimates are closer to the truth, by a lot, and the formula says exactly which segments could lose.
Where is it used?
The justification for hierarchical models, ridge regression, empirical-Bayes rankings in sports and ads, small-area estimation by statistics offices, and the bias–variance view of every regularizer, including the Laplace prior on your forecasting model's changepoint slopes.
How is it used?
Simulate from a fitted hierarchical model (or a fake-data version, Chapter 6.5) and compare total squared error of raw versus pooled estimates. Report pooled estimates; flag segments whose raw value is far outside the population as "possibly unusual, check for a missing covariate".
"Partial pooling makes every segment's estimate more accurate."
It makes the total (or average) error smaller. Segments that truly sit far from the centre (beyond $\sqrt{2\tau^2 + s^2}$) are worse off on average, and in any single sample a sizeable share of the dots (often 20% to 40%) can be worse just by luck.
"Shrinkage adds bias, so it should be avoided in rigorous work."
Bias is one part of the error; the MSE is bias² + variance. Accepting a small bias to remove a large variance is what ridge regression, Laplace priors and hierarchical models all do on purpose.
"The gains only hold if the segment values are really Normal around μ."
The exact numbers use the Normal model, but Stein's result shows that shrinking 3 or more means by a data-chosen amount lowers total expected squared error for any set of true means (with known Normal noise). A badly wrong population model (for example, a missing covariate) does reduce the gain, which is why exchangeability matters.
"Shrinkage is a fudge that makes small segments look average."
Shrinkage is the posterior mean under a model in which segments come from a common population; it minimizes expected squared error under that model, and it pulls each segment in proportion to how much of its raw average is noise.
Model answer: "MSE equals bias squared plus variance. A small segment's raw average is unbiased but very noisy. Pulling it toward the population mean adds a little bias and removes most of the noise. With $w = \tau^2/(\tau^2 + \sigma^2/n)$, the average MSE drops from $\sigma^2/n$ to $w\,\sigma^2/n$. The segments that can lose are the rare ones truly far from the mean, and across all segments the total error is much lower."
This is why partial pooling is P0 for your experimentation framework: segments with little data produce the noisiest raw lifts, and those are exactly the ones that get noticed. Shrinkage makes the segment estimates more accurate overall and turns "a segment with 40 users shows +30%" into an honest "probably close to the overall lift, uncertain". The same bias–variance logic runs your forecasting model's Laplace prior on changepoint slopes: it shrinks most slope changes toward zero to cut variance, at the cost of under-reacting to the occasional real, large trend change.
$\text{MSE}_{\text{partial}}(d) = w^2 s^2 + (1 - w)^2 d^2$; averaged over $d \sim N(0, \tau^2)$: $w s^2 \lt s^2$ (gain $(1-w)s^2$).
Losers: $|d| \gt \sqrt{2\tau^2 + s^2}$ (a minority). Total error falls; most of the gain is in small segments.
Trap: not every segment improves; the claim is about total/average error.
Quick check: $n = 1$, $\sigma = 12$, $\tau = 4$. What are the average MSEs of no pooling and partial pooling, and how far out must a segment be to lose?
$s^2 = 144$, $w = 16/(16 + 144) = 0.1$. No pooling: 144. Partial: $w s^2 = 14.4$, a 90% cut. Break-even $|d| = \sqrt{2 \times 16 + 144} = \sqrt{176} \approx 13.3 = 3.3\tau$: only about 0.1% of segments are that far out.
Partial pooling for conversion rates, and the winner's curse
For conversion rates the same idea comes with pseudo-counts, exactly as in Chapter 6.3. If segment rates come from a population $\text{Beta}(\kappa\phi, \kappa(1-\phi))$, each segment's estimate is its own data plus a "starter pack" of $\kappa$ imaginary visitors converting at the overall rate $\phi$. Big segments drown the starter pack; tiny segments are mostly starter pack.
This matters most when you rank segments. If you crown the segment with the best raw rate, you almost always crown a small, lucky one, and its raw rate overstates its true rate. This is the winner's curse: picking the maximum of noisy numbers picks the noise.
Three ways to say it:
- Picture: every segment gets the same starter pack of imaginary visitors; it only changes the small segments.
- Numbers: 2 of 5 visitors (40%) with a starter pack of 10 of 100 becomes $12/105 = 11.4\%$; 1 300 of 10 000 (13%) becomes $1310/10100 = 12.97\%$.
- Slogan: the best-looking segment is usually the luckiest small one.
Overall rate $\phi = 10\%$; the population is "worth" $\kappa = 100$ visitors, so the starter pack is 10 conversions among 100 visitors.
| Segment | conversions / visitors | raw rate | weight $\frac{n}{n + 100}$ | pooled $\frac{k + 10}{n + 100}$ |
|---|---|---|---|---|
| S1 | 2 / 5 | 40.0% | 0.048 | $12/105 = 11.4\%$ |
| S2 | 30 / 200 | 15.0% | 0.667 | $40/300 = 13.3\%$ |
| S3 | 1300 / 10000 | 13.0% | 0.990 | $1310/10100 = 12.97\%$ |
| S4 | 0 / 20 | 0.0% | 0.167 | $10/120 = 8.3\%$ |
- Raw ranking: S1 (40%) > S2 (15%) > S3 (13%) > S4 (0%).
- Pooled ranking: S2 (13.3%) > S3 (12.97%) > S1 (11.4%) > S4 (8.3%). The 5-visitor "winner" drops to third.
- The pooled estimate is a weighted average: S2 $= 0.667 \times 15\% + 0.333 \times 10\% = 13.3\%$, the same form as the Normal case with $\sigma^2/\tau^2$ replaced by $\kappa$.
- In a simulation with 15 segments of 20 to 5 000 visitors whose true rates vary around 10% ($\text{Beta}(20, 180)$), the top raw segment was one of the 5 smallest in 67% of runs, and its raw rate overstated its true rate by 5.3 points on average. The top pooled segment overstated its truth by about 0 points, and it was also truly better (13.1% vs 12.0% on average).
Beta-Binomial hierarchy: $p_g \sim \text{Beta}(\kappa\phi, \kappa(1-\phi))$, $\;k_g \mid p_g \sim \text{Binomial}(n_g, p_g)$. Given $\phi$ and $\kappa$, each segment's posterior is $\text{Beta}(\kappa\phi + k_g,\ \kappa(1-\phi) + n_g - k_g)$, with mean
$$\hat p_g = \frac{k_g + \kappa\phi}{n_g + \kappa} = w_g\,\frac{k_g}{n_g} + (1 - w_g)\,\phi, \qquad w_g = \frac{n_g}{n_g + \kappa}.$$- $\kappa$ plays the role of $\sigma^2/\tau^2$: large $\kappa$ = segments very similar = strong pooling. It is learned from how much the segment rates spread beyond Binomial noise (empirical Bayes: maximize the product of beta-binomial probabilities; full Bayes: hyperpriors on $\phi$ and $\kappa$).
- Winner's curse (a selection bias): the largest of many noisy estimates is, on average, too high, and the noisiest (smallest) segments win most often. If the population model is right, posterior means already account for that luck, so selecting on them does not systematically overstate.
Why do we need it?
Segment breakdowns, "where did the treatment work best?", top-k lists and leaderboards all rank noisy rates. Raw rates of very unequal segments put the luckiest tiny segments on top and overstate them.
Where is it used?
Segment analysis in A/B testing, click-through-rate estimates for new ads and items (cold start), "best-performing store" reports, batting averages early in a season, and Bayesian-average product ratings.
How is it used?
Estimate $\phi$ and $\kappa$ (or fit the hierarchical model), report $\hat p_g = (k_g + \kappa\phi)/(n_g + \kappa)$ with its Beta interval, and rank by pooled estimates or by the posterior probability of being best. Never rank raw rates of segments whose sizes differ by orders of magnitude.
"Adding +1 conversion and +2 visitors (Laplace smoothing) is the same as partial pooling."
That is a fixed, tiny starter pack (a Beta(1, 1) prior worth 2 visitors) centred on 50%. Partial pooling centres the starter pack on the overall rate $\phi$ and learns its size $\kappa$ from how much the segments really differ, often hundreds of visitors.
"The segment with the highest raw rate is the best segment."
With unequal sizes it is usually the luckiest small segment, and its raw rate overstates its truth (winner's curse). Rank by pooled estimates or by the posterior probability of being best.
"Pooled estimates are biased, so the top pooled segment must be overstated too."
Each pooled estimate is pulled toward $\phi$ by just the amount the luck deserves under the model. If the population model is right, the top pooled segment is not systematically overstated (about 0 points in the simulation).
Your A/B framework's Beta-Binomial models are exactly this setting, one level up: with segments, the segment rates (or segment lifts) share a population, and partial pooling is the defence against "the treatment worked best in a tiny segment". When someone asks "which segment did the treatment help most?", answer with pooled lifts and with $P(\text{lift}_g \gt 0 \mid D)$ from the hierarchical model, not with the largest raw lift. If your framework uses a logit-Normal population instead of a Beta one (Chapter 6.5), the same pseudo-count intuition holds approximately.
$p_g \sim \text{Beta}(\kappa\phi, \kappa(1-\phi))$ ⇒ $\hat p_g = \dfrac{k_g + \kappa\phi}{n_g + \kappa} = w_g\frac{k_g}{n_g} + (1-w_g)\phi$, $w_g = \dfrac{n_g}{n_g + \kappa}$ ($\kappa \leftrightarrow \sigma^2/\tau^2$).
2/5 → 11.4%, 30/200 → 13.3%, 1300/10000 → 12.97% with $\phi = 10\%$, $\kappa = 100$.
Trap: winner's curse: the top raw segment is usually a small lucky one and is overstated; rank pooled estimates.
Quick check: $\phi = 5\%$, $\kappa = 400$. A segment shows 6 conversions among 40 visitors (15%). What is the pooled estimate, and how much weight does the segment's own rate get?
Starter pack: $400 \times 0.05 = 20$ conversions among 400 visitors. Pooled $= (6 + 20)/(40 + 400) = 26/440 \approx 5.9\%$. Weight on its own rate $= 40/440 \approx 0.09$. Forty visitors are little evidence against a population that is worth 400.
Recap, cheat sheet and practice
- Complete pooling (one value for all, $\tau = 0$), no pooling (each group alone, $\tau = \infty$) and partial pooling (groups share a population) are three settings of one model.
- Partial pooling = posterior mean = precision-weighted average: $\hat\theta_g = w_g\bar y_g + (1-w_g)\mu$, $w_g = \frac{n_g/\sigma^2}{n_g/\sigma^2 + 1/\tau^2} = \frac{n_g}{n_g + \sigma^2/\tau^2}$; precisions add.
- The population is worth $\sigma^2/\tau^2$ observations (for rates: $\kappa$ visitors). Small groups shrink more; shrinkage = noise share × distance to the centre.
- $\mu$ and $\tau$ are learned from all groups, so every estimate depends on every group. Empirical Bayes plugs in $\hat\tau$ (overconfident with few groups); full Bayes averages over $\tau$ (eight schools: 7.7 ± 4.1 vs 11.4 ± 8.3 for school A).
- Shrinkage lowers total error: average MSE $w\,\sigma^2/n$ instead of $\sigma^2/n$. Only groups truly beyond $\sqrt{2\tau^2 + \sigma^2/n}$ lose on average.
- Ranking raw rates triggers the winner's curse; rank pooled estimates.
Cheat sheet
| Idea | Formula | Plain words / remember |
|---|---|---|
| Complete pooling | $\hat\theta_g = \sum n_h\bar y_h/\sum n_h$ | everyone the same; $\tau = 0$ |
| No pooling | $\hat\theta_g = \bar y_g$ | each alone; $\tau = \infty$; unbiased, noisy |
| Partial pooling | $w_g\bar y_g + (1 - w_g)\mu$ | posterior mean of the hierarchical model |
| Weight | $\frac{n/\sigma^2}{n/\sigma^2 + 1/\tau^2} = \frac{\tau^2}{\tau^2 + \sigma^2/n} = \frac{n}{n + \sigma^2/\tau^2}$ | precision is weight |
| Posterior variance | $1/(n/\sigma^2 + 1/\tau^2) = w\,\sigma^2/n$ | precisions add; narrower than the raw SE |
| Shrinkage factor | $B = 1 - w = \sigma^2/(\sigma^2 + n\tau^2)$ | = noise share of the raw average |
| Centre | $\hat\mu = \sum v_g\bar y_g/\sum v_g$, $v_g = 1/(\tau^2 + \sigma^2/n_g)$ | grand mean at $\tau = 0$; plain average at $\tau = \infty$ |
| Effective parameters | $\sum w_g + \sum(1-w_g)v_g/\sum v$ | from 1 to $G$ |
| Error | $w^2 s^2 + (1-w)^2 d^2$; average $w s^2$ | losers: $|d| \gt \sqrt{2\tau^2 + s^2}$ |
| Full Bayes variance | $E[Var(\theta \mid \tau)] + Var(E[\theta \mid \tau])$ | EB drops the second term |
| Rates | $(k + \kappa\phi)/(n + \kappa)$, $w = n/(n + \kappa)$ | starter pack of κ visitors at rate φ |
import numpy as np
import jax
import numpyro
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS
# 1) The partial-pooling formula: a precision-weighted average
def partial_pool(ybar, n, sigma, mu, tau):
data_prec = n / sigma**2 # precision of the group's own average
prior_prec = 1 / tau**2 # precision of the population N(mu, tau^2)
w = data_prec / (data_prec + prior_prec) # weight on the group's own average = n / (n + sigma^2/tau^2)
return w, w * ybar + (1 - w) * mu, np.sqrt(1 / (data_prec + prior_prec))
n = np.array([1, 9, 36, 81]); ybar = np.array([58.0, 46.0, 38.0, 40.0])
w, est, sd = partial_pool(ybar, n, sigma=12.0, mu=40.0, tau=4.0)
print(w, est, np.round(sd, 2)) # [0.1 0.5 0.8 0.9] [41.8 43. 38.4 40. ] [3.79 2.83 1.79 1.26]
print((n * ybar).sum() / n.sum()) # 40.0 complete pooling: every segment gets the grand mean
# 2) Why partial pooling lowers total error: simulate many worlds (mu, tau, sigma known)
rng = np.random.default_rng(1)
sizes = np.array([2, 4, 8, 15, 30, 60, 120, 250]); R = 20000
theta = rng.normal(40, 4, (R, 8)) # true segment values
yb = rng.normal(theta, 12 / np.sqrt(sizes)) # observed segment averages
w, part, _ = partial_pool(yb, sizes, 12.0, 40.0, 4.0)
comp = (yb * sizes).sum(1, keepdims=True) / sizes.sum() # complete pooling (grand mean)
for name, e in [("no pooling", yb), ("complete", comp), ("partial", part)]:
print(name, round(((e - theta) ** 2).sum(1).mean(), 1))
# no pooling 144.7 / complete 141.9 / partial 45.9 (theory: 144.6, 142.1, 46.1)
print(np.round((((part - theta) ** 2) - ((yb - theta) ** 2)).mean(0), 1)) # gain per segment, smallest first
# [-59.2 -25. -9.4 -3.6 -1.1 -0.3 -0.1 -0. ] almost all of the gain is in the small segments
# 3) Eight schools (Rubin 1981): estimated coaching effects y and their standard errors s
y = np.array([28., 8, -3, 7, -1, 1, 18, 12]); s = np.array([15., 10, 16, 11, 9, 11, 10, 18])
taus = np.linspace(0, 60, 6001)
V = s**2 + taus[:, None]**2; v = 1 / V
mu_hat = (v * y).sum(1) / v.sum(1); V_mu = 1 / v.sum(1) # flat prior on mu, integrated out
logp = 0.5 * np.log(V_mu) - 0.5 * np.log(V).sum(1) - 0.5 * ((y - mu_hat[:, None])**2 / V).sum(1) # flat prior on tau
p = np.exp(logp - logp.max()); p /= p.sum()
print(taus[p.argmax()], taus[np.searchsorted(p.cumsum(), 0.5)]) # 0.0 5.23: mode at 0, median 5.2
wt = np.where(taus[:, None] > 0, s**-2 / (s**-2 + np.maximum(taus[:, None], 1e-12)**-2), 0.0)
E_cond = wt * y + (1 - wt) * mu_hat[:, None] # E[theta_j | tau, y]
print(np.round((p[:, None] * E_cond).sum(0), 1)) # full Bayes: [11.4 7.9 6.1 7.6 5.1 6.1 10.7 8.5]
print(round(mu_hat[0], 2), round(np.sqrt(V_mu[0]), 2)) # 7.69 4.07: empirical Bayes plugs in tau = 0 -> all schools 7.7
# 4) The same model in NumPyro with the classic priors (non-centered form: Chapter 6.7 explains why)
def eight_schools(s, y=None):
mu = numpyro.sample("mu", dist.Normal(0, 5))
tau = numpyro.sample("tau", dist.HalfCauchy(5))
with numpyro.plate("schools", len(s)):
z = numpyro.sample("z", dist.Normal(0, 1))
theta = numpyro.deterministic("theta", mu + tau * z)
numpyro.sample("y", dist.Normal(theta, s), obs=y)
mcmc = MCMC(NUTS(eight_schools), num_warmup=1000, num_samples=4000, progress_bar=False)
mcmc.run(jax.random.PRNGKey(0), s=s, y=y)
post = mcmc.get_samples()
print(round(float(post["mu"].mean()), 1)) # 4.4: the N(0, 5) prior pulls mu toward 0
print(np.round(np.asarray(post["theta"], dtype=float).mean(0), 1)) # [6.2 4.9 3.9 4.8 3.7 4.1 6.3 4.9]
print(np.round(np.quantile(np.asarray(post["tau"]), [0.05, 0.5, 0.95]), 1)) # [0.2 2.8 9.5]
# Different (reasonable) priors, noticeably different answers: with 8 noisy schools, check prior sensitivity (Chapter 6.8)
# 5) Conversion rates: Beta-Binomial partial pooling = pseudo-counts (Chapter 6.3)
k = np.array([2, 30, 1300, 0]); nv = np.array([5, 200, 10000, 20])
phi, kappa = 0.10, 100 # population mean rate; "worth 100 visitors"
print(np.round(k / nv, 3), np.round((k + kappa * phi) / (nv + kappa), 3))
# raw [0.4 0.15 0.13 0.] -> pooled [0.114 0.133 0.13 0.083]: the 2-of-5 segment drops from first to third
# 6) The winner's curse: the best-looking raw segment is usually a small, lucky one
sz = np.array([20, 30, 45, 70, 100, 150, 230, 350, 500, 750, 1100, 1700, 2500, 3500, 5000])
pt = rng.beta(20, 180, (20000, 15)); kk = rng.binomial(sz, pt)
raw, pooled = kk / sz, (kk + 20) / (sz + 200)
i_raw, i_pool = raw.argmax(1), pooled.argmax(1); rows = np.arange(20000)
print(round((i_raw < 5).mean(), 2), round(100 * (raw[rows, i_raw] - pt[rows, i_raw]).mean(), 1)) # 0.67 5.3
print(round((i_pool < 5).mean(), 2), round(100 * (pooled[rows, i_pool] - pt[rows, i_pool]).mean(), 1)) # 0.1 0.0
1. $\sigma = 12$, $\tau = 6$, and a segment has $n = 4$ orders. What weight does its own average get?
2. In the hierarchical model, what happens to the segment estimates as $\tau \to 0$?
3. Two segments have raw averages the same distance above the centre. One has 10 users, the other 1 000. Which one moves further toward the centre?
4. In eight schools, empirical Bayes gives every school 7.7 ± 4.1. Why?
5. Which statement about shrinkage is correct?
6. Overall rate 10%, population worth $\kappa = 100$ visitors. A segment converts 3 of 10. Its pooled rate is…
Practice problems
A. $\sigma = 20$, $\tau = 5$, $\mu = 20$. A segment has $n = 16$ observations averaging 30. Find the precisions, the weight, the estimate and its posterior SD.
- Noise variance of the average: $400/16 = 25$; data precision $1/25 = 0.04$.
- Population precision $1/\tau^2 = 1/25 = 0.04$.
- $w = 0.04/(0.04 + 0.04) = 0.5$; estimate $= 0.5 \times 30 + 0.5 \times 20 = 25$.
- Posterior SD $= 1/\sqrt{0.08} \approx 3.54$ (raw SE 5).
B. Show that $\dfrac{n/\sigma^2}{n/\sigma^2 + 1/\tau^2} = \dfrac{\tau^2}{\tau^2 + \sigma^2/n} = \dfrac{n}{n + \sigma^2/\tau^2}$.
Multiply top and bottom of the first form by $\sigma^2\tau^2/n$: top $\tau^2$, bottom $\tau^2 + \sigma^2/n$. That is the second form. Multiply top and bottom of the first form by $\sigma^2$ instead: top $n$, bottom $n + \sigma^2/\tau^2$. That is the third form.
C. (Interview) A product manager sees a +25% lift in a 60-user segment (overall lift +2%) and wants to launch the feature only there. What do you say?
"With 60 users that lift is mostly noise: its standard error is probably larger than the true differences between segments. Among many segments, the biggest raw lift is usually a small, lucky one (winner's curse). Our hierarchical model shrinks each segment's lift toward the overall lift in proportion to its noise; let's look at that pooled lift and at $P(\text{lift} \gt \delta \mid D)$ for this segment. If the pooled estimate is close to +2%, the evidence for a special effect is weak, and we would need a targeted follow-up test to confirm it."
D. Three segments: $n = 10, 40, 50$ with averages 30, 20, 25. What is the centre $\hat\mu$ at $\tau = 0$ and as $\tau \to \infty$?
$\tau = 0$: weights $\propto n$: $(10 \times 30 + 40 \times 20 + 50 \times 25)/100 = (300 + 800 + 1250)/100 = 23.5$. $\tau \to \infty$: equal weights: $(30 + 20 + 25)/3 = 25$. Every observation counts equally at one end; every segment counts equally at the other.
E. $n = 4$, $\sigma = 10$, $\tau = 5$ ($\mu$, $\tau$ known). Average MSE with and without pooling? Which segments lose, and what share of the population is that?
$s^2 = 100/4 = 25$, $w = 25/(25 + 25) = 0.5$. No pooling: 25. Partial: $w s^2 = 12.5$. Losers: $|d| \gt \sqrt{2 \times 25 + 25} = \sqrt{75} \approx 8.66 = 1.73\tau$. Share: $P(|Z| \gt 1.73) \approx 8.3\%$.
F. Large segments (so little noise) show conversion rates with mean 10% and SD 2 points. Estimate $\kappa$ by matching the Beta variance, then find the weight a 50-visitor segment puts on its own rate.
$Var(p_g) = \phi(1-\phi)/(\kappa + 1)$: $0.02^2 = 0.0004 = 0.09/(\kappa + 1)$, so $\kappa + 1 = 225$ and $\kappa = 224$. A 50-visitor segment: $w = 50/(50 + 224) \approx 0.18$. It keeps less than a fifth of its own rate.
Centered vs non-centered parameterization
A hierarchical model can be written down in two ways that say exactly the same thing. To the statistics they are identical. To the algorithm that has to explore the posterior they are very different: one way hides a narrow "funnel" that makes samplers stumble and variational fits shrink, the other turns the funnel into a round, easy hill. Knowing which to use, and why, is one of the most practical skills in Bayesian modeling.
- Write a hierarchical model in its centered form $\theta_g \sim N(\mu, \tau^2)$ and its non-centered form $z_g \sim N(0,1)$, $\theta_g = \mu + \tau z_g$, and explain why they are the same model
- Draw and explain the funnel: when the spread $\tau$ is small the group effects are squeezed together (a narrow neck); when it is large they spread out (a wide mouth)
- Explain, with a "walker exploring the posterior" picture, why one step size cannot fit both the neck and the mouth
- Say what a divergence is, why it appears in the neck, and why it means "your answer may be biased", not "your answer is a bit noisy"
- Explain why a Gaussian variational fit also struggles in the funnel, and why the non-centered form fixes it
- Know the nuance: with lots of data per group the centered form is usually better; with little data the non-centered form is
- Do it in NumPyro: by hand with a $z$ site, or automatically with
LocScaleReparam(centered=0), and read the divergence count
What we need from earlier chapters: the hierarchical model $\theta_g \sim N(\mu, \tau^2)$, $y_g \sim p(y\mid\theta_g)$ with population mean $\mu$, between-group spread $\tau$ and within-group noise $\sigma$ (Chapter 6.5), and partial pooling with its precision weights (Chapter 6.6); z-scores and standardization (Chapter 4.18); the Normal distribution and the fact that $a + bZ$ is Normal when $Z$ is (Chapter 4.9); the posterior (Chapter 6.1). Two words used all chapter. A parameterization is the choice of which numbers you use to describe the unknowns of a model (like describing a place by street address or by GPS coordinates). A sampler is an algorithm that walks around the space of unknowns and collects a sample of plausible values from the posterior; MCMC samplers such as Metropolis and NUTS are taught properly in Chapters 6.9 and 6.10. Here you only need the picture of a walker exploring a landscape.
Same model, two ways to write it core
Quick refresher of Chapter 6.5 and 6.6: in a hierarchical model each group $g$ (a segment, a country, a store) has its own effect $\theta_g$, and the effects are drawn from one population bell curve with centre $\mu$ and spread $\tau$; this shared curve is what lets small groups borrow strength from big ones (partial pooling).
Now think about how you would describe one person's height. You can say "182 cm". Or you can say "1.5 standard deviations above the average height". If you know the average (170 cm) and the standard deviation (8 cm), the second description gives the first: $170 + 8 \times 1.5 = 182$. Same person, same height, two ways to write it.
Group effects are the same. The centered way says "group $g$'s effect is $\theta_g$, drawn from the bell curve $N(\mu, \tau^2)$". The non-centered way says "group $g$ is $z_g$ spreads away from the average, with $z_g$ drawn from the standard bell curve $N(0, 1)$; its effect is $\theta_g = \mu + \tau z_g$". The model, the data and the posterior are identical. What changes is which numbers the computer treats as the unknowns while it explores, and that will matter a lot.
Three ways to say it:
- Picture: height in centimetres vs height as a z-score: the same fact in two coordinate systems.
- Numbers: with $\mu = 5$ and $\tau = 2$, the score $z = 1.5$ means $\theta = 5 + 2 \times 1.5 = 8$; drawing $\theta$ straight from $N(5, 2^2)$ produces 8 just as often.
- Slogan: same model, different coordinates; the statistics do not change, the algorithm's job does.
Two recipes, one bell curve. Population mean lift $\mu = 5$ (say, +5 orders per 1 000 visitors), spread between segments $\tau = 2$.
- Centered recipe: draw $\theta_g \sim N(5, 2^2)$ directly. About 95% of segments land within $5 \pm 1.96 \times 2 = 5 \pm 3.92$, that is between $1.08$ and $8.92$.
- Non-centered recipe: draw $z_g \sim N(0, 1)$, then compute $\theta_g = 5 + 2 z_g$.
- Its mean: $E[\theta_g] = 5 + 2\,E[z_g] = 5 + 2 \times 0 = 5$.
- Its variance: $Var(\theta_g) = 2^2\,Var(z_g) = 4 \times 1 = 4$, so its standard deviation is $2$.
- A constant plus a constant times a Normal is still Normal (Chapter 4.9). So $\theta_g \sim N(5, 2^2)$: exactly the same distribution as step 1.
- Some draws: $z = 1.5 \Rightarrow \theta = 8$; $\;z = -0.5 \Rightarrow \theta = 4$; $\;z = 0 \Rightarrow \theta = 5$ (a perfectly average segment).
- And backwards: a segment with $\theta = 8$ has $z = (8 - 5)/2 = 1.5$: "1.5 population spreads above average". This is just the z-score of Chapter 4.18.
A hierarchical model with groups $g = 1, \dots, G$: hyperpriors $\mu \sim p(\mu)$ and $\tau \sim p(\tau)$ with $\tau \gt 0$, group effects $\theta_g$, and data $y_g \sim p(y\mid\theta_g)$. It can be written in two parameterizations:
$$\textbf{centered:}\quad \theta_g \mid \mu, \tau \sim N(\mu, \tau^2) \qquad\qquad \textbf{non-centered:}\quad z_g \sim N(0, 1),\;\; \theta_g = \mu + \tau\, z_g .$$- Centered: the unknowns the algorithm explores are $(\mu, \tau, \theta_1, \dots, \theta_G)$. Each $\theta_g$'s prior is centred at $\mu$ and its width depends on $\tau$.
- Non-centered: the unknowns are $(\mu, \tau, z_1, \dots, z_G)$. Each $z_g$ has the fixed prior $N(0,1)$, which does not depend on $\mu$ or $\tau$. The $\theta_g$ are computed from them (a "deterministic" quantity, not a separate unknown).
- Both give the same joint distribution of $(\mu, \tau, \theta_1, \dots, \theta_G, y)$, hence the same posterior for every quantity you report. With a perfect algorithm the two give identical answers.
- The trick works for any location–scale family (a distribution that is "a standard shape, shifted by a location and stretched by a scale"): Normal, Student-t, Laplace, Cauchy: $\theta = \mu + \tau z$ with $z$ from the standard version.
- Notation: $N(\mu, \tau^2)$ is written with the variance; NumPyro's
dist.Normal(mu, tau)takes the standard deviation $\tau$.
Why do we need it?
The two forms describe the same beliefs but give the algorithm different landscapes to explore. Choosing the right form is often the difference between a fit that silently misses part of the posterior and one that is fast and trustworthy, with no change to the model itself.
Where is it used?
Every hierarchical model fitted by MCMC or variational inference: segment-level lifts in a Bayesian A/B framework, store- or region-level effects in demand forecasting, random effects in mixed models, the classic "eight schools" example, Stan and PyMC and NumPyro tutorials, and automatic tools such as NumPyro's LocScaleReparam.
How is it used?
Write the group effects either as numpyro.sample("theta", dist.Normal(mu, tau)) (centered) or as a standard-Normal site z plus theta = mu + tau * z (non-centered). Fit, check the divergence count and effective sample size, and switch form if the diagnostics complain.
"The non-centered model is a different model with a different prior."
It is the same model. $z_g \sim N(0,1)$ with $\theta_g = \mu + \tau z_g$ implies $\theta_g \sim N(\mu, \tau^2)$ exactly. Only the coordinates the algorithm moves in are different.
"Switching to non-centered changed my estimates, so one of the two must be the wrong model."
With a perfect algorithm the answers are identical. If they differ, at least one of the two fits did not explore the posterior properly; that is a computation problem, and the rest of this chapter shows how to spot which one.
"$z_g$ is the effect of group $g$."
$z_g$ is the effect measured in "population spreads away from $\mu$". The effect itself is $\theta_g = \mu + \tau z_g$. Report $\theta_g$ (NumPyro keeps it if you wrap it in numpyro.deterministic).
In an A/B framework like yours, segment-level lifts with partial pooling are exactly this model: $\theta_g$ is segment $g$'s effect, $\mu$ the overall effect, $\tau$ how much segments really differ. You can write the segment effects either way; the posterior probabilities you report, such as $P(\theta_{g} \gt 0 \mid D)$ for a segment, do not depend on the choice if the fit is good. The global scaler you use (one standardization for all groups) is a different thing: it changes the units of $y$, not which unknowns the sampler explores.
Centered: $\theta_g \sim N(\mu, \tau^2)$. Non-centered: $z_g \sim N(0,1)$, $\theta_g = \mu + \tau z_g$.
Same model, same posterior; only the coordinates the algorithm explores differ. $z_g$ = "spreads away from average".
Trap: if the two fits disagree, a fit is broken, not the model.
Quick check: $\mu = -2$, $\tau = 0.5$. A segment has $z = -3$. What is its effect, and what does $z$ tell you?
$\theta = -2 + 0.5 \times (-3) = -3.5$. The segment is 3 population spreads below the average segment: unusual, since only about 0.13% of a standard Normal lies below $-3$.
The funnel: what the posterior looks like in centered coordinates core
In a hierarchical model the spread $\tau$ and the group effects are tied together. If $\tau$ is tiny, all segments must have almost the same effect: every $\theta_g$ is squeezed close to $\mu$. If $\tau$ is large, the segments may differ a lot: the $\theta_g$ can spread far out. So the room the $\theta_g$ have depends on the value of $\tau$.
Plot one group effect $\theta_g$ (left to right) against $\log \tau$ (bottom to top; the log lets a tiny spread like 0.05 and a big one like 20 both fit on one page). The cloud of plausible values looks like a funnel standing on its tip: a narrow neck at the bottom (small $\tau$, $\theta$ squeezed) and a wide mouth at the top (big $\tau$, $\theta$ spread out).
When does the posterior look like this? When the data say little about each group (few users per segment, noisy metrics). Then the posterior keeps the shape that the prior structure $\theta_g \sim N(\mu, \tau^2)$ gives it, and that shape is a funnel.
Three ways to say it:
- Picture: a trumpet standing on its mouthpiece: thin at the bottom, wide at the top.
- Numbers: at $\log\tau = -2$ the effects live within about $\pm 0.27$ of $\mu$; at $\log\tau = +2$ within about $\pm 14.5$: a room 55 times wider.
- Slogan: the spread decides how much room the groups have, so the shape changes with height.
The toy funnel (a famous test problem from Radford Neal, written here with $\log\tau$). Take $\mu = 0$, a prior $\log\tau \sim N(0, 1.5^2)$ and one group effect $\theta \mid \tau \sim N(0, \tau^2)$, with no data yet.
- At $\log\tau = -2$: $\tau = e^{-2} \approx 0.135$. 95% of $\theta$ lies within $\pm 1.96 \times 0.135 \approx \pm 0.265$.
- At $\log\tau = 0$: $\tau = 1$, so within $\pm 1.96$.
- At $\log\tau = 2$: $\tau = e^{2} \approx 7.39$, so within $\pm 1.96 \times 7.39 \approx \pm 14.5$.
- Width ratio between the heights $+2$ and $-2$: $e^{2}/e^{-2} = e^{4} \approx 54.6$. Same distribution, two heights, a 55-fold change in width.
- How much probability sits in the neck? $P(\log\tau \lt -1) = \Phi(-1/1.5) = \Phi(-0.667) \approx 0.252$. A quarter of all plausible values live in a region that is very narrow. A method that cannot get in there misses a quarter of the answer.
For the toy model, with $s = \log\tau$ ("s" for spread), the joint density of $(\theta, s)$ is
$$p(\theta, s) = \underbrace{N(s;\, 0,\, 1.5^2)}_{\text{prior on }\log\tau}\;\cdot\;\underbrace{N(\theta;\, 0,\, e^{2s})}_{\theta \text{ given } \tau = e^{s}} .$$The conditional width of $\theta$ at height $s$ is proportional to $\tau = e^{s}$: it changes exponentially with height. That shape is the funnel.
- In a real model, $p(\mu, \tau, \theta_1, \dots, \theta_G \mid D) \propto p(\mu)\,p(\tau)\prod_g N(\theta_g;\, \mu, \tau^2)\, p(y_g\mid\theta_g)$. When each $p(y_g\mid\theta_g)$ is wide (weak data per group), the product $\prod_g N(\theta_g;\mu,\tau^2)$ dominates the shape and every pair $(\theta_g, \log\tau)$ looks like the toy funnel.
- More groups make it worse: the term $\prod_g N(\theta_g;\mu,\tau^2)$ contains $\tau^{-G}$, which pulls hard toward small $\tau$ whenever all $\theta_g$ are near $\mu$. The neck becomes a deep, thin pit.
- In non-centered coordinates $(z, s)$ the same toy model is $N(z;\, 0, 1)\cdot N(s;\, 0, 1.5^2)$: two independent bell curves, a round hill with no funnel (shown in the widget).
- Why $\log\tau$? Because $\tau \gt 0$, NumPyro (like Stan) internally moves an unconstrained version of it, $\log\tau$. The funnel is the shape in exactly those internal coordinates.
Why do we need it?
The funnel is the reason hierarchical models are hard to compute. Once you can picture it, every symptom that follows (divergences, stuck chains, too-narrow variational fits, under-estimated $\tau$) has an obvious cause and an obvious fix.
Where is it used?
As a mental model and a standard test problem: "Neal's funnel" is a benchmark in the Stan, PyMC and NumPyro documentation; the "pairs plot" of a group effect against $\log\tau$ is the standard diagnostic plot for any hierarchical model, such as segment effects in A/B tests or store effects in demand models.
How is it used?
After fitting, scatter-plot draws of one $\theta_g$ (or $z_g$) against $\log\tau$, with divergent draws highlighted (ArviZ: az.plot_pair(..., divergences=True)). A funnel with red points piled at the neck says: reparameterize.
"The neck is a thin tail with hardly any probability, so it does not matter if we miss it."
In the toy funnel a quarter of the probability sits below $\log\tau = -1$. The neck is narrow, not empty. Missing it biases $\tau$ upward and makes the groups look more different than they are (less shrinkage).
"The funnel is a bug in the model."
It is a true feature of the posterior when data per group are weak: the data really cannot rule out "all groups are nearly identical". The funnel is fine; the trouble is computing it in these coordinates.
"Looking at $\tau$ alone will show the problem."
The marginal of $\tau$ alone can look harmless. The funnel only appears in the pair plot of a group effect against $\log\tau$. Always make that plot for hierarchical models.
In your A/B framework the funnel appears when there are many segments with few users each, or when the segments truly barely differ (posterior mass near $\tau \approx 0$). That is exactly the case where partial pooling matters most, so it is also the case where you most need the fit to reach the neck: if it does not, segment estimates are shrunk too little and some segments look like winners by noise.
Centered posterior of a group effect vs $\log\tau$ = a funnel: width of $\theta_g$ at height $\log\tau$ is $\propto \tau$.
Toy: $\log\tau \sim N(0, 1.5^2)$, $\theta\mid\tau \sim N(0,\tau^2)$: widths $\pm 0.27$ vs $\pm 14.5$ at $\log\tau = \mp 2$; 25% of mass below $\log\tau = -1$.
Appears when data per group are weak. Diagnose with a pairs plot of $\theta_g$ against $\log\tau$.
Quick check: at what height $\log\tau$ is the 95% room for $\theta$ exactly $\pm 1$?
We need $1.96\,\tau = 1$, so $\tau = 1/1.96 \approx 0.51$ and $\log\tau = \log 0.51 \approx -0.67$.
Why samplers struggle: one step size cannot fit both core
Picture a sampler as a walker exploring a landscape in the fog. The height of the ground is the posterior density. The walker moves around, and the places she spends time in become your sample: lots of time on high ground (plausible values), little on low ground. Two kinds of walker matter for us:
- Random-walk Metropolis takes a random step of typical length $\varepsilon$ (the step size). If the new spot is higher, she moves; if it is lower, she moves only sometimes, otherwise she stays put. (Chapter 6.9 builds it from scratch.)
- Hamiltonian Monte Carlo (HMC), and its self-tuning version NUTS that NumPyro uses, glides instead: like a marble rolling over the landscape, it follows the slope through many small moves, each of length $\varepsilon$. (Chapter 6.10.)
Both have one step size. In the funnel's wide mouth a good step is large, several units. In the neck a good step is tiny, a few hundredths. A step that suits the mouth jumps straight out of the neck, so in the neck almost every move is refused and the walker barely gets in. A step that suits the neck makes the walker crawl through the mouth, needing thousands of moves to cross it. There is no single good choice.
Three ways to say it:
- Picture: one pair of shoes for a ballroom and a tightrope: either you stride and fall off the rope, or you shuffle across the ballroom all night.
- Numbers: with step 1, about 17% of moves are accepted at $\log\tau = -2$ but 96% at $\log\tau = +2$, where each move covers only a thirtieth of the room.
- Slogan: the right step size changes with height; a fixed one is wrong almost everywhere.
How often a random step is accepted. For a random-walk step of size $\varepsilon$ inside a bell curve of width (standard deviation) $\tau$, the long-run acceptance rate works out to $\frac{2}{\pi}\arctan\!\left(\frac{2\tau}{\varepsilon}\right)$ (a known result for Normal targets; the widget below and a simulation agree with it). Apply it at three heights of the funnel with $\varepsilon = 1$:
- Neck, $\log\tau = -2$, $\tau = 0.135$: $\frac{2}{\pi}\arctan(0.271) = 0.637 \times 0.264 \approx 0.17$. Five of every six moves are refused.
- Middle, $\log\tau = 0$, $\tau = 1$: $\frac{2}{\pi}\arctan(2) = 0.637 \times 1.107 \approx 0.70$.
- Mouth, $\log\tau = 2$, $\tau = 7.39$: $\frac{2}{\pi}\arctan(14.8) \approx 0.96$. Nearly every move is accepted, but the room is about $2 \times 14.5 = 29$ wide.
- A random walk with steps of size 1 drifts a distance of about $\sqrt{N}$ after $N$ steps (a rule of thumb), so crossing 29 units takes about $29^2 \approx 840$ steps. Slow.
- Shrink the step to $\varepsilon = 0.3$: the neck acceptance rises to about $0.47$, but crossing the mouth now takes about $(29/0.3)^2 \approx 9\,300$ steps.
Whatever you choose, some part of the funnel is explored badly.
A Markov chain Monte Carlo (MCMC) sampler produces a sequence of points $x^{(1)}, x^{(2)}, \dots$ (a "chain"), each made from the previous one, so that in the long run the points are spread like the posterior. Its key tuning knob is the step size $\varepsilon$.
- Acceptance rate: the fraction of proposed moves the sampler takes. Too low means steps are too big; very high often means steps are too small.
- Warmup (also called adaptation): the first part of a NUTS run, thrown away, during which NumPyro tunes $\varepsilon$ (aiming at an average acceptance of about 0.8 by default,
target_accept_prob=0.8) and one fixed scale per coordinate. After warmup these are frozen. - The ideal local step is proportional to the local width of the posterior. In the funnel the local width of $\theta$ is proportional to $\tau = e^{\log\tau}$, so it changes by a factor $e$ for every unit of height. A single tuned $\varepsilon$ fits the region where most of the warmup happened (the bulk), not the neck.
- The result is poor mixing: the chain moves slowly, the effective sample size (ESS: how many independent draws your correlated draws are worth; Chapter 6.10) is small, and the neck is under-visited.
Why do we need it?
To understand why a model that is perfectly correct can still be fitted badly. The problem is not too few iterations or a bad seed; it is a landscape whose right step size changes from place to place, which no single tuned step size can follow.
Where is it used?
Tuning and debugging MCMC in NumPyro, Stan and PyMC: reading acceptance rates, step sizes and ESS for hierarchical models of segments, stores or regions; explaining in an interview why NUTS reports low ESS for $\tau$ in a centered model.
How is it used?
After a NUTS run, look at the adapted step size, the acceptance statistics, the ESS of $\tau$ (or $\log\tau$) and the trace of $\log\tau$. A trace that rarely dips low, with low ESS, while the group effects look fine, is the funnel signature.
"If the chain misses the neck, just run it longer."
Longer runs help only if the walker can get into the neck at all. When the step size is too big for the neck, HMC refuses almost every move into it, so extra iterations add more draws of the same biased picture.
"A high acceptance rate means the sampler is working well."
Tiny steps are almost always accepted and still crawl. Judge mixing by the effective sample size and the trace plots, not by acceptance alone.
"NUTS adapts its step size, so it handles the funnel automatically."
NUTS tunes one step size (and one scale per coordinate) during warmup, then freezes them. It cannot use a big step in the mouth and a tiny one in the neck. Changing the coordinates is what removes the problem.
Sampler = walker with one step size $\varepsilon$. Ideal step ∝ local width; in the funnel the width ∝ $\tau$, so it varies exponentially with $\log\tau$.
Random step in a bell curve of width $\tau$: acceptance $\frac{2}{\pi}\arctan(2\tau/\varepsilon)$: tiny in the neck, ≈ 1 (but slow) in the mouth.
Trap: NUTS adapts one global $\varepsilon$; it cannot follow a width that changes with height.
Quick check: warmup tuned the step size to $\varepsilon = 0.5$. Roughly below what height $\log\tau$ do random-walk proposals start to be mostly refused (acceptance under 0.5)?
Acceptance $\frac{2}{\pi}\arctan(2\tau/\varepsilon) = 0.5$ when $\arctan(2\tau/\varepsilon) = \pi/4$, so $2\tau/\varepsilon = 1$, $\tau = \varepsilon/2 = 0.25$, $\log\tau \approx -1.39$. Below that height most proposals are refused, and the toy funnel has about $\Phi(-1.39/1.5) \approx 18\%$ of its probability there.
Divergences: the sampler's warning light core
HMC moves by simulating physics: a frictionless puck slides over the landscape, speeding up downhill and slowing uphill. The computer cannot follow the puck continuously, so it moves it in small jumps of length $\varepsilon$. On a gentle slope the jumps follow the true path closely. There is also a built-in honesty check: the puck's total energy (height plus speed) should stay the same along the path. If the computed energy stays nearly constant, the simulation is good.
In the funnel's neck the walls are extremely steep and close together. A jump of length $\varepsilon$ overshoots the narrow valley and lands high up on the opposite wall; the next jump overshoots even more; within a few jumps the computed puck flies off to infinity and the energy explodes. NumPyro notices the exploding energy and calls the move a divergence (a "divergent transition").
A divergence is the sampler saying: "I tried to go there, my simulation broke, so I did not go." The region where it happens is visited too little, and everything you compute from the draws is biased toward the rest of the posterior.
Three ways to say it:
- Picture: a car taking a hairpin bend at motorway speed: it does not follow the road, it flies off it.
- Numbers: in a valley of width 0.1, steps of 0.3 send the energy from 0.5 to 13, 621, then 29 160 in three jumps.
- Slogan: divergences are not noise; they are a map of where the sampler could not go.
One HMC path, step by step. Take a bell curve of width $\tau = 0.1$ (like the funnel's neck at $\log\tau \approx -2.3$). Its "height" is $U(\theta) = \theta^2/(2 \times 0.1^2) = 50\,\theta^2$; the energy is $U(\theta) + p^2/2$, where $p$ is the puck's momentum (speed). Start at $\theta = 0.1$ with $p = 0$: energy $= 50 \times 0.01 = 0.5$. The leapfrog jumps (the integrator HMC uses, Chapter 6.10) give:
- Step size $\varepsilon = 0.3$: positions $0.1 \to -0.35 \to 2.35 \to -16.1 \to 110$.
- Energies along the way: $0.5 \to 13.2 \to 621 \to 29\,160 \to 1.37$ million. It should have stayed at 0.5.
- After the third jump the energy error is over 1 000: NumPyro marks the transition as divergent and does not move there.
- Step size $\varepsilon = 0.05$: positions $0.1 \to 0.0875 \to 0.0531 \to 0.0055 \to -0.0436$, energies $0.5 \to 0.493 \to 0.478 \to 0.469 \to 0.475$. The error stays below $0.04$: a healthy path.
- The rule behind it: for a bell curve of width $\tau$, leapfrog jumps are stable only if $\varepsilon \lt 2\tau$. Here $2\tau = 0.2$: $0.3$ breaks, $0.05$ is fine. In the toy funnel a step of $0.3$ is stable at the mouth ($2\tau \approx 15$ at $\log\tau = 2$) but breaks wherever $\tau \lt 0.15$, that is below $\log\tau \approx -1.9$.
In HMC and NUTS, a divergent transition is one where the simulated energy $H = -\log p(\text{position}) + \tfrac12|\text{momentum}|^2$ grows by more than a threshold during the trajectory (NumPyro's default max_delta_energy is 1000, as in Stan) or becomes infinite or NaN.
- Cause: the step size is too large for the local curvature (how sharply the landscape bends). In the funnel, the curvature in the $\theta$ direction is $1/\tau^2$, enormous in the neck.
- Meaning: the region is under-explored, so estimates are biased, not just noisy. Even a few divergences deserve attention.
- Where they occur is informative: plot the draws of a group effect against $\log\tau$ and highlight divergent draws. A pile at the neck says "funnel".
- In NumPyro:
mcmc.run(key, ..., extra_fields=("diverging",)), thenmcmc.get_extra_fields()["diverging"].sum();mcmc.print_summary()also prints "Number of divergences". - Fixes, in order: reparameterize (non-centered, next section); then, if needed, a smaller step via
NUTS(model, target_accept_prob=0.95)(slower, and it often only pushes the problem deeper); stronger priors or more data.
Why do we need it?
Divergences are one of the few diagnostics that tell you an MCMC answer is wrong, not just imprecise. Without them, a funnel-shaped posterior would quietly give you a too-large $\tau$ and too little shrinkage.
Where is it used?
Every NUTS run in NumPyro, Stan and PyMC reports them; ArviZ plots them on pair plots; they are the first thing to check after fitting a hierarchical model of segments, stores or schools, and a standard interview question about MCMC diagnostics.
How is it used?
Read the divergence count after every run. If it is above zero, plot divergent draws against $\log\tau$ (or the suspected parameter). If they cluster at the neck, switch to the non-centered form and refit; confirm the count drops to zero and the ESS of $\tau$ rises.
"Only 12 divergences out of 1 000 draws: 1.2%, so the answer is 98.8% right."
The count says how often the sampler hit the wall, not how much probability is behind the wall. In the eight-groups example below, a centered fit with a few dozen divergences never visits $\tau \lt 0.7$, while the correct posterior puts about 19% of its probability on $\tau \lt 1$.
"Divergences are random glitches; rerun with another seed."
The geometry is the same for every seed, so another seed gives the same kind of failure (maybe with a different count). Change the parameterization, not the seed.
"Raise target_accept_prob to 0.99 and the problem is solved."
A smaller step reaches a bit deeper into the neck, at a higher cost per draw, and the deepest part can still be missed. It is a second-line fix. Reparameterize first.
"I got some divergences, but R-hat is 1.00 and ESS is fine, so the fit is OK."
R̂ compares chains with each other; if every chain misses the same neck, they agree with each other and R̂ looks perfect. Divergences are a separate, direct signal that part of the posterior was not explored.
Model answer: "A divergence means the leapfrog simulation's energy error blew up, usually because the step size is too big for the local curvature. In hierarchical models that happens in the funnel's neck at small τ. The draws are then biased toward large τ. I would look at where the divergences sit, switch the group effects to a non-centered parameterization, and refit until the divergence count is zero."
Divergence = energy error $\gt 1000$ (NumPyro default) during an HMC/NUTS trajectory: the simulation broke, the move is refused.
Leapfrog on width $\tau$ is stable only if $\varepsilon \lt 2\tau$; the neck has tiny $\tau$, so divergences pile up there.
Divergences mean bias (an unexplored region), not noise. Read them with extra_fields=("diverging",); fix by reparameterizing.
Quick check: NUTS adapted $\varepsilon = 0.12$. In the toy funnel, below roughly what $\log\tau$ do you expect divergences?
Stability needs $\varepsilon \lt 2\tau$, so trouble starts when $\tau \lt 0.06$, that is $\log\tau \lt \log 0.06 \approx -2.8$. In the toy funnel only about $\Phi(-2.8/1.5) \approx 3\%$ of the probability lies there, so a small step reaches most of the neck, but each trajectory then needs many more jumps to cross the mouth.
The non-centered fix: let a formula do the squeezing core
The funnel exists because, in centered coordinates, the room that $\theta_g$ has depends on $\tau$. The fix is to explore a quantity whose room does not depend on $\tau$: the standard score $z_g$. Its prior is $N(0, 1)$ at every height, so in the coordinates $(z_g, \log\tau)$ the toy posterior is a plain round hill.
The squeezing has not disappeared; it has moved into the formula $\theta_g = \mu + \tau z_g$. The walker now takes steps of the same size in $z$ everywhere, and the formula turns each step into a $\theta$-step of size $\tau \times$ (step in $z$): tiny in the neck, big in the mouth. That is exactly the "different step size at every height" that no tuned sampler could provide.
Three ways to say it:
- Picture: draw the funnel on a rubber sheet and stretch the neck until it is as wide as the mouth; now one stride fits everywhere.
- Numbers: a $z$-step of 0.5 is a $\theta$-step of $0.5 \times 0.135 = 0.068$ at $\log\tau = -2$ and $0.5 \times 7.39 = 3.7$ at $\log\tau = +2$.
- Slogan: non-centering moves the dependence on $\tau$ out of the landscape and into a formula.
The same three heights, in $z$ coordinates. Toy model, $\mu = 0$.
- At every height, $z \sim N(0, 1)$: 95% of $z$ lies within $\pm 1.96$. The width ratio between $\log\tau = +2$ and $-2$ is now $1.96/1.96 = 1$, instead of 55.
- Leapfrog stability needs $\varepsilon \lt 2 \times (\text{width})$. In the $z$ direction the width is 1 everywhere, so any $\varepsilon \lt 2$ is stable at every height. In the $\log\tau$ direction the width is 1.5, so $\varepsilon \lt 3$. A step of 0.3 is safe everywhere: no divergences.
- Map back: the point $(z, \log\tau) = (1.2, -2)$ is $\theta = 1.2 \times 0.135 = 0.162$; the point $(1.2, 2)$ is $\theta = 1.2 \times 7.39 = 8.87$. Same $z$, very different $\theta$: the formula does the squeezing.
- Where did the funnel go? Changing variables from $\theta$ to $z$ multiplies the density by $|d\theta/dz| = \tau$, and $N(\theta;\, \mu, \tau^2) \times \tau = N(z;\, 0, 1)$. The factor that made the neck deep and thin cancels exactly.
Non-centered reparameterization: replace the unknowns $\theta_g$ by $z_g = (\theta_g - \mu)/\tau$. For fixed $\mu, \tau$ this is a one-to-one map with $d\theta_g/dz_g = \tau$, and
$$p(\mu, \tau, z_{1:G}\mid D) \;\propto\; p(\mu)\,p(\tau)\prod_{g=1}^{G} N(z_g;\, 0, 1)\; p\big(y_g \mid \theta_g = \mu + \tau z_g\big).$$- The prior part $\prod_g N(z_g; 0, 1)$ no longer involves $\tau$: no funnel from the prior.
- $\tau$ now enters only through the likelihood. When the data per group are weak, the likelihood is nearly flat and the posterior is close to round: easy for samplers and for Gaussian variational guides.
- When the data per group are strong, the likelihood forces $\mu + \tau z_g \approx$ (the group's data), which bends the posterior in $(z_g, \log\tau)$ into a thin curved ridge. That is the other side of the trade-off (two sections ahead).
- You still report $\theta_g = \mu + \tau z_g$, computed for every draw. Its posterior is the same as in the centered model.
Why do we need it?
It removes the funnel without changing the model: divergences disappear, the effective sample size of $\tau$ goes up, and the sampler (or a Gaussian guide) reaches the small-$\tau$ region where strong pooling lives.
Where is it used?
It is the default way to write random effects in the Stan, PyMC and NumPyro documentation (eight schools, radon by county, item response models), in hierarchical A/B models of segments, and inside automatic tools such as NumPyro's LocScaleReparam.
How is it used?
Inside the group plate, sample z = numpyro.sample("z", dist.Normal(0, 1)) and compute theta = numpyro.deterministic("theta", mu + tau * z). Use theta in the likelihood exactly as before. Refit and compare divergences and ESS with the centered version.
"Non-centering removes the shrinkage, because the groups no longer depend on each other."
The groups still share $\mu$ and $\tau$ through $\theta_g = \mu + \tau z_g$. The model, the pooling and the shrinkage are identical; only the coordinates changed.
"To non-center I must rewrite the whole model."
Only the group-effect sites change: one standard-Normal site and one line of arithmetic. NumPyro can even do it for you (last section of this chapter).
"Non-centered is always the better choice."
It is better when the data per group are weak. With strong data per group it creates its own thin, curved ridge and the centered form wins. See "When centered is better" below.
In your A/B framework, the segment effects become z = numpyro.sample("z", dist.Normal(0, 1)) inside the segment plate and theta = mu + tau * z. The same trick works for any parameter whose prior scale is itself a parameter. For example, in your forecasting model the changepoint slope changes have $\delta_j \sim Laplace(0, b)$; if $b$ is a learned parameter rather than a fixed setting (check which it is in your code), writing $\delta_j = b\,u_j$ with $u_j \sim Laplace(0, 1)$ would be the non-centered version (Laplace is also a location–scale family).
Non-centered: sample $z_g \sim N(0,1)$, compute $\theta_g = \mu + \tau z_g$. In $(z, \log\tau)$ the weak-data posterior is round.
Why it works: $N(\theta;\mu,\tau^2)\cdot\tau = N(z;0,1)$: the $\tau$ that shaped the funnel cancels; the formula rescales each step by $\tau$.
Trap: same model and shrinkage; and not always better (strong data per group → centered).
Quick check: in the non-centered toy, the sampler takes a $z$-step of 0.4 at $\log\tau = -3$. How big is the step in $\theta$?
$\tau = e^{-3} \approx 0.0498$, so the $\theta$-step is $0.4 \times 0.0498 \approx 0.02$. In centered coordinates the sampler would have needed a step that small on its own, at that height only.
Variational inference meets the funnel too
Stochastic variational inference (SVI), which both your projects use, does not walk around the posterior. It picks a simple shape, the guide $q$ (for example a Gaussian), and tunes its centre and width until it matches the posterior as well as that shape can (the full story is in Chapters 6.11–6.13). NumPyro's automatic guides fit a Gaussian in the internal, unconstrained coordinates, which for $\tau$ means $\log\tau$.
A Gaussian in $(\theta, \log\tau)$ is an ellipse: its $\theta$-width is the same at every height. A funnel's width is not. If the ellipse reached down into the neck, it would put lots of probability where the funnel has almost none (its $\theta$-width is far too big there), and the fit's error measure punishes that heavily. So the best ellipse stays away from the neck and from the top of the mouth: it hugs the middle and reports a $\log\tau$ that is far too certain.
In non-centered coordinates the toy posterior is itself a round Gaussian, so a Gaussian guide can match it exactly.
Three ways to say it:
- Picture: trying to cover a trumpet with an oval sticker: whatever you do, it covers the middle and misses both ends.
- Numbers: the true spread of $\log\tau$ is 1.5; the best Gaussian guide says 0.64, and gives $P(\log\tau \lt -1) = 0.06$ instead of $0.25$.
- Slogan: SVI raises no alarm in a funnel; it just becomes overconfident about $\tau$.
The best Gaussian for the toy funnel. Let $q(\theta, \log\tau)$ be Gaussian. Minimizing the mismatch $KL(q\,\|\,p)$ (Chapter 6.11) can be done exactly for this toy; the steps below give the answer (checked by numerical optimization, including a full-rank Gaussian, which does no better here).
- Best guide: $\log\tau \sim N(0, 0.64^2)$ and $\theta \sim N(0, 0.66^2)$, independent. The truth is $\log\tau \sim N(0, 1.5^2)$.
- Probability of the neck: guide $P(\log\tau \lt -1) = \Phi(-1/0.64) \approx 0.06$; truth $0.25$.
- 90% interval for $\tau$: guide $e^{\pm 1.645 \times 0.64} = [0.35,\; 2.86]$; truth $e^{\pm 1.645 \times 1.5} = [0.085,\; 11.8]$. Too narrow at both ends.
- The remaining mismatch is $KL \approx 0.85$: it cannot go lower with a Gaussian in these coordinates.
- Non-centered coordinates: the posterior is $N(z; 0,1)\,N(\log\tau; 0, 1.5^2)$, a Gaussian, so the best guide is the posterior: $KL = 0$.
- On the real eight-groups data (Code-it below), NumPyro's
AutoNormalguide on the centered model gives a standard deviation of $\log\tau$ of about 0.23 and $P(\tau \lt 1) \approx 0$; on the non-centered model about 0.7 and $P(\tau \lt 1) \approx 0.13$; a long NUTS run gives about 1.17 and 0.20. Better, though still not perfect.
In SVI the guide $q_\phi$ is fitted by maximizing the ELBO, which is the same as minimizing $KL(q_\phi \,\|\, p(\cdot\mid D))$. Because $\log p(D) = \text{ELBO} + KL$ (Chapter 6.12):
- A Gaussian guide (mean-field
AutoNormal, or full-rankAutoMultivariateNormal) has the same width in every direction at every location. It can tilt (correlation) but not widen with height. A funnel is out of its reach in centered coordinates. - $KL(q\,\|\,p)$ punishes $q$ for putting mass where $p$ has little, so the fitted $q$ shrinks inward: too narrow, especially for $\tau$ (this "mode-seeking" behaviour is explained in Chapter 6.11).
- The two parameterizations have the same $\log p(D)$ (it is the same model), so the one whose fitted guide reaches the higher ELBO has the smaller KL: a fair, practical way to compare them. On eight groups: about $-33.4$ (centered) vs $-31.6$ (non-centered).
- SVI has no divergence alarm. The symptom is quiet: a narrow posterior for $\tau$, and group effects shrunk too little or too much.
Why do we need it?
Because SVI is fast and reports no divergences, it is easy to believe its answer. In a centered hierarchical model it can be confidently wrong about how different the groups are, which changes every pooled estimate and every per-group decision.
Where is it used?
Any hierarchical model fitted with SVI and an automatic Gaussian guide: segment-level effects in an A/B framework, store or region effects in demand models, large models where NUTS is too slow and SVI is the only practical option.
How is it used?
Write group effects non-centered before fitting with an autoguide; compare final ELBOs of the two forms (higher is better, same model); and when possible, validate the SVI posterior of $\tau$ against a NUTS run on the same model or a subset (Chapter 6.15).
"SVI showed no divergences, so the funnel did not affect it."
SVI has no divergence check at all. In a funnel it returns a guide that is too narrow for $\tau$, without any warning.
"A full-rank guide captures correlations, so it fixes the funnel."
A full-rank Gaussian can tilt, but its width still cannot change with height. In the toy funnel the best full-rank fit is the same as the best mean-field fit. Change the coordinates instead (or use a more flexible guide).
"The non-centered SVI answer is now exact."
Only in the toy. With real data the non-centered posterior is not exactly Gaussian, so the guide is better but still somewhat too narrow (0.7 vs 1.17 for the spread of $\log\tau$ on eight groups). Check important conclusions against NUTS.
Your A/B framework fits its hierarchical models with SVI. If segment effects are written centered and the true between-segment spread can be small, the guide will be overconfident about $\tau$: segments get the wrong amount of shrinkage and per-segment probabilities such as $P(\theta_g \gt 0 \mid D)$ come out too extreme or too timid, with nothing in the ELBO trace to warn you. Writing them non-centered, and comparing the final ELBO of both forms, is cheap insurance. Your forecasting loop's ELBO-based early stopping and best-state checkpointing (Chapter 6.14) would happily converge to either answer: convergence of the ELBO says nothing about which parameterization is better.
A Gaussian guide has the same width at every height, so in centered coordinates it cannot follow the funnel: it under-states uncertainty in $\tau$ (toy: sd 0.64 vs 1.5; KL ≥ 0.85).
Non-centered toy: the posterior is Gaussian, KL = 0. Real data: better, not perfect.
Same model ⇒ same $\log p(D)$ ⇒ the form with the higher final ELBO has the smaller KL.
Quick check: the centered fit ends at ELBO = −33.4 and the non-centered one at −31.6. Which guide is closer to its posterior, and by how much KL?
Both models have the same $\log p(D)$, and $KL = \log p(D) - \text{ELBO}$. So the non-centered guide has the smaller KL, by about $-31.6 - (-33.4) = 1.8$ nats. (Use ELBOs estimated with many particles: single-step SVI ELBOs are noisy.)
When centered is better: lots of data per group core
Non-centering is not a free win. Imagine each segment has a million users. Then the data pin each $\theta_g$ down tightly on their own, whatever $\tau$ is. In centered coordinates that is easy: $\theta_g$ sits in a narrow vertical strip of the same width at every height of $\log\tau$.
In non-centered coordinates it becomes hard. The data say "$\mu + \tau z_g$ must be about 12". If $\tau$ is 2, $z_g$ must be about $(12 - \mu)/2$; if $\tau$ is 4, about $(12 - \mu)/4$. Every value of $\tau$ needs a different $z_g$, so the plausible region is a thin curved ridge: a funnel turned on its side. Now the non-centered sampler needs tiny steps and moves slowly.
So the choice depends on how loudly each group's own data speak compared with the spread between groups. That is the same comparison as the shrinkage weight of Chapter 6.6.
Three ways to say it:
- Picture: if each group is nailed down by its own data, describing it "relative to the population" makes every description wobble whenever the population's spread wobbles.
- Numbers: on eight groups with standard errors 9–18, non-centered wins (0 vs 54 divergences); with standard errors ten times smaller, centered wins (effective sample size about 1 040 vs 126).
- Slogan: center when the data speak loudly per group; non-center when they whisper.
Two versions of the eight groups. Each group has an estimated effect with a standard error $SE_g$ (the classic "eight schools" numbers: $SE_g$ from 9 to 18). The posterior spread between groups is about $\tau \approx 3.6$.
- Weight on a group's own data (from Chapter 6.6): $w_g = \dfrac{1/SE_g^2}{1/SE_g^2 + 1/\tau^2} = \dfrac{\tau^2}{\tau^2 + SE_g^2}$.
- Real data, $SE = 15$: $w = 13/(13 + 225) \approx 0.055$. The group's own data barely matter: the posterior is mostly the population curve, i.e. a funnel. Non-centered is the right form.
- NumPyro NUTS (500 warmup, 1 000 draws, seed 0): centered gives 54 divergences, ESS of $\tau$ about 73 and never visits $\tau \lt 0.7$; non-centered gives 0 divergences and ESS about 559.
- Now pretend each group had 100 times more users: standard errors 10 times smaller ($SE = 1.5$), and the posterior $\tau$ becomes about 11: $w = 121/(121 + 2.25) \approx 0.98$. Each group's data dominate. Centered is the right form.
- NumPyro on this version: no divergences in either, but the ESS of $\tau$ is about 1 040 (centered) vs 126 (non-centered), and the non-centered chains needed 3 to 4 times more leapfrog jumps per draw.
A useful rule of thumb (not a theorem): compare each group's standard error $SE_g = \sigma/\sqrt{n_g}$ (how precisely the group's own data measure $\theta_g$) with the between-group spread $\tau$.
- $SE_g \gg \tau$ (weight $w_g$ near 0, data weak): the posterior inherits the funnel; use non-centered.
- $SE_g \ll \tau$ (weight near 1, data strong): the non-centered posterior is a thin curved ridge; use centered.
- In between, or with a mix of big and small groups: either can work; try both and compare divergences and ESS (or ELBO for SVI). You can also non-center only the small groups (two sites), or use partial centering (next section).
- "Strong data" is about $SE_g$ relative to $\tau$, not about the total size of the dataset. A million users spread over 100 000 segments is still weak data per segment.
Why do we need it?
Blindly non-centering everything can make well-identified models slow. Knowing which side of the trade-off you are on tells you which form to start with and how to read a low effective sample size.
Where is it used?
Large A/B tests with a few big segments (often centered is better) versus many small segments (non-centered); store-level demand models with a few big stores and many small ones (mixed); Stan's and NumPyro's guidance on "centered vs non-centered depends on the data".
How is it used?
Estimate $SE_g$ for typical groups (metric standard deviation divided by $\sqrt{n_g}$) and a plausible $\tau$. Start with the form the rule suggests; fit; if the diagnostics complain (divergences or very low ESS for $\tau$), try the other form or partial centering.
"Always non-center hierarchical models; it is the professional default."
It is the right default when data per group are weak, which is common. With strong data per group it slows the sampler down. Check the diagnostics, not a habit.
"My dataset has millions of rows, so I have strong data and should center."
What matters is the standard error of each group's estimate compared with $\tau$. Millions of rows spread over many tiny segments, or a metric with tiny true differences between segments, is still weak data per group.
"No divergences in the non-centered fit, so it is the better one."
With strong data both forms can be divergence-free; then compare effective sample sizes and run time. In the rich-data example the centered fit gives about 8 times the ESS for $\tau$.
In an A/B framework like yours the answer can differ by metric. A conversion metric with tiny true differences between segments ($\tau$ of a fraction of a percentage point) needs thousands of users per segment before centered becomes preferable; a revenue metric with large segment differences may already be "strong data". If some segments are huge and some tiny, a single global choice is a compromise: watch the divergence count and the ESS of $\tau$ on real data.
Weight on group data $w_g = \tau^2/(\tau^2 + SE_g^2)$, $SE_g = \sigma/\sqrt{n_g}$.
Rule of thumb: $w$ small (weak data) → non-centered; $w$ near 1 (strong data) → centered (non-centered becomes a thin curved ridge).
Eight groups: ×1 data: NC wins (0 vs 54 divergences); ×100 data: centered wins (ESS ≈ 1 040 vs 126).
Quick check: a revenue metric has per-user sd $\sigma = 40$, segments differ by about $\tau = 2$, and each segment has 2 500 users. Which form would you start with?
$SE = 40/\sqrt{2500} = 40/50 = 0.8$. $w = 4/(4 + 0.64) \approx 0.86$: the segments' own data dominate, so start centered, and confirm with the diagnostics.
Doing it in NumPyro: by hand or with LocScaleReparam
There are two ways to non-center in NumPyro. By hand: write a standard-Normal site $z$ and compute $\theta = \mu + \tau z$ yourself. Automatically: keep the centered model exactly as it is and wrap it in a reparam handler that tells NumPyro "rewrite the site called theta for me". The handler swaps the site for a hidden standard-Normal site called theta_decentered and computes theta from it, the same arithmetic you would have written.
The automatic version also offers a dial: centered between 0 (fully non-centered) and 1 (unchanged). Values in between give partial centering, a funnel that opens only part of the way.
Three ways to say it:
- Picture: the same recipe card with one line rewritten by an editor, so the cook (the sampler) sees easier instructions.
- Numbers: on the eight groups, centered: 54 divergences; by hand: 0;
LocScaleReparam(centered=0): 0. - Slogan: one config line, same model, easier geometry.
The real runs (the Code-it block below; NUTS, 500 warmup, 1 000 draws, seed 0; reference for $P(\tau \lt 1)$ from a long run: about 0.20).
- Centered: 54 divergences, ESS of $\tau$ about 73, $P(\tau \lt 1) \approx 0.07$, smallest $\log\tau$ visited $-0.3$.
- Non-centered by hand: 0 divergences, ESS about 559, $P(\tau \lt 1) \approx 0.19$, smallest $\log\tau$ about $-5.9$.
reparam(centered, config={"theta": LocScaleReparam(centered=0)}): 0 divergences, ESS about 704, $P(\tau \lt 1) \approx 0.19$. The same model as step 2, written by the library.- The centered run is not "slightly noisy": it underestimates $P(\tau \lt 1)$ by almost a factor of 3 and never sees the neck.
numpyro.handlers.reparam(model, config={"theta": LocScaleReparam(centered=c)}) rewrites the latent site theta with dist.Normal(loc, scale) (or any location–scale distribution with real support) as
- $c = 0$: $\theta_{\text{dec}} \sim N(0, 1)$ and $\theta = \mu + \tau\,\theta_{\text{dec}}$: fully non-centered. $c = 1$: unchanged (centered).
- The new latent site is named
theta_decentered;thetabecomes a deterministic site, still returned bymcmc.get_samples(). centered=None(the default) creates a learnable parameter in $[0, 1]$ starting at 0.5. SVI can learn it; NUTS does not learn parameters, so under NUTS it stays at 0.5. For MCMC, passcentered=0explicitly.- Works for latent sites only (not observed ones), and only for real-valued location–scale distributions.
- Reading divergences:
mcmc.run(key, ..., extra_fields=("diverging",))thenmcmc.get_extra_fields()["diverging"].sum(); ESS:numpyro.diagnostics.effective_sample_sizeormcmc.print_summary().
Why do we need it?
Large models have many hierarchical sites; rewriting each by hand is error-prone. The handler changes the geometry with one line, keeps the original model code readable, and lets you switch forms (or try partial centering) per site.
Where is it used?
NumPyro's own eight-schools and hierarchical tutorials, production NumPyro models with many random effects, and SVI workflows that learn the centering (centered=None), following the "automatic reparameterisation" paper by Gorinova, Moore and Hoffman.
How is it used?
Wrap the model: model_nc = reparam(model, config={"theta": LocScaleReparam(centered=0)}), pass model_nc to NUTS or SVI, read theta from the samples as before, and compare the divergence count and ESS with the original.
LocScaleReparam(centered=0) does to one site. Your model code does not change; the handler substitutes a standard-Normal latent and computes the original site from it."LocScaleReparam() with no argument non-centers my model for NUTS."
With no argument, centered=None creates a learnable parameter starting at 0.5. NUTS does not learn it, so you get a half-centered model. Write LocScaleReparam(centered=0).
"After reparam, theta disappears from the samples."
theta is kept as a deterministic site and appears in get_samples(); the new latent is theta_decentered.
"Reparameterizing changes the priors, so I must re-check my prior predictive."
The joint distribution is unchanged, so prior and posterior predictive checks (Chapter 6.2, Chapter 6.8) give the same results. What changes is the latent site names your guide or initialization refers to.
By hand: z = sample("z", Normal(0, 1)); theta = deterministic("theta", mu + tau * z).
Automatic: reparam(model, config={"theta": LocScaleReparam(centered=0)}) → latent theta_decentered.
Trap: default centered=None = learnable 0.5 (SVI only). Count divergences with extra_fields=("diverging",).
Quick check: with centered=0.5, $\mu = 0$, how does the width of $\theta_{\text{dec}}$ change between $\log\tau = -2$ and $\log\tau = 2$?
Its scale is $\tau^{0.5} = e^{0.5\log\tau}$, so the ratio is $e^{0.5 \times 4} = e^2 \approx 7.4$: a milder funnel than the centered ratio $e^4 \approx 55$, but not round.
Recap, cheat sheet and practice
- A hierarchical model can be written centered ($\theta_g \sim N(\mu, \tau^2)$) or non-centered ($z_g \sim N(0,1)$, $\theta_g = \mu + \tau z_g$). Same model, same posterior; different coordinates for the algorithm.
- In centered coordinates, a group effect plotted against $\log\tau$ forms a funnel: the room for $\theta_g$ is proportional to $\tau$. It appears when data per group are weak.
- A sampler has one step size. The ideal step changes with height, so in the funnel steps are too big in the neck (refused moves, divergences) or too small in the mouth (crawling).
- A divergence is an HMC/NUTS trajectory whose energy error exploded (over 1000 in NumPyro). It marks a region the sampler could not explore: the estimates are biased, typically toward too-large $\tau$.
- A Gaussian variational guide cannot follow a funnel either; it silently under-states the uncertainty in $\tau$. Same model means same $\log p(D)$, so the form with the higher ELBO has the smaller KL.
- Non-centering moves the dependence on $\tau$ into the formula $\theta_g = \mu + \tau z_g$: the weak-data posterior becomes round.
- With strong data per group ($SE_g \ll \tau$) the non-centered posterior becomes a thin curved ridge and the centered form is better. Rule of thumb via the weight $w_g = \tau^2/(\tau^2 + SE_g^2)$.
- NumPyro: by hand with a
zsite, orreparam(model, config={"theta": LocScaleReparam(centered=0)}); read divergences withextra_fields=("diverging",).
Cheat sheet
| Idea | Formula / code | Remember |
|---|---|---|
| Centered | $\theta_g \sim N(\mu, \tau^2)$ | sampler moves $\theta_g$; prior width depends on $\tau$ |
| Non-centered | $z_g \sim N(0,1)$, $\theta_g = \mu + \tau z_g$ | sampler moves $z_g$; report $\theta_g$ |
| Funnel width | 95% room for $\theta$ at height $\log\tau$: $\pm 1.96\,\tau$ | toy: $\pm 0.27$ at $-2$, $\pm 14.5$ at $+2$ (ratio $e^4 \approx 55$) |
| Random-walk acceptance | $\frac{2}{\pi}\arctan(2\tau/\varepsilon)$ | tiny in the neck, near 1 but slow in the mouth |
| Leapfrog stability | $\varepsilon \lt 2\tau$ | violated in the neck → divergences |
| Divergence | energy error $\gt 1000$ (NumPyro default) | bias, not noise; look where they cluster |
| Gaussian guide in the toy funnel | best: $\log\tau \sim N(0, 0.64^2)$, KL ≈ 0.85 | truth $N(0, 1.5^2)$; non-centered: KL = 0 |
| Which form? | $w_g = \tau^2/(\tau^2 + SE_g^2)$ | small $w$ → non-centered; $w$ near 1 → centered (rule of thumb) |
| NumPyro | LocScaleReparam(centered=0) | latent theta_decentered; default None = learnable 0.5 |
| Partial centering | $\theta_{\text{dec}} \sim N(c\mu, \tau^{2c})$, $\theta = \mu + \tau^{1-c}(\theta_{\text{dec}} - c\mu)$ | $c = 1$ centered, $c = 0$ non-centered |
import numpy as np
import jax
import jax.numpy as jnp
import numpyro
import numpyro.distributions as dist
from numpyro.handlers import reparam
from numpyro.infer import MCMC, NUTS, SVI, Trace_ELBO
from numpyro.infer.autoguide import AutoNormal
from numpyro.infer.reparam import LocScaleReparam
from numpyro.diagnostics import effective_sample_size
from numpyro.optim import Adam
# Eight groups: an estimated effect y_g and its standard error sigma_g for each
# (the classic "eight schools" numbers; read them as 8 segment lifts with their SEs)
y = jnp.array([28., 8., -3., 7., -1., 1., 18., 12.])
sigma = jnp.array([15., 10., 16., 11., 9., 11., 10., 18.])
# 1) CENTERED: the sampler moves theta_g ~ Normal(mu, tau) directly
def centered(y, sigma):
mu = numpyro.sample("mu", dist.Normal(0.0, 5.0))
tau = numpyro.sample("tau", dist.HalfCauchy(5.0))
with numpyro.plate("groups", y.shape[0]):
theta = numpyro.sample("theta", dist.Normal(mu, tau)) # Normal(loc, scale): scale = sd
numpyro.sample("obs", dist.Normal(theta, sigma), obs=y)
# 2) NON-CENTERED by hand: the sampler moves z_g ~ Normal(0, 1); theta_g is computed
def noncentered(y, sigma):
mu = numpyro.sample("mu", dist.Normal(0.0, 5.0))
tau = numpyro.sample("tau", dist.HalfCauchy(5.0))
with numpyro.plate("groups", y.shape[0]):
z = numpyro.sample("z", dist.Normal(0.0, 1.0))
theta = numpyro.deterministic("theta", mu + tau * z) # same theta, new coordinates
numpyro.sample("obs", dist.Normal(theta, sigma), obs=y)
# 3) NON-CENTERED automatically: rewrite the "theta" site of the centered model.
# centered=0 -> fully non-centered. (centered=None would create a learnable
# parameter starting at 0.5, which NUTS does not learn: say 0 explicitly.)
reparam_model = reparam(centered, config={"theta": LocScaleReparam(centered=0)})
def run(model, y, sigma, seed=0):
mcmc = MCMC(NUTS(model), num_warmup=500, num_samples=1000, progress_bar=False)
mcmc.run(jax.random.PRNGKey(seed), y, sigma, extra_fields=("diverging",))
tau = np.asarray(mcmc.get_samples()["tau"])
n_div = int(mcmc.get_extra_fields()["diverging"].sum())
ess = float(effective_sample_size(tau[None, :])) # input shape: (chains, draws)
return n_div, round(ess), round(float((tau < 1).mean()), 2), round(float(np.log(tau).min()), 1)
# printed: (divergences, ESS of tau out of 1000 draws, P(tau < 1), smallest log tau visited)
for name, m in [("centered", centered), ("non-centered", noncentered), ("LocScaleReparam", reparam_model)]:
print(name, run(m, y, sigma))
# centered (54, 73, 0.07, -0.3) <- never visits the neck; a long run gives P(tau < 1) ~ 0.20
# non-centered (0, 559, 0.19, -5.9)
# LocScaleReparam (0, 704, 0.19, -6.0)
# 4) Lots of data per group (standard errors 10x smaller, like 100x more users per group)
for name, m in [("centered", centered), ("non-centered", noncentered)]:
print("rich data,", name, run(m, y, sigma / 10))
# rich data, centered (0, 1039, 0.0, 1.5) <- now centered is the efficient one
# rich data, non-centered (0, 126, 0.0, 1.7)
# 5) SVI with an automatic Gaussian guide feels the funnel too (and raises no alarm)
for name, m in [("centered", centered), ("non-centered", noncentered)]:
guide = AutoNormal(m) # Gaussian in (theta or z, mu, log tau)
svi = SVI(m, guide, Adam(0.02), Trace_ELBO(num_particles=4))
res = svi.run(jax.random.PRNGKey(0), 6000, y, sigma, progress_bar=False)
elbo = -Trace_ELBO(num_particles=4000).loss(jax.random.PRNGKey(1), res.params, m, guide, y, sigma)
tau = np.asarray(guide.sample_posterior(jax.random.PRNGKey(2), res.params, sample_shape=(20000,))["tau"])
print(name, round(float(elbo), 1), round(float(np.log(tau).std()), 2), round(float((tau < 1).mean()), 2))
# centered -33.4 0.23 0.0 (ELBO, sd of log tau, P(tau < 1)); NUTS reference: sd 1.17, P 0.20
# non-centered -31.6 0.73 0.13 higher ELBO = smaller KL (same model, same log p(D))
1. Which statement about the centered and non-centered forms of a hierarchical model is true?
2. A NUTS fit of a segment model reports 30 divergences out of 1 000 draws, piled up at small $\tau$. What is the best first action?
3. In the toy funnel ($\mu = 0$), where does 95% of $\theta$ lie at height $\log\tau = -3$?
4. Every segment's own estimate has standard error 0.5, and segments differ by about $\tau = 5$. Which form would you start with?
5. You wrap a model with reparam(model, config={"theta": LocScaleReparam()}) and run NUTS. What did you get?
centered=None creates a parameter in $[0,1]$ initialized at 0.5. SVI can learn it; MCMC does not learn parameters. Write centered=0 for a fully non-centered model.6. You fit a centered hierarchical model with SVI and an AutoNormal guide, and the true posterior has a funnel. What should you expect?
Practice problems
A. Show that $\theta = \mu + \tau z$ with $z \sim N(0,1)$ has the distribution $N(\mu, \tau^2)$. Then compute $\theta$ for $\mu = -1$, $\tau = 3$, $z = 0.4$, and the $z$ of a group with $\theta = 5$.
- Mean: $E[\theta] = \mu + \tau E[z] = \mu$. Variance: $Var(\theta) = \tau^2 Var(z) = \tau^2$.
- A constant plus a constant times a Normal is Normal, so $\theta \sim N(\mu, \tau^2)$.
- $\theta = -1 + 3 \times 0.4 = 0.2$.
- $z = (5 - (-1))/3 = 2$: that group is 2 population spreads above the average.
B. Toy funnel ($\mu = 0$, $\log\tau \sim N(0, 1.5^2)$). Find the 95% room for $\theta$ at $\log\tau = -1$ and $+1$, their ratio, and $P(\log\tau \lt -2)$.
- $\log\tau = -1$: $\tau = 0.368$, room $\pm 1.96 \times 0.368 = \pm 0.72$.
- $\log\tau = +1$: $\tau = 2.718$, room $\pm 5.33$.
- Ratio $e^{1}/e^{-1} = e^2 \approx 7.4$.
- $P(\log\tau \lt -2) = \Phi(-2/1.5) = \Phi(-1.33) \approx 0.091$: 9% of the probability is in a region where the room is under $\pm 0.27$.
C. NUTS adapted a step size $\varepsilon = 0.2$ on the toy funnel. Below what height are leapfrog jumps unstable, and how much of the toy funnel's probability is there?
Stability needs $\varepsilon \lt 2\tau$, so $\tau \gt 0.1$; unstable below $\log\tau = \log 0.1 \approx -2.30$. The probability there is $\Phi(-2.30/1.5) = \Phi(-1.54) \approx 0.06$. Six percent of the posterior lies in a region where trajectories diverge, so it is under-visited, and every summary of $\tau$ is biased upward.
D. (Interview) "What is the funnel in hierarchical models, and why does the non-centered parameterization help?"
"In a hierarchical model the group effects are drawn from $N(\mu, \tau^2)$. When the data per group are weak, the posterior inherits that structure: if $\tau$ is small, all $\theta_g$ must be close to $\mu$; if it is large, they can spread out. Plotting a $\theta_g$ against $\log\tau$ gives a funnel whose width changes exponentially with height. HMC and NUTS use one tuned step size, which is too big for the neck, so trajectories diverge there and the sampler under-explores small $\tau$, biasing the results; a Gaussian variational guide also under-states the uncertainty in $\tau$. The non-centered form samples $z_g \sim N(0,1)$ and computes $\theta_g = \mu + \tau z_g$; the prior of $z_g$ does not depend on $\tau$, so the posterior becomes roughly round, and the formula rescales every step by $\tau$ automatically. It is the same model. With strong data per group the trade-off flips and the centered form is better, so I check divergences and ESS for both."
E. Suppose that in your A/B framework the biggest segments have $SE_g \approx 0.2$ and the smallest $SE_g \approx 5$, with $\tau \approx 1$. What do you do?
Weights: big segments $w = 1/(1 + 0.04) \approx 0.96$ (strong data, centered-friendly); small segments $w = 1/(1 + 25) \approx 0.04$ (weak data, funnel-prone). A single choice is a compromise. Options: non-center everything and accept a slower sampler for the big segments (usually the safe default when many segments are small); split the segments into two sites (centered for big, non-centered for small); or use partial centering (LocScaleReparam with $0 \lt c \lt 1$, or centered=None under SVI). Then compare divergences, ESS of $\tau$ or the final ELBO.
F. Two SVI fits of the same segment model end with (many-particle) ELBOs of $-120.4$ (centered) and $-118.9$ (non-centered). The centered guide says $sd(\log\tau) = 0.2$, the non-centered one $0.6$. Which do you trust more, and why can you compare ELBOs here?
$\log p(D) = \text{ELBO} + KL(q\,\|\,p)$, and $\log p(D)$ is the same for both because it is the same model. So the non-centered guide has the smaller KL, by about $1.5$ nats, and is the better approximation. Its wider $\log\tau$ is consistent with the centered guide having been squeezed by the funnel. ELBOs of different models could not be compared this way for approximation quality, because their $\log p(D)$ would differ.
Checking a Bayesian model: posterior predictive checks, prior sensitivity, identifiability
A fitted model always gives an answer. This chapter is about the three questions you ask before you trust it. Can the fitted model produce data that look like yours? Would a reasonable colleague with a different prior reach the same conclusion? And can the data actually tell your parameters apart, or is the model splitting the credit by guesswork?
- Run a posterior predictive check (PPC): simulate replicated datasets from the fitted model and compare them with the real data, by eye and with test statistics (mean, variance, max, number of zeros)
- Compute and read a posterior predictive p-value as a diagnostic, and know why it is not a classical p-value
- Choose test statistics that target what your decisions depend on, and avoid statistics the model fits automatically
- Do a prior sensitivity analysis: refit with reasonable alternative priors and report whether the conclusion changes, especially for small data, Laplace changepoint scales, dispersion parameters and sparse holiday effects
- Recognize non-identifiability ("only $a + b$ is learned") and weak identification, and their symptoms in samplers and guides
- Explain, for your forecasting model, which component gets the credit for a bump (trend, seasonality, holiday, regressor), and how priors, constraints and better data act as identification aids
What we need from earlier chapters: the posterior predictive distribution $p(\tilde y\mid D) = \int p(\tilde y\mid\theta)\,p(\theta\mid D)\,d\theta$ and how to simulate it (Chapter 6.1); prior predictive checks, which simulate data from the prior before fitting (Chapter 6.2: this chapter is their "after fitting" partner); the Poisson and Negative Binomial distributions and overdispersion (Chapter 4.8); Beta posteriors and $P(\theta_B \gt \theta_A\mid D)$ (Chapters 6.3, 6.4); Laplace priors (Chapter 5.3); regression and collinearity (Chapter 5.13); classical p-values (Chapter 5.6). Checks specific to forecasts (seasonality, autocorrelation, horizons) are in Chapter 7.14. Notation: $D$ (or $y$) is the observed data; $y^{\text{rep}}$ ("y rep") is a replicated dataset of the same size simulated from the fitted model; $T(\cdot)$ is a test statistic.
Posterior predictive checks: can the fitted model fake your data? core
An art expert checks a forger by mixing one real painting among the forger's copies. If the expert can pick out the real one at a glance, the copies are not good enough. A posterior predictive check does the same with your model: after fitting, ask the model to make fake datasets, and compare them with your real data.
If the model has captured how your data arise, the real dataset should look like just another one of the fakes. If the real data stand out (far more zero days than any fake, a much bigger spread, a weekly pattern the fakes never show), the model is missing something, and you have learned what it is missing.
"Fake datasets from the fitted model" are draws from the posterior predictive of Chapter 6.1, made with the same size and design as the real data. They are called replicated data, $y^{\text{rep}}$.
Three ways to say it:
- Picture: a police lineup: hide the real dataset among fakes; if you can spot it, the model fails.
- Numbers: a Poisson model fitted to 30 days of orders produces fake months with 0 to 2 zero-order days; the real month had 7.
- Slogan: a good model can produce data that look like yours.
Orders per day at a small store, 30 days: 11, 0, 1, 8, 8, 2, 9, 11, 0, 1, 0, 0, 6, 0, 0, 13, 2, 7, 2, 5, 2, 1, 2, 11, 0, 6, 1, 4, 2, 5. Total 120, so the mean is 4.0; the variance is 16.4; there are 7 days with zero orders; the largest day is 13.
- Model: each day's count $y_t \sim \text{Poisson}(\lambda)$, prior $\lambda \sim \text{Gamma}(2, \text{rate } 0.5)$ (a weak prior with mean 4).
- Fit: Gamma-Poisson is conjugate (Chapter 6.3): posterior $\lambda \sim \text{Gamma}(2 + 120,\; 0.5 + 30) = \text{Gamma}(122, 30.5)$, mean $122/30.5 = 4.0$, sd $\sqrt{122}/30.5 \approx 0.36$.
- Replicate: draw one $\lambda$ from the posterior, then 30 Poisson($\lambda$) counts. That is one fake month. Repeat a few thousand times.
- Compare zero days: for $\lambda \approx 4$, $P(y = 0) = e^{-4} \approx 0.018$, so a fake month has on average $30 \times 0.018 \approx 0.55$ zero days (90% of fake months have 0 to 2). The real month has 7.
- Compare spread: a Poisson's variance equals its mean, so fake months have variance around 4 (90% between about 2.3 and 6.1). The real month has 16.4.
- Verdict: the real month is easy to pick out of the lineup. The Poisson model misses the overdispersion (variance bigger than the mean; Chapter 4.8). A Negative Binomial model, whose variance $\mu + \mu^2/\alpha$ can exceed the mean, is the natural next try.
A posterior predictive check compares the observed data $y$ with replicated datasets drawn from the posterior predictive distribution
$$p(y^{\text{rep}}\mid y) = \int p(y^{\text{rep}}\mid\theta)\,p(\theta\mid y)\,d\theta .$$Recipe, for $s = 1, \dots, S$:
- take a posterior draw $\theta^{(s)}$ (from MCMC, SVI, or a conjugate formula);
- simulate $y^{\text{rep}(s)} \sim p(y\mid\theta^{(s)})$ with exactly the same size and structure as $y$ (same number of days, same groups, same covariates);
- compare: plot $y$ among a few $y^{\text{rep}}$ (lineups, overlaid histograms or densities), or compute a summary for each (next section).
- Each $y^{\text{rep}}$ uses a different $\theta^{(s)}$, so the comparison includes parameter uncertainty, not just noise.
- Prior predictive checks (Chapter 6.2) use $\theta \sim p(\theta)$: "do my assumptions make sense before seeing data?". Posterior predictive checks use $\theta \sim p(\theta\mid y)$: "after learning from the data, can the model reproduce them?".
- A PPC checks misfit: features of the data the model cannot produce. It does not prove a model correct.
Why do we need it?
A posterior always exists, even for a terrible model. The PPC is the standard way to find out whether the likelihood you chose (Poisson, Normal, Negative Binomial, Student-t) is capable of producing data like yours, and which feature it gets wrong.
Where is it used?
The standard Bayesian workflow in Stan, PyMC and NumPyro (Predictive), ArviZ's plot_ppc, checking count models for overdispersion and excess zeros, checking forecast models for tails and seasonality (Chapter 7.14), and checking A/B metric models before trusting $P(\theta_B \gt \theta_A\mid D)$.
How is it used?
Fit the model; call Predictive(model, posterior_samples)(key, ...) without the observed data to get y_rep of shape (draws, n); plot the real data among a handful of replicates, then compute test statistics. If the real data stand out, change the likelihood or add the missing structure, and check again.
"The PPC passed, so the model is correct."
It only shows that the model can reproduce the features you looked at. Two very different models can pass the same checks; a model can pass every check and still predict the future badly (for that, use held-out data, Chapter 7.15).
"To replicate, plug in the best-fit parameter and simulate."
Use a different posterior draw for each replicate. Plug-in replicates ignore parameter uncertainty and look too tidy, making honest data seem surprising.
"A PPC is the same as the prior predictive check from Chapter 6.2."
Same simulation machinery, different question. Prior predictive (before data): are my assumptions plausible? Posterior predictive (after data): can the fitted model reproduce what I saw?
Both projects have count data where this exact check matters. In your forecasting model, a Negative Binomial likelihood is the usual answer to daily demand being overdispersed; a PPC on the training history (zero-demand days, variance, the biggest days) is how you justify NB over Poisson, or Student-t over Normal for continuous demand. In an A/B framework like yours, a Poisson likelihood for a count metric (orders per user, sessions per user) assumes variance = mean; a PPC on the control group shows quickly whether that holds before you trust $P(\theta_B \gt \theta_A\mid D)$.
PPC: draw $\theta^{(s)} \sim p(\theta\mid y)$, simulate $y^{\text{rep}(s)} \sim p(y\mid\theta^{(s)})$ (same size and design), compare with $y$.
Real data stand out among the fakes ⇒ the model misses that feature (e.g. Poisson: 0–2 zero days per fake month vs 7 real).
Trap: passing a PPC ≠ correct model; use a new posterior draw per replicate.
Quick check: why would the Negative Binomial model produce more zero days than the Poisson with the same mean of 4?
The NB is a Poisson whose rate varies from day to day (a gamma mixture, Chapter 4.8). On low-rate days zeros are common. With $\mu = 4$ and $\alpha = 1$, $P(0) = (\alpha/(\alpha + \mu))^{\alpha} = 1/5 = 0.2$, against $e^{-4} \approx 0.018$ for the Poisson: about 6 zero days per 30 instead of 0.5.
Test statistics and the posterior predictive p-value core
Eyeballing lineups is useful but subjective. To be precise, choose one number that describes a feature you care about: the average, the spread, the biggest day, the number of zero days. In statistics a number computed from a dataset is called a statistic (not the school subject: just "a number made from data"), and a statistic used to check a model is a test statistic $T$.
Compute $T$ for each fake dataset: you get a histogram of what the model expects. Mark $T$ of the real data on it. If the real value sits comfortably inside, the model reproduces that feature. If it is far out in a tail, it does not. The share of fake datasets whose $T$ is at least as large as the real one is the posterior predictive p-value (ppp). Values near 0 or near 1 mean "the real data are extreme for this model"; values in the middle mean "typical".
Three ways to say it:
- Picture: a histogram of what the model expects, and a purple line for what happened.
- Numbers: Poisson: ppp for the variance 0.000, for zero days 0.000; Negative Binomial: 0.50 and 0.38.
- Slogan: pick a feature, ask "is my data typical for the model on this feature?".
The 30 days of orders, checked with four statistics (NumPyro NUTS fits and 2 000 replicated months each; the Code-it block reproduces these numbers).
- Observed: mean 4.0, variance 16.4, zero days 7, max 13.
- Poisson model: fake months have mean about 4.0 (ppp 0.50), variance around 4 (ppp 0.000: none of 2 000 fakes reached 16.4), zero days around 0.6 (ppp 0.000), max around 8.5 (ppp 0.007).
- Read it: the mean is fine, everything about spread and extremes is badly off. The Poisson cannot produce this much variation.
- Negative Binomial model ($y_t \sim NB2(\mu, \alpha)$, $Var = \mu + \mu^2/\alpha$): the posterior gives $\mu \approx 4.1$ and $\alpha \approx 1.1$, so the implied variance is about $4.1 + 4.1^2/1.1 \approx 19$.
- Its ppp values: mean 0.48, variance 0.50, zero days 0.38, max 0.74. All comfortably inside: the NB reproduces every feature we checked.
A test statistic (or discrepancy) $T(y)$ is any number computed from a dataset; it may also depend on the parameters, $T(y, \theta)$. The posterior predictive p-value is
$$p_{B} = P\big(T(y^{\text{rep}}) \ge T(y)\;\big|\; y\big) \;\approx\; \frac{1}{S}\sum_{s=1}^{S} \mathbf{1}\!\left[T(y^{\text{rep}(s)}) \ge T(y)\right],$$where $\mathbf{1}[\cdot]$ is 1 when the condition holds and 0 otherwise.
- It is a diagnostic of misfit: values very close to 0 or 1 say the real data are extreme for the model on this feature. Read both ends (a ppp of 0.99 for the variance means the model is too spread out).
- For counts like "zero days", ties are common; this chapter counts ties as "at least as large" ($\ge$). Say which convention you use.
- It uses the data twice (to fit, then to check), so it tends to sit closer to 0.5 than a classical p-value would: it is conservative. Treat a moderately small value (say 0.05) as a real warning, not as "fine because it is not 0.01".
- It is not the probability that the model is true, and it is not a classical hypothesis test with a guaranteed 5% false alarm rate (Chapter 5.6).
Why do we need it?
Plots can be argued about; a statistic and its ppp turn "the fakes look a bit different" into "zero of 2 000 replicated months had a variance as large as the real one", which is precise, repeatable and easy to report.
Where is it used?
Model checking in the Bayesian workflow (Gelman et al.'s Bayesian Data Analysis), az.plot_bpv in ArviZ, choosing between Poisson and Negative Binomial likelihoods for count metrics and demand, and checking tails before capacity or inventory decisions.
How is it used?
From y_rep of shape (draws, n), compute T(y_rep) along the data axis, compare with T(y), and report (T(y_rep) >= T(y)).mean() next to the histogram. Do it for a handful of statistics chosen in advance.
"ppp = 0.38 means there is a 38% chance the Negative Binomial model is right."
It means 38% of fake months from the fitted NB model had at least as many zero days as the real month. It describes how typical the data are under the model, for one feature. It says nothing about the probability of the model.
"The mean passes with ppp = 0.50, so the Poisson fits."
The Poisson has a free rate $\lambda$ that is fitted to the mean, so the mean almost always passes. Use statistics the model does not fit directly (spread, zeros, extremes).
"ppp above 0.05 means accept, below means reject."
It is a diagnostic, not a decision rule, and it is conservative (it tends toward 0.5). Look at how far the real value is from the bulk and whether it matters for your decisions.
"The posterior predictive p-value is the Bayesian version of the p-value, so it means the same thing."
A classical p-value is computed under a fixed null hypothesis and, when that hypothesis is true, is uniformly distributed (5% of the time below 0.05). A ppp averages over the posterior, which was fitted to the same data, so it is not uniform and tends to be conservative. It measures how surprising one feature of the data is for the fitted model.
Model answer: "I use posterior predictive p-values as a misfit diagnostic: for a chosen statistic, the fraction of replicated datasets at least as extreme as the observed one. Values near 0 or 1 tell me which feature the model fails to reproduce. I do not read them as error rates or as the probability that the model is correct."
$p_B = P(T(y^{\text{rep}}) \ge T(y)\mid y) \approx$ share of replicates at least as extreme. Near 0 or 1 ⇒ misfit on that feature.
Orders example: Poisson ppp (mean, variance, zeros, max) = 0.50, 0.000, 0.000, 0.007; NB = 0.48, 0.50, 0.38, 0.74.
Traps: not P(model true); conservative; statistics fitted directly (the mean) always pass.
Quick check: a model's ppp for the variance is 0.995. What does that say?
99.5% of replicated datasets have a variance at least as large as the real one: the real data are less spread out than the model expects. The model is too wide (for example, too much dispersion or too heavy tails). Both ends of the ppp signal misfit.
Choosing what to check
A PPC only finds the problems you look for. A model with a free mean will always reproduce the mean; a model with a free noise scale will always reproduce the overall spread. Those checks pass even when the model is badly wrong somewhere else. Good test statistics have two properties:
- they describe a feature your decisions depend on (the big days for capacity planning, the zero days for a count model, the weekly pattern for staffing);
- the model does not fit them directly through one of its parameters, so they can actually fail.
Three ways to say it:
- Picture: a smoke detector in the kitchen finds kitchen fires; put the detectors where the fires you fear would start.
- Numbers: a one-level model of daily orders passes the mean (ppp 0.50) and the standard deviation (0.50), but fails the weekend-minus-weekday gap (0.000).
- Slogan: test what matters and what the model was not tuned to match.
Part 1: the mean can never fail. Normal data with known $\sigma$ and a flat prior on $\mu$.
- Posterior: $\mu\mid y \sim N(\bar y, \sigma^2/n)$.
- A replicate's mean, given $\mu$: $\bar y^{\text{rep}} \sim N(\mu, \sigma^2/n)$.
- Averaging over the posterior: $\bar y^{\text{rep}}\mid y \sim N(\bar y, 2\sigma^2/n)$, centred exactly at the observed mean.
- So $P(\bar y^{\text{rep}} \ge \bar y\mid y) = 0.5$ for every dataset, however wrong the model is. This statistic cannot detect anything.
Part 2: a targeted statistic. Eight weeks of simulated daily orders: weekdays around 100, weekends around 120, noise sd 8. The model wrongly assumes one level for all days.
- Observed: weekend mean minus weekday mean = 19.9.
- One-level model, ppp for the overall mean 0.50 and for the standard deviation 0.50: both "pass" (the model's $\mu$ and $\sigma$ were fitted to exactly these).
- ppp for the gap: replicated gaps fall between about $-7.2$ and $+7.6$ (95%), centred at 0; none reaches 19.9: ppp 0.000.
- A model with separate weekday and weekend levels gives a gap ppp of 0.51. The targeted statistic found the missing structure; the generic ones did not.
A useful menu of test statistics (choose a few before looking at the results):
- Location and spread: mean, median, standard deviation, variance-to-mean ratio (counts).
- Tails and extremes: maximum, minimum, 95th or 99th percentile, number of values beyond a threshold.
- Shape of counts: number or share of zeros, number of very large counts.
- Structure: means by group (segments, weekdays), spread of group means, trend in the residuals, lag-1 autocorrelation of residuals (Chapter 7.17), seasonal pattern; forecast-specific checks are collected in Chapter 7.14.
- Avoid (or do not trust a pass on) statistics that a single parameter fits directly: the mean in a model with a free mean, the overall sd in a model with a free $\sigma$.
- Checking many statistics makes a few small ppp values likely by chance; prioritize the ones tied to your decisions.
Why do we need it?
The value of a PPC depends entirely on the statistic. Generic checks give false comfort; statistics aimed at the decision (peaks for capacity, zeros for stock-outs, weekday patterns for staffing) find the failures that would actually hurt.
Where is it used?
Checking demand-forecast models (weekly pattern, peak days, zero-demand days), A/B metric models (spread of segment means, share of heavy users), insurance and risk models (tail percentiles), and any NumPyro model through Predictive plus a few NumPy reductions.
How is it used?
Write down 3–6 statistics tied to how the model will be used, including at least one the model does not fit directly. Compute them for the data and for each replicate, plot the histograms with the observed line, and act on the ones that are extreme.
"Check every statistic you can think of and keep changing the model until all of them pass."
With many statistics, a few extreme ppp values appear by chance, and tuning a model until every check passes is a form of overfitting. Choose a handful in advance, tied to how the model will be used.
"Checking the model on the data it was fitted to is cheating; only held-out data count."
The two answer different questions. PPCs find features the model cannot reproduce (criticism). Held-out data measure predictive accuracy (Chapter 7.15). A good workflow uses both.
For your forecasting model, statistics worth checking on the training history are: number of zero-demand days (NB vs Normal), the largest days and the 99th percentile (Student-t or NB tails), the weekday pattern (Fourier and seasonal terms), behaviour around holidays, and the lag-1 autocorrelation of residuals (structure that trend + seasonality miss). In the A/B framework, a useful hierarchical check is the spread of the observed segment rates: can the fitted model produce segments that differ as much (and as little) as the real ones?
Good $T$: tied to decisions, and not fitted directly by a parameter. Flat-prior Normal: ppp for the mean is exactly 0.5 for any data.
Example: one-level model passes mean (0.50) and sd (0.50), fails the weekend gap (0.000).
Pick a few statistics in advance; PPC = criticism, held-out data = accuracy.
Quick check: a Poisson model of orders per user has a free rate. Name one statistic that is useless for checking it and one that is useful.
Useless: the mean number of orders (the rate is fitted to it, so it passes). Useful: the share of users with zero orders, the variance-to-mean ratio, or the share of users with 10 or more orders: features a Poisson ties to the mean and cannot adjust freely.
Prior sensitivity: would a reasonable colleague reach the same conclusion? core
A prior is a choice, and another sensible person could have chosen a slightly different one. Prior sensitivity analysis asks: if I swap my prior for other reasonable priors, does my conclusion change? If it does not, the data have spoken and you can report the result with confidence. If it does, the data alone are not strong enough to settle the question, and the honest report says so.
It is not about trying absurd priors; it is about the range of priors a thoughtful colleague might defend: a bit wider, a bit narrower, a different family with similar meaning.
Three ways to say it:
- Picture: ask four sensible friends with different hunches; if after seeing the data they all agree, the data decided.
- Numbers: A converts 3 of 20, B 8 of 20: $P(\theta_B \gt \theta_A\mid D)$ ranges from 0.77 to 0.96 across four reasonable priors; with 10 times the data, all four give about 1.000.
- Slogan: a robust conclusion survives every reasonable prior.
A small A/B test: A converts 3 of 20 visitors, B 8 of 20. Decision rule: ship B if $P(\theta_B \gt \theta_A\mid D) \gt 0.95$. Four priors for each rate, all defensible (Beta-Binomial updating from Chapter 6.3):
- Flat Beta(1, 1): posteriors A ~ Beta(4, 18), B ~ Beta(9, 13) (means 0.18 and 0.41). $P(B \gt A) = 0.957$: ship.
- Jeffreys Beta(0.5, 0.5): $P(B \gt A) = 0.963$: ship.
- Weakly informative Beta(2, 18) ("about 10%, worth 20 visitors"): A ~ Beta(5, 35), B ~ Beta(10, 30) (means 0.125 and 0.25). $P = 0.931$: do not ship.
- Informative Beta(20, 180) ("about 10%, worth 200 visitors", from past tests): A ~ Beta(23, 197), B ~ Beta(28, 192) (means 0.105 and 0.127). $P = 0.774$: do not ship.
- Verdict: the decision flips with the prior. Report it as "B looks better, but 20 visitors per arm cannot settle it; the answer depends on how strongly we expect rates near 10%".
- Same rates with 10 times the data (30 of 200 vs 80 of 200): all four priors give $P(B \gt A) \approx 1.000$. The conclusion is now robust.
Prior sensitivity analysis:
- List the quantities you will report or act on (a decision probability, an interval, a forecast quantile).
- Choose, before seeing the results, a small set of reasonable alternative priors: scales a factor 2–3 wider and narrower, a heavier-tailed family (Student-t instead of Normal), a different but defensible centre.
- Refit with each (for small changes you can instead reweight the existing posterior draws by the ratio of new to old prior, importance sampling from Chapter 6.9).
- Report the range of the quantities and whether the decision changes.
- Large sensitivity means the data are weak for this question (small samples, rare events, parameters the data barely touch), or that prior and data disagree (prior–data conflict).
- "Flat" is a prior too, and not neutral on every scale (Chapter 6.2); include it as one option, not as the reference truth.
Why do we need it?
Stakeholders and reviewers will ask "isn't the answer just your prior?". A sensitivity analysis answers with evidence: either "no, every reasonable prior agrees", or "yes, for this question, and here is how much more data would settle it".
Where is it used?
Bayesian A/B testing with small segments, clinical trials (regulators often require it), hierarchical models with few groups, forecasting models with hand-set prior scales (changepoints, holidays, dispersion), and any report that leads to a decision.
How is it used?
Loop over a short list of prior settings, refit (or reweight draws), and tabulate the key outputs side by side. Put the table in the report; if the decision flips, say so and recommend more data or a pre-registered prior.
"Try priors until one gives the result we want, then report that one."
That is prior hacking. Choose the alternatives before seeing the results, and report all of them.
"If the answer is sensitive to the prior, the analysis is broken."
It is a finding: the data are too weak to settle this question on their own. Report the range and how much more data would make it robust.
"Use a flat prior and there is nothing to be sensitive about."
A flat prior is one of the priors; it is not neutral on every scale (flat on a rate is not flat on log-odds), and in small samples it can drive the answer as much as any other.
In an A/B framework like yours, small segments and early looks are where priors matter: a Beta prior "worth" 200 visitors dominates a segment with 20. A sensitivity table (flat, weak, informative-from-history) next to each segment-level $P(\theta_B \gt \theta_A\mid D)$ shows which segment conclusions are robust. Partial pooling makes the prior on $\tau$ part of this list too: with few segments, how much they shrink depends on it (Chapter 6.6).
Sensitivity analysis: refit with a pre-chosen set of reasonable priors; report the range of the key outputs and whether the decision flips.
3/20 vs 8/20: $P(B \gt A)$ = 0.957 / 0.963 / 0.931 / 0.774 (flat, Jeffreys, weak, strong): fragile. 10× data: all ≈ 1.000.
Trap: never pick the prior after seeing which answer it gives.
Quick check: why does the strong prior Beta(20, 180) pull $P(B \gt A)$ down so much with 20 visitors per arm?
It is worth 200 pseudo-visitors at 10%, ten times the real data per arm. Both posteriors are pulled toward 10% (means 0.105 and 0.127 instead of 0.18 and 0.41), so they overlap a lot. With 200 real visitors per arm the data weigh as much as the prior, and with 2 000 they dominate.
Where priors bite: scales, dispersion and sparse effects
Priors matter most where the data say least. Some parameters are learned from thousands of points (the overall level of demand); others from a handful (the effect of a holiday that happened twice; whether a slope change is real). Some parameters are scale parameters that decide how flexible the model is: the Laplace scale $b$ on changepoint slope changes, the prior scale on holiday effects, the between-group spread $\tau$, the dispersion $\alpha$ of a Negative Binomial. Changing their priors changes how much the rest of the model is allowed to move.
A sensitivity curve makes this visible: plot your answer against the prior scale over a range of reasonable values. A flat curve means the data decide; a steep curve means the prior does.
Three ways to say it:
- Picture: a quiet witness is easy to talk over; a loud one is not. Few data points are a quiet witness.
- Numbers: a holiday seen twice: its estimated effect moves from 0.2 to 16.5 units as the prior scale goes from 1 to 100; seen 20 times, from 14.0 to 17.8 as the scale goes from 5 to 100.
- Slogan: few data plus a scale prior means the prior is doing the talking.
A holiday effect under a Laplace prior. Each occurrence of the holiday shows a bump measured with noise sd 20 (units of daily demand); the average bump observed is $\hat d = 18$. Prior: $\delta \sim Laplace(0, b)$ (a Laplace prior pulls small effects toward 0; Chapter 5.3). Posterior computed on a fine grid.
- Seen $k = 2$ times: standard error $20/\sqrt 2 \approx 14.1$. Posterior mean of $\delta$: $b = 1$: 0.2; $b = 5$: 3.2; $b = 20$: 11.4; $b = 100$: 16.5. $P(\delta \gt 0)$ goes from 0.54 to 0.89.
- That is a steep sensitivity curve: the answer is mostly the prior's choice. Report it as uncertain, or set $b$ from other holidays or past years (an informative prior you can defend).
- Seen $k = 20$ times: standard error $4.5$. Posterior mean: $b = 5$: 14.0; $b = 20$: 17.0; $b = 100$: 17.8, and $P(\delta \gt 0) \ge 0.999$ from $b = 5$ on. A flat curve for all reasonable $b$: the data decide.
- Dispersion, same idea: for the 30 days of orders, four reasonable priors on the NB concentration $\alpha$ give posterior medians from 0.88 to 1.08, and the predicted chance of a day with 15 or more orders from 0.037 to 0.046. Mild sensitivity with 30 days; with a week of data it would be much larger.
A sensitivity curve plots a reported quantity (posterior mean, interval, decision probability, forecast quantile) against a prior hyperparameter over a reasonable range, usually on a log scale for scales. Places where it is typically steep:
- Small data: few observations overall, or per group, or per event.
- Scale and complexity parameters: the Laplace scale $b$ of changepoint slope changes $\delta_j$ (it decides how many "real" trend changes the model allows; Chapter 7.10), prior scales of holiday and regressor effects, the hierarchical spread $\tau$ with few groups.
- Dispersion and tail parameters: the NB concentration $\alpha$, the Student-t degrees of freedom $\nu$; they are learned from the few extreme observations.
- Sparse effects: holidays seen once or twice, rare events, segments with little traffic.
Where the curve is steep, either collect information (more history, pooling across similar holidays), use a prior you can justify from outside data, or report the range.
Why do we need it?
Scale priors are often set once (or left at a library default) and forgotten, yet they decide how much trend change, holiday lift or overdispersion the model can express. A sensitivity curve shows whether that hidden choice is driving your forecasts.
Where is it used?
Prophet-style forecasting models (changepoint and holiday prior scales are exposed as tuning knobs), Negative Binomial demand models (dispersion priors), hierarchical models with few groups ($\tau$ priors), and sparse regression (shrinkage scales).
How is it used?
Pick 4–8 values of the scale spanning a factor of 10 or more around your choice, refit (SVI is quick enough), and plot the key outputs against the scale. Choose the scale with holdout performance (Chapter 7.15) and report how much the outputs moved.
"The library's default prior scale is neutral, so I do not need to think about it."
A default scale is a modelling choice made for a typical dataset, on a particular scaling of $y$. For your data it may be far too tight or too loose. Check it with a sensitivity curve or holdout performance.
"The curve is flat between $b = 10$ and $b = 100$, so the prior does not matter."
It does not matter within that range. If someone could reasonably argue for $b = 2$, check there too: the small-$b$ end is where shrinkage is strongest.
Your forecasting model has three textbook examples. The Laplace scale $b$ on the changepoint slope changes $\delta_j \sim Laplace(0, b)$ decides how many of the grid and PELT candidate changepoints get real slope changes (Chapter 7.10). Holiday effects for holidays seen once or twice in the history are mostly prior. And the Negative Binomial concentration (and the Student-t $\nu$), if learned, is mostly informed by a few extreme days, which drives the width of your upper prediction intervals. For each, a short sensitivity table belongs next to the forecast.
Sensitivity curve: key output vs prior scale (log axis). Steep = the prior decides; flat = the data decide.
Holiday, Laplace prior, $\hat d = 18$: seen twice, mean 0.2 → 16.5 for $b$ = 1 → 100; seen 20 times, 14.0 → 17.8 for $b$ = 5 → 100.
Watch: small data, scale priors ($b$, $\tau$, holiday scales), dispersion/tails ($\alpha$, $\nu$), sparse effects.
Quick check: why is a Negative Binomial dispersion parameter more prior-sensitive than its mean?
The mean is estimated from every observation, roughly with standard error $\sqrt{Var/n}$. The dispersion is about how far the extreme days are from the mean, so it is learned mainly from the few largest and smallest values; with little data those few points give weak information, and the prior fills the gap.
Identifiability: when the data cannot tell parameters apart core
Two friends split a restaurant bill, and you only see the total: 60. Did Ana pay 30 and Ben 30? Ana 50 and Ben 10? Ana 70 and Ben −10? The receipt fits all of them equally well. No amount of staring at the total will tell you the split.
A model has the same problem when the data only ever "see" a combination of parameters. In the model $y_i \sim N(a + b, 1)$, every pair $(a, b)$ with the same sum produces exactly the same data. The data can teach you $a + b$ very precisely and $a - b$ not at all. The parameters $a$ and $b$ are not identifiable; only their sum is.
A Bayesian model still produces a posterior (the priors keep it finite), but along the direction the data cannot see, the posterior is simply the prior. The fit runs, nothing crashes, and the split you read off is the prior's opinion.
Three ways to say it:
- Picture: a long, thin ridge in the posterior along the line $a + b = $ constant.
- Numbers: 100 observations with mean 3: $a + b$ is known to $\pm 0.1$, but $a$ alone only to $\pm 7.1$, exactly what the prior said; the correlation of $a$ and $b$ is $-0.9999$.
- Slogan: the data can only teach what changes the data.
The $a + b$ model. $y_i \sim N(a + b, 1^2)$ for $i = 1, \dots, 100$, with $\bar y = 3$; priors $a \sim N(0, 10^2)$ and $b \sim N(0, 10^2)$. Everything is Normal, so the posterior can be written exactly.
- The data give information only about $s = a + b$: its likelihood has standard error $1/\sqrt{100} = 0.1$.
- Posterior of $a + b$: sd $\approx 0.1$ (the prior on $a + b$, with sd $\sqrt{200} \approx 14$, adds almost nothing).
- Posterior of $a - b$: sd $\sqrt{200} \approx 14.1$, identical to its prior. The data said nothing about it.
- So $sd(a) = sd(b) \approx 7.07$, the correlation between $a$ and $b$ is $-0.9999$, and the posterior means are $a = b = 1.5$: the sum 3 split evenly only because the two priors are equal.
- Shrink both prior sds to 1: $sd(a)$ becomes 0.71, the correlation $-0.99$. The prior now constrains the split, but it is still the prior doing it.
A model is identifiable if different parameter values always give different distributions of the data: $p(y\mid\theta_1) = p(y\mid\theta_2)$ for all $y$ implies $\theta_1 = \theta_2$. If some different values give the same distribution, the parameters are non-identifiable, and there are directions in parameter space along which the likelihood is perfectly flat.
- Weakly identified: technically identifiable, but the likelihood changes very little along some direction (near-flat), so in practice the data barely constrain it. This is the common case: collinear regressors, components that look alike over a short history.
- With proper priors the posterior still exists, but along flat directions it is (close to) the prior. Read these as "the data do not tell".
- Symptoms: posterior correlations near $\pm 1$; posterior sd of a parameter close to its prior sd (little contraction, $1 - Var_{\text{post}}/Var_{\text{prior}}$ near 0); slow MCMC with high autocorrelation along the ridge; mean-field SVI guides that look over-confident or change with the seed or initialization; results that move a lot when you change a prior scale.
- What remains well determined are the identified combinations (here $a + b$), and any prediction that depends only on them.
Why do we need it?
A Bayesian fit of a non-identified model looks normal: no errors, sensible-looking numbers. Without the concept you would report the prior's split as a finding, or blame the sampler for a ridge that is really in the model.
Where is it used?
Additive forecasting models (trend vs seasonality vs holidays vs regressors), regressions with collinear features, models with an intercept plus a full set of category effects, latent-factor models (sign and rotation), mixtures (label switching), and any model where two parts can explain the same pattern.
How is it used?
After fitting, look at posterior correlations and at prior vs posterior sds for each parameter; report identified combinations; and if a parameter you care about is not identified, change the design, add constraints or justified priors (see the last section of this chapter).
"The posterior mean of $a$ is 1.5, so $a$ is about 1.5."
That number is the prior's way of splitting the sum. The data support $a = 1.5$ exactly as much as $a = -10$ (with $b = 13$). Report $a + b = 3.0 \pm 0.1$, and say that $a$ and $b$ are not separately identified.
"If the model were not identifiable, NumPyro would raise an error."
With proper priors everything runs. The warning signs are quiet: correlations near $\pm 1$, posteriors that look like priors, slow mixing, seed-dependent SVI fits.
"More data will fix it."
Only data that separate the parameters help. A million more observations of $a + b$ pin the sum down further and leave $a - b$ exactly as unknown.
Identifiable: different θ ⇒ different data distributions. Non-identified directions = flat likelihood; posterior there = prior.
$y \sim N(a + b, 1)$, $n = 100$, priors $N(0, 10^2)$: sd(a + b) = 0.1, sd(a − b) = 14.1 (= prior), corr = −0.9999.
Symptoms: correlations near ±1, little contraction, slow chains, seed-dependent SVI. Report identified combinations.
Quick check: in the $a + b$ model with priors sd 10 for $a$ and sd 5 for $b$, and $\bar y = 3$ from many data, roughly how is the sum split?
In proportion to the prior variances: $a$ gets $100/(100 + 25) = 80\%$ of the sum (about 2.4) and $b$ gets 20% (about 0.6). The data fix the sum; the priors divide it.
Which component gets the credit? Identifiability in an additive forecasting model core
Your demand jumps in late November. Was it the trend turning upward, the yearly seasonality, the holiday, or the promotion (an external regressor) that ran the same week? In an additive model $y_t = g(t) + s(t) + h(t) + X_t\beta + \epsilon_t$, each of these components could produce that bump. If, in your history, two of them always happened together (the promotion always ran on the holiday), the data can measure only their combined effect. It is the $a + b$ problem wearing business clothes.
The model will still hand out credit: by the priors, and by small accidental differences. The forecast for next year's holiday-with-promotion may be fine. But the moment the two separate (a promotion without a holiday), the forecast relies on a split the data never measured.
Three ways to say it:
- Picture: two singers always sing together; you can hear the duet but cannot say who is louder.
- Numbers: holiday-and-promotion days were 20 above normal; equal priors give each 9.8; priors with sds 20 and 5 give 18.6 and 1.2. Same data.
- Slogan: if two components always move together, the data measure the sum and the priors write the split.
Holiday vs promotion. Demand above the usual level is $y_t = \beta_h h_t + \beta_x x_t + \epsilon_t$, with $h_t = 1$ on holidays, $x_t = 1$ on promotion days, noise sd 4. In the history there were 4 holidays, and the promotion ran on exactly those 4 days. On them, demand was 20 above normal on average.
- On every informative day $h_t = x_t = 1$, so the data see only $u = \beta_h + \beta_x$, measured as 20 with standard error $4/\sqrt 4 = 2$.
- Priors $\beta_h \sim N(0, 10^2)$, $\beta_x \sim N(0, 10^2)$. For Normal priors and a Normal measurement of the sum, $E[\beta_h\mid D] = \dfrac{s_h^2}{s_h^2 + s_x^2 + 2^2}\, u = \dfrac{100}{204} \times 20 \approx 9.80$, and the same for $\beta_x$. Correlation $-0.96$.
- Priors with sd 20 for the holiday and 5 for the promotion: $E[\beta_h\mid D] = \dfrac{400}{400 + 25 + 4} \times 20 \approx 18.6$ and $E[\beta_x\mid D] = \dfrac{25}{429} \times 20 \approx 1.2$.
- Same data, different priors, very different stories ("it's the holiday" vs "half and half"). The sum is about 20 in both (19.6 and 19.8).
- Next year the promotion runs alone on an ordinary week. The first model forecasts +9.8, the second +1.2. Neither number was measured.
- Fix: if the promotion had also run on a few ordinary days, those days measure $\beta_x$ alone and the split becomes identified (try it in the widget).
In an additive regression-style model, components compete for the same pattern when their columns in the design matrix are (nearly) linearly dependent: collinearity (Chapter 5.13), here called confounding between components. Typical pairs in a forecasting model:
- Trend vs yearly seasonality with one or two years of history: a rise at the end of the series could be either.
- Holiday vs regressor when a promotion or price change always coincides with the holiday.
- Regressor vs trend when the regressor (marketing spend, number of stores) grows steadily over time.
- Changepoint vs holiday when a candidate changepoint sits right at a holiday: a slope change can absorb the holiday bump, and vice versa.
- Intercept vs a full set of category effects (next section).
What stays well estimated: the total fitted value on the days the components co-occur, and forecasts for future days with the same combination. What does not: each component's separate effect, and forecasts when the components decouple.
Why do we need it?
Stakeholders ask "how much did the promotion add?". If the promotion never ran without the holiday, the honest answer is "we cannot separate them". Knowing this prevents confident but invented attributions, and bad forecasts when next year's calendar is different.
Where is it used?
Prophet-style and other additive forecasting models, marketing-mix models (channels whose spend moves together), price-and-promotion models, and any regression with collinear features where a coefficient is read as an effect.
How is it used?
Before trusting a component, look at the posterior correlation between components and at their prior vs posterior sds; plot the components separately; check whether the history contains periods where they vary independently; and state attributions with that caveat.
"The forecast fits the history well, so the component estimates are trustworthy."
A good fit only shows the sum of the components is right on the days seen. Individual components can be prior-driven and still produce a perfect fit.
"The model says the holiday adds 18.6 units."
"Holiday plus promotion together add about 20. The split is set mainly by our priors, because the promotion never ran without the holiday."
"Our model attributes most of the November lift to the holiday, so the promotion was not worth it."
If the promotion always coincided with the holiday, the model cannot separate them; the attribution reflects the priors (and, for SVI, possibly the initialization). Check the posterior correlation between the two effects and whether any period has one without the other.
Model answer: "In an additive model, components that always co-occur in the history are not separately identified; only their sum is. The posterior then shows a strong negative correlation between the two effects and posterior sds close to the priors. I would report the combined effect, test the sensitivity of the split to the prior scales, and, if the separate effect matters for a decision, run the promotion on some non-holiday days or bring in outside information as an informative prior."
This is the identifiability question in your forecasting model. A late-year increase can be absorbed by a changepoint slope change $\delta_j$ (especially if a grid or PELT candidate sits nearby), by the yearly Fourier terms (with one or two years of history they are poorly separated from the trend), by a holiday effect, or by an exogenous regressor that moves with the season. The Laplace prior on $\delta_j$ and the prior scales on holidays and $\beta$ decide the split where the data cannot. Check the posterior correlations between these components, and avoid placing changepoint candidates right on holidays unless you mean to.
Components that always co-occur ⇒ only their sum is identified (collinear design columns).
Holiday + promo always together, sum 20 (se 2): priors 10/10 → 9.8/9.8; priors 20/5 → 18.6/1.2. $E[\beta_h\mid D] = \frac{s_h^2}{s_h^2+s_x^2+se^2}\,u$.
Fix with data where they vary separately; report sums; beware forecasts when components decouple.
Quick check: in the widget, why does adding a single ordinary-day promotion change the ellipse so much?
That day has $x_t = 1$ and $h_t = 0$, so it measures $\beta_x$ by itself (with sd 4). Together with the holiday days, which measure $\beta_h + \beta_x$, the two effects are now separately identified: the ridge is cut across, and the posterior is no longer free to slide along it.
Weak identification: symptoms and fixes (priors as identification aids)
Once you have found a direction the data cannot see, you have four honest options. Get data that see it (a promotion on an ordinary day, another year of history). Remove the duplicate by a constraint that says what each parameter means (weekday effects that average to zero, so the overall level lives only in the intercept). Use a prior you can defend from outside knowledge (a prior is an identification aid, not a cheat, if it comes from real information). Or merge the components and report only what is identified.
The classic small example: a model with an intercept $m$ and a separate effect for each of the 7 weekdays. Adding 10 to $m$ and subtracting 10 from every weekday effect gives exactly the same predictions, so $m$ and the weekday effects are not separately identified. Requiring the weekday effects to sum to zero removes that duplicate direction.
Three ways to say it:
- Picture: a dial that turns without changing anything; a constraint glues it in place.
- Numbers: 4 weeks of data: without the constraint the level's sd is 6.0 (prior sd 10); with it, 0.38.
- Slogan: say what each parameter means, and the data can measure it.
Intercept plus weekday effects, 28 days (4 of each weekday), noise sd 2. Priors: level $m \sim N(0, 10^2)$, each weekday effect $N(0, 20^2)$.
- What the data measure: each weekday's level $m + s_d$, from 4 days each, with standard error $2/\sqrt 4 = 1$. Posterior sd of $m + s_{\text{Mon}}$: 1.00.
- Free model: posterior sd of $m$ is 6.0 and of $s_{\text{Mon}}$ 6.1, with correlation $-0.99$. The level has contracted only from prior variance 100 to 36 (contraction 0.64): the split between "level" and "weekday" is mostly prior.
- Sum-to-zero model (the 7 effects are constrained to average 0): sd of $m$ = $2/\sqrt{28} \approx 0.38$, sd of $s_{\text{Mon}} \approx 0.93$, correlation 0.
- The identified quantity $m + s_{\text{Mon}}$ still has sd 1.00 in both versions: predictions do not change, only the meaning of the parameters (now $m$ = the average day).
- The same holds for Fourier seasonality over whole periods: sines and cosines average to zero over a full cycle, so they do not compete with the level. Over a short, partial history they partly do.
Ways to deal with non- or weak identification:
- Design / data: collect observations where the confounded components vary separately (experiments, longer history, holdout weeks without promotions).
- Constraints: sum-to-zero (or a reference level) for category effects; centring regressors (subtract their mean) so they do not compete with the intercept; fixing signs or orderings in latent-factor models. In NumPyro, a simple sum-to-zero is
s = s_raw - s_raw.mean(). - Priors as identification aids: informative priors from outside data (last year's holiday lift, a pilot test of the promotion) legitimately pin down a direction the current data cannot. Document their source; run a sensitivity analysis.
- Reparameterize to identified combinations (sum and difference) so that samplers and guides are not fighting a ridge.
- Simplify: merge components that cannot be separated and report their joint effect.
- Posterior contraction $c = 1 - Var_{\text{post}}/Var_{\text{prior}}$ is a quick diagnostic: near 1 = learned from data; near 0 = still the prior.
Why do we need it?
Unidentified directions waste computation (slow MCMC along ridges, unstable SVI) and invite wrong interpretations. Fixing them makes parameters mean something, speeds up inference, and makes the remaining uncertainty honest.
Where is it used?
Category and seasonal effects with an intercept (PyMC's ZeroSumNormal, reference levels in regression), centred regressors in forecasting models, marketing-mix models with informative priors from experiments, and factor models with sign constraints.
How is it used?
Compute prior vs posterior sds (contraction) and posterior correlations for every parameter; for each weak direction choose one fix (constraint, centring, outside-information prior, more data, or merging), refit, and confirm that predictions did not change while the parameters became stable.
"A sum-to-zero constraint changes the model's predictions."
It removes a direction the data could not see; the identified combinations (each weekday's level) and therefore the predictions are essentially unchanged. What changes is the meaning of the parameters ($m$ = the average day).
"Using an informative prior to identify a parameter is cheating."
It is legitimate when the prior comes from real outside information (a past experiment, last year's lift) and is documented and checked for sensitivity. It is cheating only when it is chosen to produce a desired answer.
"Mean-field SVI handles ridges fine because it does not need to mix."
A mean-field guide cannot represent a correlation of $-0.99$; it typically reports each parameter far too precisely along the ridge (Chapter 6.13). Removing the ridge helps SVI as much as it helps NUTS.
Two direct links. In the forecasting model, Fourier seasonal terms average to zero over full periods, which keeps them from competing with the trend's level; centring (or globally standardizing) regressors keeps $X_t\beta$ from competing with the intercept, and informative priors from past promotions or holidays are legitimate identification aids. In the A/B framework, the syllabus warning about per-group normalization is an identifiability issue: if each group is standardized separately, every group's mean becomes 0, so the group effect is no longer in the data at all. One global scaler keeps the differences between groups visible to the model.
Fixes: more separating data; constraints (sum-to-zero, centring, reference level); defensible informative priors; reparameterize to identified combinations; merge.
Weekday example: free sd(m) = 6.0 (contraction 0.64) → sum-to-zero 0.38; sd(m + s_Mon) = 1.00 either way (predictions unchanged).
Contraction $1 - Var_{\text{post}}/Var_{\text{prior}}$ near 0 = still the prior.
Quick check: you add a regressor "number of stores open", which grew steadily over two years, to a model with a linear trend. What do you expect, and what could you do?
The regressor and the trend are nearly collinear, so their coefficients will be strongly negatively correlated and weakly identified: the model cannot tell "more stores" from "time passing". Options: centre and scale the regressor, use per-store demand instead, put an informative prior on the store effect from openings data, or accept and report only their combined effect.
Recap, cheat sheet and practice
- A posterior predictive check simulates replicated datasets $y^{\text{rep}}$ from the fitted model (a new posterior draw for each) and compares them with the real data. Real data that stand out reveal a feature the model cannot produce. Prior predictive checks (Chapter 6.2) do the same before fitting.
- A test statistic $T$ (mean, variance, max, zero count, group spread, weekday gap…) turns the comparison into a number; the posterior predictive p-value $P(T(y^{\text{rep}}) \ge T(y)\mid y)$ near 0 or 1 flags misfit. It is a conservative diagnostic, not the probability that the model is true.
- Choose statistics tied to your decisions and not fitted directly by a parameter: the mean always passes a model with a free mean.
- Prior sensitivity analysis: refit under a pre-chosen set of reasonable priors and report whether the conclusion changes. Sensitivity is largest with small data, scale parameters (Laplace $b$, holiday scales, $\tau$), dispersion and tail parameters, and sparse effects; a sensitivity curve shows it.
- Identifiability: if different parameter values give the same data distribution, the data cannot separate them; along that direction the posterior is the prior. Symptoms: correlations near ±1, little contraction, slow chains, unstable SVI.
- In additive forecasting models, components that always co-occur (holiday and promotion; trend and yearly season over a short history; a changepoint at a holiday) are only identified as a sum: the priors write the split.
- Fixes: data where components vary separately, constraints (sum-to-zero, centring), defensible informative priors (identification aids), reparameterizing, or merging and reporting the sum.
Cheat sheet
| Idea | Formula / code | Remember |
|---|---|---|
| Replicated data | $\theta^{(s)} \sim p(\theta\mid y)$, $y^{\text{rep}(s)} \sim p(y\mid\theta^{(s)})$ | same size and design as $y$; NumPyro Predictive(model, samples) |
| Posterior predictive p-value | $P(T(y^{\text{rep}}) \ge T(y)\mid y) \approx$ share of replicates | near 0 or 1 = misfit; conservative; not P(model true) |
| Count-model checks | variance, zero days, max | Poisson: var = mean; NB2: var = μ + μ²/α |
| Useless statistic | flat-prior Normal: ppp(mean) = 0.5 always | avoid statistics a parameter fits directly |
| Prior sensitivity | refit under reasonable alternatives | 3/20 vs 8/20: P(B > A) 0.77–0.96; 10× data: all ≈ 1 |
| Sensitivity curve | output vs prior scale (log axis) | holiday seen 2×: 0.2 → 16.5; seen 20×: 14.0 → 17.8 |
| Identifiable | $p(y\mid\theta_1) = p(y\mid\theta_2)\ \forall y \Rightarrow \theta_1 = \theta_2$ | else a flat likelihood direction; posterior there = prior |
| $a + b$ model | sd(a + b) = 0.1, sd(a − b) = 14.1, corr = −0.9999 | report the identified combination |
| Credit split (Normal priors) | $E[\beta_h\mid D] = \frac{s_h^2}{s_h^2+s_x^2+se^2}\,u$ | the priors write the split of the sum $u$ |
| Contraction | $1 - Var_{\text{post}}/Var_{\text{prior}}$ | near 0 = still the prior |
| Sum-to-zero | s = s_raw - s_raw.mean() | same predictions, identified level |
import numpy as np
from scipy import stats
import jax
import jax.numpy as jnp
import numpyro
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS, Predictive
# 30 days of orders for a small store (overdispersed: variance about 4x the mean)
y = np.array([11, 0, 1, 8, 8, 2, 9, 11, 0, 1, 0, 0, 6, 0, 0,
13, 2, 7, 2, 5, 2, 1, 2, 11, 0, 6, 1, 4, 2, 5])
print(y.mean(), round(y.var(ddof=1), 2), (y == 0).sum(), y.max()) # 4.0 16.41 7 13
def poisson_model(n, y=None):
lam = numpyro.sample("lam", dist.Gamma(2.0, 0.5)) # Gamma(concentration, RATE)
with numpyro.plate("days", n):
numpyro.sample("y", dist.Poisson(lam), obs=y)
def nb_model(n, y=None):
mu = numpyro.sample("mu", dist.Gamma(2.0, 0.5))
alpha = numpyro.sample("alpha", dist.Gamma(2.0, 0.5)) # NB2: Var = mu + mu^2 / alpha
with numpyro.plate("days", n):
numpyro.sample("y", dist.NegativeBinomial2(mu, alpha), obs=y)
# 1) Posterior predictive check: fit, replicate the dataset, compare test statistics
T = {"mean": lambda d: d.mean(-1), "variance": lambda d: d.var(-1, ddof=1),
"zero days": lambda d: (d == 0).sum(-1), "max": lambda d: d.max(-1)}
def ppc(model, seed=0):
mcmc = MCMC(NUTS(model), num_warmup=500, num_samples=2000, progress_bar=False)
mcmc.run(jax.random.PRNGKey(seed), n=len(y), y=jnp.array(y))
# y=None in the call below, so the "y" site is SIMULATED: one fake month per posterior draw
yrep = Predictive(model, mcmc.get_samples())(jax.random.PRNGKey(seed + 1), n=len(y))["y"]
yrep = np.asarray(yrep) # shape (2000, 30)
return {k: round(float((f(yrep) >= f(y)).mean()), 3) for k, f in T.items()} # ppp = P(T(y_rep) >= T(y) | D)
print("Poisson ppp:", ppc(poisson_model))
# Poisson ppp: {'mean': 0.503, 'variance': 0.0, 'zero days': 0.0, 'max': 0.007}
print("NB ppp: ", ppc(nb_model))
# NB ppp: {'mean': 0.484, 'variance': 0.502, 'zero days': 0.376, 'max': 0.741}
# 2) Prior sensitivity: the same A/B data under four reasonable priors (Monte Carlo P(B > A))
def p_b_beats_a(a0, b0, kA, nA, kB, nB, draws=400_000, seed=0):
rng = np.random.default_rng(seed)
tA = rng.beta(a0 + kA, b0 + nA - kA, draws)
tB = rng.beta(a0 + kB, b0 + nB - kB, draws)
return round(float((tB > tA).mean()), 3)
priors = {"flat (1,1)": (1, 1), "Jeffreys (.5,.5)": (0.5, 0.5), "weak (2,18)": (2, 18), "strong (20,180)": (20, 180)}
for data in [(3, 20, 8, 20), (30, 200, 80, 200)]:
print(data, {name: p_b_beats_a(*ab, *data) for name, ab in priors.items()})
# (3, 20, 8, 20) {'flat (1,1)': 0.957, 'Jeffreys (.5,.5)': 0.962, 'weak (2,18)': 0.93, 'strong (20,180)': 0.775}
# (30, 200, 80, 200) {... all 1.0} <- with 10x the data the prior no longer matters
# 3) A sensitivity curve: holiday effect under a Laplace(0, b) prior, seen k times
grid = np.linspace(-100, 160, 10_401)
def holiday_posterior_mean(dhat, k, b, noise_sd=20.0):
logw = -np.abs(grid) / b + stats.norm.logpdf(dhat, grid, noise_sd / np.sqrt(k))
w = np.exp(logw - logw.max())
return float((grid * w).sum() / w.sum())
for k in [2, 20]:
print(k, [round(holiday_posterior_mean(18.0, k, b), 1) for b in [1, 5, 20, 100]])
# 2 [0.2, 3.2, 11.4, 16.5] <- steep: the prior scale decides
# 20 [2.5, 14.0, 17.0, 17.8] <- flat for any b >= 5: the data decide
# 4) Identifiability: y ~ Normal(a + b, 1). The data pin down a + b, not a and b.
def ab_model(y):
a = numpyro.sample("a", dist.Normal(0.0, 10.0))
b = numpyro.sample("b", dist.Normal(0.0, 10.0))
numpyro.sample("y", dist.Normal(a + b, 1.0), obs=y)
y_ab = 3.0 + jax.random.normal(jax.random.PRNGKey(5), (100,))
mcmc = MCMC(NUTS(ab_model), num_warmup=500, num_samples=2000, progress_bar=False)
mcmc.run(jax.random.PRNGKey(0), y_ab)
s = mcmc.get_samples()
a, b = np.asarray(s["a"]), np.asarray(s["b"])
print(round(a.std(), 1), round(np.corrcoef(a, b)[0, 1], 4), round((a + b).std(), 2), round((a - b).std(), 1))
# 7.0 -0.9999 0.1 14.1 <- sd(a) and sd(a - b) are the PRIOR's; only a + b was learned
1. What is a posterior predictive check?
2. A Poisson model of daily orders gives ppp = 0.50 for the statistic "mean". What can you conclude?
3. For a count model, the ppp for "number of zero days" is 0.002. What does it suggest?
4. Under a flat prior $P(\theta_B \gt \theta_A\mid D) = 0.97$; under an informative prior from past tests it is 0.81. Your rule ships at 0.95. What should you report?
5. In the model $y_i \sim N(a + b, \sigma^2)$ with 10 000 observations, what do the data identify?
6. In your history, a promotion always ran on the holiday. The fitted model says the holiday adds 18 and the promotion 2. Which statement is right?
Practice problems
A. For 30 days with mean 4, how many zero days do you expect under Poisson(4), and under NB2 with $\mu = 4$, $\alpha = 1$? The real month had 7. What do you conclude?
- Poisson: $P(0) = e^{-4} \approx 0.018$; expected zero days $30 \times 0.018 \approx 0.55$.
- NB2: $P(0) = \left(\frac{\alpha}{\alpha + \mu}\right)^{\alpha} = \frac{1}{5} = 0.2$; expected zero days $30 \times 0.2 = 6$.
- 7 zero days are wildly unlikely under the Poisson and typical under the NB: the data need the extra day-to-day variation the NB allows.
B. Of 2 000 replicated months, 14 had a maximum day at least as large as the real maximum (13). Compute the ppp and interpret it.
$ppp = 14/2000 = 0.007$. Fewer than 1% of fake months reach the real peak: the model's tail is too light for this data. For capacity planning that matters, because the model would under-state the chance of very busy days. A heavier-tailed likelihood (NB for counts, Student-t for continuous data) is the usual next step.
C. Show that with Normal data, known $\sigma$ and a flat prior on $\mu$, the posterior predictive p-value for $T(y) = \bar y$ equals 0.5 for every dataset.
- Posterior: $\mu\mid y \sim N(\bar y, \sigma^2/n)$.
- Replicated mean given $\mu$: $\bar y^{\text{rep}} \sim N(\mu, \sigma^2/n)$.
- Combining: $\bar y^{\text{rep}}\mid y \sim N(\bar y, 2\sigma^2/n)$, symmetric around $\bar y$.
- So $P(\bar y^{\text{rep}} \ge \bar y\mid y) = 0.5$ exactly. The check can never fail, which is why statistics fitted directly by a parameter are useless.
D. A converts 3 of 20. Find the posterior mean of $\theta_A$ under Beta(1, 1) and under Beta(20, 180), and explain the gap.
Beta(1 + 3, 1 + 17) = Beta(4, 18): mean $4/22 \approx 0.182$. Beta(20 + 3, 180 + 17) = Beta(23, 197): mean $23/220 \approx 0.105$. The informative prior is worth 200 pseudo-visitors at 10%, ten times the real data, so it pulls the estimate most of the way to 10%. With 2 000 real visitors the two means would be close: that is how sensitivity fades with data.
E. In $y_i \sim N(a + b, 1)$ with a very large $n$ and $\bar y = 5$, the priors are $a \sim N(0, 3^2)$ and $b \sim N(0, 1^2)$. What are the posterior means of $a$ and $b$, and roughly what is $sd(a)$?
The data fix $a + b = 5$. Along that line the prior decides: the split is proportional to the prior variances, $a \approx 5 \times 9/10 = 4.5$ and $b \approx 0.5$. The prior, conditioned on the sum, leaves $sd(a) = \frac{3 \times 1}{\sqrt{9 + 1}} \approx 0.95$, and the correlation of $a$ and $b$ is close to $-1$. All of this is prior; the data contributed only the 5.
F. (Interview) "How would you check whether your forecasting model's likelihood and priors are adequate?"
"First a prior predictive check, to make sure the trend, changepoint and seasonality priors produce plausible demand. After fitting, posterior predictive checks on the training history with statistics tied to the decisions: number of zero-demand days and the variance for the count likelihood, the largest days and high quantiles for the tails, the weekday pattern, behaviour around holidays, and autocorrelation of residuals; ppp values near 0 or 1 tell me what is missing, and forecast-specific checks follow on held-out periods. Then a prior sensitivity analysis on the parameters the data say least about: the Laplace scale on changepoint slope changes, holiday prior scales for rare holidays, and the dispersion prior, plotting key forecasts against each scale. Finally, I look at posterior correlations between trend, seasonality, holidays and regressors; where two components always co-occur I report their combined effect rather than the split."
Why approximate inference? MCMC from scratch
Chapters 6.1–6.8 told you what to compute: a posterior. This chapter starts the computation half of the guide. You will see why the posterior of a real model cannot be computed exactly, meet the four standard ways around the problem, and then build the most important one, Markov chain Monte Carlo, from nothing: a three-state chain, then the Metropolis algorithm, then the practical skills of warmup, autocorrelation and reading trace plots. Chapter 6.10 upgrades the engine to Hamiltonian Monte Carlo and NUTS, the sampler NumPyro actually runs.
- Explain why the posterior of a real model is intractable: the evidence is an integral, and a grid over $d$ parameters costs $k^d$ evaluations
- Use a bag of posterior draws to answer any question (means, probabilities, intervals) and know that the error shrinks like $1/\sqrt{S}$
- Name the menu of approximations: MCMC, variational inference, the Laplace approximation and importance sampling, with what each returns and how each fails
- Build a Laplace approximation (a bell curve at the peak, width from the curvature) and an importance sampler (weights, weight degeneracy) by hand
- Watch a Markov chain forget its start and settle into its stationary distribution
- Write the Metropolis algorithm from scratch, explain why it targets the posterior without ever knowing $p(D)$, and tune its step size
- Handle burn-in / warmup, read autocorrelation, see through the thinning myth, and read a trace plot
What we need from earlier chapters: the posterior, the evidence $p(D)$ and the grid approximation (Chapter 6.1, especially the posterior and the evidence); conjugate models, which are the rare exact cases (Chapter 6.3); the Law of Large Numbers and the Central Limit Theorem (Chapter 4.13); standard errors (Chapter 5.5); MAP estimation (Chapter 5.2); second derivatives and the Hessian (Calculus 2.10) and Taylor series (Calculus 2.11); eigenvectors (Linear Algebra 1.11); autocorrelation of a series (Chapter 7.3, short refresher below). Words used everywhere in this chapter: a draw (or sample) is one random value produced by a computer from a distribution; a target is the distribution we want draws from (here always the posterior); approximate inference means computing an answer that is close to the exact posterior answer, with an error we can control or at least check.
Why the exact posterior is usually out of reach core
Bayes' theorem says the posterior is "likelihood × prior, divided by the evidence". The top part is easy: for any parameter value θ you choose, a computer can multiply the likelihood and the prior in a fraction of a millisecond. The bottom part, the evidence $p(D)$, is the hard part. It is the total area (or volume) under likelihood × prior, added up over every possible θ.
Think of a mountain range. You can stand at any GPS point and read the height. But someone asks for the total volume of rock. With one coordinate (walking along a line) you can measure every 10 metres and add up. With two coordinates you need a grid of squares. Your forecasting model can easily have dozens of "coordinates" (parameters), and a grid with 100 points per axis in 40 or more dimensions has more points than there are atoms in the universe (about $10^{80}$). That is what statisticians mean when they say the posterior is intractable: we can write it down, but we cannot compute it exactly in any reasonable time.
Three ways to say it:
- Picture: you can read the height of the posterior landscape anywhere, but you cannot measure the volume of the whole landscape.
- Numbers: 100 grid points per parameter: 1 parameter needs 100 evaluations; 10 parameters need $10^{20}$, about 3 000 years at a billion evaluations per second.
- Slogan: the shape is cheap, the total is expensive; approximate inference avoids computing the total.
How fast does a grid explode? Use $k = 100$ grid values per parameter (a coarse grid) and a very fast computer that evaluates likelihood × prior a billion ($10^9$) times per second.
- $d = 1$ parameter (one conversion rate, Chapter 6.1): $100^1 = 100$ evaluations. Instant.
- $d = 2$ (two rates, A and B): $100^2 = 10\,000$. Still instant.
- $d = 3$: $100^3 = 10^6$, one millisecond.
- $d = 10$ (eight segments plus a population mean and spread, Chapter 6.5): $100^{10} = 10^{20}$ evaluations. At $10^9$ per second that is $10^{11}$ seconds; one year is about $3.16 \times 10^7$ seconds, so $10^{11}/(3.16\times 10^7) \approx 3\,170$ years.
- A Prophet-style forecasting model with the Prophet defaults has 25 changepoint slopes, 20 yearly Fourier coefficients (order 10) and 6 weekly ones (order 3), plus the base slope, offset and noise scale: already $25 + 20 + 6 + 3 = 54$ parameters before any holidays or regressors. $100^{54} = 10^{108}$. No computer will ever do this.
- And "1 evaluation" is generous: each one sums a log-likelihood over every data point (every day of history, every user).
Write $\tilde p(\theta) = p(D\mid\theta)\,p(\theta)$ for the unnormalized posterior ("p tilde"). Then
$$p(\theta\mid D) = \frac{\tilde p(\theta)}{p(D)}, \qquad p(D) = \int \tilde p(\theta)\,d\theta = \int p(D\mid\theta)\,p(\theta)\,d\theta .$$- What we can compute cheaply: $\log \tilde p(\theta) = \log p(D\mid\theta) + \log p(\theta)$ at any θ we choose, and (with automatic differentiation, Calculus 2.9) its gradient $\nabla_\theta \log\tilde p(\theta)$.
- What we actually want is almost always a posterior expectation, an average over the posterior: $E[f(\theta)\mid D] = \int f(\theta)\,p(\theta\mid D)\,d\theta$. The posterior mean ($f(\theta) = \theta$), a probability such as $P(\theta_B \gt \theta_A\mid D)$ ($f$ = 1 when $\theta_B \gt \theta_A$, else 0), and the posterior predictive (Chapter 6.1) are all of this form.
- Intractable means: the integral has no formula (no closed form) and brute-force numerical integration is far too slow. A grid with $k$ points per axis needs $k^d$ evaluations in $d$ dimensions: exponential growth, the "curse of dimensionality".
- Exceptions: conjugate models (Beta-Binomial, Dirichlet-Multinomial, Gamma-Poisson, Normal-Normal; Chapter 6.3) have exact formulas. Almost everything else (hierarchical models, Student-t or Negative Binomial likelihoods, any model with many parameters) does not.
Why do we need it?
If you do not see why exact computation fails, MCMC and SVI look like mysterious rituals. Once you see that the only obstacle is one giant integral, every method in this half of the guide becomes "a clever way to avoid that integral".
Where is it used?
Every non-conjugate Bayesian model: hierarchical A/B segment models, logistic and Poisson regressions with priors, Prophet-style forecasting models, Bayesian neural networks, topic models. NumPyro, PyMC and Stan exist because of this problem.
How is it used?
Check first whether your model is conjugate (then use the exact formula). Otherwise give the software only $\log\tilde p(\theta)$ (NumPyro builds it from your sample statements) and let an approximate method (NUTS, SVI) produce draws or a fitted distribution.
"Intractable means the posterior cannot be written down."
We can write the formula $p(\theta\mid D) = \tilde p(\theta)/p(D)$ perfectly well and evaluate $\tilde p$ anywhere. What we cannot do is compute the integral $p(D)$, or any posterior average, exactly in a reasonable time.
"A finer grid, or a faster computer, will fix it."
The cost $k^d$ grows exponentially in $d$. A computer a million times faster buys you about 3 extra parameters at $k = 100$ ($100^3 = 10^6$). Methods whose cost grows slowly with $d$ are the only way out.
"Without $p(D)$ we know nothing about the posterior."
We know its shape exactly: ratios such as $p(\theta_1\mid D)/p(\theta_2\mid D) = \tilde p(\theta_1)/\tilde p(\theta_2)$ need no $p(D)$. MCMC is built entirely on such ratios.
In an A/B framework like yours, a single conversion metric with a Beta prior is conjugate: the posterior Beta$(\alpha+k, \beta+n-k)$ is exact and no approximation is needed. As soon as you add hierarchical partial pooling across segments, a Student-t likelihood for revenue, or a Poisson model with a non-conjugate prior, the evidence integral has no formula. Your forecasting model is the extreme case: with changepoint slopes $\delta_j$, Fourier coefficients, holiday effects and regressor coefficients it can easily have dozens of parameters. That is the usual reason to reach for SVI (Chapters 6.11–6.14), with NUTS (Chapter 6.10) as the tool for checking it.
"We use MCMC because Bayes' theorem cannot be applied to complex models."
Bayes' theorem applies to every model. The obstacle is computational: the normalizing integral $p(D)$ and posterior averages are high-dimensional integrals without closed forms.
Model answer: "The posterior is likelihood times prior divided by the evidence. The numerator is cheap to evaluate pointwise, but the evidence and every posterior expectation are integrals over all parameters, and grid integration costs $k^d$. So we either draw samples whose distribution approaches the posterior (MCMC) or fit a tractable distribution by optimization (VI), and neither needs the evidence."
$p(\theta\mid D) = \tilde p(\theta)/p(D)$, $\tilde p = $ likelihood × prior (cheap anywhere), $p(D) = \int\tilde p\,d\theta$ (expensive).
Grid cost $k^d$: 100 points, 10 parameters → $10^{20}$ evaluations ≈ 3 000 years at $10^9$/s.
We want expectations $E[f(\theta)\mid D]$; approximate inference gets them without $p(D)$. Trap: "intractable" ≠ "unknown formula".
Quick check: a model has 6 parameters and you use a grid of 50 values per parameter. How many evaluations is that, and how long at a million per second?
$50^6 = 15\,625\,000\,000 \approx 1.56 \times 10^{10}$ evaluations. At $10^6$ per second: $1.56\times 10^{4}$ seconds ≈ 4.3 hours. Feasible but painful, and one more parameter multiplies it by 50 (about 9 days).
A bag of draws is as good as the posterior: Monte Carlo estimates core
If you cannot compute the posterior curve, carry around a big bag of values drawn from it instead. Values in the high part of the curve appear often in the bag, values in the tails appear rarely. Now every question becomes counting: "what is the average rate?" → average the values in the bag; "how likely is it that the rate is above 10%?" → count the share of values above 10%; "what is a 90% interval?" → sort the bag and read off the 5th and 95th percent positions.
This trick is called Monte Carlo (after the casino: it uses randomness). It turns integrals into averages. Its accuracy depends on the number of draws $S$, not on the number of parameters. MCMC (later in this chapter) and NumPyro's mcmc.get_samples() give you exactly such a bag.
Three ways to say it:
- Picture: a posterior is a crowd; a bag of draws is a random survey of that crowd.
- Numbers: with 4 000 independent draws, a posterior mean is pinned down to about ±2 × sd/√4000 ≈ ±3% of a posterior sd.
- Slogan: integrals become averages; error shrinks like 1/√S.
Ten draws from the posterior Beta(5, 35) (Chapter 6.1's checkout example: prior Beta(2, 18), 3 buyers in 20 visitors). A computer gave: 0.226, 0.084, 0.159, 0.150, 0.144, 0.206, 0.146, 0.185, 0.048, 0.054.
- Posterior mean ≈ average of the draws $= 1.402/10 = 0.140$. (Exact: $5/40 = 0.125$.)
- $P(\theta \gt 0.10\mid D)$ ≈ share of draws above 0.10 $= 7/10 = 0.7$. (Exact: 0.650.)
- Sorted: 0.048, 0.054, 0.084, 0.144, 0.146, 0.150, 0.159, 0.185, 0.206, 0.226. Posterior median ≈ $(0.146 + 0.150)/2 = 0.148$.
- How far off can 10 draws be? The posterior sd is 0.0516, so the standard error of the average is $0.0516/\sqrt{10} = 0.016$. Our miss, $0.140 - 0.125 = 0.015$, is about one standard error: normal bad luck.
- With $S = 4\,000$ draws the standard error is $0.0516/\sqrt{4000} = 0.0008$: the estimate would be $0.125 \pm 0.0016$ (two standard errors).
Given draws $\theta^{(1)}, \dots, \theta^{(S)}$ from $p(\theta\mid D)$, the Monte Carlo estimate of a posterior expectation is the plain average:
$$E[f(\theta)\mid D] = \int f(\theta)\,p(\theta\mid D)\,d\theta \;\approx\; \frac{1}{S}\sum_{s=1}^{S} f\big(\theta^{(s)}\big).$$- Why it works: the Law of Large Numbers (Chapter 4.13) says the average approaches the expectation as $S$ grows.
- How accurate: for independent draws, the Central Limit Theorem gives the Monte Carlo standard error $\text{sd}(f(\theta))/\sqrt{S}$. For a probability $P(A\mid D)$ estimated by a share $\hat q$, it is $\sqrt{\hat q(1-\hat q)/S}$.
- The error does not depend on the number of parameters $d$: this is why sampling beats grids in high dimensions.
- MCMC draws are not independent (each draw depends on the previous one). They still work, but they carry less information per draw; Chapter 6.10 replaces $S$ by the effective sample size.
Why do we need it?
Every sampler, MCMC included, returns draws, not formulas. Monte Carlo averaging is how those draws become the numbers you report: posterior means, $P(\theta_B \gt \theta_A\mid D)$, credible intervals and forecast quantiles.
Where is it used?
$P(B \gt A)$ in Bayesian A/B testing (Chapter 6.4), fan charts of forecasts (draws of future paths), NumPyro's Predictive, ArviZ summaries, Monte Carlo ELBO estimates in SVI (Chapter 6.12), and risk numbers in finance.
How is it used?
Get draws (mcmc.get_samples() or draws from a fitted guide), apply the function you care about to every draw, and average: np.mean(theta_B > theta_A), np.quantile(theta, [0.05, 0.95]). Report a Monte Carlo error alongside when precision matters.
"More draws make the posterior narrower."
The posterior's width is set by the data and the prior. More draws make your estimate of the posterior (its mean, its quantiles) more precise. The histogram gets smoother, not narrower.
"Monte Carlo error is negligible, so report 6 decimals."
With 1 000 draws a probability like 0.65 has a standard error of about 0.015. Report $P(B \gt A) \approx 0.65$, not 0.6512, unless you have many more (effective) draws.
"MCMC draws are independent, so sd/√S is the error."
MCMC draws are correlated with their neighbours; the honest error is sd/√ESS with the effective sample size from Chapter 6.10, which is usually much smaller than $S$.
$E[f(\theta)\mid D] \approx \frac1S\sum_s f(\theta^{(s)})$; probabilities = shares; intervals = sorted draws.
Independent draws: error = sd/√S (probability: $\sqrt{q(1-q)/S}$), whatever the number of parameters.
Trap: more draws ≠ narrower posterior; MCMC needs ESS instead of S.
Quick check: with 2 500 independent draws you estimate $P(\theta_B \gt \theta_A\mid D) = 0.90$. What is its Monte Carlo standard error?
$\sqrt{0.9 \times 0.1 / 2500} = \sqrt{0.000036} = 0.006$. So the probability is $0.90 \pm 0.012$ (two standard errors), which is accurate enough for most decisions.
The Laplace approximation: a bell curve at the peak
Zoom in on the top of any smooth hill and it looks like an upside-down bowl. In the language of logs: near its peak, the log posterior looks like an upside-down parabola, and a parabola in the log is exactly a Normal (bell) curve. So: find the peak, measure how sharply the log posterior bends there, and use a Normal with that centre and the matching width. A sharp peak (strong bending) means a narrow bell; a gentle peak means a wide one.
Finding the peak is an optimization problem (it is the MAP estimate of Chapter 5.2), and the bending is a second derivative (in many dimensions, the Hessian matrix; Calculus 2.10).
Three ways to say it:
- Picture: put a tent on the summit, with its slope matched to the summit's roundness.
- Numbers: for Beta(5, 35) the peak is at 0.105 and the bending is 403, so the bell has sd $1/\sqrt{403} = 0.050$ (exact sd 0.052).
- Slogan: centre = MAP, width = 1/√curvature.
Laplace for the checkout posterior Beta(5, 35). The log posterior (up to a constant) is $\ell(\theta) = 4\log\theta + 34\log(1-\theta)$.
- Slope: $\ell'(\theta) = 4/\theta - 34/(1-\theta)$. Set it to 0: $4(1-\theta) = 34\theta$, so $\hat\theta = 4/38 = 0.1053$ (the mode).
- Bending: $\ell''(\theta) = -4/\theta^2 - 34/(1-\theta)^2$. At $\hat\theta$: $4/0.1053^2 = 361.0$ and $34/0.8947^2 = 42.5$, so $\ell''(\hat\theta) = -403.5$.
- Width: $\sigma = 1/\sqrt{403.5} = 0.0498$. Laplace says $\theta \approx N(0.1053, 0.0498^2)$.
- Compare with the truth: exact mean 0.125 (Laplace 0.105), exact sd 0.0516 (Laplace 0.0498), exact $P(\theta \gt 0.15) = 0.284$ (Laplace 0.184). The bell is symmetric but the posterior has a long right tail, so the right tail is badly underestimated.
- Worse: the bell puts $P(\theta \lt 0) = \Phi(-0.1053/0.0498) = \Phi(-2.11) = 0.017$ on negative conversion rates.
- Same method on the log-odds scale $\eta = \log\frac{\theta}{1-\theta}$, which runs over all real numbers. Here the peak is at $\theta = 5/40 = 0.125$ ($\eta = -1.946$) and the sd of $\eta$ is $1/\sqrt{40 \times 0.125 \times 0.875} = 0.478$. Transformed back: $P(\theta \gt 0.15) = 0.329$ (error 0.045 instead of 0.100) and zero mass below 0. Same method, better coordinates, better answer.
Let $\ell(\theta) = \log\tilde p(\theta) = \log p(D\mid\theta) + \log p(\theta)$ and let $\hat\theta$ be its maximum (the MAP). A second-order Taylor expansion (Calculus 2.11) around $\hat\theta$, where the gradient is zero, gives
$$\ell(\theta) \approx \ell(\hat\theta) - \tfrac12 (\theta-\hat\theta)^\top A\,(\theta-\hat\theta), \qquad A = -\nabla^2\ell(\hat\theta),$$which is the log of a Gaussian density. The Laplace approximation is
$$p(\theta\mid D) \approx N\big(\hat\theta,\; A^{-1}\big), \qquad \log p(D) \approx \ell(\hat\theta) + \tfrac{d}{2}\log(2\pi) - \tfrac12\log\det A .$$- $A$ is the negative Hessian ("curvature matrix") at the peak; its inverse is the approximate posterior covariance. In one dimension, variance $= 1/(-\ell''(\hat\theta))$.
- Exact when the posterior is Gaussian; good when there is a lot of data (the posterior becomes close to Normal, a result known as the Bernstein–von Mises theorem); poor for skewed, bounded or multi-peaked posteriors.
- It depends on the coordinates: fitting on an unconstrained scale (log for positive parameters, log-odds for rates) usually works better. NumPyro's
AutoLaplaceApproximationworks on the unconstrained scale. - The evidence formula above is the root of the BIC (Bayesian information criterion).
Why do we need it?
It turns an optimizer (which you already have) into an approximate posterior in one extra step. It is the cheapest way to get error bars around a MAP fit, and a good first look before running anything expensive.
Where is it used?
NumPyro's AutoLaplaceApproximation guide, INLA (integrated nested Laplace approximations, popular in spatial statistics), Gaussian-process classification, the derivation of BIC, and the "standard errors from the Hessian" printed by many MLE routines.
How is it used?
Optimize $\log\tilde p$ to get $\hat\theta$; compute the Hessian there (autodiff or finite differences); invert its negative to get a covariance; draw from $N(\hat\theta, A^{-1})$ on the unconstrained scale and transform back. Check against NUTS when it matters.
"Laplace gives the posterior mean."
It is centred at the peak (the mode). For skewed posteriors the mean is elsewhere: 0.125 versus 0.105 in the example.
"The Laplace approximation is the same whichever parameterization I use."
The bell is fitted in the coordinates you choose. On the rate scale it puts mass on negative rates; on the log-odds scale it cannot. Always fit on an unconstrained scale.
"A small Hessian-based standard error means the parameter is well determined."
The curvature only describes the immediate neighbourhood of the peak. A posterior with a second peak, a long tail or a curved ridge (the banana) can be much wider than the Hessian suggests.
Laplace: $\hat\theta = \arg\max \log\tilde p$, $A = -\nabla^2\log\tilde p(\hat\theta)$, posterior $\approx N(\hat\theta, A^{-1})$. 1-D: sd $= 1/\sqrt{-\ell''(\hat\theta)}$.
Beta(5, 35): $\hat\theta = 0.105$, curvature 403.5, sd 0.050 (exact mean 0.125, sd 0.052).
Traps: centred at the mode, not the mean; depends on the scale (use log-odds/log); blind to skew, tails, second peaks.
Quick check: a posterior's log density near its peak is $\ell(\theta) \approx 3 - 50(\theta - 2)^2$. What is the Laplace approximation?
The peak is at $\hat\theta = 2$. The second derivative is $\ell'' = -100$, so the variance is $1/100$ and the sd is $0.1$: $N(2, 0.1^2)$. (Check: a Normal with sd 0.1 has log density $-(\theta-2)^2/(2 \times 0.01) = -50(\theta-2)^2$ plus a constant.)
Importance sampling: draw from the wrong distribution, then re-weight
You want the average height of adults in a city, but the only people you can survey are at a basketball club. Tall people are over-represented there. If you know how much more common each height is at the club than in the city, you can fix the survey: count each tall person a little less and each short person a little more. That correction factor is an importance weight.
Importance sampling does exactly this with distributions. Draw from an easy distribution $q$ (the "club", called the proposal), and give each draw the weight "posterior height ÷ proposal height". Weighted averages then estimate posterior averages. It works beautifully when $q$ looks like the posterior, and badly when it does not: a few draws get almost all the weight and the rest are wasted.
Three ways to say it:
- Picture: survey the wrong crowd, then re-weight every answer by how over- or under-represented that kind of person is.
- Numbers: four draws with normalized weights 0.23, 0.59, 0.17, 0.01 are worth about $1/(0.23^2 + 0.59^2 + 0.17^2 + 0.01^2) \approx 2.3$ equally weighted draws.
- Slogan: weight = target ÷ proposal; watch the weights, not just the estimate.
Four draws for the checkout posterior. Target $\tilde p(\theta) = \theta^4(1-\theta)^{34}$ (Beta(5, 35) without its constant). Proposal: Uniform(0, 0.4), height $q = 2.5$ everywhere on $[0, 0.4]$. The draws were 0.05, 0.10, 0.20, 0.30.
- Target heights ×$10^6$: $0.05^4 \times 0.95^{34} \times 10^6 = 1.093$; for 0.10: $2.781$; for 0.20: $0.811$; for 0.30: $0.044$.
- Weights are target ÷ proposal; the proposal height 2.5 is the same for every draw, so it cancels after normalizing. Total $= 1.093 + 2.781 + 0.811 + 0.044 = 4.729$.
- Normalized weights $\bar w$: $1.093/4.729 = 0.231$, $0.588$, $0.172$, $0.009$ (they sum to 1).
- Weighted mean: $0.231 \times 0.05 + 0.588 \times 0.10 + 0.172 \times 0.20 + 0.009 \times 0.30 = 0.01155 + 0.0588 + 0.0344 + 0.0027 = 0.10745 \approx 0.107$ (exact 0.125; four draws is very few).
- Weight effective sample size: $1/\sum \bar w^2 = 1/(0.0534 + 0.3457 + 0.0296 + 0.0001) = 1/0.4288 = 2.33$. Four draws are worth about 2.3 equally weighted ones: the draw at 0.30 is nearly wasted.
For any proposal density $q(\theta)$ that is positive wherever the posterior is positive:
$$E_p[f(\theta)] = \int f(\theta)\frac{p(\theta\mid D)}{q(\theta)}\,q(\theta)\,d\theta \;\approx\; \frac{\sum_s w_s f(\theta^{(s)})}{\sum_s w_s}, \qquad \theta^{(s)} \sim q, \quad w_s = \frac{\tilde p(\theta^{(s)})}{q(\theta^{(s)})}.$$- This is self-normalized importance sampling: dividing by $\sum w_s$ cancels the unknown $p(D)$, so the unnormalized $\tilde p$ is enough. (As a bonus, $\frac1S\sum_s w_s$ estimates $p(D)$ itself.)
- Normalized weights $\bar w_s = w_s/\sum_j w_j$. Weight effective sample size: $\text{ESS}_w = 1/\sum_s \bar w_s^2 = (\sum w_s)^2/\sum w_s^2$; it lies between 1 (one draw has all the weight) and $S$ (equal weights).
- Degeneracy: when $q$ is a poor match, or the dimension is high, a handful of weights dominate and $\text{ESS}_w \ll S$. If $q$ has lighter tails than $p$, the weights can have infinite variance and the estimate never settles.
- Rule for a safe proposal: centred near the posterior and somewhat wider than it (heavier tails), never narrower.
Why do we need it?
It re-uses draws from one distribution to answer questions about another. That makes it the tool for "what if" questions (a different prior, leaving out one data point) without refitting, and for correcting or checking an approximation such as a VI fit.
Where is it used?
PSIS-LOO cross-validation (az.loo in ArviZ, Pareto-smoothed importance sampling), checking a variational fit with the Pareto $\hat k$ diagnostic, particle filters and sequential Monte Carlo, re-weighting posterior draws for prior sensitivity (Chapter 6.8), rare-event simulation.
How is it used?
Draw from $q$, compute $\log w_s = \log\tilde p(\theta^{(s)}) - \log q(\theta^{(s)})$ in logs, subtract the maximum before exponentiating (log-sum-exp), normalize, and always report $\text{ESS}_w$ next to the estimate. A tiny $\text{ESS}_w$ means "do not trust this".
"Any proposal works if I take enough draws."
Only in theory. If $q$ is narrower than the target (lighter tails), the weights can have infinite variance: the estimate jumps whenever a rare tail draw arrives and never settles. Always use a proposal at least as wide as the target.
"A large weight ESS proves the estimate is right."
A too-narrow proposal can give even weights and a decent ESS while never visiting the tails: the estimate is confidently biased. ESS measures weight balance among the draws you have, not what you missed.
"Importance sampling is a general-purpose posterior sampler."
In many dimensions the weights degenerate exponentially (the second widget). It is a great correction or checking tool on top of a good approximation, not a replacement for MCMC or VI.
$E_p[f] \approx \sum_s w_s f(\theta^{(s)})/\sum_s w_s$, $\theta^{(s)} \sim q$, $w_s = \tilde p/q$ (compute in logs).
Weight ESS $= 1/\sum\bar w_s^2$, between 1 and $S$; collapses exponentially with dimension.
Traps: proposal must be wider than the target; a decent ESS does not prove you visited the tails.
Quick check: three normalized weights are 0.98, 0.01, 0.01. What is the weight ESS and what does it mean?
$1/(0.98^2 + 0.01^2 + 0.01^2) = 1/(0.9604 + 0.0001 + 0.0001) = 1/0.9606 = 1.04$. The three draws are worth about one: the estimate is essentially the value of a single draw. The proposal is a poor match.
Markov chains: a random walk that forgets where it started core
Suppose a shop's daily traffic is Low, Normal or High, and tomorrow's level depends only on today's level (not on last week). After a Low day, the next day is Low 60% of the time, Normal 30%, High 10%; and so on for the other states. A process like this, where the next step depends only on the present state, is a Markov chain.
Start it on a Low day and look at the chances for each later day. Day 1 is very likely Low. But the chain forgets: by day 10, it no longer matters where you started; the chances have settled at 30% Low, 50% Normal, 20% High, and they stay there forever. That settled distribution is the stationary distribution. And one long run of the chain spends 30% of its days in Low, 50% in Normal, 20% in High.
That is the whole idea of MCMC: design a chain whose stationary distribution is your posterior, run it for a long time, and record the states it visits. The record is a bag of posterior draws.
Three ways to say it:
- Picture: a walker hopping between rooms with fixed door probabilities ends up spending a fixed share of time in each room, whatever room it started in.
- Numbers: starting from Low: day 1 (0.60, 0.30, 0.10), day 2 (0.43, 0.42, 0.15), day 3 (0.357, 0.468, 0.175), … → (0.3, 0.5, 0.2).
- Slogan: build a chain whose long-run time shares equal the posterior; then just watch it.
Transition probabilities (row = today, column = tomorrow; each row sums to 1):
| today ↓ / tomorrow → | Low | Normal | High |
|---|---|---|---|
| Low | 0.6 | 0.3 | 0.1 |
| Normal | 0.2 | 0.7 | 0.1 |
| High | 0.1 | 0.3 | 0.6 |
- Day 0: we know today is Low: $\pi_0 = (1, 0, 0)$.
- Day 1: the Low row: $\pi_1 = (0.6, 0.3, 0.1)$.
- Day 2, P(Low) = P(Low on day 1)×0.6 + P(Normal)×0.2 + P(High)×0.1 $= 0.6 \times 0.6 + 0.3 \times 0.2 + 0.1 \times 0.1 = 0.36 + 0.06 + 0.01 = 0.43$. Similarly P(Normal) $= 0.6\times0.3 + 0.3\times0.7 + 0.1\times0.3 = 0.42$ and P(High) $= 0.15$.
- Day 3: $(0.357, 0.468, 0.175)$; day 5: $(0.311, 0.495, 0.194)$; day 10: $(0.300, 0.500, 0.200)$.
- Check that $(0.3, 0.5, 0.2)$ does not change: P(Low) $= 0.3\times0.6 + 0.5\times0.2 + 0.2\times0.1 = 0.18 + 0.10 + 0.02 = 0.30$; P(Normal) $= 0.09 + 0.35 + 0.06 = 0.50$; P(High) $= 0.03 + 0.05 + 0.12 = 0.20$. It is stationary.
A Markov chain is a sequence of random states $X_0, X_1, X_2, \dots$ with the Markov property: the next state depends only on the current one, $P(X_{t+1} = j\mid X_t = i, X_{t-1}, \dots, X_0) = P_{ij}$. The matrix $P$ (rows sum to 1) is the transition matrix.
- The distribution of the state evolves as $\pi_{t+1} = \pi_t P$ (row vector times matrix).
- A stationary distribution satisfies $\pi = \pi P$: once there, the chain stays there (in distribution; the walker itself keeps moving). It is a left eigenvector of $P$ with eigenvalue 1 (Linear Algebra 1.11).
- If the chain can reach every state from every state (irreducible) and does not cycle with a fixed period (aperiodic), then $\pi_t \to \pi$ from any start, and the long-run share of time spent in each state equals $\pi$ (the ergodic theorem: time averages → posterior averages). The speed of forgetting is set by the second-largest eigenvalue of $P$ (here 0.5: the distance to $\pi$ roughly halves each day).
- Detailed balance: if $\pi_i P_{ij} = \pi_j P_{ji}$ for all pairs (the flow from $i$ to $j$ equals the flow back), then $\pi$ is stationary, because summing over $i$ gives $\sum_i \pi_i P_{ij} = \pi_j\sum_i P_{ji} = \pi_j$. Metropolis is built to satisfy detailed balance. (The traffic chain above is stationary without detailed balance: it is a sufficient condition, not a necessary one.)
- The same ideas hold for continuous states (a parameter θ): $P$ becomes a "transition kernel", a rule for drawing $\theta_{t+1}$ given $\theta_t$.
Why do we need it?
Markov chains are the "MC" in MCMC. Their two properties (forgetting the start, and time shares equal to the stationary distribution) are exactly why a sampler's recorded path can be used as posterior draws, and why the first part of the path must be discarded.
Where is it used?
Every MCMC sampler (Metropolis, Gibbs, HMC, NUTS), PageRank (the stationary distribution of a random web surfer), hidden Markov models for regimes in time series, customer journey and churn models, and text generation with n-gram models.
How is it used?
In MCMC you never write $P$ down. You design a step rule that leaves the posterior stationary (detailed balance), start somewhere, run, throw away the early part (warmup) and keep the rest. Diagnostics (Chapter 6.10) check that the chain has forgotten its start.
"The chain converges to the most likely state and stays there."
The distribution converges; the walker never stops moving. It keeps visiting every state, each in proportion to its stationary probability. That is exactly what a sampler must do.
"Consecutive states of a Markov chain are independent draws from π."
Each state depends on the previous one (the strip shows runs of the same colour). Long-run shares are right, but neighbouring draws are correlated. This is the autocorrelation you will measure later in this chapter.
"Stationary means the chain has stopped changing."
It means the probabilities no longer change from day to day. Individual days still change.
Markov property: next state depends only on the current one. $\pi_{t+1} = \pi_t P$; stationary: $\pi = \pi P$.
Irreducible + aperiodic ⇒ forgets the start; time shares → π (ergodic theorem). Detailed balance $\pi_i P_{ij} = \pi_j P_{ji}$ ⇒ π stationary.
MCMC = a chain designed to have the posterior as π. Trap: draws are correlated, not independent.
Quick check: start the traffic chain on a Normal day. What are the chances for day 1 and for Low on day 2?
Day 1 is the Normal row: (0.2, 0.7, 0.1). Day 2, P(Low) $= 0.2\times0.6 + 0.7\times0.2 + 0.1\times0.1 = 0.12 + 0.14 + 0.01 = 0.27$. Already close to the stationary 0.3.
The Metropolis algorithm from scratch core
Picture a blindfolded hiker on the posterior landscape (high ground = high posterior). The hiker can feel the height under their feet and at any spot they test, but cannot see the map. Each step:
- Propose a random step to a nearby spot.
- If the new spot is higher, always go there.
- If it is lower, go there only sometimes: with probability "new height ÷ current height". A spot half as high is accepted half the time; a spot 100 times lower almost never.
- If the move is refused, stay, and write down the current spot again.
Because only the ratio of two heights is used, the unknown evidence $p(D)$ cancels: we only need $\tilde p$ = likelihood × prior. And the "sometimes go downhill" rule is tuned exactly so that, in the long run, the hiker spends time in each region in proportion to its posterior probability.
Three ways to say it:
- Picture: uphill always, downhill sometimes, otherwise stay and count this spot again.
- Numbers: from θ = 0.10, a proposal at 0.15 has height ratio 0.725 (accepted 72.5% of the time); a proposal at 0.02 has ratio 0.029.
- Slogan: accept with probability min(1, ratio); the evidence cancels in the ratio.
A. Two Metropolis steps for the checkout posterior, $\tilde p(\theta) = \theta^4(1-\theta)^{34}$, current value $\theta = 0.10$.
- Proposal $\theta' = 0.15$. Ratio $r = \tilde p(0.15)/\tilde p(0.10) = (0.15/0.10)^4 \times (0.85/0.90)^{34} = 5.0625 \times 0.1432 = 0.725$.
- Draw $u$ from Uniform(0, 1), say $u = 0.41$. Since $0.41 \lt 0.725$: accept. The chain records $0.15$.
- Next proposal from 0.15, say $\theta' = 0.02$: $r = (0.02/0.15)^4 \times (0.98/0.85)^{34}$. In logs: $4\ln(0.1333) + 34\ln(1.1529) = -8.06 + 4.84 = -3.22$, so $r = e^{-3.22} = 0.040$.
- Draw $u = 0.63 \gt 0.040$: reject. The chain records $0.15$ again. The repeat is not a bug: it is how the chain spends more time in high-probability places.
B. Why it targets the right distribution (three states, target $\pi = (0.3, 0.5, 0.2)$, propose one of the two other states with probability ½ each).
- From Low (0.3): to Normal, ratio $0.5/0.3 \gt 1$, always accepted: $P_{LN} = \tfrac12 \times 1 = 0.5$. To High, ratio $0.2/0.3$: $P_{LH} = \tfrac12 \times \tfrac23 = 0.333$. Stay: $1 - 0.5 - 0.333 = 0.167$.
- From Normal (0.5): to Low $\tfrac12 \times 0.6 = 0.3$; to High $\tfrac12 \times 0.4 = 0.2$; stay 0.5. From High (0.2): to Low 0.5, to Normal 0.5, stay 0.
- Detailed balance, Low ↔ Normal: $0.3 \times 0.5 = 0.15$ and $0.5 \times 0.3 = 0.15$. ✓ Low ↔ High: $0.3 \times 0.333 = 0.1$ and $0.2 \times 0.5 = 0.1$. ✓ Normal ↔ High: $0.5 \times 0.2 = 0.1$ and $0.2 \times 0.5 = 0.1$. ✓
- So $(0.3, 0.5, 0.2)$ is stationary for this chain: we built the right transition matrix using only ratios of the target.
Random-walk Metropolis for a target $p(\theta) \propto \tilde p(\theta)$:
- Start at some $\theta_0$. For $t = 0, 1, 2, \dots$:
- Propose $\theta' = \theta_t + \varepsilon z$, $z \sim N(0, I)$. The step size $\varepsilon$ is the only tuning knob; this proposal is symmetric: going from $a$ to $b$ is as likely as from $b$ to $a$.
- Acceptance probability $\alpha = \min\!\big(1, \tilde p(\theta')/\tilde p(\theta_t)\big)$, computed in logs: $\log r = \log\tilde p(\theta') - \log\tilde p(\theta_t)$.
- Draw $u \sim$ Uniform(0, 1). If $\log u \lt \log r$: $\theta_{t+1} = \theta'$ (accept). Otherwise $\theta_{t+1} = \theta_t$ (reject, and record the current value again).
- Why it works: the probability of moving from θ to θ′ is $q(\theta'\mid\theta)\min(1, p(\theta')/p(\theta))$. Multiply by $p(\theta)$: $q(\theta'\mid\theta)\min(p(\theta), p(\theta'))$, which is symmetric in θ and θ′. That is detailed balance, so $p$ is stationary; with a proposal that can reach everywhere, the chain converges to it.
- Metropolis–Hastings allows a non-symmetric proposal $q$ by using $r = \frac{\tilde p(\theta')\,q(\theta_t\mid\theta')}{\tilde p(\theta_t)\,q(\theta'\mid\theta_t)}$.
- Rejected steps are part of the output; dropping them would bias the sample toward low-probability regions.
Why do we need it?
It is the simplest algorithm that produces posterior draws from nothing but $\log\tilde p$. Every modern sampler (HMC, NUTS) keeps its accept/reject step, so understanding Metropolis is understanding the safety net under all of them.
Where is it used?
Physics simulations (it was invented for them in 1953), Gibbs and Metropolis-within-Gibbs steps for discrete parameters, PyMC's Metropolis step, the accept/reject inside HMC and NUTS, simulated annealing in optimization, and teaching.
How is it used?
Write log_p(theta); loop: propose, compare logs, accept or stay, append. Tune the step size until the acceptance rate is roughly 0.2–0.5. In practice for continuous models you use NUTS instead, which tunes itself.
"When a proposal is rejected, nothing is recorded for that step."
The current value is recorded again. Dropping repeats would over-represent the places the chain leaves quickly (low-probability regions) and bias every estimate.
"Metropolis needs the normalized posterior."
Only the ratio $\tilde p(\theta')/\tilde p(\theta)$ is used; $p(D)$ cancels. That is the whole reason MCMC exists.
"Metropolis climbs to the peak like an optimizer."
It accepts downhill moves with probability equal to the height ratio, so it keeps wandering through the whole posterior. An optimizer would stop at the top and report one point.
"Proposals outside the allowed range (a rate below 0) break the algorithm."
They have $\tilde p = 0$, so they are always rejected. That is valid, just wasteful; samplers like NUTS avoid it by working on an unconstrained scale (log-odds, log).
"MCMC samples from the posterior because it computes the posterior first."
MCMC never computes the posterior or the evidence. It builds a Markov chain whose stationary distribution is the posterior, using only ratios of likelihood × prior.
Model answer: "Metropolis proposes a move and accepts it with probability min(1, p̃(θ′)/p̃(θ)). The normalizing constant cancels in that ratio. The rule satisfies detailed balance, so the posterior is the chain's stationary distribution, and after the chain forgets its start, the visited states are (correlated) draws from the posterior."
Metropolis: $\theta' = \theta + \varepsilon z$; accept if $\log u \lt \log\tilde p(\theta') - \log\tilde p(\theta)$; else record θ again.
Works because $p(\theta)q(\theta'\mid\theta)\min(1, p(\theta')/p(\theta))$ is symmetric (detailed balance); $p(D)$ cancels.
Traps: rejections are recorded; it is not an optimizer; proposals outside the support are simply rejected.
Quick check: the current value has $\log\tilde p = -10.0$ and the proposal has $\log\tilde p = -10.7$. What is the acceptance probability?
$\log r = -10.7 - (-10.0) = -0.7$, so $r = e^{-0.7} = 0.497$. The move is accepted about half the time. If the proposal had $\log\tilde p = -9.5$ (higher), $r \gt 1$ and it would always be accepted.
Tuning the step size: too small crawls, too big gets stuck
Metropolis has one knob, the step size ε. With tiny steps, almost every proposal is accepted, but each move is tiny: the hiker shuffles and needs ages to cross the landscape. With huge steps, proposals land far away in low ground and are almost all refused: the hiker stands still for long stretches. In between there is a "just right" size where moves are both bold and often accepted.
A high acceptance rate is therefore not a sign of a good sampler. What matters is how quickly the chain produces new, different information, which we measure by the effective sample size (ESS): roughly, how many independent draws your correlated chain is worth.
Three ways to say it:
- Picture: a shuffling walker, a frozen walker, and a striding walker in between.
- Numbers: on a 1-D standard Normal, step 0.1 accepts about 97% of proposals, yet 1 000 steps are worth fewer than 10 independent draws; step 2.4 accepts about 44% and 1 000 steps are worth about 200.
- Slogan: tune for effective draws per second, not for acceptance.
Why both extremes waste effort (target: a standard Normal, sd 1).
- Step $\varepsilon = 0.1$: a typical proposal changes $\log\tilde p$ by very little, so $r \approx 1$ and nearly everything is accepted. But a random walk of $n$ steps of size 0.1 travels only about $0.1\sqrt{n}$; to cross the bulk of the target (about 4 units, from −2 to 2) takes about $(4/0.1)^2 = 1\,600$ steps.
- Step $\varepsilon = 20$: most proposals land 10+ sds away where $\tilde p$ is about $e^{-50}$ of the current height: refused. The chain repeats the same value for dozens of steps.
- Theory for Gaussian targets (a rule of thumb, not a law): in 1-D the best step is about $2.4\times$ the posterior sd, with acceptance near 0.44; in many dimensions about $2.38/\sqrt{d}$ per sd, with acceptance near 0.234.
- The $1/\sqrt{d}$ means random-walk Metropolis needs smaller steps as the number of parameters grows, so it slows down. Correlated or curved posteriors make it worse: the step must fit the narrowest direction. Chapter 6.10's Hamiltonian Monte Carlo fixes exactly this by using gradients.
- Acceptance rate: the share of proposals accepted over a run.
- Effective sample size (ESS): the number of independent draws that would give the same Monte Carlo error as your $n$ correlated draws; ESS $= n/\tau$ where $\tau = 1 + 2\sum_{k\ge1}\rho_k$ adds up the autocorrelations $\rho_k$ (defined in the autocorrelation section below; computed properly in Chapter 6.10).
- Optimal scaling (Roberts, Gelman and Gilks, 1997, for Gaussian-like targets): best step $\approx 2.38\,\sigma/\sqrt{d}$, acceptance $\approx 0.234$ for large $d$; about 0.44 in one dimension. Treat these as starting points.
- Adaptive tuning: adjust ε during warmup until the acceptance rate hits a target, then freeze it (NUTS does this automatically, Chapter 6.10).
Why do we need it?
The same algorithm can be useless or efficient depending on one number. Knowing the shape of the trade-off tells you what a bad acceptance rate (0.98 or 0.02) is telling you, and why modern samplers adapt their step size during warmup.
Where is it used?
Tuning any random-walk sampler, PyMC's step adaptation, the step-size adaptation of HMC and NUTS (target acceptance 0.8 in NumPyro), and judging sampler output: NumPyro reports the adapted step size and mean acceptance probability.
How is it used?
Run short pilot chains with a few step sizes, look at acceptance and ESS per second, pick the best, and discard the pilot runs. For a model with many correlated parameters, stop tuning Metropolis and switch to NUTS.
"An acceptance rate of 95% means the sampler works great."
It usually means the steps are too small and the chain is crawling. Look at ESS (or the trace), not acceptance alone. (NUTS is different: its target of 0.8 is about step accuracy along a trajectory, Chapter 6.10.)
"0.234 is the correct acceptance rate for every problem."
It is an asymptotic result for idealized Gaussian-like targets in high dimension. For one dimension it is about 0.44, and for real posteriors anything in roughly 0.15–0.5 can be fine. A rule of thumb, not a law.
"I can keep tuning the step size while I collect the draws I report."
Changing the step rule based on the chain's own history can break the stationary distribution. Tune during warmup, freeze, then collect.
Small ε: high acceptance, tiny moves, high autocorrelation. Large ε: few acceptances, flat trace. Best: in between.
Rule of thumb (Gaussian targets): ε ≈ 2.38 σ/√d, acceptance ≈ 0.234 (≈ 0.44 in 1-D). Judge by ESS per second.
Trap: high acceptance ≠ good mixing; tune only during warmup.
Quick check: a 20-parameter posterior with all sds about 0.5. What Metropolis step size does the rule of thumb suggest, and what acceptance rate?
$2.38 \times 0.5/\sqrt{20} = 1.19/4.47 \approx 0.27$, with acceptance near 0.234. If the parameters are strongly correlated, even this will mix slowly, because the step must fit the narrowest direction.
Burn-in and warmup: throw away the walk to the posterior
A chain has to start somewhere, and that "somewhere" is usually arbitrary: a random guess, or a value far out in the tail. The first part of the run is the chain travelling toward the high-probability region. Those early values are not typical posterior values; they are leftovers of the starting point. If you keep them, they drag every average toward the start.
So we discard the beginning. The old name is burn-in (like letting a new engine run for a while). Modern samplers call it warmup, because they also use that period to tune themselves (step size, and in HMC the "mass matrix", Chapter 6.10). Warmup draws are never part of the reported sample.
Three ways to say it:
- Picture: wait for the thermometer to settle before you read it.
- Numbers: 100 early draws averaging 0.35 mixed with 900 good draws averaging 0.125 give an overall mean of 0.1475, an 18% overestimate.
- Slogan: forget the start, then start counting.
How much damage can a start do? The checkout posterior has mean 0.125. A chain started at θ = 0.70 with small steps needs about 100 steps to reach the bulk.
- Suppose the first 100 draws average 0.35 (they slide down from 0.70) and the next 900 average 0.125.
- Mean of all 1 000: $(100 \times 0.35 + 900 \times 0.125)/1000 = (35 + 112.5)/1000 = 0.1475$.
- Error: $0.1475 - 0.125 = 0.0225$, which is 18% of the true value and about 13 times the Monte Carlo error of 900 independent draws ($0.0516/\sqrt{900} = 0.0017$).
- Discard the first 100 (burn-in): mean $\approx 0.125$, up to Monte Carlo noise.
- With 100 000 good draws after them, the same 100 bad draws would shift the mean by only $100 \times (0.35 - 0.125)/100\,100 \approx 0.0002$: burn-in matters most for short runs, but it is free and always done.
- Burn-in: the first $B$ iterations of a chain, discarded because their distribution still depends on the starting point.
- Warmup (NumPyro, Stan; "tune" in PyMC): the burn-in period plus automatic tuning of the sampler during it. Because the sampler changes during warmup, warmup draws are not valid posterior draws and are always dropped. In NumPyro,
MCMC(NUTS(model), num_warmup=500, num_samples=1000)runs 1 500 iterations per chain andget_samples()returns only the last 1 000. - There is no formula for the right $B$. Convergence cannot be proven from the output, only checked: with trace plots, and by running several chains from different starts and comparing them ($\hat R$, Chapter 6.10). A common habit is to spend as many iterations on warmup as on sampling, or half as many.
Why do we need it?
Without it, every posterior summary carries a bias toward wherever the chain happened to start, and the bias is biggest exactly when runs are short or the start is poor.
Where is it used?
Every MCMC run: num_warmup in NumPyro's MCMC, iter_warmup in Stan (CmdStanPy), tune in PyMC's pm.sample. Also in simulations of any Markov process (queues, traffic) when you want steady-state behaviour.
How is it used?
Choose a warmup length (500–1 000 is common for NUTS on small models), run several chains from dispersed starts, and confirm with trace plots and $\hat R$ that all chains reached the same region before the kept part begins. If not, lengthen warmup or fix the model.
"Burn-in fixes a chain that is mixing badly."
Burn-in only removes the effect of the starting point. A chain that is stuck, too slow, or trapped in one of several modes stays broken however much you discard.
"Warmup draws are just extra posterior draws I can add back for free."
During warmup the sampler is still changing its own settings and the chain may not have reached the posterior; those draws do not come from the posterior. NumPyro drops them for you.
"My chain looks flat after 50 steps, so it has converged."
One chain can sit in one region and look settled while missing another region entirely. Run several chains from different starts and compare them (Chapter 6.10).
If you validate your SVI fits with NUTS (Chapter 6.15), num_warmup is one of the first settings to report. A good habit: run 4 chains, start with num_warmup equal to num_samples, and check that all chains agree before trusting $P(\theta_B \gt \theta_A\mid D)$ or a forecast interval. Warmup has a cousin in your SVI loop: the first ELBO values are also far from the answer, which is one reason relative-ELBO early stopping needs patience (and, in many loops, a minimum number of steps) (Chapter 6.14).
Burn-in = discard the first $B$ draws (they still remember the start). Warmup = burn-in + automatic tuning; never reported.
Bias from a bad start shrinks as the run gets longer, but removing it is free. Convergence is checked (traces, several chains, $\hat R$), never proven.
Trap: burn-in cannot fix slow mixing or a missed mode.
Quick check: NumPyro runs MCMC(NUTS(model), num_warmup=300, num_samples=700, num_chains=4). How many iterations run in total, and how many draws come back?
Each chain runs $300 + 700 = 1\,000$ iterations, so $4\,000$ in total. get_samples() returns $4 \times 700 = 2\,800$ draws per parameter; the $1\,200$ warmup iterations are discarded.
Autocorrelation, and why thinning does not help
An MCMC chain moves in small, connected steps, so each draw is similar to the one before it. Think of a video: 1 000 frames of a 40-second clip contain far less information than 1 000 photos taken on different days, because neighbouring frames are almost the same picture. Autocorrelation measures that sameness: how correlated each draw is with the draw $k$ steps later.
A popular reaction is thinning: keep only every 10th draw "to remove the autocorrelation". The kept draws are indeed less correlated with each other, but you threw away nine draws that each carried some new information. Thinning never makes your estimates more precise. It only saves memory and disk.
Three ways to say it:
- Picture: deleting 9 of every 10 video frames does not turn the video into independent photos; it just gives you fewer frames.
- Numbers: 10 000 draws with lag-1 correlation 0.9 are worth about 526 independent draws; keeping every 10th leaves 1 000 draws worth about 483.
- Slogan: autocorrelation costs information; thinning costs more.
A chain with lag-1 autocorrelation ρ = 0.9 (each draw is 0.9 × the previous plus fresh noise: an AR(1) process, Chapter 7.3). Its autocorrelation at lag $k$ is $\rho^k$.
- Lags: $\rho_1 = 0.9$, $\rho_5 = 0.9^5 = 0.59$, $\rho_{10} = 0.349$, $\rho_{30} = 0.042$. It takes about 30 steps to "forget".
- Effective sample size: ESS $= n/(1 + 2\sum_k \rho^k) = n\,\frac{1-\rho}{1+\rho}$. With $n = 10\,000$: $10\,000 \times 0.1/1.9 = 526$.
- Thin by 10: 1 000 draws whose lag-1 correlation is $0.9^{10} = 0.349$. ESS $= 1\,000 \times (1 - 0.349)/(1 + 0.349) = 1\,000 \times 0.651/1.349 = 483$.
- $483 \lt 526$: thinning lost about 8% of the information, and every estimate became slightly less precise.
- Monte Carlo error of a posterior mean (posterior sd 1): unthinned $1/\sqrt{526} = 0.044$; thinned $1/\sqrt{483} = 0.046$; and the naive (wrong) $1/\sqrt{10\,000} = 0.010$, which pretends the draws are independent.
- Autocorrelation at lag $k$: $\rho_k = \text{Corr}(\theta_t, \theta_{t+k})$ along the chain, estimated as for a time series. The ACF plot shows $\rho_0 = 1, \rho_1, \rho_2, \dots$. A good sampler's ACF drops to about 0 within a few lags.
- Effective sample size: $\text{ESS} = n/\tau$ with $\tau = 1 + 2\sum_{k\ge1}\rho_k$, the "integrated autocorrelation time" (how many steps one independent draw costs). For AR(1), $\tau = (1+\rho)/(1-\rho)$. Chapter 6.10 shows how software estimates it from finite chains.
- Thinning by $m$: keep $\theta_m, \theta_{2m}, \dots$. For positively autocorrelated chains, ESS(thinned) ≤ ESS(all): thinning never improves Monte Carlo precision. It is justified only for storage, or when each kept draw is expensive to process later (for example, a full forecast simulation per draw).
- Autocorrelation does not make MCMC wrong; it makes it less efficient. Averages still converge to the right answer; they just need more draws.
Why do we need it?
The raw number of draws overstates how much you know. Autocorrelation is what turns "4 000 draws" into "worth 600 independent draws", which decides how many digits of $P(\theta_B \gt \theta_A\mid D)$ you can trust.
Where is it used?
The n_eff column of NumPyro's print_summary, ess_bulk/ess_tail in ArviZ, the thinning argument of NumPyro's MCMC (and thin in Stan), and the ACF plots used for forecast residuals (Chapter 7.17).
How is it used?
Look at the ACF and the ESS of every important quantity. If ESS is too low, run longer or use a better sampler (NUTS, a reparameterization). Thin only when memory or downstream cost forces you to, and then report the ESS of what you kept.
"Thinning removes autocorrelation, so it makes the sample better."
It makes the kept draws less correlated, but there are fewer of them, and the total information (ESS) goes down or at best stays the same. Use all draws unless storage forces you to thin.
"High autocorrelation means the MCMC answer is wrong."
Averages of an autocorrelated chain still converge to the correct posterior values; they need more draws to reach a given precision. The real dangers are a chain that has not converged or that misses part of the posterior.
"Lag-1 autocorrelation of 0.2 means the draws are basically independent."
Look at the whole ACF: a small lag-1 value followed by a long, slow tail can still cut ESS a lot. ESS sums all lags.
$\rho_k = \text{Corr}(\theta_t, \theta_{t+k})$; ESS $= n/(1 + 2\sum_k\rho_k)$; AR(1): ESS $= n(1-\rho)/(1+\rho)$.
ρ = 0.9, n = 10 000 → ESS 526; thin by 10 → ESS 483 (smaller!).
Traps: thinning never raises precision; autocorrelation = inefficiency, not bias.
Quick check: a chain of 8 000 draws has $\tau = 1 + 2\sum\rho_k = 20$. What is its ESS, and what is the Monte Carlo error of a posterior mean if the posterior sd is 0.3?
ESS $= 8\,000/20 = 400$. Error $= 0.3/\sqrt{400} = 0.015$. Treating the draws as independent would claim $0.3/\sqrt{8000} = 0.0034$, four and a half times too optimistic.
Trace plots: the chain's heartbeat
A trace plot is the simplest and most useful MCMC picture: iteration number on the horizontal axis, the value of one parameter on the vertical axis, a line through the draws. It is like a heart monitor: you learn to recognise a healthy rhythm at a glance.
Healthy looks like a "fuzzy caterpillar": a dense, flat band that jumps up and down quickly with no trend. Unhealthy patterns each have a cause: a slow wander (steps too small), flat stairs (steps too big, many refusals), a slope at the start (burn-in not discarded), jumps between two levels or no jumps at all (two modes), long flat stretches in one region (the sampler cannot enter a narrow region, like the neck of a funnel).
Three ways to say it:
- Picture: a fuzzy caterpillar is good; snakes, stairs and slopes are bad.
- Numbers: in a healthy trace the chain crosses its own mean every few iterations; in a stuck one it can stay on one side for hundreds.
- Slogan: look at the trace before you trust any number.
Diagnosing from a trace and two numbers (1-D standard Normal target, 2 000 Metropolis steps).
- Acceptance 0.98, trace drifts slowly across the range over hundreds of steps: steps too small. The ACF stays high for many lags; ESS is tiny. Fix: larger steps.
- Acceptance 0.03, trace is a staircase of flat segments dozens of steps long: steps too big. Fix: smaller steps.
- Acceptance 0.44, dense band centred on 0, no trend: healthy (ESS a few hundred).
- Trace starts at 12 and slides to 0 over the first few dozen steps, then looks healthy: discard the slide as burn-in.
- Trace sits near −3 for the whole run (or jumps to +3 once and stays): two modes, and the chain rarely crosses between them. Its histogram may show only one bump: the most dangerous case, because a single chain looks fine.
- Trace plot: $\theta_t$ against $t$ for one parameter (one line per chain when there are several).
- What to look for: (1) stationarity: no trend or drift in the kept part; (2) mixing: rapid up-and-down movement across the whole range; (3) agreement: several chains overlap (Chapter 6.10).
- Companions: the marginal histogram of the same draws, the ACF, and with several chains the rank plot (histograms of the ranks of each chain's draws among all draws; flat = good). ArviZ:
az.plot_trace,az.plot_rank. - A trace plot can show problems; it cannot prove there are none.
Why do we need it?
Numbers like a mean or even an ESS can look plausible for a broken chain. Two seconds of looking at a trace catches burn-in, stuck chains, bad step sizes and multiple modes before they reach a report.
Where is it used?
The first plot in every Bayesian workflow: az.plot_trace on NumPyro, PyMC and Stan output, MCMC dashboards, and (the same idea) training-loss curves in machine learning and the ELBO trace of your SVI loop.
How is it used?
Get draws per chain (mcmc.get_samples(group_by_chain=True)), plot each key parameter's trace with one colour per chain, and check: flat, fuzzy, overlapping. Investigate any parameter that does not look like a caterpillar before using its numbers.
"The trace looks like a fuzzy caterpillar, so the chain has converged."
A healthy-looking single chain can be exploring only one mode. Compare several chains started in different places; if they sit at different levels, the posterior has not been explored (Chapter 6.10's $\hat R$ turns this into a number).
"I only need to look at the trace of the parameter I report."
Problems often show first in hyperparameters (a hierarchical scale τ, a noise σ, a Negative Binomial concentration). Look at the traces of these "structural" parameters too.
Trace = value vs iteration. Good: flat, fuzzy, overlapping chains. Bad: wander (ε too small), stairs (ε too big), initial slope (burn-in), level jumps or a missing level (modes), long sticky stretches (funnel).
Pair it with the histogram and ACF; with several chains, rank plots.
Trap: one good-looking chain proves nothing about modes it never visited.
Quick check: a trace has an acceptance rate of 0.99 and slowly drifts upward for the whole run. What is wrong and what do you change?
The steps are much too small: nearly all moves are accepted but they are tiny, so the chain has not even finished travelling (it may still be in its burn-in). Increase the step size (or let NUTS adapt it), run longer, and discard the early part.
Recap, cheat sheet and practice
- The posterior $p(\theta\mid D) = \tilde p(\theta)/p(D)$ is intractable for real models: $\tilde p$ (likelihood × prior) is cheap anywhere, but the evidence $p(D) = \int\tilde p\,d\theta$ and every posterior average are high-dimensional integrals, and a grid costs $k^d$.
- A bag of draws answers everything by averaging and counting; for independent draws the error is sd/√S, whatever the dimension.
- The menu: MCMC (correlated draws, asymptotically exact), VI (a fitted $q_\phi$, fast, family-limited; Chapters 6.11–6.14), Laplace ($N(\hat\theta, A^{-1})$ from the peak and its curvature; crude, scale-dependent), importance sampling (weights $\tilde p/q$; weight ESS collapses with dimension).
- A Markov chain with a stationary distribution π forgets its start and spends a share π of its time in each state. MCMC designs a chain whose π is the posterior.
- Metropolis: propose $\theta' = \theta + \varepsilon z$, accept with probability $\min(1, \tilde p(\theta')/\tilde p(\theta))$, otherwise record θ again. $p(D)$ cancels; detailed balance makes the posterior stationary.
- Step size trades acceptance against move length; judge by ESS, not acceptance (rule of thumb: 0.234 in high dimension, 0.44 in 1-D). Random-walk steps must shrink like $1/\sqrt d$: the motivation for HMC (Chapter 6.10).
- Burn-in / warmup discards the walk from the start (and, in modern samplers, the tuning phase). Autocorrelation makes draws worth less: ESS $= n/(1 + 2\sum\rho_k)$. Thinning never raises precision.
- Trace plots: fuzzy caterpillars good; wander, stairs, slopes, missing modes and sticky stretches bad. One good chain proves nothing about modes it never visited.
Cheat sheet
| Idea | Formula / rule | In words |
|---|---|---|
| Unnormalized posterior | $\tilde p(\theta) = p(D\mid\theta)p(\theta)$; $p(\theta\mid D) = \tilde p/p(D)$ | the shape is cheap, the total is not |
| Grid cost | $k^d$ evaluations | 100 per axis, 10 parameters → $10^{20}$ |
| Monte Carlo | $E[f\mid D] \approx \frac1S\sum f(\theta^{(s)})$, error $\text{sd}/\sqrt S$ | integrals become averages |
| Laplace | $N(\hat\theta, A^{-1})$, $A = -\nabla^2\log\tilde p(\hat\theta)$ | a bell at the peak; fit on log / log-odds scale |
| Importance sampling | $w_s = \tilde p(\theta^{(s)})/q(\theta^{(s)})$; $\text{ESS}_w = 1/\sum\bar w_s^2$ | re-weight draws from an easy $q$; $q$ wider than $p$ |
| Markov chain | $\pi_{t+1} = \pi_t P$; stationary $\pi = \pi P$ | forgets the start; time shares → π |
| Detailed balance | $\pi_i P_{ij} = \pi_j P_{ji}$ ⇒ π stationary | flow from i to j = flow back |
| Metropolis | accept if $\log u \lt \log\tilde p(\theta') - \log\tilde p(\theta)$ | uphill always, downhill sometimes, else repeat θ |
| Step size (rule of thumb) | $\varepsilon \approx 2.38\sigma/\sqrt d$, acceptance ≈ 0.234 (0.44 in 1-D) | tune in warmup, then freeze |
| Effective sample size | $n/(1 + 2\sum_k\rho_k)$; AR(1): $n(1-\rho)/(1+\rho)$ | how many independent draws the chain is worth |
| Thinning | ESS(every $m$-th) ≤ ESS(all) | saves memory, never precision |
| NumPyro | MCMC(NUTS(model), num_warmup=…, num_samples=…, num_chains=4), get_samples(group_by_chain=True), numpyro.diagnostics.effective_sample_size | warmup draws are dropped for you |
import numpy as np
from scipy import stats, optimize
from numpyro.diagnostics import effective_sample_size # the ESS that NumPyro reports (Chapter 6.10)
rng = np.random.default_rng(0)
# The checkout posterior: prior Beta(2, 18), 3 buyers out of 20 visitors -> exactly Beta(5, 35)
a, b = 5, 35
exact = stats.beta(a, b)
def log_p(t):
"""unnormalized log posterior: log(likelihood x prior) = 4 log t + 34 log(1 - t); -inf outside (0, 1)"""
t = np.asarray(t, dtype=float)
out = np.full(t.shape, -np.inf)
ok = (t > 0) & (t < 1)
out[ok] = (a - 1) * np.log(t[ok]) + (b - 1) * np.log1p(-t[ok])
return out if out.ndim else float(out)
print(f"exact: mean {exact.mean():.4f} sd {exact.std():.4f} P(theta > 0.10) {exact.sf(0.10):.3f}")
# exact: mean 0.1250 sd 0.0516 P(theta > 0.10) 0.650
# 1) Why grids die: 100 points per parameter, 1e9 evaluations per second
for d in (1, 2, 3, 10, 54):
n_eval = 100.0 ** d
print(f"d = {d:2d}: {n_eval:.0e} evaluations = {n_eval / 1e9 / 3.156e7:.1e} years")
# d = 1: 1e+02 evaluations = 3.2e-15 years
# d = 2: 1e+04 evaluations = 3.2e-13 years
# d = 3: 1e+06 evaluations = 3.2e-11 years
# d = 10: 1e+20 evaluations = 3.2e+03 years <- about 3 000 years
# d = 54: 1e+108 evaluations = 3.2e+91 years
# 2) Monte Carlo: a bag of independent draws answers every question
draws = exact.rvs(4000, random_state=rng)
se = draws.std(ddof=1) / np.sqrt(draws.size)
print(f"MC (4000 draws): mean {draws.mean():.4f} +- {2 * se:.4f} P(theta > 0.10) {np.mean(draws > 0.10):.3f}")
# MC (4000 draws): mean 0.1243 +- 0.0016 P(theta > 0.10) 0.650
# 3) Laplace approximation: peak + curvature, on the rate scale
res = optimize.minimize_scalar(lambda t: -log_p(t), bounds=(1e-6, 1 - 1e-6), method="bounded")
mode, h = res.x, 1e-5
curv = -(log_p(mode + h) - 2 * log_p(mode) + log_p(mode - h)) / h**2 # -(second derivative)
lap = stats.norm(mode, 1 / np.sqrt(curv))
print(f"Laplace: mode {mode:.4f} curvature {curv:.1f} sd {lap.std():.4f} "
f"P(theta > 0.15) {lap.sf(0.15):.3f} (exact {exact.sf(0.15):.3f}) P(theta < 0) {lap.cdf(0):.3f}")
# Laplace: mode 0.1053 curvature 403.5 sd 0.0498 P(theta > 0.15) 0.184 (exact 0.284) P(theta < 0) 0.017
# 4) Importance sampling with a Normal proposal: weights and weight ESS
for m, s in [(0.13, 0.07), (0.30, 0.05)]:
th = rng.normal(m, s, 2000)
logw = log_p(th) - stats.norm(m, s).logpdf(th) # log(target / proposal)
w = np.exp(logw - logw.max()); w /= w.sum() # normalize (the unknown p(D) cancels)
print(f"IS proposal N({m}, {s}^2): estimate {np.sum(w * th):.4f} weight ESS {1 / np.sum(w**2):7.1f} of 2000")
# IS proposal N(0.13, 0.07^2): estimate 0.1242 weight ESS 1689.7 of 2000
# IS proposal N(0.3, 0.05^2): estimate 0.1863 weight ESS 41.3 of 2000 <- off centre: biased, few useful draws
# 5) The traffic Markov chain forgets its start
P = np.array([[0.6, 0.3, 0.1], [0.2, 0.7, 0.1], [0.1, 0.3, 0.6]])
pi = np.array([1.0, 0.0, 0.0]) # start: a Low day
for day in range(10):
pi = pi @ P
vals, vecs = np.linalg.eig(P.T)
stat = np.real(vecs[:, np.argmax(np.real(vals))]); stat /= stat.sum()
print("day 10:", pi.round(4), " stationary:", stat.round(4), " eigenvalues:", np.sort(np.real(vals)).round(2))
# day 10: [0.3002 0.4999 0.1998] stationary: [0.3 0.5 0.2] eigenvalues: [0.4 0.5 1. ]
# 6) Metropolis from scratch
def metropolis(log_p, x0, n, step, rng):
x, lp, acc = x0, log_p(x0), 0
out = np.empty(n)
for i in range(n):
y = x + step * rng.standard_normal()
ly = log_p(y)
if np.log(rng.uniform()) < ly - lp: # accept with probability min(1, p(y)/p(x))
x, lp, acc = y, ly, acc + 1
out[i] = x # a rejection records the current value AGAIN
return out, acc / n
chain, acc = metropolis(log_p, 0.30, 20000, 0.05, rng)
kept = chain[1000:] # burn-in: drop the walk from the start
ess = float(effective_sample_size(kept[None, :])) # shape (chains, draws)
print(f"Metropolis: acceptance {acc:.2f} mean {kept.mean():.4f} (all draws {chain.mean():.4f}) "
f"ESS {ess:.0f} of {kept.size} MCSE {kept.std() / np.sqrt(ess):.4f}")
# Metropolis: acceptance 0.70 mean 0.1251 (all draws 0.1254) ESS 1973 of 19000 MCSE 0.0012
# 7) Step size: acceptance vs effective draws
for step in (0.005, 0.05, 0.12, 0.5):
c, ac = metropolis(log_p, 0.125, 10000, step, rng)
print(f"step {step:5.3f}: acceptance {ac:.2f} ESS {float(effective_sample_size(c[None, :])):6.0f} of 10000")
# step 0.005: acceptance 0.97 ESS 14 of 10000 <- crawls
# step 0.050: acceptance 0.70 ESS 1058 of 10000
# step 0.120: acceptance 0.44 ESS 1920 of 10000 <- about 2.4 x the posterior sd (0.052): best
# step 0.500: acceptance 0.13 ESS 806 of 10000
# 8) The thinning myth: keep every 10th draw
thin = kept[::10]
print(f"ESS all {ess:.0f} ({kept.size} draws) vs ESS thinned {float(effective_sample_size(thin[None, :])):.0f} ({thin.size} draws)")
# ESS all 1973 (19000 draws) vs ESS thinned 1220 (1900 draws) <- thinning lost information
1. What makes the posterior of a 10-parameter hierarchical model intractable?
2. In random-walk Metropolis a proposal is rejected. What is added to the chain for that iteration?
3. Which two methods become arbitrarily accurate if you simply give them more computation?
4. Four importance draws have normalized weights 0.5, 0.5, 0, 0. What is the weight effective sample size?
5. You thin a positively autocorrelated chain by keeping every 10th draw. What happens to the effective sample size?
6. A Metropolis run reports an acceptance rate of 0.98. What is the most likely situation?
Practice problems
A. A model has 8 parameters. You try a grid with 30 values per parameter on a machine doing $10^8$ evaluations per second. How long does it take? What if you add two more parameters?
- $30^8 = 6.56 \times 10^{11}$ evaluations.
- Time: $6.56\times10^{11}/10^8 = 6\,561$ seconds ≈ 1.8 hours.
- Two more parameters multiply the work by $30^2 = 900$: about 68 days. Exponential growth: this is why grids are only for 1–3 parameters.
B. Metropolis on a standard Normal target, current value 1.0. Compute the acceptance probability for proposals 0.5 and 2.5.
- $\log\tilde p(\theta) = -\theta^2/2$. At 1.0: $-0.5$.
- Proposal 0.5: $\log\tilde p = -0.125$; $\log r = -0.125 - (-0.5) = 0.375$; $r = e^{0.375} = 1.45 \gt 1$: always accepted (uphill).
- Proposal 2.5: $\log\tilde p = -3.125$; $\log r = -3.125 + 0.5 = -2.625$; $r = e^{-2.625} = 0.072$: accepted about 7% of the time.
C. A Poisson rate has posterior $\tilde p(\lambda) = \lambda^9 e^{-3\lambda}$ (a Gamma(10, rate 3)). Find its Laplace approximation on the λ scale and on the log λ scale, and compare with the exact mean 3.33 and sd 1.05.
- λ scale: $\ell(\lambda) = 9\ln\lambda - 3\lambda$; $\ell' = 9/\lambda - 3 = 0$ gives $\hat\lambda = 3$; $\ell'' = -9/\lambda^2 = -1$ at 3, so sd $= 1$. Laplace: $N(3, 1^2)$, with $P(\lambda \lt 0) = \Phi(-3) = 0.0013$ on impossible values.
- log scale, $u = \ln\lambda$: the density of $u$ gains a factor λ (the Jacobian), so $\ell(u) = 10u - 3e^u$; $\ell' = 10 - 3e^u = 0$ gives $\lambda = e^{\hat u} = 10/3 = 3.33$ (the exact mean); $\ell'' = -3e^{u} = -10$, so sd$(u) = 1/\sqrt{10} = 0.316$.
- Transformed back, $\lambda = e^u$ is log-normal: always positive and right-skewed like the truth. The λ-scale bell is centred at the mode 3 and is symmetric. Same method, better coordinates.
D. A two-state chain has transition matrix rows (0.9, 0.1) and (0.3, 0.7). Find its stationary distribution using detailed balance, and check it.
- Detailed balance: $\pi_1 \times 0.1 = \pi_2 \times 0.3$, so $\pi_1 = 3\pi_2$.
- With $\pi_1 + \pi_2 = 1$: $\pi_2 = 0.25$, $\pi_1 = 0.75$.
- Check $\pi = \pi P$: first entry $0.75 \times 0.9 + 0.25 \times 0.3 = 0.675 + 0.075 = 0.75$. ✓ (Any two-state chain satisfies detailed balance at its stationary distribution.)
E. A chain of 5 000 draws behaves like AR(1) with ρ = 0.8, and the posterior sd is 2. Compute the ESS and the Monte Carlo error of the posterior mean. How wrong is the naive error?
- ESS $= 5\,000 \times (1 - 0.8)/(1 + 0.8) = 5\,000 \times 0.2/1.8 = 556$.
- Monte Carlo standard error $= 2/\sqrt{556} = 0.085$.
- Naive $2/\sqrt{5\,000} = 0.028$: three times too small. Reporting the mean to two decimals would be overclaiming.
F. (Interview) "How does MCMC get posterior samples if it never computes the normalizing constant? And what can go wrong in practice?"
"MCMC builds a Markov chain whose stationary distribution is the posterior. In Metropolis, I propose a move and accept it with probability min(1, p̃(θ′)/p̃(θ)); the evidence cancels in that ratio, so I only need likelihood × prior. The rule satisfies detailed balance, so once the chain has forgotten its starting point, the states it visits are draws from the posterior, correlated with each other. In practice three things go wrong: the chain has not converged yet (so I discard warmup and run several chains from different starts), it mixes slowly (high autocorrelation, so the effective sample size is much smaller than the number of draws, and I report Monte Carlo error with ESS, not n), or it misses part of the posterior, such as a second mode or a narrow funnel neck. That is why I check trace plots, R-hat and ESS, and why modern samplers like NUTS use gradients and adapt their step size during warmup."
Hamiltonian Monte Carlo, NUTS and MCMC diagnostics
Random-walk Metropolis (Chapter 6.9) stumbles around like a blindfolded hiker. Hamiltonian Monte Carlo gives the hiker a skateboard and a map of the slopes: it uses the gradient of the log posterior to glide far across the landscape in one move and still be accepted almost every time. NUTS is the version that chooses its own trajectory length and step size, and it is what NumPyro runs when you write MCMC(NUTS(model)). The second half of the chapter is the toolkit for deciding whether to trust the output: trace plots, R̂, effective sample size, Monte Carlo standard error, divergences and warmup.
- Explain why random-walk proposals are slow (distance grows like √n) and how gradients help
- Describe HMC with the physics picture: position = parameters, momentum = a random kick, potential energy $U = -\log\tilde p$, total energy $H = U + K$
- Run the leapfrog integrator by hand, and explain why its energy error stays small and what happens when the step is too big
- Compute an HMC acceptance probability $\min(1, e^{-\Delta H})$ and explain why HMC moves far with high acceptance
- Explain how NUTS doubles its trajectory until it makes a U-turn, and how warmup adapts the step size (target acceptance 0.8) and the mass matrix
- Read trace plots of several chains, compute split R̂, ESS and the Monte Carlo standard error sd/√ESS
- Recognise divergences (the funnel), know what to do about them, and read NumPyro's
print_summarylike a checklist
What we need from earlier chapters: Metropolis, acceptance ratios, step size, warmup, autocorrelation and the first look at effective sample size (Chapter 6.9); the funnel and the non-centered fix (Chapter 6.7, especially divergences and the fix); gradients (Calculus 2.4) and how autodiff computes them (Calculus 2.9); gradient descent and step sizes (Optimization 3.3); covariance matrices and their scales (Chapter 5.15). Words: the gradient $\nabla\log\tilde p(\theta)$ is the vector of slopes of the log posterior: it points uphill, toward more probable values. A trajectory is the path a simulated particle follows. Energy here is just a number we compute; no real physics is involved.
Why random walks are slow, and what the gradient adds core
A random-walk sampler forgets which way it was going after every step. Step left, step right, step left… after $n$ steps it has typically moved only about $\sqrt n$ step-lengths from where it started, because many steps cancel. To cross a posterior that is 40 step-lengths wide, it needs about $40^2 = 1\,600$ steps. And Chapter 6.9 showed that its steps must be small whenever the posterior is narrow in some direction or has many parameters.
Now give the walker two things. First, a map of the slopes: the gradient of $\log\tilde p$, which says which way is uphill (toward more probable values). Second, momentum: once it is moving, it keeps moving in the same direction until the landscape turns it around. Like a skateboarder in a bowl, it can glide all the way across the posterior in one go, following the curve of the valley instead of bumping into the walls. That is Hamiltonian Monte Carlo (HMC).
Three ways to say it:
- Picture: a drunk walker versus a skateboarder who feels the slope.
- Numbers: to cross 4 units with steps of 0.1: a random walk needs about $(4/0.1)^2 = 1\,600$ steps; a straight glide needs about $4/0.1 = 40$.
- Slogan: random walks diffuse ($\sqrt n$); momentum travels ($n$).
A correlated posterior with correlation ρ = 0.99 between two parameters (unit sds), the shape of a long thin cigar.
- The covariance matrix has eigenvalues $1 + \rho = 1.99$ and $1 - \rho = 0.01$, so the cigar has sd $\sqrt{1.99} = 1.41$ along its length and $\sqrt{0.01} = 0.1$ across it.
- Random-walk Metropolis must use steps about the size of the narrow sd, say 0.1, or nearly everything is rejected.
- To travel the cigar's length (about $\pm 2$ sds $= 2 \times 2 \times 1.41 \approx 5.6$ units) it needs about $(5.6/0.1)^2 \approx 3\,100$ steps, and that is just one crossing.
- A gliding trajectory with step 0.1 that follows the cigar needs about $5.6/0.1 = 56$ gradient steps for the same distance, and it is one HMC iteration.
- Price: every step needs the gradient, which autodiff (JAX) computes at a cost of a small multiple of one $\log\tilde p$ evaluation.
- Diffusive behaviour: for a random walk with step size $\varepsilon$, the typical distance after $n$ steps is $\varepsilon\sqrt n$, so crossing a distance $D$ costs about $(D/\varepsilon)^2$ steps.
- For random-walk Metropolis on a $d$-dimensional Gaussian-like target the good step size shrinks like $d^{-1/2}$ (Chapter 6.9), so the number of $\log\tilde p$ evaluations per effective draw grows roughly like $d$. For HMC the number of gradient evaluations grows roughly like $d^{1/4}$ (theoretical scaling results for idealized targets; real models vary).
- Gradient information: $\nabla_\theta\log\tilde p(\theta) = \nabla_\theta\log p(D\mid\theta) + \nabla_\theta\log p(\theta)$. It does not need $p(D)$ either (the constant disappears when you differentiate a log).
- HMC needs continuous parameters with a differentiable log density. Discrete latent variables must be summed out (NumPyro can enumerate some) or sampled with other methods.
Why do we need it?
Real posteriors have tens to thousands of correlated parameters. Random-walk samplers become hopelessly slow there; gradient-based samplers stay practical. This is why every modern probabilistic programming language ships HMC/NUTS as its default.
Where is it used?
NUTS is the default sampler in Stan, PyMC and NumPyro (numpyro.infer.NUTS); HMC variants are used for Bayesian neural networks, hierarchical A/B models, epidemiology models and Gaussian processes. The gradients come from autodiff (JAX in NumPyro).
How is it used?
You never write the gradient by hand: write the model, NumPyro differentiates $\log\tilde p$ with JAX, and NUTS uses it. Your job is to keep the model differentiable (no hard if on parameters), continuous, and in good coordinates (non-centered, unconstrained).
"HMC uses the gradient to climb to the peak, like an optimizer."
The gradient acts like gravity on a frictionless puck: it bends the path and changes the speed, but the puck never stops at the bottom. It keeps rolling through the whole posterior, so HMC samples rather than optimizes.
"HMC works for any model."
It needs a differentiable log density over continuous parameters. Discrete parameters (a changepoint index, a mixture label) must be marginalized out or handled by another sampler.
Random walk: distance ≈ ε√n, crossing D costs ≈ (D/ε)² steps. Momentum: ≈ D/ε steps.
HMC uses $\nabla\log\tilde p$ (from autodiff; no $p(D)$ needed) to travel far in one iteration.
Trap: the gradient bends the path; it does not make HMC an optimizer. Needs continuous, differentiable parameters.
Quick check: a posterior is 3 units wide and a random-walk sampler can only use steps of 0.05. Roughly how many steps per crossing? With momentum?
Random walk: $(3/0.05)^2 = 60^2 = 3\,600$ steps. Momentum: about $3/0.05 = 60$ steps, sixty times fewer.
The physics picture: position, momentum, energy core
HMC pretends the parameters are the position of a puck on an ice landscape whose height is $U(\theta) = -\log\tilde p(\theta)$. Probable values are low ground (valleys); improbable values are high ground. At the start of each iteration we give the puck a random shove: a momentum $p$, one number per parameter, drawn from a standard Normal. Then we let physics run for a while.
Physics keeps the total energy constant: height energy $U$ plus motion energy $K = \tfrac12|p|^2$. Rolling downhill turns height into speed; rolling uphill turns speed back into height. So a puck with total energy $H$ can never climb above height $H$, and it keeps sliding back and forth through the valley, visiting high-probability regions. Where it is after a while is the proposal. Because the next iteration draws a fresh kick, the energy level changes from iteration to iteration, so the puck visits both the valley floor and the higher slopes.
Three ways to say it:
- Picture: a frictionless puck in a bowl that gets a random kick each round.
- Numbers: on $U = \theta^2/2$, a puck at θ = 0 kicked with $p = 1.5$ has $H = 0 + 1.5^2/2 = 1.125$, so it swings between θ = −1.5 and +1.5 (where $U = 1.125$, $K = 0$).
- Slogan: position = parameters, momentum = random kick, energy = −log posterior + ½|p|².
A puck on the standard Normal landscape $U(\theta) = \theta^2/2$ (that is, $\tilde p(\theta) = e^{-\theta^2/2}$).
- Start at $\theta = 0$ (the bottom) with momentum $p = 1.5$: $U = 0$, $K = 1.5^2/2 = 1.125$, $H = 1.125$.
- As it climbs, $K$ falls and $U$ rises with their sum fixed. It stops (for an instant) where $U = 1.125$: $\theta^2/2 = 1.125$, $\theta = \pm 1.5$.
- The exact motion here is $\theta(t) = 1.5\sin t$, $p(t) = 1.5\cos t$: a full swing takes $2\pi \approx 6.28$ time units.
- Start instead at $\theta = 2$ with $p = 0$: $H = U = 2$; it slides down, passes 0 at full speed ($K = 2$, $p = -2$) and climbs to $\theta = -2$ on the other side.
- The gradient is what moves it: $dp/dt = -U'(\theta) = \frac{d}{d\theta}\log\tilde p(\theta) = -\theta$. At θ = 2 the puck is pushed back toward 0 with force 2.
HMC adds an auxiliary momentum variable $p$ (same dimension as θ) and defines the Hamiltonian (total energy)
$$H(\theta, p) = U(\theta) + K(p), \qquad U(\theta) = -\log\tilde p(\theta), \qquad K(p) = \tfrac12\,p^\top M^{-1}p .$$- $M$ is the mass matrix; with $M = I$, $K = \tfrac12|p|^2$ and $p \sim N(0, I)$. (Warmup adapts $M$; see the warmup section.)
- Hamilton's equations describe the motion: $\dfrac{d\theta}{dt} = M^{-1}p$, $\quad\dfrac{dp}{dt} = -\nabla U(\theta) = \nabla\log\tilde p(\theta)$.
- Exact Hamiltonian motion keeps $H$ constant, is reversible (flip the momentum and it retraces its path) and preserves volume in $(\theta, p)$ space. These three properties are what make the HMC proposal valid.
- The joint density $\propto e^{-H(\theta, p)} = \tilde p(\theta)\,e^{-K(p)}$: θ and $p$ are independent, θ has exactly the posterior as its distribution, and $p$ is just a Normal we throw away after each iteration.
Why do we need it?
The energy picture explains every HMC behaviour you will see: why acceptance is high (energy is nearly conserved), why the step size matters (simulation error changes the energy), and why narrow regions cause divergences (huge forces).
Where is it used?
Inside NUTS in NumPyro, Stan and PyMC; NumPyro's potential_fn (the $U$ of a model) and the energy extra field of an MCMC run; the same equations drive molecular-dynamics simulations in chemistry and physics; HMC itself was invented for lattice simulations in particle physics.
How is it used?
You never set the physics up yourself: NumPyro turns your model into $U(\theta)$ on an unconstrained scale, draws $p$, and integrates the motion. Reading "energy" diagnostics (and divergences) is how you see when the physics breaks down.
"The momentum is a parameter of the model."
It is an auxiliary variable invented by the sampler: drawn fresh from $N(0, M)$ each iteration and thrown away afterwards. Your model never sees it.
"Since energy is conserved, the puck always stays at the same height."
Total energy is conserved; height $U$ goes up and down as speed changes. Different iterations get different random kicks and so explore different energy levels.
$U(\theta) = -\log\tilde p(\theta)$, $K(p) = \tfrac12 p^\top M^{-1}p$, $H = U + K$; $p \sim N(0, M)$ fresh each iteration.
$d\theta/dt = M^{-1}p$, $dp/dt = \nabla\log\tilde p(\theta)$; exact motion conserves $H$, is reversible, preserves volume.
Trap: momentum is not a model parameter; $U$ changes along the path, $H$ does not.
Quick check: on $U(\theta) = \theta^2/2$ a puck starts at θ = 1 with $p = -1$. What is $H$, and how far left can it go?
$H = 1^2/2 + (-1)^2/2 = 1$. It turns around where $U = 1$: $\theta^2/2 = 1$, $\theta = -\sqrt 2 \approx -1.41$.
The leapfrog integrator: simulating the slide on a computer core
A computer cannot follow the puck continuously; it moves it in small time steps of size ε. The obvious recipe (move with the current speed, then update the speed) slowly adds energy at every step: the simulated puck spirals outward and flies off. The leapfrog recipe fixes this with a symmetric pattern: a half-step kick to the momentum, a full step for the position, another half-step kick. Kick, drift, kick.
That symmetry makes the leapfrog path reversible (run it backwards and you return exactly) and keeps its energy error small and bounded: it wobbles a little up and down but does not grow with the number of steps, as long as ε is small enough. If ε is too big for the narrowest part of the posterior, the simulation becomes unstable and the energy explodes. That explosion is what NumPyro calls a divergence.
Three ways to say it:
- Picture: half kick, full drift, half kick; repeat $L$ times.
- Numbers: on $U = \theta^2/2$ from θ = 1, ε = 0.5: four leapfrog steps keep $H$ between 0.469 and 0.500; four naive (Euler) steps push it from 0.5 to 1.22.
- Slogan: leapfrog energy errors wobble, they do not drift, until ε is too big.
One leapfrog step by hand on $U(\theta) = \theta^2/2$ (so $\nabla\log\tilde p(\theta) = -\theta$), start θ = 1, $p = 0$, ε = 0.5. Start energy $H = 0.5$.
- Half kick: $p \leftarrow p + \tfrac{\varepsilon}{2}(-\theta) = 0 - 0.25 \times 1 = -0.25$.
- Drift: $\theta \leftarrow \theta + \varepsilon p = 1 + 0.5 \times (-0.25) = 0.875$.
- Half kick with the new gradient: $p \leftarrow -0.25 - 0.25 \times 0.875 = -0.46875$.
- New energy: $0.875^2/2 + 0.46875^2/2 = 0.3828 + 0.1099 = 0.4927$. Error $-0.0073$. (The exact motion gives θ = cos 0.5 = 0.8776.)
- Three more steps: $H = 0.4776, 0.4688, 0.4747$: the error wobbles but stays below 0.04.
- Naive Euler instead ($\theta \leftarrow \theta + \varepsilon p$, then $p \leftarrow p - \varepsilon\theta_{\text{old}}$): $H = 0.625, 0.781, 0.977, 1.221$. Each step multiplies $H$ by $1 + \varepsilon^2 = 1.25$: the error grows without limit.
One leapfrog step of size ε (unit mass):
$$p_{\frac12} = p + \tfrac{\varepsilon}{2}\nabla\log\tilde p(\theta), \qquad \theta' = \theta + \varepsilon\, p_{\frac12}, \qquad p' = p_{\frac12} + \tfrac{\varepsilon}{2}\nabla\log\tilde p(\theta').$$- $L$ steps give a trajectory of length (in time) $\varepsilon L$ and cost $L$ gradient evaluations (consecutive half kicks merge, so one new gradient per step).
- Reversible: starting from $(\theta', -p')$ and stepping retraces the path. Volume-preserving ("symplectic"). These make the HMC proposal valid without any extra correction terms.
- Energy error $\Delta H = H(\theta_L, p_L) - H(\theta_0, p_0)$ is of order $\varepsilon^2$ and bounded over long trajectories when ε is below the stability limit.
- Stability limit: in a direction where the posterior is roughly Normal with sd σ, leapfrog is stable only if $\varepsilon \lt 2\sigma$ (with unit mass). The narrowest direction sets the limit; beyond it the energy grows exponentially (a divergence).
Why do we need it?
The exact physics cannot be computed for a real posterior. Leapfrog is the cheap, stable approximation whose small, non-accumulating energy error makes long trajectories (and therefore far moves) affordable.
Where is it used?
Every HMC and NUTS implementation (NumPyro's velocity_verlet, Stan, PyMC), molecular dynamics (where it is called velocity Verlet), and planetary orbit simulations. NumPyro reports the number of leapfrog steps per draw as num_steps.
How is it used?
Automatically: NUTS picks ε during warmup and $L$ per iteration. You meet it through its symptoms: many leapfrog steps per draw (slow sampling: a hard geometry) or divergences (ε too large for some region).
"A smaller step size is always better."
Smaller ε means a more accurate simulation but more gradient evaluations for the same distance. The goal is the largest ε that keeps the energy error small, which is exactly what warmup adaptation looks for.
"The energy error grows the longer the trajectory."
With leapfrog below its stability limit, the error oscillates and stays bounded; that is why HMC can afford long trajectories. Above the limit (ε more than about 2 × the narrowest sd) it grows exponentially.
Leapfrog: $p \mathrel{+}= \tfrac\varepsilon2\nabla\log\tilde p(\theta)$; $\theta \mathrel{+}= \varepsilon p$; $p \mathrel{+}= \tfrac\varepsilon2\nabla\log\tilde p(\theta)$. $L$ steps = $L$ gradients, time $\varepsilon L$.
Reversible + volume-preserving; energy error $O(\varepsilon^2)$, bounded. Stable only if ε ≲ 2 × the narrowest sd.
Trap: Euler drifts (×(1+ε²) per step on a Gaussian); too-big ε explodes = divergence.
Quick check: a posterior has sds 5 and 0.2 along its two principal directions. With unit mass, roughly what is the largest stable leapfrog step, and why is that a problem?
About $2 \times 0.2 = 0.4$, set by the narrow direction. Crossing the wide direction (sd 5, so about 20 units) then takes about $20/0.4 = 50$ or more steps. Rescaling the coordinates (the mass matrix, adapted in warmup) removes this mismatch.
One HMC iteration: kick, glide, check core
Put the pieces together. Each HMC iteration: (1) draw a fresh random momentum (the kick); (2) run $L$ leapfrog steps (the glide); (3) compare the total energy at the end with the energy at the start. If the simulation were exact, the energy would be identical and the end point would always be accepted. Leapfrog is slightly off, so we accept the end point with probability $e^{-\Delta H}$ (capped at 1): a Metropolis check that exactly cancels the simulation error. If rejected, the chain records the current point again, as in Chapter 6.9.
Here is HMC's superpower: because the energy error stays small even after many steps, the end point can be far away and still have acceptance near 1. Random-walk Metropolis can never do that: a far random jump almost always lands somewhere improbable.
Three ways to say it:
- Picture: kick the puck, let it slide, keep the landing spot if the energy books balance.
- Numbers: ΔH = 0.1 → accepted with probability 0.905; ΔH = 2 → 0.135; ΔH = −0.3 → always.
- Slogan: far moves, high acceptance, because physics nearly conserves energy.
Acceptance from the energy error.
- From the leapfrog example: start θ = 1, $p = 0$ ($H_0 = 0.5$), four steps of ε = 0.5 end at θ = −0.436 with $H = 0.4747$. $\Delta H = -0.0253 \lt 0$, so $e^{-\Delta H} = 1.026 \gt 1$: accepted for sure. The chain moved 1.44 units (1.44 posterior sds) in one iteration.
- A trajectory with $\Delta H = 0.1$: acceptance $e^{-0.1} = 0.905$.
- $\Delta H = 0.5$: $e^{-0.5} = 0.607$. $\Delta H = 2$: $e^{-2} = 0.135$. $\Delta H = 5$: $0.0067$.
- Why $e^{-\Delta H}$: the joint density of (θ, p) is proportional to $e^{-H}$, so the Metropolis ratio of end to start is $e^{-H_1}/e^{-H_0} = e^{-\Delta H}$. The proposal is symmetric because leapfrog is reversible and volume-preserving.
Hamiltonian Monte Carlo (step size ε, $L$ steps, mass $M$), from the current θ:
- Draw $p \sim N(0, M)$. Compute $H_0 = U(\theta) + K(p)$.
- Run $L$ leapfrog steps to get $(\theta^*, p^*)$; compute $H_1 = U(\theta^*) + K(p^*)$.
- Accept θ* with probability $\min\!\big(1, e^{-(H_1 - H_0)}\big)$; otherwise keep θ. Discard $p$.
- The θ-draws have the posterior as their stationary distribution; momentum resampling lets the chain move between energy levels.
- Tuning: ε controls accuracy (acceptance); the trajectory length $\varepsilon L$ controls how far each iteration travels. Plain HMC needs you to set both; NUTS sets them automatically.
- NumPyro's
HMCkernel is this algorithm (with atrajectory_lengthsetting);NUTSchooses the length per iteration.
Why do we need it?
It is the combination that matters: gradients for direction, momentum for distance, and the accept step for exactness. Remove the accept step and you get a biased sampler; remove the momentum and you are back to a random walk.
Where is it used?
numpyro.infer.HMC, NUTS (HMC with automatic trajectory lengths) in NumPyro, Stan and PyMC, Bayesian neural network samplers, and lattice simulations in physics where HMC was invented ("hybrid Monte Carlo", 1987).
How is it used?
For real models you run NUTS. You read HMC's health from its outputs: the mean acceptance probability (near the target 0.8 after warmup), the number of leapfrog steps per draw, the step size, and divergences.
"HMC's high acceptance rate means it is exploring well."
Acceptance only says the simulation was accurate. A trajectory that is too short moves very little even with 100% acceptance; one that is too long can loop back near its start. Exploration is measured by ESS.
"HMC does not need the Metropolis accept step because physics is exact."
The leapfrog simulation is approximate. The accept step with $e^{-\Delta H}$ is exactly what removes the discretization error and keeps the posterior as the stationary distribution.
"HMC is faster than Metropolis because it accepts more proposals."
HMC is more efficient because gradient-guided trajectories make distant proposals that are still accepted, so successive draws are far less correlated (higher ESS per gradient evaluation), especially in high dimension and with correlated parameters.
Model answer: "HMC augments the parameters with a momentum, simulates Hamiltonian dynamics on U = −log posterior with the leapfrog integrator, and accepts the end point with probability min(1, e^(−ΔH)). Because leapfrog nearly conserves energy, the end point can be far from the start and still be accepted, which suppresses the random-walk behaviour. Each iteration costs L gradient evaluations, which JAX provides by autodiff."
HMC iteration: $p \sim N(0, M)$; $L$ leapfrog steps; accept with $\min(1, e^{-\Delta H})$, $\Delta H = H_1 - H_0$.
ΔH = 0.1 → 0.905, 0.5 → 0.607, 2 → 0.135. Far moves with high acceptance because the energy error stays small.
Trap: acceptance measures simulation accuracy, not exploration; judge by ESS per gradient.
Quick check: an HMC trajectory starts with $H_0 = 12.40$ and ends with $H_1 = 12.95$. What is the acceptance probability? What if $H_1 = 12.10$?
$\Delta H = 0.55$: acceptance $e^{-0.55} = 0.577$. If $H_1 = 12.10$, $\Delta H = -0.30 \lt 0$ and $e^{0.30} \gt 1$, so it is always accepted.
NUTS: grow the trajectory until it makes a U-turn core
Plain HMC needs you to choose the number of leapfrog steps $L$, and the right choice is hard. Too few steps and each iteration barely moves: a random walk again. Too many and the trajectory swings past the far side of the posterior and starts coming back, like a skateboarder rolling up the other wall of the bowl and down again; those extra gradients are wasted, and the end point may land near the start. Worse, the best length differs from one region of the posterior to another.
The No-U-Turn Sampler (NUTS) chooses the length itself, every iteration. It starts with one leapfrog step and keeps doubling the trajectory, each time extending it forward or backward in time at random. After each doubling it checks the two ends: if they have started moving toward each other, the path is making a U-turn, so it stops. Then it picks the next draw from the points along the whole trajectory (favouring points whose energy stayed close to the start), in a careful way that keeps the posterior as the stationary distribution.
Three ways to say it:
- Picture: keep rolling until you notice you are heading back where you came from.
- Numbers: trajectories of 1, 3, 7, 15, 31, … steps; NumPyro's default
max_tree_depth=10caps it at $2^{10} - 1 = 1\,023$ steps. - Slogan: NUTS = HMC that picks its own path length (and, in warmup, its own step size).
The U-turn check with numbers. The trajectory's backward end is at $q^- = (0, 0)$ with momentum $p^- = (1, 0)$; its forward end is at $q^+ = (2, 1)$ with momentum $p^+ = (-0.5, 0.8)$.
- Vector from the backward end to the forward end: $q^+ - q^- = (2, 1)$.
- Is the forward end still moving away? $(q^+ - q^-)\cdot p^+ = 2 \times (-0.5) + 1 \times 0.8 = -1 + 0.8 = -0.2 \lt 0$: no, it has turned back.
- The backward end: $(q^+ - q^-)\cdot p^- = 2 \times 1 + 1 \times 0 = 2 \gt 0$. (In backward time this end moves along $-p^-$, away from $q^+$.)
- One of the two dot products is negative, so a U-turn has started: stop doubling.
- Sizes: after the 1st, 2nd, 3rd, 4th doubling the trajectory has $1, 3, 7, 15$ leapfrog steps ($2^j - 1$). If the U-turn appears at doubling 5, NUTS spent 31 gradient evaluations on this iteration. NumPyro records this per draw as
num_steps.
- Doubling: at tree depth $j$ the trajectory is extended by $2^j$ leapfrog steps from its forward end (forward in time) or its backward end (backward in time, using −ε), chosen at random. Total steps after $j$ doublings: $2^j - 1$.
- U-turn criterion (simple form): stop when $(q^+ - q^-)\cdot p^- \lt 0$ or $(q^+ - q^-)\cdot p^+ \lt 0$ (with a mass matrix, the momenta are first multiplied by $M^{-1}$). NumPyro and Stan use a refined version based on sums of momenta and also check every sub-tree; the idea is the same.
- Also stop if the energy error explodes (a divergence) or if the depth reaches
max_tree_depth(default 10 in NumPyro). Hitting the maximum often is a sign of a hard geometry (very different scales or strong correlations). - Choosing the draw: the next state is sampled from all points of the trajectory with weights proportional to $e^{-H}$ (NumPyro uses this "multinomial" version, with a bias toward the newest half of the tree). This preserves the posterior as the stationary distribution.
- NUTS also needs a step size ε and a mass matrix; both are learned during warmup (next section). Reference: Hoffman and Gelman (2014), "The No-U-Turn Sampler".
Why do we need it?
It removes the hardest tuning choice of HMC (the trajectory length) and adapts it per iteration, so one sampler works out of the box on most continuous models. That is why it is the default everywhere.
Where is it used?
numpyro.infer.NUTS, Stan's default sampler, pm.sample() in PyMC, BlackJAX, TensorFlow Probability. Every "fit the model with MCMC" step in modern Bayesian workflows, including validating an SVI fit (Chapter 6.15).
How is it used?
MCMC(NUTS(model), num_warmup=500, num_samples=1000, num_chains=4).run(key, data). Check divergences and num_steps (via extra_fields); if trajectories keep hitting 1 023 steps, rescale or reparameterize the model rather than raising max_tree_depth blindly.
"NUTS has no tuning parameters."
It chooses the trajectory length per iteration and, in warmup, the step size and mass matrix. You still choose the warmup length, the number of chains, target_accept_prob and max_tree_depth, and the model's parameterization matters a lot.
"NUTS takes the end point of the trajectory as the next draw."
It samples a point from the whole trajectory, weighted by $e^{-H}$; taking the end point would break the guarantee that the posterior is stationary.
"If NUTS hits the maximum tree depth, raise max_tree_depth."
Long trajectories usually mean very different scales or strong correlations. Standardize inputs, reparameterize, or use a dense mass matrix first; a deeper tree only makes each iteration slower.
NUTS: double the trajectory (random direction) until $(q^+ - q^-)\cdot p^- \lt 0$ or $(q^+ - q^-)\cdot p^+ \lt 0$; then sample a point from the trajectory ∝ $e^{-H}$.
Steps after $j$ doublings: $2^j - 1$; max_tree_depth=10 → ≤ 1 023 per draw (num_steps).
Trap: NUTS still has settings (warmup, chains, target_accept_prob) and still depends on parameterization.
Quick check: ends $q^- = (1, 1)$, $p^- = (0, 1)$, $q^+ = (3, 0)$, $p^+ = (1, 0.5)$. Has a U-turn started?
$q^+ - q^- = (2, -1)$. $(2, -1)\cdot(1, 0.5) = 2 - 0.5 = 1.5 \gt 0$; $(2, -1)\cdot(0, 1) = -1 \lt 0$. One product is negative, so yes: the backward end has turned toward the other end. Stop doubling.
Warmup: how NUTS tunes its step size and mass matrix
Chapter 6.9 introduced warmup as "throw away the walk from the start". For NUTS it does a second job: learning its own settings. Two settings matter most.
- Step size ε. Too small wastes gradients; too big causes large energy errors and rejections (or divergences). During warmup NUTS nudges ε after every iteration: if the acceptance probability was above the target (NumPyro default 0.8) it makes ε bigger, if below, smaller. The nudges shrink over time and settle on a value.
- Mass matrix. If one parameter has posterior sd 10 and another 0.01, a single ε cannot suit both (the leapfrog stability limit is set by the narrowest direction). The mass matrix rescales the parameters so that every direction has a similar width. NUTS estimates the posterior variances from warmup draws and uses them.
Three ways to say it:
- Picture: before the race, the skateboarder tries a few runs to choose a stride and reshapes the bowl so it is round.
- Numbers: target acceptance 0.8 → the step size settles where about 80% of trajectories are accepted.
- Slogan: warmup = forget the start + learn ε + learn the scales.
Why the mass matrix matters: an independent posterior with sd 4 for $x$ and 0.1 for $y$.
- Unit mass: leapfrog is stable only for $\varepsilon \lt 2 \times 0.1 = 0.2$. Crossing $x$'s range (about $\pm 2$ sds = 16 units) at under 0.2 per step takes more than 80 steps, and the $y$ direction oscillates wildly meanwhile.
- Rescale: $z_x = x/4$, $z_y = y/0.1$. Both now have sd 1, the stability limit is about 2, and a trajectory of a handful of steps of ε ≈ 0.4 crosses the bulk.
- Using an inverse mass matrix $M^{-1} = \text{diag}(4^2, 0.1^2) = \text{diag}(16, 0.01)$ is exactly equivalent to this rescaling. NumPyro estimates these variances from warmup draws.
- In the widget below the effective draws per 1 000 gradients rise from about 5 (unit mass) to about 250 (adapted mass): roughly 50 times faster for the same posterior.
- Step-size adaptation (dual averaging): after warmup iteration $m$ with acceptance statistic $\alpha_m$, NumPyro (like Stan) updates $\bar H_m = (1 - \tfrac{1}{m + t_0})\bar H_{m-1} + \tfrac{1}{m+t_0}(\delta - \alpha_m)$ and $\log\varepsilon_m = \mu - \tfrac{\sqrt m}{\gamma}\bar H_m$, with target $\delta$ =
target_accept_prob(default 0.8), $\gamma = 0.05$, $t_0 = 10$, $\mu = \log(10\varepsilon_0)$. A running weighted average $\log\bar\varepsilon_m$ (weights $m^{-0.75}$) is the step size used after warmup. - Mass matrix adaptation: $M^{-1}$ ≈ estimated posterior covariance from warmup draws; diagonal by default in NumPyro (
dense_mass=False: only the variances), full withdense_mass=True(also captures correlations; costs $O(d^2)$ memory). - Warmup windows (Stan's schedule, used by NumPyro): a fast first window of 75 iterations (step size only), then "slow" windows of 25, 50, 100, … iterations that each end with a mass-matrix update (and a step-size restart), and a final fast window of 50. For
num_warmup=1000: 0–74, 75–99, 100–149, 150–249, 250–449, 450–949, 950–999. - After warmup, ε and $M$ are frozen; only then are draws recorded. Higher
target_accept_prob(e.g. 0.95) gives smaller ε: more accurate, slower trajectories.
Why do we need it?
Without adaptation every model would need hand-tuned step sizes and rescaling, and badly scaled parameters would make HMC crawl. Warmup is what makes NUTS work out of the box.
Where is it used?
NumPyro's NUTS(adapt_step_size=True, adapt_mass_matrix=True, dense_mass=False, target_accept_prob=0.8) defaults, Stan's adaptation, PyMC's tune phase. The adapted step size appears in mcmc.last_state.adapt_state.step_size.
How is it used?
Give enough warmup (several hundred iterations), standardize your inputs so parameters have similar scales (the global scaler in your A/B framework helps here too), raise target_accept_prob when you see a few divergences, and consider dense_mass=True for small models with strong correlations.
"The step size NUTS reports is a property of my model."
It is a tuning result: it depends on the target acceptance, the mass matrix, the warmup length and the random seed (the four chains in the Code-it adapted 0.19 to 0.25). A very small adapted ε is a hint that some region of the posterior is narrow or badly scaled.
"A diagonal mass matrix fixes correlations."
A diagonal $M$ only rescales each parameter. Strong correlations need dense_mass=True (feasible for small models) or a reparameterization.
"I can shorten warmup to save time; the samples will still be fine."
Too short a warmup leaves ε and $M$ poorly tuned (slow, or divergent) and may not even reach the posterior. Check R̂ and ESS; if they are bad, warmup is one of the first things to lengthen.
Warmup tunes ε by dual averaging toward target_accept_prob (0.8) and $M^{-1}$ ≈ posterior variances (diagonal by default), then freezes both.
Mass matrix = rescaling: $M^{-1} = \text{diag}(\sigma^2)$ ⇔ sampling $z = \theta/\sigma$. Stability limit ≈ 2 × narrowest sd in those coordinates.
Trap: diagonal $M$ does not fix correlations; tiny adapted ε signals a hard geometry.
Quick check: during warmup, an iteration's acceptance statistic is 0.95 while the target is 0.8. Which way does dual averaging push ε, and why?
Upward. The term $\delta - \alpha = 0.8 - 0.95 \lt 0$ lowers $\bar H$, and $\log\varepsilon = \mu - \frac{\sqrt m}{\gamma}\bar H$ then increases. Acceptance above target means the simulation is more accurate than needed, so bigger (cheaper) steps are allowed.
Several chains and their trace plots core
Chapter 6.9 showed that one chain can look perfectly healthy while missing half the posterior. The fix is simple: run several chains (four is the usual number), each from a different, spread-out starting point, and compare them. If they have all forgotten their starts and are exploring the same region with the same spread, their trace plots overlap like four intertwined caterpillars. If one chain sits somewhere else, wanders slowly, or gets stuck, the overlay shows it immediately.
It is like asking four people to explore a city independently and draw a map: if the four maps agree, you trust them; if one map shows a different city, someone got lost (or the city has two parts).
Three ways to say it:
- Picture: four coloured caterpillars on top of each other = good; four separate snakes = bad.
- Numbers: four chain means 0.07, 0.01, 0.00, 0.01 (agree, posterior sd 1) versus −1.28, 2.65, −0.71, 0.15 (disagree wildly).
- Slogan: one chain can lie; four chains from different starts rarely tell the same lie.
Reading four traces (a parameter with posterior mean 0 and sd 1, 400 kept draws per chain).
- All four bands overlap, each chain crosses the others hundreds of times, chain means 0.07, 0.01, 0.00, 0.01: the differences are a few hundredths of a posterior sd. Healthy.
- Four slow snakes with means −1.28, 2.65, −0.71, 0.15: each chain covers only part of the range in 400 draws. Not converged: run much longer or use a better sampler.
- Two chains near −2.5, two near +2.5, never crossing: two modes. Pooling them would give a mean near 0 where the posterior has almost no mass.
- Chains that agree most of the time but sometimes freeze for dozens of iterations at one level: sticky regions (a funnel neck), usually accompanied by divergences.
- Multiple chains: independent runs of the same sampler with different random seeds and dispersed initial values. NumPyro's default
init_to_uniformdraws starts uniformly in (−2, 2) on the unconstrained scale;MCMC(..., num_chains=4)runs four. - On a CPU, call
numpyro.set_host_device_count(4)before JAX starts so chains run in parallel; otherwise NumPyro falls back to running them one after another (chain_method="sequential";"vectorized"is another option). mcmc.get_samples(group_by_chain=True)returns arrays of shape (chains, draws, …), which is the input every multi-chain diagnostic needs.- What to check by eye: overlap (same level and spread), no trends, quick back-and-forth. ArviZ:
az.plot_trace, andaz.plot_rank(rank plots: each chain's ranks should be uniform).
Why do we need it?
Comparing chains is the only practical way to detect non-convergence: a chain cannot tell you about regions it has never visited, but another chain that started elsewhere can.
Where is it used?
Every serious MCMC run: num_chains=4 in NumPyro, four chains by default in Stan's CmdStanPy interface, and several chains by default in PyMC's pm.sample (check your version's default); R̂ (next section) is computed from them.
How is it used?
Run 4 chains, keep group_by_chain=True, overlay their traces per parameter (one colour per chain), then compute R̂ and ESS. Investigate any parameter whose chains do not overlap before using its numbers.
"Four chains is overkill; one long chain is the same."
One long chain cannot reveal a mode it never found, and its R̂ can only compare its own halves. Several dispersed chains are the main protection against confidently wrong answers, and on parallel hardware they cost no extra wall-clock time.
"If the chains disagree I can just pool them; the average is still fine."
Disagreeing chains mean the sampler has not explored the posterior. Pooling two modes gives a mean where there may be no posterior mass at all. Fix the problem first.
Run 4 chains from dispersed starts (num_chains=4, group_by_chain=True); overlay their traces.
Good: overlapping, flat, fast-mixing. Bad: separate levels (modes / non-convergence), slow snakes, frozen stretches.
Trap: pooling disagreeing chains hides the problem.
Quick check: three chains overlap nicely; the fourth sits at a level 3 posterior sds higher and never crosses the others. What do you conclude?
The sampler has not converged to a single picture of the posterior: either the fourth chain is stuck in a second mode (which might be real), or the other three are stuck and it is right. Do not report pooled numbers. Investigate the model (multimodality, identifiability, label switching) and run longer or reparameterize.
R̂ ("R-hat"): do the chains agree? core
R̂ turns "do the traces overlap?" into one number. It compares two spreads: the spread of draws inside each chain, and the spread you get when you pool all chains together. If every chain has explored the whole posterior, pooling adds nothing: both spreads are the same and R̂ = 1. If the chains sit in different places, pooling adds the gaps between them, the pooled spread is bigger, and R̂ is above 1.
The modern version first splits each chain into a first and a second half and treats them as separate chains. That way a single chain that is still drifting (its first half differs from its second half) also raises R̂.
Three ways to say it:
- Picture: four caterpillars on top of each other (R̂ ≈ 1) vs four caterpillars at different heights (R̂ ≫ 1).
- Numbers: half-chains whose means differ by ±0.5 posterior sds give R̂ = 1.15; means differing by ±0.1 give R̂ = 1.002.
- Slogan: R̂ = √(total spread ÷ within-chain spread); aim for ≤ 1.01.
Split R̂ by hand. After splitting, we have 4 half-chains of $n = 100$ draws, each with within-chain variance 1.
- Within: $W = $ average of the within variances $= 1$.
- Half-chain means 0.5, −0.5, 0.5, −0.5. Their average is 0; their variance (dividing by $4 - 1 = 3$) is $(0.25 \times 4)/3 = 0.333$. Between: $B = n \times 0.333 = 33.3$.
- Pooled estimate of the posterior variance: $\hat V = \frac{n-1}{n}W + \frac{B}{n} = 0.99 \times 1 + 33.3/100 = 0.99 + 0.333 = 1.323$.
- $\hat R = \sqrt{\hat V/W} = \sqrt{1.323} = 1.150$: far above 1.01, not converged.
- Same but with means 0.1, −0.1, 0.1, −0.1: variance of the means $0.04/3 = 0.0133$, $B = 1.33$, $\hat V = 0.99 + 0.0133 = 1.0033$, $\hat R = 1.0017$. Differences of ±0.1 sd with 100 draws per half are about what pure chance produces, so R̂ is essentially 1.
Take $m$ chains, split each in half: $2m$ sequences of length $n$. With $\bar\theta_j$ and $s_j^2$ the mean and variance of sequence $j$ and $\bar\theta$ the overall mean:
$$W = \frac{1}{2m}\sum_j s_j^2, \qquad B = \frac{n}{2m-1}\sum_j(\bar\theta_j - \bar\theta)^2, \qquad \hat V = \frac{n-1}{n}W + \frac{B}{n}, \qquad \hat R = \sqrt{\hat V / W}.$$- $\hat R \to 1$ as the chains converge to the same distribution. It is computed per parameter (and per quantity you care about).
- Rule of thumb (a convention, not a test): $\hat R \le 1.01$ (Vehtari, Gelman, Simpson, Carpenter and Bürkner, 2021); the older threshold 1.1 is now considered too lenient.
- NumPyro's
r_hatcolumn is this split R̂ (numpyro.diagnostics.split_gelman_rubin). ArviZ'saz.rhat/az.summaryuse a rank-normalized split R̂, more robust for heavy tails; values are usually close. - R̂ near 1 is necessary, not sufficient: if all chains are stuck in the same wrong region (for example all started in one mode), R̂ can be 1.00.
Why do we need it?
With hundreds of parameters (segments, changepoint slopes, Fourier coefficients) you cannot eyeball every trace. R̂ is the automatic screen that flags which parameters' chains disagree.
Where is it used?
The r_hat column of mcmc.print_summary(), numpyro.diagnostics.split_gelman_rubin, az.rhat and az.summary in ArviZ, Stan's and PyMC's summaries, and automated checks in production MCMC pipelines.
How is it used?
Compute R̂ for every parameter (and key derived quantities such as a lift or $P(\theta_B \gt \theta_A)$); sort by R̂; inspect the traces of the worst ones; if any exceed 1.01, run longer, reparameterize, or fix the model before using the results.
"R̂ = 1.00, so the posterior is correct."
R̂ only says the chains agree with each other. All four could miss the same mode, or all be blocked from the same funnel neck (divergences would hint at that). Check divergences, ESS, traces and posterior predictive checks too.
"R̂ below 1.1 is fine."
That was the old convention. Current guidance is 1.01, together with enough effective draws; 1.05 on a key parameter deserves a longer run.
"I only need R̂ for the parameters I report."
Check hyperparameters (τ, σ, dispersion) too: if they have not converged, nothing that depends on them has.
Split each chain in half; $W$ = mean within variance, $B = n\,\text{Var}(\text{half means})$, $\hat V = \frac{n-1}{n}W + \frac Bn$, $\hat R = \sqrt{\hat V/W}$.
Target ≤ 1.01 (old: 1.1). NumPyro: split R̂; ArviZ: rank-normalized split R̂.
Trap: R̂ ≈ 1 is necessary, not sufficient (chains can share the same blind spot).
Quick check: 4 half-chains of 50 draws, within variances all 2.0, means 1.0, 1.2, 0.8, 1.0. Compute split R̂.
$W = 2$. Mean of means 1.0; variance of means $= (0 + 0.04 + 0.04 + 0)/3 = 0.0267$; $B = 50 \times 0.0267 = 1.333$. $\hat V = 0.98 \times 2 + 1.333/50 = 1.96 + 0.0267 = 1.9867$. $\hat R = \sqrt{1.9867/2} = 0.997$, essentially 1: these differences are what chance produces with 50 draws per half (standard error of a half-mean ≈ $\sqrt{2/50} = 0.2$).
Effective sample size and the Monte Carlo standard error core
Converged chains still produce correlated draws, so 4 000 draws are worth fewer than 4 000 independent ones. The effective sample size (ESS, Chapter 6.9) is that "worth". It answers the practical question: how precisely do my draws pin down a posterior summary? The answer is the Monte Carlo standard error (MCSE) = posterior sd ÷ √ESS: the typical size of the error that comes only from using a finite number of draws.
MCSE is not the posterior uncertainty. The posterior sd says how uncertain the parameter is given the data; the MCSE says how uncertain your computer's estimate of a posterior summary is. Running longer shrinks the MCSE, never the posterior sd.
Three ways to say it:
- Picture: the posterior is a blurry photo; MCSE is the grain from printing it with too few dots.
- Numbers: sd 0.160 and ESS 924 give MCSE $0.160/\sqrt{924} = 0.0053$; pretending the 4 000 draws were independent would claim 0.0025.
- Slogan: MCSE = sd/√ESS, not sd/√n.
From the Code-it run below: the average log-odds μ of the segment model, 4 chains × 1 000 draws.
print_summaryreports mean −2.11, std 0.16,n_eff923.57.- MCSE of the posterior mean $= 0.160/\sqrt{924} = 0.160/30.4 = 0.0053$.
- So the posterior mean is $-2.113 \pm 0.011$ (two MCSE): two decimals are trustworthy, the third is not.
- The naive $0.160/\sqrt{4000} = 0.0025$ would be twice too optimistic.
- For a probability such as $P(\theta_B \gt \theta_A\mid D) = 0.92$ estimated from draws with ESS 1 000: MCSE $\approx\sqrt{0.92 \times 0.08/1000} = 0.0086$. Reporting "0.92" is fine; "0.9214" is not.
- Rule of thumb: total ESS of at least about 400 (100 per chain with 4 chains) before trusting R̂ and summaries; more for tail quantities such as a 95% interval endpoint.
- Multi-chain ESS (Stan, NumPyro): estimate the autocorrelations $\hat\rho_k = 1 - (W - \bar\gamma_k)/\hat V$ from all chains ($\bar\gamma_k$ = average autocovariance at lag $k$), add them in pairs while the pair sums stay positive and decreasing (Geyer's initial monotone sequence), and set $\text{ESS} = mn/\hat\tau$ with $\hat\tau = -1 + 2\sum(\text{pairs})$. This is NumPyro's
n_eff(numpyro.diagnostics.effective_sample_size). - Bulk-ESS (ArviZ
ess_bulk): the same computed after replacing draws by their normalized ranks; it measures how reliable the centre (mean, median) is. Tail-ESS (ess_tail): ESS of the indicators "draw ≤ 5% quantile" and "draw ≤ 95% quantile"; it measures how reliable interval endpoints are. Tails mix more slowly, so tail-ESS is often smaller. - Monte Carlo standard error of a posterior mean: $\text{MCSE} = \text{sd}/\sqrt{\text{ESS}}$. For a probability $q$: $\sqrt{q(1-q)/\text{ESS}}$. (ArviZ:
az.mcse.) - Rule of thumb (Vehtari et al., 2021): bulk-ESS and tail-ESS ≥ 100 per chain (≥ 400 for 4 chains); then make the MCSE small compared with the precision your decision needs.
Why do we need it?
It decides how many digits you may report and whether you need more draws. Without it, a sticky chain's 4 000 draws look as convincing as 4 000 independent ones.
Where is it used?
The n_eff column of NumPyro's print_summary, ArviZ's ess_bulk, ess_tail, mcse_mean in az.summary, Stan's summaries, and stopping rules for long MCMC runs.
How is it used?
Read ESS for every parameter and derived quantity; compute MCSE = sd/√ESS; if MCSE is too big for the decision (for example, $P(B \gt A)$ near your 0.95 threshold), run more draws or improve the sampler, then report numbers rounded to their MCSE.
"MCSE is the uncertainty about the parameter."
The posterior sd is the uncertainty about the parameter. MCSE is the extra, computational error in your estimate of a posterior summary; it shrinks with more draws, the posterior sd does not.
"ESS is the same for every quantity."
Each parameter and each summary has its own ESS. Means usually mix fastest; tail quantiles (interval endpoints) and variance parameters are often much slower. Check tail-ESS before trusting a 95% interval.
"An ESS larger than the number of draws is a bug."
It can happen: NUTS often produces anti-correlated draws (a trajectory that ends on the other side of the posterior), and then ESS for the mean can exceed $n$, so ESS can approach or exceed the number of draws. In the Code-it the z rows reach 2 000–3 100 effective draws out of 4 000; on simple, well-behaved targets values above 4 000 can occur.
In an A/B framework like yours, the decision quantity is a probability such as $P(\theta_B - \theta_A \gt \delta\mid D)$. If you ever compute it from MCMC draws (for example when validating SVI with NUTS), its Monte Carlo error is $\sqrt{q(1-q)/\text{ESS}}$: with ESS 400 and $q = 0.95$ that is 0.011, so "0.95 vs 0.94" is not a real difference and a decision threshold at 0.95 needs more draws. The same thinking applies to draws from an SVI guide: there the draws are independent, so ESS equals the number of draws, but the guide's own approximation error is not measured by any MCSE (Chapter 6.15).
ESS = $mn/\hat\tau$ (NumPyro n_eff); bulk-ESS for centres, tail-ESS for interval ends (ArviZ).
MCSE(mean) = sd/√ESS; MCSE(probability) = $\sqrt{q(1-q)/\text{ESS}}$. Rule of thumb: ESS ≥ 400 total (100 per chain).
Trap: MCSE ≠ posterior sd; never use sd/√n for MCMC draws.
Quick check: a forecast parameter has posterior sd 2.4 and ESS 144. What is the MCSE of its posterior mean, and how should you report a mean of 17.326?
MCSE $= 2.4/\sqrt{144} = 2.4/12 = 0.2$. Two MCSE is 0.4, so report about 17.3 (±0.4 from Monte Carlo error alone); the digits "26" are noise. For a tighter estimate, run longer: quadrupling ESS halves the MCSE.
Divergences: when the simulation blows up core
A divergence is a trajectory whose energy exploded: the leapfrog simulation went unstable because it entered a region where the posterior curves far more sharply than the step size can handle. The classic place is the narrow neck of the funnel that a centered hierarchical model creates when the group spread τ is small (Chapter 6.7). With a step size tuned for the wide mouth, the puck that enters the neck is thrown out violently.
Why it matters: the sampler cannot get into that region, so it quietly under-samples it. Your posterior summaries are then biased (for a hierarchical model, typically too few draws with small τ), and nothing else in the output may look wrong. NUTS flags each such transition for you. Treat divergences as a warning light on the dashboard, not as a detail.
Three ways to say it:
- Picture: a skateboarder with long strides cannot ride down a narrow chute; he bounces out, so the chute is never explored.
- Numbers: the centered eight-schools model gives 260 divergences out of 4 000 draws; asking for smaller steps (target 0.95) still leaves 123; non-centering gives 4, then 0.
- Slogan: divergences mean biased draws; fix the geometry, not just the step size.
The eight-schools funnel (four chains × 1 000 draws, from the Code-it below).
- Centered,
target_accept_prob=0.8: 260 divergences, R̂(τ) = 1.071, ESS(τ) = 124. Clearly broken. - Centered, 0.95 (smaller steps): 123 divergences, R̂ 1.064, ESS 115. Smaller steps help a little but cannot fit both the neck and the mouth.
- Non-centered ($\theta_j = \mu + \tau z_j$, $z_j \sim N(0, 1)$), 0.8: 4 divergences, R̂ 1.001, ESS 3 353.
- Non-centered, 0.95: 0 divergences, R̂ 1.000, ESS 3 753.
- Lesson: change the coordinates first (Chapter 6.7); raise
target_accept_probto mop up a few remaining divergences.
- Divergent transition: an iteration whose energy error exceeds a threshold; NumPyro uses $\Delta H \gt 1000$ (
max_delta_energy), Stan the same. NUTS stops building that trajectory and the iteration is flagged (extra_fields=("diverging",); also printed as "Number of divergences" byprint_summary). - Cause: curvature too high for the step size somewhere the posterior has mass (funnels, sharp ridges, near hard boundaries).
- Consequence: that region is under-explored, so estimates are biased, and the bias does not shrink by running longer.
- Responses, in order: (1) reparameterize (non-centered hierarchical effects, standardize predictors, unconstrained scales); (2) raise
target_accept_prob(0.9, 0.95, 0.99) for smaller steps; (3) reconsider priors (a weakly informative prior on τ keeps it away from the extreme neck). Never just delete divergent draws. - Locate them: plot divergent iterations in pairs of parameters (for example log τ against a group effect); they cluster where the trouble is.
Why do we need it?
Divergences are the one diagnostic that points at bias rather than just imprecision, and they are specific to the gradient-based samplers you use. A run with divergences can have R̂ = 1.00 and large ESS and still be wrong.
Where is it used?
NumPyro's diverging extra field and the "Number of divergences" line, Stan's divergence warnings, ArviZ's az.plot_pair(..., divergences=True). Hierarchical A/B models, random-effects models and Gaussian-process length scales are the usual suspects.
How is it used?
Always check that the divergence count is 0 (a handful is still a warning). If not, find where they cluster, reparameterize, then raise target_accept_prob, and rerun until there are none.
"A few divergences out of thousands of draws are harmless."
Even a handful means some region could not be explored, and that region may matter (small τ = strong pooling). Treat any divergence as a reason to investigate; aim for zero.
"Just raise target_accept_prob to 0.99 and the problem is solved."
Smaller steps reduce divergences but cost much more time and often cannot remove them (123 remained at 0.95 in the centered model). The real fix is usually a better parameterization.
"No divergences, so the posterior is right."
No divergences is necessary for trust, not sufficient. You still need R̂, ESS, traces and posterior predictive checks.
Your A/B framework pools segments hierarchically. When the data say segments are similar, the posterior of the between-segment spread τ piles up near 0, which is exactly the funnel neck. If you validate the model with NUTS and see divergences, switch to the non-centered form $\theta_g = \mu + \tau z_g$ (Chapter 6.7). SVI does not report divergences: a Gaussian guide can sit happily in the mouth of the funnel and simply miss the neck, so the absence of warnings from SVI is not evidence that the geometry is easy (Chapter 6.7).
Divergence: ΔH > 1000 in a trajectory (NumPyro diverging). Cause: curvature too high for ε (funnel neck).
Effect: bias (region under-sampled), not just noise. Fix: reparameterize first, then raise target_accept_prob.
Eight schools: centered 260 → 123 (0.95); non-centered 4 → 0. Trap: never ignore or delete them.
Quick check: a run has R̂ = 1.00 for every parameter, ESS above 2 000, and 37 divergences. Can you report it?
No. Divergences indicate that part of the posterior could not be explored, so the estimates may be biased even though the chains agree and mix well (they can agree on the same blind spot). Find where the divergences occur, reparameterize, and rerun until there are none.
Putting it together: reading print_summary like a checklist core
A pilot does not take off because the plane "looks fine"; they go through a checklist in a fixed order. After every NUTS run, do the same. The order matters, because some failures make later numbers meaningless: if there are divergences, a perfect R̂ means nothing; if R̂ is bad, ESS is meaningless; only when the sampler is healthy do Monte Carlo errors and the posterior itself deserve attention.
NumPyro prints most of what you need in one table, mcmc.print_summary(): one row per parameter with the posterior mean, sd, median, a 90% interval, n_eff (ESS) and r_hat, followed by the number of divergences.
Three ways to say it:
- Picture: a pre-flight checklist: divergences, R̂, ESS, traces, MCSE, then the model itself.
- Numbers: pass if divergences = 0, every r_hat ≤ 1.01, every n_eff ≥ 400 (4 chains), and MCSE small for the decision.
- Slogan: sampler first, model second, conclusions last.
Reading the Code-it's summary (segment model, 4 chains × 1 000 draws).
- Divergences: "Number of divergences: 0". Pass.
- R̂: μ 1.01, τ 1.00, every $z$ 1.00. μ is right at the threshold: acceptable for a first look, and a longer run is cheap if μ drives a decision.
- ESS: μ 924, τ 944, the $z$'s 2 000–3 100. All far above 400. Pass.
- MCSE: μ's mean $-2.11$ has MCSE $0.16/\sqrt{924} = 0.005$; report μ ≈ −2.11 (sd 0.16).
- Columns 5.0% and 95.0%: the 90% highest posterior density interval (HPDI: the shortest interval that holds 90% of the draws;
prob=0.9), not the 5th and 95th percentiles. For τ it is [0.00, 0.58]: its posterior piles up near 0, so the shortest interval starts at the boundary. - Sampler verdict: healthy. Next: posterior predictive checks (Chapter 6.8) before trusting the model's conclusions.
MCMC diagnostics checklist (stop at the first failure and fix it):
- Divergences = 0 (
extra_fields=("diverging",)). Otherwise reparameterize / raisetarget_accept_prob. - R̂ ≤ 1.01 for every parameter, hyperparameter and key derived quantity. Otherwise run longer, check for modes, reparameterize.
- ESS ≥ about 400 in total (bulk and tail in ArviZ). Otherwise run longer or improve the geometry.
- Traces / rank plots look like overlapping caterpillars;
num_stepsis not stuck at the maximum ($2^{10} - 1 = 1\,023$). - MCSE = sd/√ESS small compared with the precision your decision needs.
- Only then: model checks (posterior predictive checks, prior sensitivity; Chapter 6.8) and conclusions.
print_summary(prob=0.9)columns:mean,std,median, the HPDI bounds labelled5.0%and95.0%,n_eff(classic multi-chain ESS),r_hat(split R̂). ArviZ'saz.summaryaddsmcse_mean,ess_bulk,ess_tailand uses rank-normalized R̂.- Passing every check means "no evidence of sampling problems", not "exact posterior".
Why do we need it?
MCMC fails silently: it always returns numbers. A fixed checklist catches the common failures before they reach a decision, and gives you a crisp, defensible answer when someone asks "how do you know it converged?".
Where is it used?
After every mcmc.run: mcmc.print_summary(), numpyro.diagnostics, ArviZ (az.summary, az.plot_trace, az.plot_rank, az.plot_pair(divergences=True)), and automated checks in production pipelines that refit Bayesian models on a schedule.
How is it used?
Run 4 chains with extra_fields=("diverging", "num_steps"); read divergences, then r_hat, then n_eff; plot the traces of the worst parameters; compute MCSE for the decision quantity; then run posterior predictive checks. Log the settings and the diagnostics with the result.
print_summary (first rows of the Code-it output) with each part labelled. The interval columns are the bounds of the 90% highest posterior density interval; n_eff is the effective sample size and r_hat the split R̂."All diagnostics passed, so the model is right."
The diagnostics check the sampler: whether the draws represent the posterior of the model you wrote. Whether that model fits the data is a separate question for posterior predictive checks and prior sensitivity (Chapter 6.8).
"The 5.0% and 95.0% columns are the 5th and 95th percentiles."
In NumPyro's print_summary they are the bounds of the 90% highest posterior density interval (prob=0.9). For skewed posteriors (like τ near 0) they differ from the equal-tailed percentiles.
"NUTS gives the exact posterior."
NUTS is asymptotically exact: its draws converge to the posterior as the run grows, but any finite run has Monte Carlo error (MCSE = sd/√ESS) and may not have converged (R̂, divergences). Chapter 6.15 develops this comparison with SVI.
Model answer: "I'd say NUTS targets the exact posterior, and with a healthy run (no divergences, R̂ ≤ 1.01, enough ESS) its summaries are accurate up to a Monte Carlo error I can quantify. That is different from VI, whose error comes from the guide family and does not shrink with more computation. But a finite MCMC run is never exact, and the diagnostics can only reveal problems, not prove their absence."
Both of your projects run SVI, so where does this checklist live in them? In validation. A defensible workflow: on a smaller version of the forecasting model (fewer changepoints, a shorter history) or of the hierarchical A/B model, run NUTS with 4 chains, pass the checklist, and compare its posterior means, intervals and decision probabilities with your SVI fit (Chapter 6.15). If NUTS itself fails the checklist (for example, divergences around the hierarchical τ or around Laplace-prior changepoint slopes $\delta_j$ concentrated near 0), that is a geometry problem your SVI guide is facing too, silently.
Checklist: divergences = 0 → R̂ ≤ 1.01 → ESS ≥ 400 (bulk, tail) → traces / tree depth → MCSE vs decision → posterior predictive checks.
print_summary: mean, std, median, 90% HPDI (5.0%/95.0%), n_eff, r_hat, "Number of divergences".
Say it right: NUTS is asymptotically exact, never "exact"; passing diagnostics ≠ correct model.
Quick check: a summary shows 0 divergences, r_hat 1.00 everywhere, but n_eff = 85 for a variance parameter with 4 × 1 000 draws. What do you do?
The sampler is healthy but inefficient for that parameter: 85 effective draws give a large MCSE and unreliable tails. Run longer (or improve the parameterization, e.g. sample its log, use a dense mass matrix) until ESS is at least about 400, especially if intervals for that parameter matter.
Recap, cheat sheet and practice
- Random walks diffuse (distance ≈ ε√n). HMC uses the gradient $\nabla\log\tilde p$ and a random momentum to glide far in one iteration.
- Physics picture: $U(\theta) = -\log\tilde p(\theta)$, $K(p) = \tfrac12 p^\top M^{-1}p$, $H = U + K$; exact motion conserves $H$.
- Leapfrog (half kick, drift, half kick) is reversible and volume-preserving; its energy error is small and bounded until ε exceeds about 2 × the narrowest sd, where it explodes (a divergence).
- One HMC iteration: draw $p$, run $L$ leapfrog steps, accept with $\min(1, e^{-\Delta H})$. Far moves with high acceptance; judge efficiency by ESS per gradient.
- NUTS doubles the trajectory until it makes a U-turn (at most $2^{10} - 1 = 1\,023$ steps by default) and samples the next point from the trajectory. Warmup adapts ε by dual averaging toward
target_accept_prob= 0.8 and a (diagonal) mass matrix from the posterior variances. - Diagnostics: several chains and their trace plots; split R̂ ≤ 1.01; ESS ≥ ~400 (bulk and tail); MCSE = sd/√ESS; divergences = 0 (bias, not noise; fix the geometry).
- Passing the checklist means the draws represent your model's posterior up to Monte Carlo error. NUTS is asymptotically exact, never "exact".
Cheat sheet
| Idea | Formula / rule | In words |
|---|---|---|
| Energy | $H(\theta, p) = -\log\tilde p(\theta) + \tfrac12 p^\top M^{-1}p$, $p \sim N(0, M)$ | height + motion; momentum is auxiliary |
| Leapfrog | $p \mathrel{+}= \tfrac\varepsilon2\nabla\log\tilde p$; $\theta \mathrel{+}= \varepsilon M^{-1}p$; $p \mathrel{+}= \tfrac\varepsilon2\nabla\log\tilde p$ | kick, drift, kick; $L$ gradients per trajectory |
| Stability | ε ≲ 2 × narrowest sd (in mass-scaled units) | beyond it the energy explodes |
| Acceptance | $\min(1, e^{-\Delta H})$ | ΔH = 0.1 → 0.905; 2 → 0.135 |
| U-turn | stop when $(q^+ - q^-)\cdot p^\pm \lt 0$ | trajectory sizes $2^j - 1$; max depth 10 |
| NumPyro NUTS defaults | target_accept_prob=0.8, max_tree_depth=10, dense_mass=False, init_to_uniform | warmup adapts ε and diagonal $M$ |
| Warmup schedule | 75 fast + slow windows 25, 50, 100, … + 50 fast | ε and $M$ frozen afterwards |
| Split R̂ | $\sqrt{\hat V/W}$, $\hat V = \frac{n-1}{n}W + \frac Bn$ | ≤ 1.01; necessary, not sufficient |
| ESS | $mn/\hat\tau$ (multi-chain, Geyer); bulk and tail versions in ArviZ | ≥ 400 total (100 per chain) |
| MCSE | sd/√ESS; probability: $\sqrt{q(1-q)/\text{ESS}}$ | decides the digits you may report |
| Divergence | ΔH > 1000 (max_delta_energy) | bias; reparameterize, then raise target_accept_prob |
print_summary | mean, std, median, 90% HPDI, n_eff, r_hat, divergences | checklist: div → r_hat → n_eff → traces → MCSE |
import numpyro
numpyro.set_host_device_count(4) # 4 CPU "devices", so 4 chains run in parallel (call before JAX starts)
import jax, jax.numpy as jnp, numpy as np
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS
from numpyro.diagnostics import effective_sample_size, split_gelman_rubin
# 1) A small A/B-style hierarchical model: one conversion rate per segment, partial pooling on the log-odds scale
n = np.array([120, 340, 95, 410, 60, 220, 180, 75]) # users per segment
k = np.array([10, 41, 15, 37, 3, 31, 13, 13]) # conversions per segment
def model(n, k=None):
mu = numpyro.sample("mu", dist.Normal(-2.0, 1.0)) # average log-odds
tau = numpyro.sample("tau", dist.HalfNormal(1.0)) # spread between segments
with numpyro.plate("seg", len(n)):
z = numpyro.sample("z", dist.Normal(0.0, 1.0)) # non-centered (Chapter 6.7)
theta = numpyro.deterministic("theta", jax.nn.sigmoid(mu + tau * z))
numpyro.sample("k", dist.Binomial(total_count=n, probs=theta), obs=k)
mcmc = MCMC(NUTS(model), num_warmup=500, num_samples=1000, num_chains=4, progress_bar=False)
mcmc.run(jax.random.PRNGKey(0), n=n, k=k, extra_fields=("diverging", "num_steps"))
mcmc.print_summary() # mean, std, median, 90% HPDI (5.0% / 95.0%), n_eff, r_hat
# mean std median 5.0% 95.0% n_eff r_hat
# mu -2.11 0.16 -2.11 -2.37 -1.87 923.57 1.01
# tau 0.32 0.20 0.30 0.00 0.58 944.36 1.00
# z[0] -0.43 0.77 -0.43 -1.68 0.84 3130.28 1.00
# z[1] 0.27 0.67 0.28 -0.82 1.39 2213.26 1.00
# z[2] 0.59 0.77 0.59 -0.59 1.89 2508.47 1.00
# z[3] -0.48 0.68 -0.45 -1.67 0.52 2011.33 1.00
# z[4] -0.67 0.86 -0.67 -2.10 0.72 3085.18 1.00
# z[5] 0.58 0.70 0.58 -0.49 1.78 2056.16 1.00
# z[6] -0.71 0.76 -0.71 -2.01 0.44 2325.85 1.00
# z[7] 0.68 0.84 0.70 -0.64 2.14 2379.36 1.00
#
# Number of divergences: 0
ex = mcmc.get_extra_fields()
print("divergences:", int(ex["diverging"].sum()), "| leapfrog steps per draw: median",
int(np.median(ex["num_steps"])), "max", int(ex["num_steps"].max()))
print("adapted step size per chain:", np.round(np.asarray(mcmc.last_state.adapt_state.step_size), 3))
# divergences: 0 | leapfrog steps per draw: median 15 max 63
# adapted step size per chain: [0.224 0.19 0.219 0.25 ]
# 2) Monte Carlo standard error = sd / sqrt(ESS), by hand, for mu
mu = np.asarray(mcmc.get_samples(group_by_chain=True)["mu"]) # shape (4 chains, 1000 draws)
ess = float(effective_sample_size(mu)); sd = mu.std()
print(f"mu: mean {mu.mean():.3f} sd {sd:.3f} ESS {ess:.0f} MCSE {sd / np.sqrt(ess):.4f} "
f"naive sd/sqrt(n) {sd / np.sqrt(mu.size):.4f}")
# mu: mean -2.113 sd 0.160 ESS 924 MCSE 0.0053 naive sd/sqrt(n) 0.0025 <- naive is 2x too optimistic
# 3) Divergences: the eight-schools funnel, centered vs non-centered (Chapter 6.7)
y = np.array([28., 8., -3., 7., -1., 1., 18., 12.]); s = np.array([15., 10., 16., 11., 9., 11., 10., 18.])
def centered(y, s):
mu = numpyro.sample("mu", dist.Normal(0, 5)); tau = numpyro.sample("tau", dist.HalfCauchy(5))
with numpyro.plate("J", 8):
theta = numpyro.sample("theta", dist.Normal(mu, tau))
numpyro.sample("y", dist.Normal(theta, s), obs=y)
def noncentered(y, s):
mu = numpyro.sample("mu", dist.Normal(0, 5)); tau = numpyro.sample("tau", dist.HalfCauchy(5))
with numpyro.plate("J", 8):
z = numpyro.sample("z", dist.Normal(0, 1))
numpyro.sample("y", dist.Normal(mu + tau * z, s), obs=y)
for name, m in [("centered", centered), ("non-centered", noncentered)]:
for target in (0.8, 0.95):
mc = MCMC(NUTS(m, target_accept_prob=target), num_warmup=500, num_samples=1000, num_chains=4, progress_bar=False)
mc.run(jax.random.PRNGKey(1), y=y, s=s, extra_fields=("diverging",))
tau_s = np.asarray(mc.get_samples(group_by_chain=True)["tau"])
print(f"{name:12s} target_accept {target}: divergences {int(mc.get_extra_fields()['diverging'].sum()):3d} "
f"tau R-hat {float(split_gelman_rubin(tau_s)):.3f} ESS {float(effective_sample_size(tau_s)):5.0f}")
# centered target_accept 0.8: divergences 260 tau R-hat 1.071 ESS 124
# centered target_accept 0.95: divergences 123 tau R-hat 1.064 ESS 115 <- smaller steps are not enough
# non-centered target_accept 0.8: divergences 4 tau R-hat 1.001 ESS 3353
# non-centered target_accept 0.95: divergences 0 tau R-hat 1.000 ESS 3753 <- fix the geometry first
# 4) What a failed R-hat looks like: two far-apart modes, chains start in different places
def two_modes():
x = numpyro.sample("x", dist.Normal(0, 10))
mix = jnp.logaddexp(dist.Normal(-4, 0.5).log_prob(x), dist.Normal(4, 0.5).log_prob(x)) + jnp.log(0.5)
numpyro.factor("modes", mix - dist.Normal(0, 10).log_prob(x)) # target = the 50/50 mixture
mc = MCMC(NUTS(two_modes), num_warmup=300, num_samples=500, num_chains=4, progress_bar=False)
mc.run(jax.random.PRNGKey(3))
x = np.asarray(mc.get_samples(group_by_chain=True)["x"])
print("chain means:", np.round(x.mean(axis=1), 2), " split R-hat:", round(float(split_gelman_rubin(x)), 2))
# chain means: [-3.99 4. -4. -4. ] split R-hat: 7.85 <- chains disagree: never pool these
# (whole script: about 5 seconds on a laptop CPU; exact numbers can differ slightly across JAX versions and machines)
1. In HMC, what is the momentum $p$?
2. An HMC trajectory has energy error ΔH = 0.7. What is its acceptance probability?
3. When does NUTS stop extending a trajectory?
4. A 4-chain NUTS run reports r_hat = 1.08 for τ and 1.00 for everything else, with no divergences. What is the right reaction?
5. A parameter has posterior sd 0.5 and ESS 100. What is the Monte Carlo standard error of its posterior mean?
6. A centered hierarchical model gives 12 divergences, R̂ = 1.00 and ESS above 1 000 for every parameter. What does it mean?
Practice problems
A. Leapfrog by hand on $U(\theta) = \theta^2/2$: start θ = 0, $p = 1$, ε = 0.5. Do two steps, compute the energy after each, and the acceptance probability.
- $H_0 = 0 + 1^2/2 = 0.5$.
- Step 1: half kick $p = 1 - 0.25 \times 0 = 1$; drift θ = $0 + 0.5 \times 1 = 0.5$; half kick $p = 1 - 0.25 \times 0.5 = 0.875$. $H = 0.125 + 0.3828 = 0.5078$.
- Step 2: half kick $p = 0.875 - 0.125 = 0.75$; drift θ = $0.5 + 0.375 = 0.875$; half kick $p = 0.75 - 0.25 \times 0.875 = 0.53125$. $H = 0.3828 + 0.1411 = 0.5239$.
- ΔH = 0.0239, acceptance $e^{-0.0239} = 0.976$. (Exact motion: θ = sin 1 = 0.841, p = cos 1 = 0.540.)
B. A posterior has sds 3 and 0.05 along two independent directions. With unit mass, what limits the step size, and roughly how many steps does a trajectory need to cross the wide direction? What changes with an adapted diagonal mass matrix?
- Stability needs ε ≲ 2 × 0.05 = 0.1 (the narrow direction).
- Crossing ±2 sds of the wide direction is 12 units: at least 120 leapfrog steps per trajectory.
- With $M^{-1} = \text{diag}(9, 0.0025)$ the sampler works in units where both sds are 1: ε can be near 1 (limit about 2), and a handful of steps crosses the posterior. That is a speed-up of more than 20× in gradients per iteration.
C. NUTS ends: $q^- = (0, 0)$, $p^- = (1, 1)$, $q^+ = (2, 3)$, $p^+ = (0.5, -1)$. Should it keep doubling?
$q^+ - q^- = (2, 3)$. $(2, 3)\cdot(1, 1) = 5 \gt 0$; $(2, 3)\cdot(0.5, -1) = 1 - 3 = -2 \lt 0$. The forward end has started to come back: a U-turn. Stop and sample the next draw from the trajectory.
D. Two chains are split into 4 halves of $n = 200$, each with within variance 0.5. The half means are 2.0, 2.1, 1.9, 3.0. Compute split R̂. Which half is the problem?
- $W = 0.5$. Mean of the means $= 9.0/4 = 2.25$.
- Squared deviations: $0.0625, 0.0225, 0.1225, 0.5625$, sum $0.77$; variance $0.77/3 = 0.2567$; $B = 200 \times 0.2567 = 51.3$.
- $\hat V = \frac{199}{200}\times 0.5 + 51.3/200 = 0.4975 + 0.2567 = 0.754$. $\hat R = \sqrt{0.754/0.5} = \sqrt{1.508} = 1.228$.
- The fourth half (mean 3.0) is the outlier: the second chain drifted in its second half. Splitting the chains is exactly what exposes this kind of drift.
E. From NUTS draws with ESS 600 you estimate $P(\theta_B - \theta_A \gt \delta\mid D) = 0.955$ and your decision threshold is 0.95. Can you decide? How large must ESS be for an MCSE of 0.0025?
- MCSE $= \sqrt{0.955 \times 0.045/600} = \sqrt{0.0000716} = 0.0085$.
- Two MCSE ≈ 0.017: the estimate is plausibly anywhere in [0.938, 0.972], which straddles 0.95. Not decidable from these draws.
- For MCSE 0.0025: ESS $= 0.955 \times 0.045/0.0025^2 = 0.042975/0.00000625 \approx 6\,900$. Run longer (more draws or chains).
- And remember this is only the Monte Carlo error; the decision is still only as good as the model.
F. (Interview) "NUTS gives the exact posterior, so let's replace our SVI pipeline with NUTS and stop worrying." How do you respond?
"NUTS is asymptotically exact: with enough draws from a converged run its summaries approach the true posterior of our model, and I can measure the remaining Monte Carlo error with MCSE = sd/√ESS. But a finite run is never exact, it can fail silently (divergences in hierarchical funnels, chains stuck in different modes), so every run needs the checklist: zero divergences, R̂ ≤ 1.01, enough bulk and tail ESS, healthy traces. It is also much more expensive than SVI: every draw costs up to a thousand gradient evaluations over the full dataset (standard NUTS has no minibatch mode), and it does not fit naturally into our JIT-compiled training loop with early stopping. I would keep SVI for production fits and use NUTS on a reduced model or subset to validate it: compare means, intervals and decision probabilities, and if they disagree, improve the guide (for example a full-rank or low-rank guide instead of mean-field)."
Variational inference and KL divergence
MCMC walks around the posterior and collects samples. Variational inference does something different: it picks a simple, well-understood distribution and bends it until it looks as much like the posterior as it can. That turns inference into optimization, which is fast and scales to big models like your forecasting model. The price: the answer is only as good as the shape you allowed, and the way we measure "looks like" (the KL divergence) has a strong opinion of its own. This chapter teaches both, with pictures.
- Say what variational inference (VI) is: choose a family of simple distributions $q_\phi(\theta)$ and tune the variational parameters $\phi$ until $q_\phi$ is as close as possible to the posterior
- Define the KL divergence $KL(q\,\|\,p)$ intuitively and mathematically, compute it by hand, and prove it is never negative
- Know that $KL(q\,\|\,p) \ne KL(p\,\|\,q)$, and see why minimizing $KL(q\,\|\,p)$ is mode-seeking (it can miss whole modes) while $KL(p\,\|\,q)$ is mass-covering
- Explain why VI minimizes $KL(q\,\|\,p)$ even so: it is the direction we can compute, because the unknown evidence $p(D)$ only adds a constant
- Separate the approximation error (the best $q$ in the family is still wrong) from the optimization error (we did not find the best $q$)
- See why VI uncertainty tends to come out too narrow (under-dispersed), especially for heavy tails and for correlated parameters under a mean-field guide
What we need from earlier chapters: prior, likelihood, posterior and evidence (Chapter 6.1); why the posterior is usually impossible to compute exactly and the menu of approximations (Chapter 6.9); the Normal, Student-t and Gamma distributions (Chapter 4.9, Chapter 4.10); the expected value $E[\cdot]$ and the fact that $E[g(X)] \ne g(E[X])$ in general (Chapter 4.5); the multivariate Normal and its covariance ellipse (Chapter 5.15); gradient descent and Adam (Optimization guide). The short definition of KL from t-SNE (Chapter 5.17) is extended here. Notation: θ is the vector of unknown parameters; $p(\theta\mid D)$ is the posterior; $q_\phi(\theta)$ (say "q with parameters phi") is the approximation; $\phi$ is the list of numbers that pick one $q$ from the family; $\log$ is the natural logarithm, so KL is measured in "nats". $E_q[f(\theta)]$ means "the average of $f(\theta)$ when θ is drawn from $q$".
Variational inference: turn inference into optimization core
You need a suit for tomorrow. A tailor could make one that fits you perfectly, but that takes weeks. Instead you go to a shop with a rack of ready-made suits. Each suit has a few adjustable things: the size, the sleeve length. You try them on and adjust until the fit is as good as the rack allows. It will not be perfect, but it is good, and it is ready today.
The exact posterior is the tailor-made suit: for real models it is too hard to compute (Chapter 6.1: the evidence integral has no formula, and grids explode with many parameters). Variational inference is the shop. The rack is a family of simple distributions, for example "all Normal distributions". The adjustable knobs are the variational parameters $\phi$, for example a mean and a standard deviation. We turn the knobs until the simple distribution $q_\phi$ is as close as possible to the posterior. "As close as possible" needs a score; the score is the KL divergence, which the next section defines properly. For now: it is 0 when $q$ equals the posterior and gets bigger the worse the match.
The word "variational" comes from an old branch of calculus that optimizes over whole functions. Here it simply means: we search over distributions by searching over the knobs $\phi$.
Three ways to say it:
- Picture: slide and stretch a simple bell curve until it lies on top of the posterior as well as it can.
- Numbers: a posterior over 3 candidate rates is approximated by a "two-point" distribution with one knob $w$; the best knob is $w = 0.439$ and leaves a mismatch of 0.129.
- Slogan: MCMC samples the posterior; VI fits it.
A complete VI problem with one knob. Take the posterior of Chapter 6.1: three candidate conversion rates 5%, 10%, 15%, equal prior weights, and 3 buyers among 20 visitors. The exact posterior is $p = (0.121,\ 0.386,\ 0.493)$.
- Family. Suppose our approximation is only allowed to put weight on 10% and 15%: $q_w = (0,\ w,\ 1-w)$. This is the "rack".
- Variational parameter. The single number $w$ between 0 and 1. This is the "knob": $\phi = w$.
- Score. $KL(q_w\|p) = \sum_i q_i \log(q_i/p_i)$, adding only over the candidates where $q_i \gt 0$. At $w = 0.5$: $0.5\log(0.5/0.386) + 0.5\log(0.5/0.493) = 0.5(0.2588) + 0.5(0.0141) = 0.1294 + 0.0071 = 0.136$.
- Optimize. Try other knob settings: $w = 0.3$ gives $0.170$, $w = 0.4$ gives $0.132$, $w = 0.6$ gives $0.181$. The minimum is at $w = p_2/(p_2 + p_3) = 0.386/0.879 = 0.439$: the posterior's own weights, rescaled onto the two allowed candidates.
- Leftover mismatch. At the best knob, $KL = 0.439\log(0.439/0.386) + 0.561\log(0.561/0.493) = -\log(0.879) = 0.129$. It is not 0, because no setting of $w$ can put weight on 5%. This leftover is the approximation error of the family.
Notice what the best $q$ did with the 5% candidate (posterior weight 0.121): it simply ignored it. Remember this: it is the first sign of the "mode-seeking" behaviour you will meet below.
Variational inference approximates the posterior $p(\theta\mid D)$ by the member of a chosen family that is closest to it:
$$\phi^\star = \arg\min_{\phi}\; KL\big(q_\phi(\theta)\,\big\|\,p(\theta\mid D)\big), \qquad \text{then use } q_{\phi^\star}(\theta) \text{ in place of } p(\theta\mid D).$$- Variational family $\mathcal{Q} = \{q_\phi\}$: the set of distributions we allow, for example all Normals $N(\mu, \sigma^2)$, or all products of independent Normals ("mean-field"), or Normals with a full covariance matrix (Chapter 6.13).
- Variational parameters $\phi$: the numbers that pick one member, for example $\phi = (\mu, \sigma)$. They are not the model's parameters θ. θ is what the model is about (a rate, a trend slope); $\phi$ describes our belief about θ (where we think θ is, and how unsure we are).
- Objective: a number that measures the mismatch, here $KL(q_\phi\|p)$. In practice we maximize the equivalent ELBO instead (Chapter 6.12), because it can be computed without knowing $p(D)$.
- Optimization: gradient-based (Adam, in NumPyro's SVI) or coordinate updates. The result is a whole distribution, so you get means, intervals and predictive draws from it, just as from MCMC samples.
- Two separate errors: the approximation error (the best member of the family is still not the posterior) and the optimization error (the optimizer stopped before reaching the best member, or found a worse local optimum).
Why do we need it?
For models with hundreds of parameters, exact posteriors are impossible and MCMC can be slow. Optimization is fast, uses gradients, works with minibatches and scales to large models and datasets. VI gives a full approximate posterior in seconds instead of minutes or hours.
Where is it used?
NumPyro and Pyro SVI (your forecasting model and the A/B framework), Stan's ADVI, variational autoencoders (VAEs), Bayesian neural networks, topic models (LDA), large hierarchical and time-series models where NUTS would be too slow.
How is it used?
Write the model; choose a family (in NumPyro, a "guide" such as AutoNormal); let the optimizer tune the guide's parameters; then draw samples from the fitted guide and use them like posterior samples (means, intervals, predictions). Always remember the answer is an approximation.
"VI gives the posterior, just faster."
VI gives the closest member of the family you chose. If the family cannot express the posterior's shape (skew, heavy tails, several modes, correlations), the answer is wrong in exactly those ways, however long you optimize.
"The variational parameters φ are the model's parameters."
θ (a conversion rate, a trend slope) is what the model is about. φ (a mean and a scale for θ) describes our uncertain belief about θ. A guide with 100 latent θ's and a mean-field Normal family has 200 numbers in φ.
"If the optimizer converged, the approximation is good."
Convergence only says we found the bottom of this family's valley (or a local dip). The leftover KL, the approximation error, can still be large. Check the fit (Chapter 6.15 compares SVI with NUTS on the same model).
Both of your projects use SVI, which is variational inference with stochastic gradients (Chapter 6.12). In NumPyro the guide is the family $q_\phi$: AutoNormal is "independent Normals, one per latent variable"; your forecasting model chooses between a full-rank and a low-rank multivariate Normal guide depending on model size (Chapter 6.13). The trend, changepoint, seasonality, holiday, regressor and noise parameters are θ; the guide's locations, scales and covariance factors are φ. Your training loop is the "turn the knobs" step.
VI: pick a family $q_\phi$, then $\phi^\star = \arg\min_\phi KL(q_\phi\,\|\,p(\theta\mid D))$; use $q_{\phi^\star}$ as the posterior.
φ = knobs of the approximation (e.g. μ, σ), not the model parameters θ.
Trap: best in the family ≠ the posterior. Leftover KL = approximation error.
Quick check: a model has 50 latent parameters and you use a family of independent Normals. How many variational parameters are there?
Each latent parameter gets its own mean and its own scale, so $\phi$ has $2 \times 50 = 100$ numbers. (NumPyro's AutoNormal stores them as ..._auto_loc and ..._auto_scale for each latent site; Chapter 6.13 counts the other guides.)
KL divergence: how badly does $q$ describe $p$? core
Two forecasters describe the same thing: which plan new customers pick (Basic, Pro or Team). The true shares are $p$. A second forecaster believes $q$. We want one number that says how wrong $q$ is.
Go category by category. For each one, compare the two probabilities with a ratio $q/p$. A ratio of 1 means "they agree". A ratio of 4 means "$q$ thinks this is 4 times more likely than it really is". Take the log of each ratio (so "4 times too high" and "4 times too low" become $+1.39$ and $-1.39$), then take the average, weighted by q: categories that $q$ believes in count more. That weighted average is the Kullback–Leibler (KL) divergence $KL(q\|p)$. Some single terms can be negative, but the total is never negative, and it is 0 only when $q$ and $p$ agree everywhere.
The order of the letters matters. $KL(q\|p)$ averages over $q$'s beliefs; $KL(p\|q)$ averages over $p$'s. They usually give different numbers, and that difference is the most important idea in this chapter.
Three ways to say it:
- Picture: wherever $q$ puts weight, measure how much taller $q$'s bar is than $p$'s (on a log scale), then average.
- Numbers: $q = (0.3, 0.3, 0.4)$ against $p = (0.6, 0.3, 0.1)$: $KL(q\|p) = 0.347$ but $KL(p\|q) = 0.277$.
- Slogan: KL is the average log-ratio, weighted by the distribution written first; zero only for a perfect match.
Plan choice. True shares $p = (0.6, 0.3, 0.1)$ for Basic, Pro, Team. A forecaster's shares $q = (0.3, 0.3, 0.4)$.
- Ratios $q_i/p_i$: $0.3/0.6 = 0.5$, $\;0.3/0.3 = 1$, $\;0.4/0.1 = 4$.
- Logs: $\log 0.5 = -0.693$, $\;\log 1 = 0$, $\;\log 4 = 1.386$.
- Weight each log by $q_i$: $0.3 \times (-0.693) = -0.208$, $\;0.3 \times 0 = 0$, $\;0.4 \times 1.386 = 0.555$.
- Add: $KL(q\|p) = -0.208 + 0 + 0.555 = 0.347$ nats. The big cost comes from Team, where $q$ is 4 times too high and puts a lot of weight.
- The other direction: ratios $p_i/q_i = 2, 1, 0.25$; logs $0.693, 0, -1.386$; weights $p_i$: $0.6 \times 0.693 = 0.416$, $\;0$, $\;0.1 \times (-1.386) = -0.139$. Add: $KL(p\|q) = 0.277$.
- Two different numbers for the same pair: $0.347 \ne 0.277$. KL is not symmetric.
Two Normals. For $q = N(0, 1)$ and $p = N(0, 2^2)$ the formula below gives $KL(q\|p) = \log 2 + \tfrac{1}{8} - \tfrac12 = 0.693 + 0.125 - 0.5 = 0.318$, but $KL(p\|q) = \log\tfrac12 + \tfrac{4}{2} - \tfrac12 = -0.693 + 2 - 0.5 = 0.807$. A narrow $q$ inside a wide $p$ is cheap in $KL(q\|p)$ and expensive in $KL(p\|q)$. Keep this in mind: it is why VI tends to give intervals that are too narrow.
For two distributions $q$ and $p$ over the same θ, the KL divergence of $q$ from $p$ is
$$KL(q\,\|\,p) = \sum_\theta q(\theta)\log\frac{q(\theta)}{p(\theta)} \quad\text{(discrete)}, \qquad KL(q\,\|\,p) = \int q(\theta)\log\frac{q(\theta)}{p(\theta)}\,d\theta = E_q\big[\log q(\theta) - \log p(\theta)\big] \quad\text{(continuous)}.$$- Never negative: $KL(q\|p) \ge 0$ ("Gibbs' inequality"). Proof in one line: the log is a curve that bends down, so the average of logs is at most the log of the average (Jensen's inequality): $-KL(q\|p) = E_q\big[\log\tfrac{p}{q}\big] \le \log E_q\big[\tfrac{p}{q}\big] = \log\sum_{q(\theta)\gt 0} p(\theta) \le \log 1 = 0$.
- Zero only for a perfect match: $KL(q\|p) = 0$ exactly when $q = p$ (everywhere that matters).
- Not symmetric: $KL(q\|p) \ne KL(p\|q)$ in general, and it does not obey the triangle rule. That is why it is called a divergence, not a distance.
- Can be infinite: if $q(\theta) \gt 0$ somewhere that $p(\theta) = 0$, then $\log(q/p) = +\infty$ there and $KL(q\|p) = \infty$. (The reverse case, $p \gt 0$ where $q = 0$, makes $KL(p\|q)$ infinite.)
- Unchanged by relabelling θ: if you change variables one-to-one (for example $\eta = \log\lambda$), both $q$ and $p$ change by the same Jacobian factor, so the ratio $q/p$ and the KL stay the same.
- Two Normals (a formula worth knowing): $KL\big(N(\mu_1,\sigma_1^2)\,\|\,N(\mu_2,\sigma_2^2)\big) = \log\dfrac{\sigma_2}{\sigma_1} + \dfrac{\sigma_1^2 + (\mu_1-\mu_2)^2}{2\sigma_2^2} - \dfrac12.$
- Units: with the natural log, KL is in nats; with $\log_2$ it is in bits ($1$ nat $= 1.443$ bits).
Why do we need it?
To fit one distribution to another we need a single number for "how different are they?" that is 0 for a perfect match and grows with the mismatch. KL is that number, it has a clean formula, and its gradient is easy to compute, so optimizers can use it.
Where is it used?
Variational inference and the ELBO (this chapter and the next), the KL term of a VAE's loss, t-SNE (Chapter 5.17), cross-entropy loss (cross-entropy = entropy + KL), knowledge distillation, policy updates in reinforcement learning (PPO, TRPO), and monitoring data drift between two histograms.
How is it used?
For discrete distributions: scipy.stats.entropy(q, p) returns $KL(q\|p)$ (note the order). For two Normals use the formula. For anything else, estimate $E_q[\log q - \log p]$ by averaging over draws from $q$. Always write down which distribution comes first.
"KL divergence is the distance between two distributions."
It is not a distance: $KL(q\|p) \ne KL(p\|q)$, and it does not obey the triangle rule. Read it as "how badly does the first describe the second, judged from the first one's point of view".
"Some terms are negative, so KL can be negative."
Single terms $q_i\log(q_i/p_i)$ can be negative, but the total is always $\ge 0$ when both $q$ and $p$ add up to 1. If you get a negative KL, one of them is not normalized, or a sign is flipped.
"scipy.stats.entropy(p, q) computes $KL(q\|p)$."
It computes $KL(p\|q) = \sum p\log(p/q)$, with the first argument as the weights. Swap the arguments to get the other direction.
$KL(q\|p) = E_q[\log q - \log p] = \sum q\log(q/p) \ge 0$; $= 0$ only if $q = p$.
Two Normals: $\log\frac{\sigma_2}{\sigma_1} + \frac{\sigma_1^2 + (\mu_1-\mu_2)^2}{2\sigma_2^2} - \frac12$.
Trap: not symmetric; infinite when the first puts mass where the second has none.
Quick check: $q = (0.5, 0.5)$ and $p = (0.9, 0.1)$. Compute $KL(q\|p)$.
$0.5\log(0.5/0.9) + 0.5\log(0.5/0.1) = 0.5(-0.588) + 0.5(1.609) = -0.294 + 0.805 = 0.511$ nats. (The other direction: $0.9\log 1.8 + 0.1\log 0.2 = 0.529 - 0.161 = 0.368$.)
The family and the approximation error: the best $q$ can still be wrong
Back to the suit shop. If every suit on the rack is a straight cut, then no amount of adjusting makes a straight cut fit a very curved body. The leftover misfit is not the shop assistant's fault; it is the rack's. In VI this leftover is the approximation error: the KL that remains at the best member of the family.
Two things decide how big it is. First, the shape of the family: a Normal is symmetric with light tails and one peak, so it cannot copy skew, heavy tails or two peaks. Second, where the family lives: a Normal stretches over all real numbers, so it cannot be the right shape for a rate λ that must be positive. NumPyro solves the second problem with a trick: it fits the Normal to $\log\lambda$, which can be any real number, and then maps back with $\lambda = e^{\eta}$.
Good news: with more data, many posteriors become more and more bell-shaped, and the approximation error of a Normal family shrinks.
Three ways to say it:
- Picture: the family is a region; the posterior is a point; the approximation error is the gap from the closest point of the region.
- Numbers: for a log-rate posterior after one day with 3 orders, the best Normal still has $KL = 0.021$; with 30 days of data (about 60 orders) it is $0.0014$.
- Slogan: optimization can find the best suit on the rack, but it cannot change the rack.
The best Normal for a log-rate posterior (the posterior of the widget above). Orders per day follow a Poisson with rate λ; the prior is Gamma(1, 1); one day shows 3 orders. The posterior is $\lambda\mid D \sim$ Gamma$(a, b)$ with $a = 1 + 3 = 4$ and $b = 1 + 1 = 2$ (shape and rate, Chapter 6.3). We fit $q = N(\mu, \sigma^2)$ to $\eta = \log\lambda$.
- On the log scale the posterior is $\log p(\eta) = a\log b - \log\Gamma(a) + a\eta - b e^{\eta}$ (the $a\eta$ includes the Jacobian of $\lambda = e^\eta$).
- Under $q$, $E_q[\eta] = \mu$ and $E_q[e^\eta] = e^{\mu + \sigma^2/2}$ (the mean of a log-normal, Chapter 4.10). The entropy of a Normal is $-E_q[\log q] = \tfrac12\log(2\pi e\sigma^2)$.
- So $KL(q\|p) = -\tfrac12\log(2\pi e\sigma^2) - a\log b + \log\Gamma(a) - a\mu + b\,e^{\mu + \sigma^2/2}$.
- Set the slope in μ to zero: $-a + b e^{\mu+\sigma^2/2} = 0$, so $e^{\mu + \sigma^2/2} = a/b = 2$.
- Set the slope in σ to zero: $-1/\sigma + b\sigma e^{\mu+\sigma^2/2} = -1/\sigma + a\sigma = 0$, so $\sigma^2 = 1/a = 0.25$ and $\sigma^\star = 0.5$.
- Then $\mu^\star = \log 2 - \sigma^{\star 2}/2 = 0.693 - 0.125 = 0.568$.
- Compare with the truth: the posterior of η has mean $\psi(a) - \log b = 0.563$ and standard deviation $\sqrt{\psi'(a)} = 0.533$ (ψ and ψ′ are the "digamma" and "trigamma" functions; any library computes them, e.g.
scipy.special.digamma). The best Normal is close but a bit too narrow (0.500 vs 0.533). - The leftover score is $KL(q^\star\|p) = 0.0208$ (for this target it equals $\log\Gamma(a) - (a-\tfrac12)\log a + a - \tfrac12\log 2\pi \approx \tfrac{1}{12a} = \tfrac{1}{48}$). No Normal does better: that is the approximation error.
And a Normal on λ itself? The posterior Gamma(4, 2) has mean 2 and sd 1. A Normal with that mean and sd puts probability $\Phi(-2) = 0.023$ on negative rates, where $p(\lambda) = 0$. Then $\log(q/p) = +\infty$ there, and $KL(q\|p) = \infty$. Every Normal on λ has some mass below 0, so every one has infinite KL.
For a family $\mathcal{Q} = \{q_\phi\}$ and posterior $p = p(\theta\mid D)$:
$$\text{approximation error} = \min_{\phi} KL(q_\phi\,\|\,p) \;\ge 0, \qquad \text{zero only if the posterior itself is in } \mathcal{Q}.$$- What you actually get is $KL(q_{\hat\phi}\|p)$ for the $\hat\phi$ the optimizer returned: approximation error plus optimization error (stopped early, noisy gradients, a local optimum).
- Support must match. $q$ must put no mass where $p(\theta) = 0$, or $KL(q\|p) = \infty$. NumPyro's automatic guides therefore work in an unconstrained space: positive parameters through $\log$, probabilities in (0, 1) through the logit, then they map back. A Normal on $\log\lambda$ is a log-normal on λ. KL is the same on both scales.
- Ways to shrink the approximation error: a richer family (full-rank or low-rank covariance, Chapter 6.13; mixtures; normalizing flows), a better parameterization of the model (Chapter 6.7), or more data (under regularity conditions, posteriors become close to Normal as data grow).
- Richer families cost more variational parameters and more computation per step. That trade-off is the whole of Chapter 6.13.
Why do we need it?
To judge a VI result honestly you must separate "the optimizer has not finished" from "the family cannot do better". The first is fixed by training longer; the second only by changing the family or the model's parameterization.
Where is it used?
Choosing a NumPyro guide (AutoNormal, AutoMultivariateNormal, AutoLowRankMultivariateNormal, AutoIAFNormal), the full-rank vs low-rank choice in your forecasting loop, and the constraint transforms that NumPyro applies to every bounded parameter.
How is it used?
Fit the simple guide first; then fit a richer one. If the richer guide reaches a clearly higher ELBO or noticeably wider intervals, the simple family was costing you accuracy. For a few key parameters, compare with a NUTS run (Chapter 6.15).
"If the ELBO stopped improving, the approximation error is small."
A flat ELBO says the optimizer reached the bottom of the family's valley. The height of that bottom (the approximation error) can still be large; only a richer family or a comparison with NUTS shows it.
"With enough data, a Normal guide is always fine."
Many posteriors become close to Normal with lots of data, but not all: several modes, weakly identified parameters (trend vs changepoints, Chapter 6.8), hierarchical funnels with few observations per group, and parameters near a boundary stay non-Normal.
"A Normal guide on a positive parameter can give negative draws."
Not in NumPyro's automatic guides: they fit the Normal in the unconstrained space (for example $\log\sigma$) and transform back, so draws of σ are always positive. If you write your own guide with dist.Normal for a positive site, you do get this problem.
In your forecasting model, the noise scale σ (and the Negative Binomial concentration, if that likelihood is used) must be positive. An automatic guide fits a Normal to their logs, so what you get back is a log-normal, which is right-skewed on the original scale. In an A/B framework like yours, a conversion rate θ in (0, 1) is handled on the logit scale. When you report a posterior mean of σ from the guide, remember it is the mean of a log-normal, not $e^{\text{loc}}$ (that is the median).
Approximation error $= \min_\phi KL(q_\phi\|p) \ge 0$; 0 only if the posterior is in the family.
Total error = approximation error + optimization error.
Trap: $q$ must live where $p$ lives; Normal on λ > 0 gives KL = ∞. AutoNormal fits on the unconstrained (log/logit) scale.
Quick check: after 10 days in the widget (a = 21), what are σ* and the approximation error?
$\sigma^\star = 1/\sqrt{21} = 0.218$. The approximation error is $\approx 1/(12 \times 21) = 1/252 = 0.0040$ (exact value 0.00397). The true posterior sd of $\log\lambda$ is $\sqrt{\psi'(21)} = 0.221$ (ψ′ is the "trigamma" function, the variance of a log-Gamma variable), so the best Normal is only about 1% too narrow.
Reverse vs forward KL: mode-seeking vs mass-covering core
Picture two mountain villages with a deep, empty valley between them. People (probability) live in the villages, almost nobody in the valley. You are allowed to draw one simple oval on the map to describe where people live.
- Rule 1, $KL(q\|p)$, the reverse KL: "you are punished heavily for drawing your oval where nobody lives". The average is taken over your oval ($q$), and wherever $p$ is tiny, $\log(q/p)$ is huge. Covering both villages would put the middle of your oval in the empty valley. So you choose one village and draw a tight oval around it. That is mode-seeking (also called zero-forcing: $q$ is forced to be near zero wherever $p$ is near zero).
- Rule 2, $KL(p\|q)$, the forward KL: "you are punished heavily for leaving people outside your oval". The average is taken over where people live ($p$), and wherever your oval is thin, $\log(p/q)$ is huge. So you draw one big oval around both villages, valley and all. That is mass-covering (zero-avoiding: $q$ avoids being zero wherever $p$ has mass).
Variational inference uses Rule 1. That is why VI can miss whole modes and tends to report too little uncertainty.
Three ways to say it:
- Picture: reverse KL hugs one hill tightly; forward KL throws a blanket over all hills.
- Numbers: for two modes at −3 and +3, the reverse-KL Normal is N(2.98, 1.02²) (one mode); the forward-KL Normal is N(0, 3.16²) (both modes and the valley).
- Slogan: $KL(q\|p)$: "never be where p is not". $KL(p\|q)$: "never miss where p is".
Two modes. Target $p = \tfrac12 N(-3, 1) + \tfrac12 N(3, 1)$: two equal hills at −3 and +3. Family: one Normal $N(\mu, \sigma^2)$.
- Forward KL, best fit. It matches the target's mean and variance (shown in the definition below). Mean: $\tfrac12(-3) + \tfrac12(3) = 0$. Variance: the average spread within a hill (1) plus the spread of the hill centres ($3^2 = 9$), so $1 + 9 = 10$ and $\sigma = \sqrt{10} = 3.16$. The fit is $N(0, 3.16^2)$ with $KL(p\|q) = 0.462$. Its peak sits in the empty valley.
- Reverse KL, best fit. Found numerically: $N(2.98, 1.02^2)$, one hill (or the mirror image at −2.98). $KL(q\|p) = 0.689$.
- Why about 0.69? Where $q$ lives (near +3), the target is half of a $N(3, 1)$ hill: $p(\theta) \approx \tfrac12 q(\theta)$. So $\log(q/p) \approx \log 2 = 0.693$ everywhere $q$ has mass, and the average is about 0.693. The reverse KL calmly pays "half the mass is missing" and nothing more.
- Score each fit with the other rule. The blanket $N(0, 3.16^2)$ under the reverse KL: $0.938$ (worse than 0.689, because it puts a lot of mass in the valley). The one-hill fit under the forward KL: $7.86$ (terrible, because a whole hill of people is left outside).
- A trap for optimizers. The reverse KL also has a "compromise" resting point $N(0, 2.74^2)$ with $KL = 0.841$. An optimizer that starts exactly in the middle can stop there. Which answer VI returns depends on where it starts.
For an approximation $q$ and a target $p$:
$$\underbrace{KL(q\,\|\,p) = E_q\Big[\log\frac{q(\theta)}{p(\theta)}\Big]}_{\text{reverse, "exclusive": used by VI}} \qquad\qquad \underbrace{KL(p\,\|\,q) = E_p\Big[\log\frac{p(\theta)}{q(\theta)}\Big]}_{\text{forward, "inclusive"}}$$- Reverse averages over $q$. If $q \gt 0$ where $p \approx 0$, the cost is huge, so $q$ stays inside the regions where $p$ is large. Behaviour: mode-seeking / zero-forcing; with one simple $q$ and several modes, it usually locks onto one mode; its spread tends to be too small (under-dispersed).
- Forward averages over $p$. If $q \approx 0$ where $p \gt 0$, the cost is huge, so $q$ stretches over all of $p$'s mass. Behaviour: mass-covering / zero-avoiding; it puts mass between modes; its spread tends to be too large there.
- Forward KL with a Normal $q$ = moment matching. $KL(p\|q) = E_p[\log p] - E_p[\log q]$, and the first part does not depend on $q$. For $q = N(\mu,\sigma^2)$: $E_p[\log q] = -\tfrac12\log(2\pi\sigma^2) - \dfrac{Var_p(\theta) + (E_p[\theta]-\mu)^2}{2\sigma^2}$, which is largest at $\mu = E_p[\theta]$ and $\sigma^2 = Var_p(\theta)$.
- The reverse KL has no such shortcut; it can have several local minima (one per mode, sometimes a compromise as well). The forward KL with a Normal family has a single best answer.
- "Forward" and "reverse" are just names from the VI literature: "forward" is $KL(p\|q)$, the true distribution first.
Why do we need it?
Knowing which direction you minimize tells you which mistakes to expect. VI's reverse KL gives answers that are confident and possibly blind to other modes; a forward-KL method gives answers that are cautious and possibly smeared across modes.
Where is it used?
Reverse KL: NumPyro/Pyro SVI, ADVI in Stan, VAEs. Forward KL: maximum likelihood (fitting a model to data minimizes KL(data ‖ model)), expectation propagation (EP, locally), cross-entropy training, importance-sampling-based methods such as reweighted wake-sleep.
How is it used?
When you fit with SVI, assume the reverse KL's habits: run from several random starts and compare the answers (different modes show up as different results), and compare key intervals with a NUTS run on a smaller version of the model.
"Mode-seeking means VI finds the highest mode."
It finds a mode: usually the one nearest to where the optimizer started, and sometimes a wide compromise. Run SVI from several random starts; if the answers disagree, the posterior has several modes and a single Gaussian guide is not enough.
"Forward KL is the better objective, so VI is doing it wrong."
Forward KL needs averages over the posterior $p$, which is exactly what we cannot compute (next section). And its answer is not always better: with two modes it puts its peak in an empty valley.
"If the reverse KL is small, nothing important is missing."
Missing a whole mode with half of the mass costs only about $\log 2 = 0.69$ nats. The reverse KL barely notices what $q$ never visits.
"KL divergence is symmetric, so it does not matter which way round we write it."
"$KL(q\|p) \ne KL(p\|q)$. VI minimizes the reverse, $KL(q\|p)$, which averages over $q$: it punishes $q$ for putting mass where the posterior has little, so it is mode-seeking and tends to under-state uncertainty. The forward $KL(p\|q)$ averages over $p$ and is mass-covering."
Model answer: "In SVI we minimize $KL(q\|p)$, because we can average over our own $q$ but not over the unknown posterior. The consequence is zero-forcing behaviour: the fitted $q$ sits inside a high-density region, can miss other modes, and its intervals are usually too narrow. For multimodal or strongly correlated posteriors I check against NUTS or use a richer guide."
In your forecasting model, two explanations can fit the same history almost equally well: a changepoint in the trend or a stronger seasonal swing (identifiability, Chapter 6.8). If the posterior really has two such modes, SVI with a single Gaussian guide will usually report one of them, confidently. Running the training loop from a few different seeds and comparing the fitted components is a cheap check. In an A/B framework like yours, a hierarchical posterior can also be lopsided (the funnel of Chapter 6.7); mode-seeking there shows up as a guide that collapses the group effects too much.
Reverse $KL(q\|p)$ (VI): average over $q$; mode-seeking, zero-forcing, under-dispersed, can miss modes; several local optima.
Forward $KL(p\|q)$: average over $p$; mass-covering; for a Normal $q$ it matches the mean and variance.
Trap: $KL(q\|p) \ne KL(p\|q)$; missing half the mass costs only $\log 2 \approx 0.69$.
Quick check: target $\tfrac12 N(-4, 1) + \tfrac12 N(4, 1)$. Give the forward-KL best Normal, and guess the reverse-KL best one.
Forward: moment matching, mean 0, variance $1 + 4^2 = 17$, so $N(0, 4.12^2)$. Reverse: one hill, about $N(\pm 4, 1)$, with $KL(q\|p) \approx \log 2 = 0.693$ (the modes are so far apart that the other hill adds almost nothing where $q$ lives).
Why VI minimizes $KL(q\|p)$: the unknown evidence drops out
If the reverse KL has such bad habits, why does every SVI library use it? Because it is the only direction we can actually compute.
Remember the posterior: $p(\theta\mid D) = p(\theta, D)/p(D)$. The top, prior × likelihood $= p(\theta, D)$, is easy: plug in any θ. The bottom, the evidence $p(D)$, is the impossible integral. Now look at what each KL needs.
- $KL(q\|p)$ is an average over $q$, a distribution we chose ourselves and can draw from. Inside the average, $\log p(\theta\mid D) = \log p(\theta, D) - \log p(D)$. The $\log p(D)$ part is the same number for every θ, so it comes out of the average as a constant. A constant shifts the score up or down but does not move the location of the minimum. So we can find the best $q$ without ever knowing $p(D)$.
- $KL(p\|q)$ is an average over the posterior. To compute it we would need draws from the posterior, which is the very thing we were trying to get.
Three ways to say it:
- Picture: the score we can compute and the true KL are the same curve, just shifted up; both bottom out at the same knob setting.
- Numbers: in the three-candidate example, at $w = 0.5$ the computable score is 1.943 and the true KL is 0.136; they differ by exactly $-\log p(D) = 1.807$, at every $w$.
- Slogan: we cannot compute the KL, but we can compute "KL minus a constant", and that is enough to find the minimum.
The three candidates again (prior weights $\tfrac13$, 3 buyers of 20). Prior × likelihood: $\tilde p = (0.0199,\ 0.0634,\ 0.0809)$; evidence $p(D) = 0.1642$, $\log p(D) = -1.807$. Family $q_w = (0, w, 1-w)$.
- The computable score uses only $\tilde p$: $J(w) = E_q[\log q - \log\tilde p] = \sum_i q_i(\log q_i - \log\tilde p_i)$.
- At $w = 0.5$: $0.5(\log 0.5 - \log 0.0634) + 0.5(\log 0.5 - \log 0.0809) = 0.5(-0.693 + 2.758) + 0.5(-0.693 + 2.514) = 1.033 + 0.911 = 1.943$.
- The true KL at $w = 0.5$ (from the first section): $0.136$. Difference: $1.943 - 0.136 = 1.807 = -\log p(D)$.
- At $w = 0.439$: $J = 1.936$, $KL = 0.129$, difference again $1.807$.
- So $J(w) = KL(q_w\|p) - \log p(D)$ for every $w$. Both are smallest at $w = 0.439$. We found the best $q$ without dividing by $p(D)$.
- The forward direction would need $KL(p\|q_w) = \sum_i p_i\log(p_i/q_i)$ with the normalized posterior $p_i$. Here it is even infinite: $p_1 = 0.121 \gt 0$ but $q_1 = 0$.
Start from the definition and split the log of the posterior:
$$\begin{aligned} KL\big(q\,\|\,p(\theta\mid D)\big) &= E_q\big[\log q(\theta) - \log p(\theta\mid D)\big] \\ &= E_q\big[\log q(\theta) - \log p(\theta, D) + \log p(D)\big] \\ &= \underbrace{E_q\big[\log q(\theta) - \log p(\theta, D)\big]}_{\text{computable: } -\text{ELBO}} + \underbrace{\log p(D)}_{\text{constant in } q}. \end{aligned}$$- $p(\theta, D) = p(D\mid\theta)\,p(\theta)$ is the joint (likelihood × prior): one evaluation per θ, no integral.
- The negative of the first term, $\text{ELBO}(q) = E_q[\log p(\theta, D) - \log q(\theta)]$, is the evidence lower bound. Rearranged: $\log p(D) = \text{ELBO}(q) + KL(q\|p(\theta\mid D))$. Minimizing the KL is the same as maximizing the ELBO. Chapter 6.12 derives this carefully and shows how to estimate the ELBO with random draws.
- The averages are over $q$, so they can be estimated by drawing θ from $q$, which is easy because we picked a simple $q$.
Why do we need it?
It explains the whole design of VI: we accept the reverse KL's mode-seeking habits because it is the only direction that needs neither $p(D)$ nor posterior samples, just the joint $p(\theta, D)$ and draws from our own $q$.
Where is it used?
Every ELBO-based method: NumPyro's Trace_ELBO, Pyro, Stan's ADVI, the VAE loss, Bayesian neural networks. The same "the constant drops out" idea makes MCMC work too: Metropolis only uses ratios of $p(\theta, D)$ (Chapter 6.9).
How is it used?
You never compute $p(D)$. You write the model (which gives $\log p(\theta, D)$), choose a guide $q$, and let SVI maximize the ELBO. The number it reports is $-$ELBO, which equals $KL - \log p(D)$, so only its changes and its minimum are meaningful.
"The loss that SVI prints is the KL divergence to the posterior."
It is $-\text{ELBO} = KL(q\|p) - \log p(D)$: the KL plus an unknown constant. It can be any size, even negative. Only its changes during training and its minimum carry meaning.
"Just estimate $p(D)$ first, then compute the real KL."
Estimating $p(D)$ is an integral over all θ, the same hard problem as the posterior itself. VI is designed to never need it. (The ELBO is a lower bound on $\log p(D)$, which is the next chapter's story.)
$KL(q\|p(\theta\mid D)) = E_q[\log q - \log p(\theta, D)] + \log p(D) = -\text{ELBO} + \log p(D)$.
$\log p(D)$ is constant in $q$, so minimizing KL ⇔ maximizing the ELBO. Needs only draws from $q$ and the joint.
Trap: the forward KL needs posterior draws, so VI cannot use it directly.
Quick check: for two different guides on the same model and data, SVI reports final losses 1 520.3 and 1 498.7. Which guide is closer to the posterior in KL, and by how much?
Loss $= -\text{ELBO} = KL - \log p(D)$, and $\log p(D)$ is the same for both (same model, same data). So the difference of losses is the difference of KLs: the second guide is closer by about $1520.3 - 1498.7 = 21.6$ nats (up to the Monte Carlo noise in each loss estimate). The losses themselves do not tell you either KL.
Under-dispersion: why VI intervals often come out too narrow
"Dispersion" just means spread. An approximation is under-dispersed if it is narrower than the real posterior: its intervals are too short, and it sounds more certain than the data allow.
The reverse KL causes this whenever $q$ cannot copy the posterior's shape. Its rule is "never be where $p$ is not". When the shapes cannot match, $q$ must decide between two mistakes: leaving some of $p$'s mass uncovered (cheap under the reverse KL) or spreading into places where $p$ is thin (expensive). It picks the cheap mistake, so it ends up too narrow.
Heavy tails are the clearest case. A Student-t posterior has a narrow middle and long tails. A Normal cannot have both. The reverse KL fits the middle and gives up on the tails.
Three ways to say it:
- Picture: the orange bell sits snugly on the green peak, while the green tails stick out on both sides.
- Numbers: target Student-t with ν = 3 (sd 1.73): the reverse-KL Normal has σ = 1.26; its "95%" interval holds only 91% of the real probability.
- Slogan: when in doubt, VI shrinks.
Heavy-tailed posterior. Target: Student-t with ν = 3 degrees of freedom, centre 0, scale 1. Its sd is $\sqrt{\nu/(\nu-2)} = \sqrt3 = 1.732$, and its central 95% interval is $\pm 3.18$.
- Forward-KL Normal (moment matching): $N(0, 1.732^2)$. Its 95% interval is $\pm 1.96 \times 1.732 = \pm 3.39$.
- Reverse-KL Normal (found numerically; by symmetry μ = 0): σ = 1.260, with $KL(q\|p) = 0.041$.
- Its 95% interval: $\pm 1.96 \times 1.260 = \pm 2.47$.
- How much of the real posterior is inside $\pm 2.47$? $P(|T_3| \le 2.47) = 0.910$. The interval called "95%" really covers 91%.
- For comparison, the moment-matched Normal scores $KL(q\|p) = 0.100$ under the reverse KL, worse than 0.041. The reverse KL actively prefers the narrower Normal.
- The same happened for the log-rate posterior earlier: best σ = 0.500 against a true sd of 0.533.
An approximation $q$ is under-dispersed if its spread is smaller than the posterior's: $Var_q(\theta_i) \lt Var_p(\theta_i)$, or its intervals are shorter and cover less than their nominal level under the true posterior. Over-dispersed is the opposite.
- Minimizing $KL(q\|p)$ with a family that cannot match the posterior's shape tends to under-disperse: most strongly for correlated parameters under a mean-field family (next section), for heavy tails, and for several modes (one mode is dropped).
- It is a strong tendency, not a theorem for every possible posterior, and the means can be off as well, not only the spreads. The only real proof is a check: compare with NUTS on the same model, or simulate data from known parameters and see how often the intervals cover the truth.
- Minimizing $KL(p\|q)$ (forward) tends to over-disperse instead.
Why do we need it?
Uncertainty is the whole point of a Bayesian model. If the guide is too narrow, every downstream number (intervals, probabilities like P(B > A), forecast bands) is over-confident, and decisions made from them take more risk than they think.
Where is it used?
Interval coverage checks for forecasts (Chapter 7.16), SVI vs NUTS comparisons (Chapter 6.15), simulation-based calibration of Bayesian software, and the decision to use richer guides in your forecasting model (Chapter 6.13).
How is it used?
Treat VI intervals as a lower bound on the real uncertainty unless checked. For the parameters that drive decisions, compare the guide's sd with a NUTS run on a smaller dataset; check the coverage of forecast intervals on a holdout.
"VI is over-confident, so multiply its standard deviations by some factor."
The shrinkage depends on the model, the parameter and the guide: tiny for a nearly Normal posterior, huge along strong correlations. There is no universal correction factor; compare with NUTS on the parameters that matter.
"Under-dispersion only affects the tails, so the intervals are fine."
The 95% interval is decided by the tails. Under-dispersion shows up exactly in the numbers you report: interval widths, tail probabilities, $P(\theta_B \gt \theta_A)$.
"Variational inference always underestimates the posterior variance."
"With the reverse KL and a family that cannot match the posterior, VI usually under-estimates the spread, most of all along correlations (mean-field), in heavy tails and when it drops a mode. It is a strong tendency, not a guarantee in every model, and it can also bias the means."
Model answer: "Because SVI minimizes $KL(q\|p)$, putting mass where the posterior is thin is expensive and leaving posterior mass uncovered is cheap. So when the guide cannot match the shape, it errs towards being too narrow. I check the guide's intervals against NUTS on a subset, or check interval coverage on held-out data."
In an A/B framework like yours, an under-dispersed posterior pushes $P(\theta_B \gt \theta_A\mid D)$ toward 0 or 1 too early: decisions look more certain than they are. In your forecasting model, the forecast bands come from guide draws pushed through the model; an under-dispersed guide gives bands that cover less than their nominal level on a holdout (coverage checks, Chapter 7.16). A Student-t likelihood changes the shape of the data noise; the guide's shape for the parameters is a separate choice.
Under-dispersed = $q$ narrower than the posterior; its "95%" intervals cover less than 95%.
Reverse KL + too-simple family → tends to under-disperse (heavy tails, correlations, dropped modes). Student-t ν = 3: σ 1.26 vs sd 1.73, coverage 91%.
Trap: a tendency, not a law; check against NUTS or with coverage on held-out data.
Quick check: why does the reverse KL prefer σ = 1.26 over the "honest" σ = 1.73 for the Student-t(3) target?
To reach a standard deviation of 1.73 without heavy tails, the Normal must be fat in the "shoulders": at $\theta = 2.5$ it has density 0.081 while the Student-t has only 0.039, so $\log(q/p) = 0.74$ there, and $q$ puts a lot of its own mass in that region. The reverse KL averages exactly these log-ratios over $q$, giving 0.100. The narrower Normal (σ = 1.26) stays where the Student-t is high, and the far tails, which it ignores, cost almost nothing because the reverse KL only averages over places where $q$ itself has mass: 0.041.
Posterior correlations and the mean-field problem (a preview) core
Fit a straight trend line to a few weeks of sales. If the slope were a little higher, the intercept would have to be a little lower for the line to still pass through the data. The two parameters move together: the posterior over (slope, intercept) is a long, thin, tilted ellipse, a ridge. That is a posterior correlation.
The simplest VI family, mean-field, gives every parameter its own independent bell curve. Independent bells can only draw ellipses whose axes are horizontal and vertical; they cannot tilt. Under the reverse KL, $q$ must not poke out of the thin ridge (that is where $p$ is nearly 0), so it shrinks into a small round blob in the middle. Its width matches how much one parameter can move while the other is held fixed, which is much less than how much it really varies.
Three ways to say it:
- Picture: a small coin placed inside a long, thin, tilted cigar.
- Numbers: with correlation 0.9, the mean-field sd is $\sqrt{1 - 0.9^2} = 0.44$ times the true sd, and its "95%" interval covers only 61%.
- Slogan: mean-field VI reports how uncertain each parameter would be if all the others were already known.
A correlated Normal posterior. Two parameters, each with posterior mean 0 and sd 1, correlation ρ = 0.9.
- Covariance $\Sigma = \begin{bmatrix} 1 & 0.9 \\ 0.9 & 1 \end{bmatrix}$, determinant $1 - 0.81 = 0.19$.
- Precision (inverse covariance) $\Lambda = \Sigma^{-1} = \frac{1}{0.19}\begin{bmatrix} 1 & -0.9 \\ -0.9 & 1 \end{bmatrix}$, so $\Lambda_{11} = \Lambda_{22} = 1/0.19 = 5.26$.
- The mean-field $q$ that minimizes $KL(q\|p)$ is $q_1 = N(0, 1/\Lambda_{11})$, $q_2 = N(0, 1/\Lambda_{22})$ (formula in the definition): variance $0.19$, sd $\sqrt{0.19} = 0.436$.
- The true marginal sd is 1. So $q$ is 56% too narrow. Its 95% interval, $\pm 1.96 \times 0.436 = \pm 0.854$, contains $P(|Z| \le 0.854) = 0.61$ of the real posterior.
- The variance 0.19 has a meaning: it is the conditional variance $Var(\theta_1\mid\theta_2) = 1 - \rho^2$, the wiggle room of $\theta_1$ once $\theta_2$ is pinned.
- The leftover score: $KL(q^\star\|p) = -\tfrac12\log(1-\rho^2) = -\tfrac12\log 0.19 = 0.830$.
- Sums and differences go wrong in opposite ways: $Var(\theta_1 + \theta_2)$ is really $1 + 1 + 2(0.9) = 3.8$, but $q$ says $0.19 + 0.19 = 0.38$ (10 times too small); $Var(\theta_1 - \theta_2)$ is really $2 - 1.8 = 0.2$, but $q$ says $0.38$ (too large).
A mean-field family assumes the parameters are independent under q:
$$q(\theta) = \prod_{i=1}^{d} q_i(\theta_i), \qquad \text{e.g. } q_i = N(m_i, s_i^2) \text{ (NumPyro's } \texttt{AutoNormal}\text{)}.$$- For a Normal posterior $N(\mu, \Sigma)$ with precision $\Lambda = \Sigma^{-1}$, the reverse-KL optimum is $q_i^\star = N(\mu_i,\ 1/\Lambda_{ii})$: the means are exact, the variances are the conditional variances $1/\Lambda_{ii} = Var(\theta_i\mid\theta_{-i}) \le \Sigma_{ii}$. In two dimensions $1/\Lambda_{ii} = \sigma_i^2(1-\rho^2)$ and $KL(q^\star\|p) = -\tfrac12\log(1-\rho^2)$.
- Mean-field minimizing the forward KL gives instead the true marginals $q_i = p(\theta_i)$: right marginal spreads, but no correlation, so $q$ puts mass in the empty "corners" off the ridge.
- Coordinate updates. One classic way to fit mean-field VI (coordinate ascent, CAVI) updates one factor at a time. For this Normal target the mean update is "set $m_1$ to the conditional mean of $\theta_1$ given $\theta_2 = m_2$", then the same for $m_2$. With strong correlation these steps zig-zag slowly along the ridge: each full sweep shrinks the distance to the centre by a factor $\rho^2$.
- Families that can tilt fix the problem: a full-rank Normal $N(m, LL^\top)$ or a low-rank-plus-diagonal Normal (you will meet both in Chapter 6.13).
- "Mean-field" is a name borrowed from physics: each variable feels only the average ("mean") effect of the others.
Why do we need it?
Most real posteriors have correlated parameters: slope and intercept, a hierarchical mean and its group effects, trend and changepoints, seasonality and holidays. You must know what a mean-field guide does to them before trusting its intervals.
Where is it used?
AutoNormal and AutoDiagonalNormal in NumPyro, the default ADVI in Stan (meanfield), classic CAVI for topic models (LDA) and Gaussian mixtures, and the per-weight Gaussians of most Bayesian neural networks.
How is it used?
Start with a mean-field guide because it is cheap (2 numbers per parameter). Then check: fit a full-rank or low-rank guide (or NUTS on a subset) and compare the intervals of the quantities you report, especially sums such as a forecast. If they differ, the correlations matter.
"A mean-field guide assumes the parameters are independent in the posterior."
It assumes independence in q. The posterior stays correlated; the fitted $q$ simply cannot show it, and it pays by shrinking the spreads.
"Mean-field gets the means wrong too."
For a Normal posterior the means come out exactly right and only the spreads are wrong. For skewed or multimodal posteriors the means can move as well.
"Each parameter's interval is a bit narrow, so any combination is a bit narrow too."
Combinations can be wrong by large factors in either direction: with ρ = 0.9 the variance of $\theta_1 + \theta_2$ is 10 times too small, while $\theta_1 - \theta_2$ is almost 2 times too large. A forecast is a combination of many correlated parameters.
"Mean-field VI underestimates variance because the optimizer does not converge."
"It is the family, not the optimizer. The best mean-field $q$ under $KL(q\|p)$ matches the conditional variances $1/\Lambda_{ii}$, not the marginal variances $\Sigma_{ii}$."
Model answer: "A mean-field family has no correlations. Minimizing $KL(q\|p)$ forces $q$ to stay inside the high-density ridge of a correlated posterior, so each factor's variance becomes the conditional variance. For a bivariate Normal with correlation ρ the sd shrinks by $\sqrt{1-\rho^2}$, about 0.44 at ρ = 0.9. A full-rank or low-rank Gaussian guide can represent the correlation and fixes it."
Your forecasting model is full of correlated parameters: the base slope $k$, the offset $m$ and the changepoint adjustments $\delta_j$ trade off against each other (raise an early slope, lower a later δ), and seasonality, holiday and regressor coefficients can compete for the same bumps (Chapter 6.8). A mean-field guide would report each of them as too certain and would get forecast bands, which add many correlated pieces, wrong in a direction that depends on the signs of the correlations. That is the statistical reason your loop uses a full-rank guide for small models and a low-rank guide for large ones (Chapter 6.13): both can represent correlations. In an A/B framework like yours, a hierarchical model's population mean, spread τ and group effects are correlated too (Chapters 6.5–6.7).
Mean-field: $q(\theta) = \prod_i q_i(\theta_i)$. For a Normal posterior, min $KL(q\|p)$ gives $q_i = N(\mu_i, 1/\Lambda_{ii})$: exact means, conditional variances.
2D: sd shrinks by $\sqrt{1-\rho^2}$ (ρ = 0.9: 0.44, coverage 61%); $KL = -\tfrac12\log(1-\rho^2)$.
Trap: sums and differences of parameters can be off by large factors. Fix: full-rank / low-rank guides (6.13).
Quick check: posterior correlation ρ = 0.6, marginal sds 1. What are the mean-field sds and the leftover KL?
$\sqrt{1 - 0.36} = 0.8$ for both, so 20% too narrow. $KL = -\tfrac12\log 0.64 = 0.223$. Its 95% interval $\pm 1.57$ covers $P(|Z| \le 1.57) = 0.88$ of the true marginal.
Recap, cheat sheet and practice
- Variational inference turns inference into optimization: choose a family $q_\phi$, tune the variational parameters φ to minimize $KL(q_\phi\|p(\theta\mid D))$, then use $q$ as the posterior.
- KL divergence $KL(q\|p) = E_q[\log q - \log p] \ge 0$, zero only when $q = p$; not symmetric; infinite when $q$ has mass where $p$ has none.
- The best member of the family can still be wrong: the approximation error. Total error = approximation + optimization error. Constrained parameters are handled on an unconstrained (log/logit) scale.
- Reverse $KL(q\|p)$ (VI): mode-seeking, zero-forcing, can miss modes, several local optima, tends to be under-dispersed. Forward $KL(p\|q)$: mass-covering; for a Normal $q$ it matches mean and variance.
- VI uses the reverse direction because $KL(q\|p) = -\text{ELBO} + \log p(D)$: the evidence is a constant, so we only need the joint $p(\theta, D)$ and draws from $q$.
- Mean-field $q = \prod q_i$ cannot represent correlations: for a Normal posterior it gets the means right and the variances equal to the conditional variances $1/\Lambda_{ii}$ (sd × $\sqrt{1-\rho^2}$ in 2D). Full-rank and low-rank guides fix this (Chapter 6.13).
Cheat sheet
| Idea | Formula | Plain words / remember |
|---|---|---|
| VI objective | $\phi^\star = \arg\min_\phi KL(q_\phi\,\|\,p(\theta\mid D))$ | fit, don't sample; φ = knobs of $q$, not θ |
| KL (discrete) | $\sum_i q_i\log(q_i/p_i)$ | average log-ratio, weighted by the first; $\ge 0$ |
| KL of two Normals | $\log\frac{\sigma_2}{\sigma_1} + \frac{\sigma_1^2 + (\mu_1-\mu_2)^2}{2\sigma_2^2} - \frac12$ | N(0,1) vs N(0,2²): 0.318 one way, 0.807 the other |
| Approximation error | $\min_\phi KL(q_\phi\|p)$ | log-rate example: $\approx 1/(12a)$; must respect the support |
| Reverse KL | $E_q[\log q/p]$ | "never be where p is not": one mode, too narrow |
| Forward KL | $E_p[\log p/q]$ | "never miss where p is": Normal $q$ = moment match |
| Why reverse | $KL(q\|p) = -\text{ELBO} + \log p(D)$ | $\log p(D)$ constant in $q$; no posterior draws needed |
| Two modes ±3 | reverse: N(2.98, 1.02²), KL 0.689 ≈ log 2 | forward: N(0, 3.16²); compromise trap N(0, 2.74²) |
| Mean-field, Normal posterior | $q_i = N(\mu_i, 1/\Lambda_{ii})$ | conditional variance; ρ = 0.9 → sd × 0.44, coverage 61% |
import numpy as np
from scipy import stats, optimize, special
# 1) KL between two discrete distributions, in both directions
p = np.array([0.6, 0.3, 0.1]) # true plan shares (Basic, Pro, Team)
q = np.array([0.3, 0.3, 0.4]) # an approximation
print(round(np.sum(q * np.log(q / p)), 3), round(np.sum(p * np.log(p / q)), 3)) # 0.347 0.277
print(round(stats.entropy(q, p), 3)) # 0.347 scipy: entropy(a, b) = KL(a || b), the first argument gives the weights
# 2) KL between two Normals: the formula vs a Monte Carlo average over draws from q
def kl_normal(m1, s1, m2, s2): # KL( N(m1, s1^2) || N(m2, s2^2) )
return np.log(s2 / s1) + (s1**2 + (m1 - m2)**2) / (2 * s2**2) - 0.5
rng = np.random.default_rng(0)
th = rng.normal(0, 1, 200_000) # draws from q = N(0, 1)
mc = np.mean(stats.norm.logpdf(th, 0, 1) - stats.norm.logpdf(th, 0, 2)) # E_q[log q - log p]
print(round(kl_normal(0, 1, 0, 2), 4), round(mc, 3), round(kl_normal(0, 2, 0, 1), 4)) # 0.3181 0.317 0.8069
# 3) Reverse vs forward KL: the best single Normal for two modes at -3 and +3
x = np.linspace(-15, 15, 30001); dx = x[1] - x[0]
p_x = 0.5 * stats.norm.pdf(x, -3, 1) + 0.5 * stats.norm.pdf(x, 3, 1)
def rev(v): # KL(q || p) with q = N(v[0], exp(v[1])^2)
lq = stats.norm.logpdf(x, v[0], np.exp(v[1]))
return np.sum(np.exp(lq) * (lq - np.log(p_x))) * dx
def fwd(v): # KL(p || q)
lq = stats.norm.logpdf(x, v[0], np.exp(v[1]))
return np.sum(p_x * (np.log(p_x) - lq)) * dx
r = optimize.minimize(rev, [2.0, 0.0], method="Nelder-Mead")
f = optimize.minimize(fwd, [2.0, 0.0], method="Nelder-Mead")
print(np.round([r.x[0], np.exp(r.x[1]), r.fun], 3)) # [2.984 1.023 0.689] one mode, KL close to log 2
print(np.round([f.x[0], np.exp(f.x[1]), f.fun], 3)) # [0. 3.162 0.462] moment match: N(0, 10)
r0 = optimize.minimize(rev, [0.0, 1.0], method="Nelder-Mead")
print(np.round([r0.x[0], np.exp(r0.x[1]), r0.fun], 3)) # [0. 2.744 0.841] started in the middle: a worse local optimum
# 4) Under-dispersion with heavy tails: Student-t(3) target, best Normal under KL(q || p)
z = np.linspace(-8, 8, 4001); wz = stats.norm.pdf(z) * (z[1] - z[0]) # quadrature for E_q[.]
def rev_t(log_s):
s = np.exp(log_s); th = s * z
return np.sum(wz * (stats.norm.logpdf(th, 0, s) - stats.t.logpdf(th, 3)))
s_rev = np.exp(optimize.minimize_scalar(rev_t, bounds=(-2, 2), method="bounded").x)
cover = stats.t.cdf(1.96 * s_rev, 3) - stats.t.cdf(-1.96 * s_rev, 3)
print(round(s_rev, 3), round(np.sqrt(3), 3), round(cover, 3)) # 1.26 1.732 0.91 the "95%" interval covers 91%
# 5) Mean-field fit of a correlated Normal posterior (rho = 0.9)
rho = 0.9
Sigma = np.array([[1, rho], [rho, 1]]); Lam = np.linalg.inv(Sigma)
def kl_gauss(m_q, S_q, m_p, S_p): # KL( N(m_q, S_q) || N(m_p, S_p) )
Sp_inv = np.linalg.inv(S_p); d = m_p - m_q
return 0.5 * (np.trace(Sp_inv @ S_q) - len(m_q) + d @ Sp_inv @ d + np.log(np.linalg.det(S_p) / np.linalg.det(S_q)))
obj = lambda v: kl_gauss(v[:2], np.diag(np.exp(2 * v[2:])), np.zeros(2), Sigma) # independent q: diagonal covariance
res = optimize.minimize(obj, [0.5, -0.5, 0.0, 0.0])
print(np.round(res.x[:2], 3) + 0.0, np.round(np.exp(res.x[2:]), 3)) # [0. 0.] [0.436 0.436] means exact, sds too small
print(np.round(1 / np.sqrt(np.diag(Lam)), 3), round(res.fun, 3), round(-0.5 * np.log(1 - rho**2), 3)) # [0.436 0.436] 0.83 0.83
# 6) Best Normal for eta = log(lambda) when lambda | D ~ Gamma(4, rate 2): closed form vs the truth
a, b = 4, 2
mu_star, s_star = np.log(a / b) - 1 / (2 * a), 1 / np.sqrt(a)
true_mean, true_sd = special.digamma(a) - np.log(b), np.sqrt(special.polygamma(1, a))
kl_min = special.gammaln(a) - (a - 0.5) * np.log(a) + a - 0.5 * np.log(2 * np.pi)
print(round(mu_star, 3), s_star, round(true_mean, 3), round(true_sd, 3), round(kl_min, 4)) # 0.568 0.5 0.563 0.533 0.0208
1. Which quantity does variational inference in NumPyro (SVI with Trace_ELBO) effectively minimize?
2. The posterior has two well-separated modes of equal size. You fit one Normal $q$ by minimizing $KL(q\|p)$. What do you typically get?
3. A Normal posterior has marginal sds 1 and 1 and correlation 0.8. What sd does the best mean-field $q$ (reverse KL) give each parameter?
4. Why does VI use $KL(q\|p)$ rather than $KL(p\|q)$?
5. A Normal $q$ is fitted to a parameter that must be positive (a rate λ), directly on the λ scale. What is $KL(q\|p)$?
6. What is the "approximation error" of a variational family?
Practice problems
A. Compute $KL(q\|p)$ and $KL(p\|q)$ for $q = (0.5, 0.25, 0.25)$ and $p = (0.25, 0.25, 0.5)$. Why are they equal here?
- $KL(q\|p) = 0.5\log 2 + 0.25\log 1 + 0.25\log 0.5 = 0.347 + 0 - 0.173 = 0.173$.
- $KL(p\|q) = 0.25\log 0.5 + 0.25\log 1 + 0.5\log 2 = -0.173 + 0 + 0.347 = 0.173$.
- They agree because $p$ is $q$ with the first and last categories swapped, so the two sums contain the same terms. In general (for example the plan-choice example, 0.347 vs 0.277) they differ.
B. $q = N(1, 0.5^2)$ and $p = N(0, 1)$. Compute both KLs and say which direction punishes "q is too narrow and shifted" more.
- $KL(q\|p) = \log(1/0.5) + (0.25 + 1)/(2 \times 1) - 0.5 = 0.693 + 0.625 - 0.5 = 0.818$.
- $KL(p\|q) = \log(0.5/1) + (1 + 1)/(2 \times 0.25) - 0.5 = -0.693 + 4 - 0.5 = 2.807$.
- The forward KL punishes it far more (2.81 vs 0.82): much of $p$'s mass lies where the narrow $q$ is tiny. The reverse KL, which VI uses, is relatively relaxed about a $q$ that is too narrow.
C. Target $p = 0.3\,N(-2, 1) + 0.7\,N(3, 0.5^2)$. Find the forward-KL best Normal by hand, and describe what the reverse KL does.
- Mean: $0.3(-2) + 0.7(3) = -0.6 + 2.1 = 1.5$.
- $E[\theta^2]$: $0.3(1 + 4) + 0.7(0.25 + 9) = 1.5 + 6.475 = 7.975$. Variance: $7.975 - 1.5^2 = 5.725$, sd $2.39$. Forward fit: $N(1.5, 2.39^2)$, whose peak sits where $p$ has very little mass.
- Reverse KL (numerically): starting near the big mode it finds $N(3.00, 0.50^2)$ with $KL \approx 0.356 \approx -\log 0.7$; starting near the small mode it finds $N(-1.99, 1.02^2)$ with $KL \approx 1.20 \approx -\log 0.3$, a worse local optimum. Each one-mode fit pays about $-\log$(weight of its mode).
D. A posterior has correlation −0.95 between two standardized parameters. Give the mean-field sd, its interval coverage, and what happens to $\theta_1 + \theta_2$ and $\theta_1 - \theta_2$.
- sd $= \sqrt{1 - 0.9025} = 0.312$; $KL = -\tfrac12\log 0.0975 = 1.164$.
- Its 95% interval $\pm 1.96 \times 0.312 = \pm 0.612$ covers $P(|Z| \le 0.612) = 0.46$ of the true marginal: less than half.
- True $Var(\theta_1 + \theta_2) = 1 + 1 + 2(-0.95) = 0.1$; $q$ says $0.0975 + 0.0975 = 0.195$ (too wide). True $Var(\theta_1 - \theta_2) = 2 + 1.9 = 3.9$; $q$ says 0.195 (20 times too small). With a negative correlation, the difference is the quantity that mean-field badly underestimates.
E. (Interview) "Our SVI fit says $P(\theta_B \gt \theta_A\mid D) = 0.99$, but a NUTS run on the same model says 0.93. Which do you trust, and why might they differ?"
"I would trust NUTS more here, after checking its diagnostics (R̂, ESS, divergences; Chapter 6.10). SVI minimizes $KL(q\|p)$, which tends to make $q$ too narrow, especially with a mean-field guide when parameters are correlated, as they are in a hierarchical model where group effects share a population mean and spread. A too-narrow posterior for $\theta_B - \theta_A$ pushes $P(\theta_B \gt \theta_A)$ toward 1. I would try a full-rank or low-rank guide; if its answer moves toward 0.93, the guide family was the problem. For the final decision I would report the NUTS number or a guide that matches it."
F. Three candidates with posterior $p = (0.121, 0.386, 0.493)$, prior × likelihood $\tilde p = (0.0199, 0.0634, 0.0809)$ and $\log p(D) = -1.807$. For $q = (0, 0.3, 0.7)$ compute $KL(q\|p)$ and the computable score $J = E_q[\log q - \log\tilde p]$, and check that they differ by $-\log p(D)$.
- $KL = 0.3\log(0.3/0.386) + 0.7\log(0.7/0.493) = 0.3(-0.252) + 0.7(0.351) = -0.076 + 0.245 = 0.170$.
- $J = 0.3(\log 0.3 - \log 0.0634) + 0.7(\log 0.7 - \log 0.0809) = 0.3(-1.204 + 2.758) + 0.7(-0.357 + 2.515) = 0.466 + 1.511 = 1.977$.
- $J - KL = 1.977 - 0.170 = 1.807 = -\log p(D)$. Same gap as at $w = 0.5$ and $w = 0.439$: the evidence only shifts the score.
The ELBO and stochastic variational inference
Chapter 6.11 said: make $q$ close to the posterior by minimizing $KL(q\|p)$. But that KL contains the one number we can never compute, the evidence $p(D)$. This chapter shows the way around it. One line of algebra splits $\log p(D)$ into a part we can compute, the ELBO, plus the KL. Then we see how NumPyro estimates the ELBO with random draws, gets gradients through those draws with the reparameterization trick, and climbs it with Adam, one noisy step at a time. That loop is SVI, and it is the engine inside both of your projects.
- Derive $\log p(D) = \text{ELBO}(q) + KL(q\,\|\,p(\theta\mid D))$ step by step, and explain why maximizing the ELBO is the same as minimizing the KL
- Read the ELBO two ways: $E_q[\log p(D,\theta)] + H[q]$ (fit + entropy) and $E_q[\log p(D\mid\theta)] - KL(q\,\|\,p(\theta))$ (data fit − complexity)
- Estimate the ELBO by Monte Carlo with $S$ draws ("particles") and see the noise fall like $1/\sqrt{S}$
- Explain the reparameterization trick $\theta = \mu + \sigma\varepsilon$ and why it gives low-noise gradients
- Name every piece of SVI: model, guide, latent variables, variational parameters, ELBO estimator, gradients, optimizer, and know that
svi.updatereturns the loss = −ELBO - Use minibatches correctly: scale the data part of the ELBO by $N/B$
What we need from earlier chapters: VI, the KL divergence and its two directions, and the fact that $KL(q\|p) = -\text{ELBO} + \log p(D)$ (Chapter 6.11); prior, likelihood, posterior and evidence (Chapter 6.1); the Normal-Normal update (Chapter 6.3); expectations and the law of large numbers (Chapter 4.5, Chapter 4.13); the chain rule and gradients (Calculus guide); stochastic gradient descent and Adam (Optimization guide). Notation: $p(D, \theta) = p(D\mid\theta)\,p(\theta)$ is the joint (likelihood × prior); $q_\phi(\theta)$ is the guide with variational parameters $\phi$; $H[q] = -E_q[\log q(\theta)]$ is the entropy of $q$, a measure of how spread out it is; $S$ is the number of Monte Carlo draws; $N$ is the number of data points and $B$ the minibatch size. Running example: a mean effect θ with prior $N(0, 1)$ and four observations $y = (0.5, 1.5, 0.8, 1.2)$, each $N(\theta, 1)$. Its exact posterior is $N(0.8, 0.2)$, written as $N(\text{mean}, \text{variance})$, so its sd is $\sqrt{0.2} = 0.447$ (precision $1 + 4 = 5$, mean $4 \times 1.0/5$), and $\log p(D) = -5.170$.
Deriving the ELBO: $\log p(D) = \text{ELBO} + KL$ core
Think of a bucket of a fixed size. The size of the bucket is $\log p(D)$: how probable the data is under your model. It does not depend on $q$ at all; it is fixed by the model and the data. Now pick any approximation $q$. It fills the bucket up to some water level: that level is the ELBO. The empty space above the water is the KL divergence from $q$ to the posterior.
Two facts follow at once. The water can never rise above the rim, because the empty space (a KL) is never negative: the ELBO is always below $\log p(D)$. That is its name: Evidence Lower BOund. And since the bucket's size is fixed, every bit of water you add removes the same amount of empty space: raising the ELBO is shrinking the KL. We cannot measure the empty space directly (it needs $p(D)$), but we can measure the water level. So we climb the ELBO.
Three ways to say it:
- Picture: one fixed bar, $\log p(D)$, split into a computable part (ELBO) and a gap (KL).
- Numbers: three candidate rates, uniform $q$: $-1.807 = -1.965 + 0.158$. With $q$ equal to the posterior: $-1.807 = -1.807 + 0$.
- Slogan: maximize the ELBO = minimize the KL, and the ELBO can never pass $\log p(D)$.
The three candidate rates from Chapter 6.1 (5%, 10%, 15%; prior $\tfrac13$ each; 3 buyers among 20). Prior × likelihood $= p(D, \theta) = (0.0199,\ 0.0634,\ 0.0809)$, with logs $(-3.919,\ -2.759,\ -2.514)$. Evidence $p(D) = 0.1642$, $\log p(D) = -1.807$. Posterior $(0.121,\ 0.386,\ 0.493)$. Take the uniform guess $q = (\tfrac13, \tfrac13, \tfrac13)$.
- $E_q[\log p(D,\theta)] = \tfrac13(-3.919 - 2.759 - 2.514) = \tfrac13(-9.192) = -3.064$.
- $E_q[\log q(\theta)] = \log\tfrac13 = -1.099$.
- $\text{ELBO} = E_q[\log p(D,\theta)] - E_q[\log q(\theta)] = -3.064 + 1.099 = -1.965$.
- $KL(q\|p) = \tfrac13\big[\log\tfrac{0.333}{0.121} + \log\tfrac{0.333}{0.386} + \log\tfrac{0.333}{0.493}\big] = \tfrac13(1.0136 - 0.1468 - 0.3914) = 0.158$.
- Add: $-1.965 + 0.158 = -1.807 = \log p(D)$. ✓
- Now take $q$ = the posterior itself: the ELBO becomes exactly $-1.807$ and the KL is 0. The best two-point $q$ of Chapter 6.11 ($w = 0.439$) gives ELBO $-1.936$ and KL $0.129$: again they add to $-1.807$.
Every $q$ splits the same number $-1.807$ differently. The better the $q$, the more of it is ELBO and the less is KL.
The derivation, one small step per line. Take any distribution $q(\theta)$ with $q \gt 0$ wherever the posterior is positive.
$$\begin{aligned} \log p(D) &= \textstyle\int q(\theta)\,d\theta \cdot \log p(D) && q \text{ adds up to 1} \\ &= E_q\big[\log p(D)\big] && \log p(D) \text{ does not depend on } \theta \\ &= E_q\Big[\log \frac{p(D,\theta)}{p(\theta\mid D)}\Big] && \text{Bayes: } p(\theta\mid D) = p(D,\theta)/p(D) \\ &= E_q\Big[\log \Big(\frac{p(D,\theta)}{q(\theta)}\cdot\frac{q(\theta)}{p(\theta\mid D)}\Big)\Big] && \text{multiply and divide by } q(\theta) \\ &= \underbrace{E_q\big[\log p(D,\theta) - \log q(\theta)\big]}_{\text{ELBO}(q)} + \underbrace{E_q\big[\log q(\theta) - \log p(\theta\mid D)\big]}_{KL(q\,\|\,p(\theta\mid D))} && \log(ab) = \log a + \log b \end{aligned}$$- $\text{ELBO}(q) = E_q[\log p(D,\theta) - \log q(\theta)]$ needs only the joint (likelihood × prior) and draws from $q$.
- Because $KL \ge 0$: $\;\text{ELBO}(q) \le \log p(D)$, with equality exactly when $q = p(\theta\mid D)$.
- Because $\log p(D)$ is fixed: $\;\arg\max_\phi \text{ELBO}(q_\phi) = \arg\min_\phi KL(q_\phi\,\|\,p(\theta\mid D))$.
- A second route (Jensen). $\log p(D) = \log\int p(D,\theta)\,d\theta = \log E_q\Big[\frac{p(D,\theta)}{q(\theta)}\Big] \ge E_q\Big[\log\frac{p(D,\theta)}{q(\theta)}\Big] = \text{ELBO}$, because the log of an average is at least the average of the logs (the log curve bends down). The size of that "Jensen gap" is exactly the KL.
Why do we need it?
The KL to the posterior cannot be computed (it contains $p(D)$), so we cannot minimize it directly. The ELBO differs from it only by that constant, and it can be estimated from draws of $q$. It turns an impossible objective into a computable one.
Where is it used?
Every variational method: NumPyro's Trace_ELBO and TraceMeanField_ELBO, Pyro, Stan's ADVI, the VAE loss (reconstruction − KL), Bayesian neural networks. In your forecasting model, ELBO-based early stopping watches this exact quantity.
How is it used?
Optimizers minimize, so libraries report the loss = −ELBO. You watch the loss fall, stop when it no longer improves (Chapter 6.14), and compare losses of different guides on the same model and data: the lower loss (higher ELBO) is the guide closer to the posterior.
"The ELBO is the log evidence."
It is a lower bound on the log evidence. The two are equal only when $q$ is exactly the posterior; otherwise the difference is $KL(q\|p(\theta\mid D)) \gt 0$.
"Model A has a higher ELBO than model B, so model A fits better."
Each ELBO is $\log p(D\mid\text{model}) - KL$, and the two KL gaps can be very different (and both ELBOs are noisy estimates). An ELBO difference between models is a rough hint, not a model comparison. Compare models with held-out predictive performance or posterior predictive checks (Chapter 6.8). Comparing guides on the same model and data is safe.
"My ELBO is positive, so something is broken."
With continuous data, densities can be bigger than 1, so $\log p(D)$ and the ELBO can be any sign. Only changes in the ELBO and comparisons on the same model and data carry meaning.
The syllabus notes that you implemented ELBO-based stopping. This is the quantity your loop watches. NumPyro's svi.update(state, ...) returns (state, loss) where loss = −ELBO (a one-draw estimate by default), so "improvement" means the loss going down. Because the bound can never pass $\log p(D)$, the ELBO curve flattens as $q$ approaches the best member of the guide family; how your loop decides that it has flattened (relative tolerance, patience, keeping the best state) is the subject of Chapter 6.14.
$\log p(D) = \underbrace{E_q[\log p(D,\theta) - \log q(\theta)]}_{\text{ELBO}} + KL(q\,\|\,p(\theta\mid D))$.
ELBO ≤ log p(D); equal iff $q$ = posterior. Max ELBO ⇔ min KL, because $\log p(D)$ is fixed.
Trap: ELBO ≠ evidence; comparing ELBOs across different models is unreliable. Loss = −ELBO.
Quick check: for some $q$ the ELBO is −120.4, and $\log p(D) = -118.9$ (known from a conjugate formula). What is $KL(q\|p(\theta\mid D))$?
$KL = \log p(D) - \text{ELBO} = -118.9 - (-120.4) = 1.5$ nats. If someone reported an ELBO of −118.5 for this model and data, you would know it is an estimation error: the ELBO can never be above $\log p(D)$ (apart from Monte Carlo noise in the estimate).
Two ways to read the ELBO: fit + entropy, and data fit − complexity core
The ELBO is one number, but it can be split into two forces that pull $q$ in opposite directions. Seeing the tug-of-war explains what the optimizer is really doing.
- Reading 1: fit + entropy. $E_q[\log p(D,\theta)]$ rewards $q$ for putting its draws where prior × likelihood is high. On its own it would squeeze $q$ into a single spike at the best θ (the MAP of Chapter 5.2). The entropy $H[q] = -E_q[\log q]$ rewards $q$ for being spread out. The balance between the two gives a $q$ that is both in the right place and honestly wide.
- Reading 2: data fit − complexity. $E_q[\log p(D\mid\theta)]$ rewards explaining the data. $KL(q\|p(\theta))$ charges $q$ for moving away from the prior. It is the same trade-off as a regularized loss: "fit the data, but do not wander far from what you believed without good reason".
Three ways to say it:
- Picture: a spring pulls $q$ toward the best-fitting θ and makes it narrow; a second spring keeps it wide and near the prior.
- Numbers: at the exact posterior of the running example, data fit $-4.446$, distance from the prior $0.725$: ELBO $= -4.446 - 0.725 = -5.170$.
- Slogan: ELBO = how well you explain the data − how much you had to change your mind to do it.
Running example (prior $N(0,1)$, data $0.5, 1.5, 0.8, 1.2$ each $N(\theta, 1)$) with $q = N(0.5, 0.3^2)$. For a Normal $q$, $E_q[(y - \theta)^2] = (y - \mu)^2 + \sigma^2$, and $\log N(y;\theta,1) = -0.919 - (y-\theta)^2/2$ (where $0.919 = \tfrac12\log 2\pi$).
- Data fit, one term per observation: $y = 0.5$: $-0.919 - (0^2 + 0.09)/2 = -0.964$; $y = 1.5$: $-0.919 - (1 + 0.09)/2 = -1.464$; $y = 0.8$: $-0.919 - (0.09 + 0.09)/2 = -1.009$; $y = 1.2$: $-0.919 - (0.49 + 0.09)/2 = -1.209$. Sum: $E_q[\log p(D\mid\theta)] = -4.646$.
- Prior term: $E_q[\log p(\theta)] = -0.919 - (0.5^2 + 0.09)/2 = -0.919 - 0.170 = -1.089$.
- Entropy: $H[q] = \tfrac12\log(2\pi e \sigma^2) = \tfrac12\log(2\pi e \times 0.09) = 0.215$.
- Reading 1: $E_q[\log p(D,\theta)] + H[q] = (-4.646 - 1.089) + 0.215 = -5.735 + 0.215 = -5.520$.
- KL to the prior (two-Normal formula of Chapter 6.11): $KL(N(0.5, 0.3^2)\|N(0,1)) = \log\tfrac{1}{0.3} + \tfrac{0.09 + 0.25}{2} - \tfrac12 = 1.204 + 0.170 - 0.5 = 0.874$.
- Reading 2: $E_q[\log p(D\mid\theta)] - KL(q\|p(\theta)) = -4.646 - 0.874 = -5.520$. The same number. ✓
- Compared with $\log p(D) = -5.170$, the KL to the posterior is $0.350$: this $q$ is too far left and too narrow.
Using $\log p(D,\theta) = \log p(D\mid\theta) + \log p(\theta)$:
$$\text{ELBO}(q) = \underbrace{E_q[\log p(D,\theta)]}_{\text{expected log joint (fit)}} + \underbrace{H[q]}_{\text{entropy}} = \underbrace{E_q[\log p(D\mid\theta)]}_{\text{expected log-likelihood}} - \underbrace{KL\big(q(\theta)\,\|\,p(\theta)\big)}_{\text{distance from the prior}} .$$- Second form from the first: $E_q[\log p(D\mid\theta)] + E_q[\log p(\theta) - \log q(\theta)]$, and the last average is $-KL(q\|p(\theta))$.
- Do not mix up the two KLs: the ELBO contains $KL(q\|\text{prior})$, and it differs from $\log p(D)$ by $KL(q\|\text{posterior})$.
- For the running example the ELBO is $\text{const} - \tfrac52\big((\mu - 0.8)^2 + \sigma^2\big) + \log\sigma$: the fit term wants σ → 0, the entropy $\log\sigma$ wants σ large, and the best σ is $1/\sqrt5 = 0.447$, the posterior sd.
- The entropy of a Normal has a formula, $\tfrac12\log(2\pi e\sigma^2)$, and so does the KL between Normals. NumPyro's
TraceMeanField_ELBOuses such formulas when they exist;Trace_ELBOestimates every term by sampling. - A VAE's loss is exactly reading 2 with a minus sign: reconstruction error + KL(encoder ‖ prior).
Why do we need it?
The two readings tell you what each part of the objective does. Without the entropy (or the KL to the prior), the optimizer would collapse $q$ onto a point estimate and report zero uncertainty. With them, VI behaves like regularized fitting that keeps a spread.
Where is it used?
The VAE loss (reconstruction + KL), the KL-annealing trick in deep generative models, NumPyro's TraceMeanField_ELBO (analytic KL to the prior), Bayes-by-backprop for neural networks, and debugging SVI runs by logging the two parts separately.
How is it used?
When an SVI fit looks wrong, look at the parts: a guide whose scales shrink to almost zero has lost the entropy battle (often a too-large learning rate or a bad initialization); a guide stuck at the prior has a data term that is too weak (for example, a forgotten minibatch scale, below).
"The KL in the ELBO is the KL to the posterior."
Inside the ELBO sits $KL(q\|\text{prior})$, a regularizer we can compute. The KL to the posterior is the gap between the ELBO and $\log p(D)$, which we cannot compute.
"Maximizing the expected log-likelihood alone would be better: it fits the data best."
Without the KL-to-prior (or entropy) term, $q$ shrinks to a spike at the maximum-likelihood point: no uncertainty, and the prior is ignored. The penalty is what makes the result a posterior approximation.
In your forecasting model, reading 2 says: the guide is rewarded for explaining the daily series (expected log-likelihood under your Normal, Student-t or Negative Binomial likelihood) and charged for moving the changepoint slopes $\delta_j$ away from their Laplace prior, the seasonality coefficients away from theirs, and so on. That charge is how the Laplace prior's shrinkage enters SVI. In an A/B framework like yours, the same term pulls each segment's effect toward the population level of a hierarchical prior.
$\text{ELBO} = E_q[\log p(D,\theta)] + H[q] = E_q[\log p(D\mid\theta)] - KL(q\|p(\theta))$.
Fit pulls $q$ to a narrow spike at the best θ; entropy / KL-to-prior keeps it wide and near the prior.
Trap: the ELBO contains KL(q‖prior); its gap to log p(D) is KL(q‖posterior). Two different KLs.
Quick check: in the running example, why is the best σ the same (0.447) for every μ?
The ELBO is $\text{const} - \tfrac52(\mu - 0.8)^2 - \tfrac52\sigma^2 + \log\sigma$: μ and σ appear in separate terms. Setting the σ-slope to zero, $-5\sigma + 1/\sigma = 0$, gives $\sigma^2 = 1/5$, so $\sigma = 0.447$ regardless of μ. (In models where the log joint is not quadratic, the best width does depend on where $q$ sits.)
Monte Carlo ELBO estimates: particles and noise
The ELBO is an average over every possible θ under $q$. For our tiny example there is a formula, but for a real model (your forecasting model with dozens of parameters) there is not. So we do what a pollster does: we cannot ask everyone, so we ask a few people at random and average their answers. Draw a few values $\theta_1, \dots, \theta_S$ from $q$ (they are called particles), compute $\log p(D,\theta_s) - \log q(\theta_s)$ for each, and average. That is a Monte Carlo estimate of the ELBO.
It is right on average (unbiased), but every estimate wobbles. The wobble shrinks like $1/\sqrt{S}$: four times as many particles halve it. NumPyro's default is a single particle per step, so the loss it prints is very noisy. And one surprise: if $q$ were exactly the posterior, every particle would give the same value, $\log p(D)$, and the noise would vanish.
Three ways to say it:
- Picture: a cloud of estimates around the true ELBO that tightens as you add particles.
- Numbers: for $q = N(0.5, 0.3^2)$ the one-particle estimates have sd 0.59; with 4 particles 0.30; with 16, 0.15; with 64, 0.075.
- Slogan: four times the particles, half the noise, four times the cost.
Three particles by hand. $q = N(0.5, 0.3^2)$; suppose the three standard-normal draws are $\varepsilon = -1, 0, 1$, so $\theta = 0.5 + 0.3\varepsilon = 0.2, 0.5, 0.8$.
- $\theta = 0.2$: $\log p(D,\theta) = -6.185$, $\log q(\theta) = -0.215$; term $= -6.185 + 0.215 = -5.970$.
- $\theta = 0.5$: $\log p(D,\theta) = -5.510$, $\log q(\theta) = 0.285$; term $= -5.795$.
- $\theta = 0.8$: $\log p(D,\theta) = -5.285$, $\log q(\theta) = -0.215$; term $= -5.070$.
- Average: $(-5.970 - 5.795 - 5.070)/3 = -16.834/3 = -5.611$. The exact ELBO is $-5.520$; this estimate is off by $-0.09$, which is normal for 3 particles.
- Repeating with fresh draws many times: the sd of one-particle estimates is $0.59$; of 3-particle averages $0.59/\sqrt3 = 0.34$.
- At the exact posterior $q = N(0.8, 0.2)$, every term equals $\log p(D) = -5.170$, so the estimate has no noise at all.
With particles $\theta_1, \dots, \theta_S$ drawn independently from $q_\phi$:
$$\widehat{\text{ELBO}} = \frac{1}{S}\sum_{s=1}^{S}\Big[\log p(D,\theta_s) - \log q_\phi(\theta_s)\Big], \qquad E\big[\widehat{\text{ELBO}}\big] = \text{ELBO}, \qquad sd\big(\widehat{\text{ELBO}}\big) = \frac{sd_q\big[\log p(D,\theta) - \log q(\theta)\big]}{\sqrt S}.$$- Unbiased: right on average. Noisy: the noise falls like $1/\sqrt S$, and the cost grows like $S$.
- The noise depends on how far $q$ is from the posterior: $\log p(D,\theta) - \log q(\theta) = \log p(D) + \log p(\theta\mid D) - \log q(\theta)$, which is constant only when $q = p(\theta\mid D)$.
- In NumPyro:
Trace_ELBO(num_particles=S); the default is $S = 1$. The value returned bysvi.updateis $-\widehat{\text{ELBO}}$ for that step's draws. - The gradient is estimated from the same draws, so it is noisy too: SVI is stochastic gradient ascent.
Why do we need it?
For real models the ELBO's average over $q$ has no formula. Sampling turns it into something a computer can evaluate in one pass of the model, for any model you can write in NumPyro, which is what makes SVI "black box".
Where is it used?
NumPyro's and Pyro's Trace_ELBO, Stan's ADVI (which also estimates the ELBO with draws), VAEs (usually one sample per data point), and every training curve of an SVI run, including the loss your loop monitors.
How is it used?
Keep $S = 1$ during training (many cheap noisy steps beat few precise ones), but judge progress from a smoothed or averaged loss. When you need an accurate ELBO value (to compare two guides), evaluate it once at the end with many particles.
"The loss went up on this step, so the fit got worse."
With one particle, a single step's loss can jump by more than the whole improvement of the last hundred steps. Judge progress from averages over many steps, or from a separate many-particle evaluation.
"More particles is always better."
More particles cost proportionally more per step. The noise falls only like $1/\sqrt S$, so 100 particles buy a 10-fold noise reduction at 100 times the cost. Often more steps with $S = 1$ (and a decaying learning rate) are the better deal.
This noise is the reason your custom loop cannot stop at the first step where the ELBO fails to improve. Compared against the best ELBO so far, a noisy estimate will often look worse by chance. That is why a stopping rule needs a tolerance, patience across several evaluations and a saved best state, and why the size of the noise relative to the size of the ELBO matters (Chapter 6.14 builds exactly this).
$\widehat{\text{ELBO}} = \frac1S\sum_s[\log p(D,\theta_s) - \log q(\theta_s)]$, $\theta_s \sim q$: unbiased, sd ∝ $1/\sqrt S$, cost ∝ $S$.
NumPyro: Trace_ELBO(num_particles=S), default 1. Zero noise only if $q$ = posterior.
Trap: one step's loss is mostly noise; smooth before you judge.
Quick check: with one particle the ELBO estimate has sd 2.4. How many particles do you need for sd 0.3?
sd falls like $1/\sqrt S$: $2.4/\sqrt S = 0.3$ gives $\sqrt S = 8$, so $S = 64$ particles, and each step costs about 64 times as much.
The reparameterization trick: gradients through random draws core
To climb the ELBO we need its slope with respect to the knobs φ = (μ, σ). But the ELBO is an average over draws from $q_\phi$, and turning a knob changes which θ's get drawn. How do you take the derivative of a dice roll?
The trick is to separate the randomness from the knobs. First roll the dice: draw a standard-normal number $\varepsilon \sim N(0,1)$. It has no knobs inside. Then build the particle: $\theta = \mu + \sigma\varepsilon$. For that fixed ε, θ is now an ordinary smooth function of μ and σ, so ordinary calculus works. Nudge μ and every particle slides by the same amount; nudge σ and each particle stretches away from μ in proportion to its ε. JAX's automatic differentiation can then follow the chain rule straight through the draw.
Everyday picture: instead of re-rolling the dice every time you adjust a knob, you roll once, write the numbers on cards, and then only shift and stretch the cards.
Three ways to say it:
- Picture: fixed marbles ε on a bottom rail, mapped to the top rail by "stretch by σ, shift by μ".
- Numbers: with ε = 0.2, −1, 0.8 and q = N(1, 0.5²), the slope of $E_q[\theta^2]$ in μ comes out as exactly 2.0, and in σ as 0.56 (true value 1.0; more draws get closer).
- Slogan: draw the noise first, then let the parameters move it.
A gradient by hand. Let $f(\theta) = \theta^2$ and $q = N(\mu, \sigma^2)$ with μ = 1, σ = 0.5. The exact answer is known: $E_q[\theta^2] = \mu^2 + \sigma^2 = 1.25$, so the slopes are $\partial/\partial\mu = 2\mu = 2$ and $\partial/\partial\sigma = 2\sigma = 1$.
- Draw three noise values: $\varepsilon = 0.2,\ -1.0,\ 0.8$.
- Build the particles: $\theta = 1 + 0.5\varepsilon = 1.1,\ 0.5,\ 1.4$.
- Chain rule: $\frac{\partial f(\theta)}{\partial\mu} = f'(\theta)\frac{\partial\theta}{\partial\mu} = 2\theta \times 1$ and $\frac{\partial f(\theta)}{\partial\sigma} = 2\theta \times \varepsilon$.
- Slope in μ: $(2.2 + 1.0 + 2.8)/3 = 6.0/3 = 2.0$. (Exactly right here, because these three ε average to 0.)
- Slope in σ: $(2(1.1)(0.2) + 2(0.5)(-1.0) + 2(1.4)(0.8))/3 = (0.44 - 1.0 + 2.24)/3 = 1.68/3 = 0.56$. The truth is 1.0: three draws are noisy. With 200 000 draws, JAX's
jax.gradgives (2.00, 1.01). - For the ELBO itself, $f(\theta) = \log p(D,\theta) - \log q(\theta)$. In the running example $\frac{d}{d\theta}\log p(D,\theta) = \sum_i (y_i - \theta) - \theta = 4 - 5\theta$, so one particle at θ = 0.5 gives the μ-slope $4 - 2.5 = 1.5$, which here equals the exact $-5(\mu - 0.8) = 1.5$ at μ = 0.5.
If θ can be written as a smooth function of φ and a parameter-free noise, $\theta = g_\phi(\varepsilon)$ with $\varepsilon \sim p(\varepsilon)$, then
$$\nabla_\phi\, E_{q_\phi}\big[f(\theta)\big] = E_{\varepsilon}\big[\nabla_\phi\, f(g_\phi(\varepsilon))\big] \approx \frac1S\sum_{s=1}^S \nabla_\phi f\big(g_\phi(\varepsilon_s)\big).$$- For a Normal guide: $\theta = \mu + \sigma\varepsilon$, $\varepsilon \sim N(0,1)$, so $\partial\theta/\partial\mu = 1$ and $\partial\theta/\partial\sigma = \varepsilon$. For a multivariate Normal guide: $\theta = m + L\varepsilon$ with $L$ a Cholesky factor (Chapter 6.13). For constrained parameters the draw is then pushed through a transform such as exp, which is also smooth.
- For the ELBO, $f(\theta) = \log p(D,\theta) - \log q_\phi(\theta)$; autodiff handles both parts.
- The alternative, the score-function (REINFORCE, likelihood-ratio) estimator, $\nabla_\phi E_q[f] = E_q\big[f(\theta)\,\nabla_\phi\log q_\phi(\theta)\big]$, works even for discrete θ, but it uses only the values of $f$, not its slope, and is usually far noisier. Subtracting a constant "baseline" $b$ from $f$ keeps it unbiased (because $E_q[\nabla_\phi\log q_\phi] = 0$) and reduces the noise.
- In NumPyro: a distribution with
has_rsample = Truecan be reparameterized (Normal, Laplace, Student-t, Gamma, Beta…; not Poisson or Bernoulli).Trace_ELBOsupports only reparameterized latent sites;TraceGraph_ELBOadds score-function terms for non-reparameterizable ones; discrete latent variables are often summed out instead (TraceEnum_ELBO).
Why do we need it?
Gradient-based optimizers need the slope of the ELBO with respect to the guide's parameters. The reparameterization trick gives an unbiased slope with low noise for any model that JAX can differentiate, which is what makes SVI fast and general.
Where is it used?
Every Trace_ELBO step in NumPyro and Pyro, Stan's ADVI, the VAE (where it was popularized), Bayes-by-backprop neural networks, and normalizing-flow guides such as AutoIAFNormal.
How is it used?
You do not write it: the guide draws ε internally and builds θ from its loc and scale, and jax.grad differentiates the loss through that. What you do: keep latent variables continuous (or enumerate discrete ones) so that Trace_ELBO applies.
"Sampling is random, so you cannot take gradients through it."
You cannot differentiate the dice roll, but you do not need to: draw the parameter-free noise ε first, then θ = μ + σε is an ordinary differentiable function of μ and σ.
"The score-function estimator is wrong because it ignores the slope of f."
It is unbiased (right on average); it is just much noisier, so it needs many more samples or baselines. It is the tool for discrete latent variables, which cannot be reparameterized.
In your forecasting model every latent parameter (trend, changepoint slopes, seasonality and holiday coefficients, regressor weights, noise scale) is continuous, and a full-rank or low-rank Gaussian guide draws them as $m + L\varepsilon$ (plus transforms for positive ones). So each SVI step can use the low-noise reparameterized gradient from Trace_ELBO. In an A/B framework like yours the conversion rates are continuous too (Beta priors are reparameterizable); discrete choices, if any, would need enumeration instead.
$\theta = \mu + \sigma\varepsilon$, $\varepsilon \sim N(0,1)$ ⇒ $\nabla_\phi E_q[f(\theta)] = E_\varepsilon[\nabla_\phi f(\mu + \sigma\varepsilon)]$; $\partial\theta/\partial\mu = 1$, $\partial\theta/\partial\sigma = \varepsilon$.
Low noise because it uses the slope of $f$. Score function $E_q[f\,\nabla\log q]$: also unbiased, far noisier; needed for discrete θ.
Trap: Trace_ELBO needs reparameterizable latents; discrete ones → enumeration or TraceGraph_ELBO.
Quick check: with q = N(2, 1²) and ε = 0.5, what is θ, and what are ∂θ/∂μ and ∂θ/∂σ?
θ = 2 + 1 × 0.5 = 2.5. ∂θ/∂μ = 1 and ∂θ/∂σ = ε = 0.5. If $f(\theta) = \theta^2$, this one particle gives slopes $2\theta \times 1 = 5$ in μ and $2\theta \times 0.5 = 2.5$ in σ (exact: 4 and 2).
Stochastic variational inference: model, guide, gradients, optimizer core
Now put the pieces together. SVI = Stochastic Variational Inference: "variational" because we fit a distribution $q_\phi$; "stochastic" because every step uses random draws, so the ELBO and its gradient are noisy estimates. One step: draw particles from the guide, score them with the model, get the gradient through the draws (reparameterization), and let the optimizer nudge φ uphill. Then repeat a few thousand times.
It is like a hiker climbing a hill in fog with a shaky altimeter. Each reading is noisy, and some steps even go slightly downhill, but on average the steps point up, and small steps add up to the summit.
The words NumPyro uses:
- Model: your Python function with the prior and the likelihood (
numpyro.samplestatements). It defines $p(D, \theta)$. - Latent variables: the unobserved sample sites, θ (rates, slopes, coefficients). Observed sites have
obs=data. - Guide: a second function that defines $q_\phi(\theta)$ for every latent variable. An autoguide such as
AutoNormalwrites it for you. - Variational parameters φ: the guide's
numpyro.paramvalues (locations, scales, covariance factors). These, and only these, are trained. - ELBO estimator:
Trace_ELBO(), which computes the Monte Carlo loss $-\widehat{\text{ELBO}}$. - Optimizer: usually Adam, which gives each parameter its own step size based on the recent size of its gradients.
Three ways to say it:
- Picture: an orange ellipse crawling onto the green posterior while a jittery ELBO curve climbs toward a ceiling it can never pass.
- Numbers: one hand-made step moves $q$ from $N(0, 1^2)$ to $N(0.65, 0.80^2)$ and the exact ELBO from $-7.97$ to $-5.74$ (the best possible is $-5.17$).
- Slogan: sample, score, differentiate, step, repeat.
One SVI step by hand on the running example. Guide $q = N(\mu, \sigma^2)$ with φ = (μ, ℓ), where $\sigma = e^{\ell}$ (the log keeps σ positive, as autoguides do). Start at the prior: μ = 0, ℓ = 0 (σ = 1). One particle; plain gradient ascent with step size 0.1 so you can check the arithmetic (NumPyro would use Adam).
- Sample: ε = −0.5, so θ = 0 + 1 × (−0.5) = −0.5.
- Score: $\log p(D\mid\theta) = -8.466$, $\log p(\theta) = -1.044$, so $\log p(D,\theta) = -9.510$; $\log q(\theta) = -1.044$. ELBO estimate $= -9.510 + 1.044 = -8.466$, so this step's loss is $8.466$. (The exact ELBO at this $q$ is $-7.966$: one particle is noisy.)
- Differentiate (reparameterization; $\frac{d}{d\theta}\log p(D,\theta) = 4 - 5\theta = 6.5$): slope in μ $= 6.5 \times 1 = 6.5$; slope in ℓ $= 6.5 \times \sigma\varepsilon + 1 = 6.5 \times (-0.5) + 1 = -2.25$ (the $+1$ is the entropy's slope, since $H[q] = \ell + $ const).
- Step uphill: μ ← 0 + 0.1 × 6.5 = 0.65; ℓ ← 0 + 0.1 × (−2.25) = −0.225, so σ = $e^{-0.225}$ = 0.80.
- Check: the exact ELBO went from −7.966 to −5.741. The exact slopes at the start were 4 and −4; our one-particle slopes (6.5 and −2.25) had the right signs but the wrong sizes: noise.
- Repeat. After a few hundred such steps, $q$ wobbles around $N(0.8, 0.447^2)$, the posterior.
The SVI algorithm. Given a model $p(D,\theta)$ and a guide $q_\phi(\theta)$:
- Initialize φ (NumPyro's autoguides start the locations at random values in $(-2, 2)$ on the unconstrained scale and every scale at 0.1).
- Draw $\varepsilon_1,\dots,\varepsilon_S$ and form particles $\theta_s = g_\phi(\varepsilon_s)$.
- Compute $\widehat{\text{ELBO}}(\phi) = \frac1S\sum_s[\log p(D,\theta_s) - \log q_\phi(\theta_s)]$ (minibatch-scaled if needed, next section).
- Compute its gradient $\hat g = \nabla_\phi\widehat{\text{ELBO}}$ by automatic differentiation.
- Update φ with the optimizer (Adam on the loss $-\widehat{\text{ELBO}}$).
- Repeat until the ELBO stops improving.
guide = AutoNormal(model) # q_phi: independent Normals (unconstrained space)
svi = SVI(model, guide, numpyro.optim.Adam(0.01), Trace_ELBO())
state = svi.init(jax.random.PRNGKey(0), data)
for t in range(num_steps):
state, loss = svi.update(state, data) # loss = -ELBO estimate for this step
params = svi.get_params(state) # the trained phi
# or, in one call (jit-compiled loop): result = svi.run(key, num_steps, data); result.losses
- Because $\hat g$ is unbiased, this is stochastic gradient ascent; with a step size that decreases suitably over time it converges (to a local optimum of the ELBO). With a constant step size the parameters keep wobbling around the optimum.
- The loss NumPyro reports is $-\widehat{\text{ELBO}}$. It is never meaningful on its own scale (it includes $-\log p(D)$); only its trend is.
- The randomness comes from an explicit PRNG key, so a run is reproducible when you fix the key (Chapter 6.16).
Why do we need it?
It is the practical way to fit a Bayesian model whose posterior has no formula and whose size makes MCMC slow: every step is one run of the model plus one backward pass, it works for any differentiable model, and it scales with minibatches.
Where is it used?
NumPyro's and Pyro's SVI class, your forecasting model's custom training loop, the A/B framework's SVI inference, VAEs and other deep generative models, large hierarchical models, Bayesian neural networks.
How is it used?
Write the model, pick an autoguide, pick Adam with a modest learning rate, run a few thousand steps, watch a smoothed loss, keep the best parameters, then call guide.sample_posterior or Predictive(model, guide=guide, params=params) to get draws for intervals and forecasts.
"Adam adapts its step sizes, so the learning rate does not matter."
Adam's steps are roughly the learning rate in size. Too large: the parameters jump around the optimum (or blow up); too small: thousands of wasted steps. With a constant rate the final parameters keep wobbling; a decaying rate, or keeping the best state (Chapter 6.14), calms that.
"After enough steps SVI gives the exact posterior."
At best it reaches the best member of the guide family (Chapter 6.11), and only a local optimum of the ELBO. The mean-field guide on the correlated target in the widget never gets closer than KL = 0.83.
"A good run should drive the loss to 0."
The loss is $-\text{ELBO} \ge -\log p(D)$, which can be any number. It should level off, not reach a particular value.
"svi.update returns the ELBO, so improvement means it goes up."
"It returns (svi_state, loss) with loss $= -\widehat{\text{ELBO}}$, the negative of a noisy one-step estimate. Improvement means the loss goes down (ELBO up), judged over several evaluations, not one step."
Model answer: "SVI maximizes the ELBO, $E_q[\log p(D,\theta) - \log q(\theta)]$, which equals $\log p(D) - KL(q\|p(\theta\mid D))$. Each step draws particles from the guide with the reparameterization trick, estimates the ELBO, differentiates it with JAX, and takes an Adam step. NumPyro's svi.update returns the loss, the negative ELBO estimate, so I monitor a smoothed loss and keep the parameters with the best value."
Your forecasting model's custom training loop is this algorithm with the update JIT-compiled: svi.init once, then repeated calls to a jitted svi.update, each returning (state, loss) with loss = −ELBO. On top of it you added relative-ELBO early stopping with patience and best-state checkpointing. Those choices answer exactly the problems this section shows: the trace is noisy, it flattens below a ceiling you do not know, and the last iterate is not necessarily the best one. Chapter 6.14 takes the loop apart piece by piece; Chapter 6.13 explains why the loop picks a full-rank or a low-rank guide.
SVI step: draw $\varepsilon$ → $\theta = g_\phi(\varepsilon)$ → $\widehat{\text{ELBO}}$ → $\nabla_\phi$ by autodiff → Adam step on loss $= -\widehat{\text{ELBO}}$.
Model = $p(D,\theta)$; guide = $q_\phi$; latents = unobserved sites; φ = guide params (only these are trained).
Trap: svi.update returns the loss (−ELBO); improvement = loss down, judged on a smoothed trace.
Quick check: the last 5 losses printed are 1 204.1, 1 199.8, 1 206.3, 1 201.0, 1 203.5. Has the run improved over these steps?
You cannot tell from these numbers: they go up and down by ±3 around about 1 203, which looks like one-particle noise. Compare averages over longer windows (for example the mean of the last 200 losses vs the 200 before), or evaluate the ELBO with many particles at two saved states.
Minibatches: scale the data part by $N/B$
Every ELBO estimate runs the model on all $N$ data points: one log-likelihood term per point. With millions of users or events that is slow. So we do what a city does when it estimates total water use: measure 100 random households and multiply by (number of households ÷ 100). In SVI: pick a random minibatch of $B$ points, add up their log-likelihood terms, and multiply by $N/B$. On average that equals the full sum (it is unbiased); it just adds a second source of noise.
Forget the multiplier, and the model believes it has seen only $B$ data points. The prior then counts $N/B$ times too much compared with the data, so the posterior comes out far too wide and pulled toward the prior. This is one of the most common silent bugs in hand-written SVI code.
Three ways to say it:
- Picture: look at a small random slice of the data, then blow it up to full size; leave the prior and the guide alone.
- Numbers: $N = 10\,000$, $B = 100$: multiply by 100. Without it, a posterior sd of 0.02 becomes 0.2.
- Slogan: scale the data part by $N/B$, never the prior.
A big A/B-style dataset. $N = 10\,000$ users, each with a measurement $y_i \sim N(\theta, 2^2)$; prior $\theta \sim N(0, 10^2)$; minibatch $B = 100$.
- Scale factor: $N/B = 10\,000/100 = 100$.
- At one particle θ, suppose the 100 sampled users have log-likelihood terms adding up to $-212.4$. Estimated full-data sum: $100 \times (-212.4) = -21\,240$.
- Add the prior and guide terms once, unscaled: $\widehat{\text{ELBO}} = -21\,240 + \log p(\theta) - \log q(\theta)$.
- What the correct ELBO converges to: the full-data posterior, with sd $\approx \sigma/\sqrt N = 2/\sqrt{10\,000} = 0.02$ (the prior is negligible here).
- If the factor 100 is forgotten, the objective is the ELBO of a dataset of only 100 users: sd $\approx 2/\sqrt{100} = 0.2$, ten times too wide.
If the data are independent given θ, the ELBO's data part is a sum over points, and a uniformly random minibatch $\mathcal{B}$ of size $B$ gives the unbiased estimate
$$\widehat{\text{ELBO}} = \frac{N}{B}\sum_{i\in\mathcal{B}} \log p(y_i\mid\theta_s) + \log p(\theta_s) - \log q_\phi(\theta_s), \qquad \theta_s \sim q_\phi .$$- Doubly stochastic: noise from the particle θ and noise from the choice of minibatch. Both average out; neither biases the gradient.
- Only terms that come once per data point are scaled. Global terms (priors on shared parameters, the guide of shared parameters) are not. Local latent variables that belong to one data point (inside the plate) are scaled together with their data.
- In NumPyro:
with numpyro.plate("data", N, subsample_size=B) as idx: numpyro.sample("y", dist.Normal(mu, sigma), obs=y[idx]). The plate's size must be the full $N$; NumPyro then multiplies the log-likelihood by $N/B$ for you. (numpyro.handlers.scaledoes the same by hand.) - Requirement: each term must depend on θ and its own data point only. Models that link neighbouring points (autoregressive terms, state-space models) cannot be split this way without extra care.
Why do we need it?
The cost of one SVI step grows with $N$. Minibatches make it grow with $B$ instead, so models with millions of observations can be trained, at the price of extra gradient noise.
Where is it used?
Large-scale SVI in NumPyro and Pyro (plate(..., subsample_size=B)), the original stochastic variational inference for topic models, VAEs and Bayesian deep learning, and experimentation platforms that fit user-level models.
How is it used?
Declare the plate with the full size and a subsample size, index the data with the returned indices, use a decaying learning rate to calm the extra noise, and check on a small dataset that the minibatch fit matches the full-batch fit.
"Scale the whole ELBO by N/B."
Only the per-data-point terms. The prior on shared parameters and the guide's terms for them appear once in the full ELBO, so they appear once in the estimate.
"numpyro.plate("data", B) on a minibatch is fine."
The plate's size must be the full $N$, with subsample_size=B. A plate of size $B$ tells NumPyro the dataset has $B$ points, and no scaling happens.
"Minibatch noise is the same as particle noise, so one more particle fixes it."
They are separate sources. More particles reduce the θ-noise; a larger batch reduces the data-sampling noise. Both are reduced by a decaying learning rate.
Whether you need minibatches depends on $N$. In an A/B framework like yours, a Beta-Binomial model works from aggregated counts (successes and trials per group), so there is nothing to minibatch; a user-level model with millions of rows would need it. A daily forecasting series of a few hundred or thousand days is usually fitted full-batch. If you ever minibatch days, the trick is valid only because, given the parameters, each day's likelihood term depends on that day alone ($y_t = g(t) + s(t) + h(t) + X_t\beta + \epsilon_t$ with independent noise); an autoregressive error term would break that.
Minibatch ELBO: $\frac{N}{B}\sum_{i\in\mathcal B}\log p(y_i\mid\theta) + \log p(\theta) - \log q(\theta)$: unbiased, doubly stochastic.
NumPyro: plate("data", N, subsample_size=B) (full N!) applies the N/B scale automatically.
Trap: forgetting N/B = pretending you have B points: posterior $\sqrt{N/B}$ times too wide, pulled to the prior.
Quick check: N = 50 000, B = 500. A minibatch's log-likelihood sum is −1 830. What does it contribute to the ELBO estimate, and how wrong would the posterior sd be without scaling?
Scale $N/B = 100$, so the data part is $100 \times (-1830) = -183\,000$. Without scaling the model acts as if it had 500 points instead of 50 000; for a well-identified parameter the posterior sd scales like $1/\sqrt{n}$, so it would be about $\sqrt{100} = 10$ times too wide.
Recap, cheat sheet and practice
- $\log p(D) = \text{ELBO}(q) + KL(q\,\|\,p(\theta\mid D))$ for every $q$: the ELBO is a lower bound on the log evidence, and because $\log p(D)$ is fixed, maximizing the ELBO is minimizing the KL. Derived by inserting $q/q$ inside $E_q[\log p(D)]$, or by Jensen's inequality.
- Two readings: $E_q[\log p(D,\theta)] + H[q]$ (fit + entropy) and $E_q[\log p(D\mid\theta)] - KL(q\|p(\theta))$ (data fit − distance from the prior). The ELBO contains the KL to the prior; it falls short of $\log p(D)$ by the KL to the posterior.
- The ELBO is estimated by Monte Carlo with $S$ particles: unbiased, sd ∝ $1/\sqrt S$, zero noise only at the exact posterior. NumPyro's default is one particle, so the loss trace is noisy.
- The reparameterization trick $\theta = \mu + \sigma\varepsilon$ makes the draw a smooth function of φ, giving low-noise gradients; the score-function estimator works for discrete θ but is much noisier.
- SVI: sample → score → differentiate → Adam step, repeated. Model = $p(D,\theta)$, guide = $q_\phi$, latent variables = unobserved sites, variational parameters = the guide's params.
svi.updatereturns(state, loss)with loss = −ELBO. - Minibatches: scale the per-point log-likelihood by $N/B$ (NumPyro:
plate("data", N, subsample_size=B)); never the prior. Forgetting it acts as if you had only $B$ points.
Cheat sheet
| Idea | Formula | Plain words / remember |
|---|---|---|
| The identity | $\log p(D) = \text{ELBO} + KL(q\|p(\theta\mid D))$ | fixed total = computable part + gap |
| ELBO | $E_q[\log p(D,\theta) - \log q(\theta)]$ | ≤ log p(D); = only if $q$ = posterior |
| Reading 1 | $E_q[\log p(D,\theta)] + H[q]$ | fit + entropy; entropy stops collapse to a point |
| Reading 2 | $E_q[\log p(D\mid\theta)] - KL(q\|p(\theta))$ | data fit − complexity (VAE loss) |
| Monte Carlo ELBO | $\frac1S\sum_s[\log p(D,\theta_s) - \log q(\theta_s)]$ | unbiased; sd ∝ $1/\sqrt S$; num_particles=S |
| Reparameterization | $\theta = \mu + \sigma\varepsilon$; $\nabla_\phi E_q[f] = E_\varepsilon[\nabla_\phi f(\theta)]$ | $\partial\theta/\partial\mu = 1$, $\partial\theta/\partial\sigma = \varepsilon$ |
| Score function | $E_q[f(\theta)\nabla_\phi\log q_\phi(\theta)]$ | works for discrete θ; much noisier; baselines help |
| SVI step | state, loss = svi.update(state, data) | loss = −ELBO estimate; improvement = loss down (smoothed) |
| Minibatch | $\frac NB\sum_{i\in\mathcal B}\log p(y_i\mid\theta) + \log p(\theta) - \log q(\theta)$ | plate size = full N; scale data terms only |
import numpy as np
import jax, jax.numpy as jnp
import numpyro, numpyro.distributions as dist
from numpyro.infer import SVI, Trace_ELBO
from numpyro.infer.autoguide import AutoNormal
from numpyro.optim import Adam
from scipy import special
# 1) Model: orders per day y_t ~ Poisson(lam), prior lam ~ Gamma(2, rate 1)
y = jnp.array([3.0, 1.0, 4.0, 2.0, 5.0])
def model(y):
lam = numpyro.sample("lam", dist.Gamma(2.0, 1.0)) # Gamma(concentration, rate)
with numpyro.plate("days", y.shape[0]):
numpyro.sample("y", dist.Poisson(lam), obs=y)
guide = AutoNormal(model) # q: Normal on log(lam), mapped back with exp
optim = Adam(lambda t: 0.02 / (1 + t / 500)) # a decaying step size calms the final wobble
svi = SVI(model, guide, optim, Trace_ELBO()) # Trace_ELBO(num_particles=1) by default
res = svi.run(jax.random.PRNGKey(0), 4000, y, progress_bar=False)
loc, scale = float(res.params["lam_auto_loc"]), float(res.params["lam_auto_scale"])
a, b = 2 + float(y.sum()), 1 + len(y) # exact posterior: Gamma(17, 6)
print(round(loc, 3), round(scale, 3)) # 1.016 0.251 the fitted guide (loc, scale of log lam)
print(round(np.log(a / b) - 1 / (2 * a), 3), round(1 / np.sqrt(a), 3)) # 1.012 0.243 the best possible Normal (Chapter 6.11)
print(np.round(np.asarray(res.losses[:3]), 2), np.round(np.asarray(res.losses[-3:]), 2))
# [15.85 15.88 16.05] [10.29 10.24 10.27] loss = -ELBO estimate: it falls, then wobbles
# 2) The ELBO is a lower bound on log p(D); the gap is KL(q || posterior)
log_pD = (-float(jnp.sum(jax.scipy.special.gammaln(y + 1))) # Gamma-Poisson evidence in closed form
- special.gammaln(2.0) + 2.0 * np.log(1.0) + special.gammaln(a) - a * np.log(b))
neg_elbo = Trace_ELBO(num_particles=100_000).loss(jax.random.PRNGKey(1), res.params, model, guide, y)
print(round(log_pD, 3), round(-float(neg_elbo), 3), round(log_pD + float(neg_elbo), 4)) # -10.239 -10.246 0.0071
# log p(D), the ELBO (below it), and the gap = KL(q || posterior)
# 3) Monte Carlo noise of the ELBO estimate falls like 1/sqrt(S)
for S in [1, 10, 100]:
elbo_S = Trace_ELBO(num_particles=S)
one = jax.jit(jax.vmap(lambda k: elbo_S.loss(k, res.params, model, guide, y)))
est = one(jax.random.split(jax.random.PRNGKey(2), 20_000)) # 20 000 independent loss evaluations
print(S, round(float(est.std()), 3), round(float(est.std()) * np.sqrt(S), 3)) # 1 0.124 0.124 | 10 0.04 0.126 | 100 0.013 0.126
# 4) The reparameterization trick: gradient of E_q[theta^2] through theta = mu + sigma * eps
eps = jax.random.normal(jax.random.PRNGKey(3), (200_000,))
f = lambda mu, sig: jnp.mean((mu + sig * eps) ** 2)
print([round(float(v), 3) for v in jax.grad(f, argnums=(0, 1))(1.0, 0.5)]) # [2.002, 1.012] exact: (2 mu, 2 sigma) = (2, 1)
# 5) Minibatches: plate(size=N, subsample_size=B) multiplies the log-likelihood by N/B
N, B = 2_000, 100
data = 3.0 + 2.0 * jax.random.normal(jax.random.PRNGKey(7), (N,))
def model_mb(batch, size): # size = N is right; size = B is the classic bug
mu = numpyro.sample("mu", dist.Normal(0.0, 10.0))
with numpyro.plate("obs", size, subsample_size=batch.shape[0]):
numpyro.sample("x", dist.Normal(mu, 2.0), obs=batch)
rng = np.random.default_rng(0)
for size in [N, B]:
svi_mb = SVI(model_mb, AutoNormal(model_mb), Adam(lambda t: 0.05 / (1 + t / 300)), Trace_ELBO())
state = svi_mb.init(jax.random.PRNGKey(5), data[:B], size)
update = jax.jit(svi_mb.update, static_argnums=2) # a jit-compiled update, as in a custom loop
for t in range(3000):
state, loss = update(state, data[rng.choice(N, B, replace=False)], size) # returns (state, loss = -ELBO)
print(size, round(float(svi_mb.get_params(state)["mu_auto_scale"]), 3)) # 2000 0.046 (right) | 100 0.197 (forgot N/B)
print(round(2.0 / np.sqrt(N), 3), round(2.0 / np.sqrt(B), 3)) # 0.045 0.2 exact posterior sd with N points vs "as if B points"
1. For some guide $q$, the ELBO is −52.3, and for this model and data $\log p(D) = -50.0$. What is $KL(q\|p(\theta\mid D))$?
2. What does NumPyro's svi.update(state, data) return?
(svi_state, loss), with loss $= -\widehat{\text{ELBO}}$ from that step's particles (one by default). Lower loss = higher ELBO; single values are noisy.3. Which expression equals the ELBO?
4. You go from 4 to 16 particles per ELBO estimate. The standard deviation of the estimate…
5. What does the reparameterization trick do?
6. $N = 1\,000\,000$ observations, minibatches of $B = 1\,000$, and the code forgets the $N/B$ factor. What happens to the fitted posterior?
Practice problems
A. Three candidates (5%, 10%, 15%), prior × likelihood $(0.0199, 0.0634, 0.0809)$ with logs $(-3.919, -2.759, -2.514)$, $\log p(D) = -1.807$. For $q = (0.2, 0.4, 0.4)$ compute the ELBO and the KL to the posterior.
- $E_q[\log p(D,\theta)] = 0.2(-3.919) + 0.4(-2.759) + 0.4(-2.514) = -0.784 - 1.104 - 1.006 = -2.893$.
- $E_q[\log q] = 0.2\log 0.2 + 0.4\log 0.4 + 0.4\log 0.4 = -0.322 - 0.367 - 0.367 = -1.055$.
- $\text{ELBO} = -2.893 + 1.055 = -1.838$.
- $KL = \log p(D) - \text{ELBO} = -1.807 + 1.838 = 0.031$ (direct computation: 0.0312). This $q$ is close to the posterior $(0.121, 0.386, 0.493)$, so the gap is small.
B. Running example (prior $N(0,1)$, data $0.5, 1.5, 0.8, 1.2$, each $N(\theta,1)$), guide $q = N(1.0, 0.2^2)$. Compute the ELBO with reading 2 and the KL to the posterior.
- Data terms $-0.919 - ((y-1)^2 + 0.04)/2$: $-1.064$, $-1.064$, $-0.959$, $-0.959$; sum $E_q[\log p(D\mid\theta)] = -4.046$.
- $KL(N(1, 0.2^2)\|N(0,1)) = \log 5 + (0.04 + 1)/2 - 0.5 = 1.609 + 0.520 - 0.5 = 1.629$.
- $\text{ELBO} = -4.046 - 1.629 = -5.675$.
- $KL(q\|\text{posterior}) = \log p(D) - \text{ELBO} = -5.170 + 5.675 = 0.505$ (the guide is too far right and too narrow: posterior is $N(0.8, 0.447^2)$).
C. In your training loop, a one-particle loss has sd about 3 near convergence, and you want to detect improvements of about 0.1 from a single evaluation. How many particles would that need? What would you do instead?
To get sd 0.1 from sd 3 you need $(3/0.1)^2 = 900$ particles per evaluation, which makes every step about 900 times as expensive. Instead: keep training with one particle, evaluate the ELBO less often but with many particles (or average the loss over a window of steps), and use a stopping rule with a tolerance and patience rather than single comparisons (Chapter 6.14).
D. For $q = N(\mu, \sigma^2)$ and $f(\theta) = e^\theta$, use the reparameterization trick to find $\partial E_q[f]/\partial\mu$ and $\partial E_q[f]/\partial\sigma$, and check against the formula $E_q[e^\theta] = e^{\mu + \sigma^2/2}$.
- Write $\theta = \mu + \sigma\varepsilon$: $E_q[e^\theta] = E_\varepsilon[e^{\mu + \sigma\varepsilon}]$.
- μ-slope: $E_\varepsilon[e^{\mu+\sigma\varepsilon} \cdot 1] = e^{\mu+\sigma^2/2}$.
- σ-slope: $E_\varepsilon[e^{\mu+\sigma\varepsilon}\,\varepsilon] = e^{\mu}E[\varepsilon e^{\sigma\varepsilon}] = e^\mu \cdot \sigma e^{\sigma^2/2} = \sigma e^{\mu+\sigma^2/2}$ (using $E[\varepsilon e^{\sigma\varepsilon}] = \frac{d}{d\sigma}E[e^{\sigma\varepsilon}] = \frac{d}{d\sigma}e^{\sigma^2/2}$).
- Check: $\frac{\partial}{\partial\mu}e^{\mu+\sigma^2/2} = e^{\mu+\sigma^2/2}$ and $\frac{\partial}{\partial\sigma}e^{\mu+\sigma^2/2} = \sigma e^{\mu+\sigma^2/2}$. ✓ (This is the term that pulls the log-rate guide of Chapter 6.11.)
E. (Interview) "In two minutes: what is the ELBO, why does maximizing it work, and what does your training loop actually optimize?"
"The posterior needs the evidence $p(D)$, which is an intractable integral. For any approximation $q$ we have $\log p(D) = \text{ELBO}(q) + KL(q\|p(\theta\mid D))$, where $\text{ELBO} = E_q[\log p(D,\theta) - \log q(\theta)]$ only needs the joint and draws from $q$. Since $\log p(D)$ is fixed and the KL is non-negative, the ELBO is a lower bound, and pushing it up pushes the KL down. Equivalently it is the expected log-likelihood minus the KL from $q$ to the prior: fit the data, but pay for moving away from the prior. In SVI we estimate it with a few reparameterized draws, $\theta = \mu + \sigma\varepsilon$, differentiate with JAX and step with Adam. NumPyro's svi.update returns the loss, the negative ELBO estimate, so my loop minimizes that loss; because each value is noisy, it stops on a relative improvement with patience and keeps the best state. The result is the best guide in the family, which can still under-state uncertainty."
F. $N = 1\,000\,000$ events, $B = 1\,000$. For one particle θ: the minibatch's log-likelihood sum is $-2\,310.5$, $\log p(\theta) = -4.2$ and $\log q(\theta) = -1.3$. Compute the ELBO estimate, and the estimate a buggy version without scaling would report.
- Scale $N/B = 1\,000$; data part $1\,000 \times (-2\,310.5) = -2\,310\,500$.
- Add prior and guide terms once: $-2\,310\,500 + (-4.2) - (-1.3) = -2\,310\,502.9$.
- Buggy (no scaling): $-2\,310.5 - 4.2 + 1.3 = -2\,313.4$. Here the prior term is a much larger share of the objective, which is exactly why the buggy fit stays near the prior.
Variational guides: mean-field, full-rank, low-rank
SVI needs a guide: the family of shapes it is allowed to use for the posterior. NumPyro offers three Gaussian guides that differ in one thing only, how much covariance they can store. Mean-field stores none, full-rank stores all of it, and low-rank stores a few important directions. This chapter shows exactly what each one learns (with NumPyro's real parameter counts), what each one gets wrong, how memory and compute grow with model size, and how to defend an automatic "full-rank when small, low-rank when big" choice in an interview.
- Say what a guide is in NumPyro: one Gaussian over the flattened, unconstrained latent vector of length $d$, and count $d$ for your own models
- Explain the mean-field guide, derive why its variances shrink to $1/\Lambda_{ii}$ when the posterior is correlated, and show when sums and differences come out too narrow or too wide
- Explain the full-rank guide $N(\boldsymbol\mu, LL^\top)$: why it learns a Cholesky factor $L$, why that costs $d + d(d+1)/2$ numbers, and where $O(d^2)$ and $O(d^3)$ really come from
- Explain the low-rank guide $WW^\top + D$: the $d(r+2)$ numbers, what $r$ directions can capture, what they cannot (chains, many separate pairs), and how to pick $r$
- Compute parameter counts, memory and per-step work for each guide and any $d$, $r$ (Module 68)
- Choose the family, write an illustrative size-based selection rule, and validate the choice
What we need from earlier chapters: variational inference and the KL divergence (Chapter 6.11); the ELBO and SVI (Chapter 6.12: SVI maximizes the ELBO, a noisy Monte Carlo estimate of how well the guide matches the posterior); covariance matrices, the multivariate Normal, eigenvectors and a first look at guide structures (Chapter 5.15); the Cholesky factorization and low-rank approximation (Linear Algebra, Chapter 1.13). Notation: $d$ = number of unconstrained latent numbers; $\boldsymbol\theta\in\mathbb{R}^d$ = the latent vector; $\Sigma$ = the posterior covariance and $\Lambda = \Sigma^{-1}$ its precision matrix; $q$ = the guide; $r$ = the rank of a low-rank guide. "Learned numbers" means the scalars the optimizer actually updates.
What a guide is: one Gaussian over the unconstrained latent vector core
In SVI you never compute the posterior itself. You pick a family of simple distributions and let the optimizer find the member that is closest to the posterior. In NumPyro this family is called the guide. Think of it as a mould: the posterior is a lump of clay with some shape, and the guide is the mould you are allowed to press it into. A round mould cannot hold a long tilted shape. A mould that can tilt can.
An autoguide (NumPyro's AutoNormal, AutoMultivariateNormal, AutoLowRankMultivariateNormal…) builds this mould for you in four moves:
- It finds every latent sample site in your model (every
numpyro.samplewithoutobs=). - It moves each one to the whole real line: a positive σ becomes $\log\sigma$, a probability $p$ becomes its log-odds, a probability vector with $K$ entries becomes $K-1$ free numbers. This is the unconstrained space.
- It flattens all of them into one long vector of $d$ numbers.
- It puts a Gaussian on that vector. The three autoguides differ only in the covariance matrix of this Gaussian.
Three ways to say it:
- Picture: line up every unknown of your model in one long row of boxes, and put one bell-shaped cloud over the whole row.
- Numbers: a forecasting model with a trend, 25 changepoints, 26 Fourier coefficients, 4 holidays, 3 regressors and a noise scale has $d = 61$; the guide is a Gaussian in 61 dimensions.
- Slogan: choosing a Gaussian guide = choosing a covariance structure for a $d$-dimensional Gaussian.
Count $d$ for two models like yours.
- Forecasting model. Trend slope $k$ and offset $m$: 2 numbers.
- 25 changepoint slope adjustments $\delta_1,\dots,\delta_{25}$: 25 numbers.
- Yearly seasonality of Fourier order 10: one sine and one cosine coefficient per order, $2\times10 = 20$; weekly order 3: $2\times3 = 6$. Together 26.
- 4 holiday effects + 3 regressor coefficients: 7.
- The noise scale $\sigma \gt 0$ becomes one unconstrained number, $\log\sigma$: 1.
- Total: $d = 2 + 25 + 26 + 7 + 1 = 61$. (NumPyro agrees: the guides below were initialized on exactly this model.)
- A/B model. A Dirichlet over 4 categories: 4 shares that must add to 1, so only $4 - 1 = 3$ free numbers. A population mean $\mu$ and spread $\tau \gt 0$: 2. Six segment effects $z_1,\dots,z_6$: 6. Total $d = 3 + 2 + 6 = 11$.
Let the model have latent sites $\theta^{(1)},\dots,\theta^{(m)}$ with supports (allowed ranges) such as "all reals", "positive" or "probability vector". For each site there is a fixed invertible map $T_j$ from the real numbers onto that support. The unconstrained latent vector is
$$\mathbf{z} = \big(T_1^{-1}(\theta^{(1)}),\ \dots,\ T_m^{-1}(\theta^{(m)})\big) \in \mathbb{R}^d,$$flattened into one vector; $d$ is the latent dimension. A Gaussian guide is
$$q_\phi(\mathbf{z}) = N(\mathbf{z};\ \boldsymbol\mu,\ \Sigma_q), \qquad \theta^{(j)} = T_j(\mathbf{z}_j),$$- $\phi$ = the variational parameters: the mean $\boldsymbol\mu\in\mathbb{R}^d$ plus whatever numbers build $\Sigma_q$. These are what SVI learns (Chapter 6.12).
- The variational family is the set of all guides you can reach by changing $\phi$: diagonal $\Sigma_q$ (mean-field), full $\Sigma_q$ (full-rank), or $WW^\top + D$ (low-rank).
- The guide is Gaussian in $\mathbf{z}$. In the original units it is a transformed Gaussian: for $\sigma = e^{z}$ it is a log-normal, skewed to the right.
- $d$ counts unconstrained scalars, not sample statements: a site of shape (25,) adds 25; a $K$-category simplex adds $K - 1$.
Why do we need it?
Every cost and every limitation of a guide is a function of $d$ and of the covariance structure. Without counting $d$ correctly you cannot predict memory, speed, or which correlations your uncertainty estimates will silently drop.
Where is it used?
All NumPyro and Pyro autoguides (AutoNormal, AutoDiagonalNormal, AutoMultivariateNormal, AutoLowRankMultivariateNormal, AutoLaplaceApproximation), Stan's ADVI (mean-field and full-rank), and your forecasting model's automatic guide choice by size.
How is it used?
Write the model, then count $d$: trace the model once and add up the unconstrained sizes of the latent sites (code at the end of the chapter), or read auto_loc.shape after svi.init. Use $d$ to pick the guide, then interpret the guide's means and spreads in unconstrained space.
"$d$ is the number of numpyro.sample statements."
$d$ counts unconstrained numbers. One statement with shape (25,) adds 25; a 4-category Dirichlet adds 3; an observed site adds nothing.
"The guide says σ has mean exp(auto_loc)."
The loc lives in unconstrained space. For a positive site NumPyro uses exactly $\sigma = e^{z}$, so $e^{\text{loc}}$ is the median of σ under the guide; the mean is larger. Other supports use other maps (log-odds for probabilities, a shifted log for $\nu \gt 1$, stick-breaking for simplexes), so use guide.sample_posterior or guide.median(params) instead of transforming by hand.
"The guide's correlations are between σ, p and δ as I wrote them."
They are correlations between the unconstrained numbers ($\log\sigma$, log-odds…), after flattening.
In your forecasting model, $d$ is the sum of everything above: every changepoint on the grid adds one $\delta_j$, every Fourier order adds two coefficients per seasonality, every holiday and regressor adds one, and the likelihood adds its scale (plus ν for Student-t, or a concentration for the Negative Binomial). This is the "model size" that an automatic full-rank vs low-rank choice looks at. In an A/B framework like yours, $d$ grows with the number of segments in a hierarchical model and with $K - 1$ for each Dirichlet.
"The guide is the posterior."
"The guide is my approximation to the posterior: a parametric family (here Gaussian in unconstrained space) whose parameters SVI tunes to maximize the ELBO."
Model answer: "NumPyro's autoguides flatten all latent sites into one unconstrained vector of size $d$ and fit a Gaussian over it. The guides differ in the covariance they allow: diagonal, low-rank plus diagonal, or full. Everything is Gaussian in unconstrained space, so positive parameters come out log-normal-like in their own units."
Guide = Gaussian $N(\boldsymbol\mu, \Sigma_q)$ over the flattened unconstrained latent vector $\mathbf{z}\in\mathbb{R}^d$; the type of guide = the structure of $\Sigma_q$.
$d$ counts unconstrained numbers (shape (25,) → 25; simplex of $K$ → $K-1$; σ → 1).
Trap: guide means/correlations live in unconstrained space; a positive parameter's guide marginal is skewed.
Quick check: a model has a Dirichlet over 5 categories, 10 segment effects, a population mean and a population sd. What is $d$?
$(5 - 1) + 10 + 1 + 1 = 16$. The Dirichlet adds 4 because the 5 shares must sum to 1; the sd is positive but still adds one unconstrained number.
The mean-field Gaussian: every parameter gets its own independent bell core
The simplest guide gives every number in $\mathbf{z}$ its own bell curve (a mean and a spread) and treats them as independent: learning one tells you nothing about another. It is cheap and fast. Its picture is an ellipse that may stretch along the axes but may never tilt.
When does that hurt? Put two almost identical regressors in a model, say temperature and "feels-like" temperature. The data pin down their sum very well, but not how to split it: "+1 on temperature, −1 on feels-like" fits almost as well. The posterior is a long, thin, tilted ellipse. The mean-field guide must use an upright ellipse, and the KL direction SVI uses, $KL(q\,\|\,p)$, punishes putting guide mass where the posterior has almost none (Chapter 6.11). So it shrinks to a small blob that fits inside the thin ellipse, and each coefficient looks far more certain than it really is.
Three ways to say it:
- Picture: an upright ellipse squeezed inside a tilted one.
- Numbers: two parameters with sd 1 and correlation 0.9: mean-field reports sd 0.44 for each.
- Slogan: mean-field reports how unsure you would be about one parameter if you already knew all the others.
A correlated posterior, done by hand. Posterior $N(\mathbf{0}, \Sigma)$ with $\Sigma = \begin{bmatrix} 1 & 0.9 \\ 0.9 & 1 \end{bmatrix}$ (both sds 1, correlation 0.9).
- Determinant: $1\cdot1 - 0.9^2 = 0.19$.
- Precision $\Lambda = \Sigma^{-1} = \frac{1}{0.19}\begin{bmatrix} 1 & -0.9 \\ -0.9 & 1 \end{bmatrix} = \begin{bmatrix} 5.263 & -4.737 \\ -4.737 & 5.263 \end{bmatrix}$.
- The best mean-field guide (derived below) uses variance $1/\Lambda_{ii} = 1/5.263 = 0.19$ for each coordinate, so sd $= \sqrt{0.19} = 0.436$. The truth is 1.
- The sum $\theta_1 + \theta_2$: true variance $1 + 1 + 2(0.9) = 3.8$ (sd 1.95); mean-field $0.19 + 0.19 = 0.38$ (sd 0.62). Ten times too small.
- The difference $\theta_1 - \theta_2$: true variance $1 + 1 - 2(0.9) = 0.2$ (sd 0.45); mean-field again $0.38$ (sd 0.62). Almost twice too large.
So mean-field is not simply "too narrow everywhere". It is too narrow along the long axis of the posterior and too wide along the short axis.
The mean-field Gaussian guide is
$$q(\mathbf{z}) = \prod_{i=1}^{d} N(z_i;\ \mu_i,\ s_i^2), \qquad \Sigma_q = \text{diag}(s_1^2, \dots, s_d^2),$$with $2d$ learned numbers: $d$ means and $d$ scales (kept positive by a softplus transform). In NumPyro, AutoNormal stores them per site (delta_auto_loc, delta_auto_scale, …) and AutoDiagonalNormal stores them as two vectors auto_loc, auto_scale of length $d$. Same family, different bookkeeping.
What it converges to when the posterior is Gaussian, $p = N(\mathbf{m}, \Sigma)$ with $\Lambda = \Sigma^{-1}$:
$$KL(q\,\|\,p) = \tfrac12\Big[\textstyle\sum_i \Lambda_{ii}s_i^2 + (\boldsymbol\mu-\mathbf{m})^\top\Lambda(\boldsymbol\mu-\mathbf{m}) - d - \sum_i \log s_i^2 + \log|\Sigma|\Big].$$- The mean appears only in the quadratic term, which is smallest (zero) at $\boldsymbol\mu = \mathbf{m}$: the means are exact.
- Differentiate in $s_i^2$: $\tfrac12\big(\Lambda_{ii} - 1/s_i^2\big) = 0$, so $s_i^2 = 1/\Lambda_{ii}$.
- $1/\Lambda_{ii}$ is the variance of $\theta_i$ given all the other parameters (a standard fact about Gaussians), which is never larger than the marginal variance $\Sigma_{ii}$. Equality holds only when $\theta_i$ is uncorrelated with the rest.
For non-Gaussian posteriors there is no closed form, but the same pull toward "too narrow along correlated directions" remains.
Why do we need it?
It is the cheapest Gaussian guide ($2d$ numbers, work that grows like $d$), so it scales to huge models and is a good first fit. Knowing exactly how it fails tells you when its intervals can be trusted and when they cannot.
Where is it used?
NumPyro AutoNormal and AutoDiagonalNormal, Stan's ADVI meanfield (its default), variational autoencoders (a diagonal Gaussian encoder per data point), Bayesian neural networks with "Bayes by backprop", and topic models (classic mean-field VI for LDA).
How is it used?
Fit it first because it is fast. Then compare its marginal sds with a full-rank or low-rank fit (or NUTS on a subset). If the sds grow a lot when correlations are allowed, mean-field was hiding correlated parameters: switch guides or reparameterize the model.
"Mean-field is over-confident about everything."
For a Gaussian posterior each marginal variance comes out as $1/\Lambda_{ii} \le \Sigma_{ii}$, so marginals are never too wide. But combinations can go either way: too narrow along the posterior's long axis, too wide along its short axis.
"Strong correlations make the mean-field means wrong."
For a Gaussian posterior the means are exact; only the spreads suffer. For skewed, heavy-tailed or multi-modal posteriors the means can shift as well.
"AutoNormal and AutoDiagonalNormal are different approximations."
Both are diagonal Gaussians in unconstrained space with $2d$ learned numbers. AutoNormal keeps one loc/scale pair per site (readable names, better support for mean-field ELBO variants); AutoDiagonalNormal uses one flat vector.
In your forecasting model, several parameters trade off: the base slope $k$ against the changepoint adjustments $\delta_j$ (a steeper start with a negative adjustment can draw the same line), a holiday effect against the weekly seasonality on the same day, two related regressors against each other. A mean-field guide reports each one as more certain than it is, and gets the spread of the forecast, which adds all components, wrong in one direction or the other. In an A/B framework like yours, the population mean and the segment effects of a hierarchical model are correlated, so a mean-field fit shrinks their intervals too.
"Mean-field VI underestimates the variance."
"Mean-field VI with reverse KL gives each parameter its conditional variance $1/\Lambda_{ii}$ (for a Gaussian posterior), so marginal variances are too small when parameters are correlated. The variance of a combination of parameters can be too small or too large, depending on the direction."
Model answer: "Mean-field assumes the posterior factorizes. Minimizing $KL(q\,\|\,p)$ is zero-forcing, so the factorized guide fits inside the posterior. For a Gaussian posterior the optimum is exact means and variances $1/\Lambda_{ii}$, the variance of each parameter given all the others. With correlation 0.9 that is $0.19$ instead of $1$. Sums along the correlated direction look ten times too certain."
Mean-field: $q = \prod_i N(\mu_i, s_i^2)$, $2d$ learned numbers (AutoNormal / AutoDiagonalNormal).
Gaussian posterior → optimum $\mu_i = m_i$, $s_i^2 = 1/\Lambda_{ii}$ = variance given all others $\le \Sigma_{ii}$.
Trap: not "too narrow everywhere": too narrow along the long axis, too wide along the short axis; forecasts (sums) can be off either way.
Quick check: two parameters, both with posterior sd 2 and correlation 0.6. What sd does the optimal mean-field guide report?
For two parameters $1/\Lambda_{ii} = \Sigma_{ii}(1 - \rho^2)$, so the sd is $2\sqrt{1 - 0.36} = 2 \times 0.8 = 1.6$ for each, instead of 2.
The full-rank Gaussian: any tilt, built from a Cholesky factor core
The full-rank guide may use any covariance matrix, so its ellipse can tilt and stretch in any direction and it can store a correlation for every pair of parameters. It does not learn the covariance directly. It learns a recipe for making correlated noise out of independent noise.
The recipe is a lower-triangular matrix $L$: start with independent standard Normal numbers $\varepsilon_1, \varepsilon_2, \dots$; the first coordinate uses only $\varepsilon_1$; the second mixes $\varepsilon_1$ and $\varepsilon_2$; the third mixes $\varepsilon_1, \varepsilon_2, \varepsilon_3$; and so on. Mixing makes the coordinates move together. This $L$ is the Cholesky factor of the covariance (Linear Algebra, Chapter 1.13): $\Sigma = LL^\top$.
Three ways to say it:
- Picture: take a round cloud of dots and stretch-and-shear it into any tilted ellipse.
- Numbers: $L = \begin{bmatrix} 2 & 0 \\ 1.5 & 1 \end{bmatrix}$ turns independent noise into a cloud with correlation 0.83.
- Slogan: full-rank stores every pairwise trade-off, and pays about $d^2/2$ numbers for it.
From $L$ to the covariance, and one sample. $\boldsymbol\mu = (0, 0)$, $L = \begin{bmatrix} 2 & 0 \\ 1.5 & 1 \end{bmatrix}$.
- $\Sigma = LL^\top$: $\Sigma_{11} = 2\cdot2 = 4$; $\Sigma_{12} = 2\cdot1.5 = 3$; $\Sigma_{22} = 1.5^2 + 1^2 = 3.25$.
- sds: $\sqrt4 = 2$ and $\sqrt{3.25} = 1.803$; correlation $3/(2 \times 1.803) = 0.832$.
- Draw $\boldsymbol\varepsilon = (1, -1)$. Then $\mathbf{z} = \boldsymbol\mu + L\boldsymbol\varepsilon = (2\cdot1,\ 1.5\cdot1 + 1\cdot(-1)) = (2, 0.5)$.
- Log-determinant for the density: $\log|\Sigma| = 2(\log 2 + \log 1) = 1.386$. Only the diagonal of $L$ is needed.
- Count for $d = 10$: 10 means + the lower triangle of a $10\times10$ matrix, $10\cdot11/2 = 55$, so 65 learned numbers. NumPyro returns
auto_scale_trilas a full $10\times10$ array (100 cells), but the 45 cells above the diagonal are always 0 and are not learned.
The full-rank Gaussian guide is
$$q(\mathbf{z}) = N(\mathbf{z};\ \boldsymbol\mu,\ LL^\top), \qquad \mathbf{z} = \boldsymbol\mu + L\boldsymbol\varepsilon,\ \ \boldsymbol\varepsilon\sim N(\mathbf{0}, I),$$with $L$ lower-triangular with a positive diagonal. Its log-density is $\log q(\mathbf{z}) = -\tfrac12\|L^{-1}(\mathbf{z}-\boldsymbol\mu)\|^2 - \sum_i \log L_{ii} - \tfrac d2\log 2\pi$.
- Learned numbers: $d + \frac{d(d+1)}{2}$.
- NumPyro (
AutoMultivariateNormal):auto_locof shape $(d,)$ andauto_scale_trilof shape $(d, d)$. The optimizer updates $d(d+1)/2$ unconstrained numbers that are mapped to $L = \text{diag}(\mathbf{s})\,L_{\text{unit}}$, where $L_{\text{unit}}$ is lower-triangular with ones on the diagonal ($d(d-1)/2$ free entries) and $\mathbf{s} \gt 0$ comes through a softplus ($d$ entries). It starts at $L = 0.1\,I$. - Why learn $L$ and not $\Sigma$: every such $L$ gives a valid covariance automatically (no "is it positive definite?" check); a sample costs one triangular matrix-vector product, about $d^2/2$ multiply-adds; $\log|\Sigma| = 2\sum_i\log L_{ii}$ costs $O(d)$; the density needs one triangular solve, again about $d^2/2$.
- Cost: $O(d^2)$ memory and $O(d^2)$ work per sample in the training loop. Work of order $d^3$ appears only when a dense $d\times d$ matrix must be factorized or multiplied out: computing a Cholesky factor of a given covariance (about $d^3/3$ operations), inverting it, forming $\Sigma = LL^\top$ explicitly, or a Laplace approximation from a Hessian.
- It is still a Gaussian in unconstrained space: no skewness, heavy tails, several modes or funnels.
Why do we need it?
When parameters trade off, only a guide that stores correlations gives honest intervals for each parameter and for combinations such as forecasts. Full-rank stores all of them, so it is the gold standard among Gaussian guides when $d$ is small enough to afford.
Where is it used?
NumPyro AutoMultivariateNormal, Stan's ADVI fullrank, Pyro AutoMultivariateNormal, the Laplace approximation (a full Gaussian at the mode), Gaussian process inference with small inducing sets, and your forecasting model whenever its size allows.
How is it used?
Use it when $d$ is up to a few hundred. Read the covariance as L @ L.T from params["auto_scale_tril"], or draw samples with guide.sample_posterior. Compare its marginal sds with mean-field to see how much the correlations mattered.
"A full-rank guide gives the exact posterior."
It is still a Gaussian in unconstrained space, fitted with a noisy optimizer. It removes the "no correlations" restriction, nothing more: skewness, heavy tails, several modes and funnels stay out of reach.
"params['auto_scale_tril'] is the covariance matrix."
It is $L$. The covariance is L @ L.T, and the marginal sds are the row norms of $L$ (jnp.linalg.norm(L, axis=-1)), not its diagonal.
"Full-rank SVI does an $O(d^3)$ Cholesky decomposition every step."
NumPyro learns $L$ directly, so a training step costs about $d^2$ per sample (a triangular product and a triangular solve). $O(d^3)$ is the price of factorizing a dense covariance, which the loop avoids. Memory still grows like $d^2$, and that is usually what stops you first.
For a forecasting model of the size of the chapter's example ($d = 61$; count your own $d$), a full-rank guide learns $61 + 61\cdot62/2 = 61 + 1\,891 = 1\,952$ numbers: cheap. In exchange it captures every trade-off at once: slope vs changepoint adjustments, holiday vs weekday, overlapping regressors. That is why full-rank is the natural choice while the model is small, and why the automatic switch only moves away from it when $d$ (and so $d^2$) gets large.
"Full-rank is $O(d^3)$ per step."
"Full-rank stores $O(d^2)$ numbers. With a learned Cholesky factor each step costs $O(d^2)$ per sample; $O(d^3)$ appears when a dense covariance must be factorized or inverted."
Model answer: "The full-rank guide is $N(\boldsymbol\mu, LL^\top)$ with $L$ lower-triangular. Parameterizing by $L$ guarantees a valid covariance, makes sampling a matrix-vector product and the log-determinant a sum over the diagonal. It learns $d + d(d+1)/2$ numbers, so memory grows quadratically. At $d = 1\,000$ that is about half a million numbers; at $d = 20\,000$ about 200 million, which is where it stops being practical."
Full-rank: $q = N(\boldsymbol\mu, LL^\top)$, $\mathbf{z} = \boldsymbol\mu + L\boldsymbol\varepsilon$; learned numbers $d + d(d+1)/2$ (auto_loc, auto_scale_tril).
Per step $O(d^2)$ work and memory; $O(d^3)$ only to factorize/invert a dense matrix.
Trap: auto_scale_tril is $L$, not Σ; full-rank ≠ exact.
Quick check: $d = 50$. How many numbers does AutoMultivariateNormal learn, and how many array cells does svi.get_params return?
Learned: $50 + 50\cdot51/2 = 50 + 1\,275 = 1\,325$. Returned: auto_loc (50) + auto_scale_tril as a $50\times50$ array (2 500) = 2 550 cells, of which the 1 225 above the diagonal are zeros.
The low-rank Gaussian: a few shared directions plus private noise core
Think of the daily sales of 30 products in a shop. They rise and fall together because of a few shared causes (the weather, payday, a promotion), and on top of that each product has its own random wobble. You do not need $30\times30$ numbers to describe how they co-move: a few "shared drivers" plus one private spread per product will do.
The low-rank guide makes exactly this bet about the posterior. It keeps $r$ directions in which many parameters move together (the columns of a $d\times r$ matrix $W$), plus a private variance for every parameter (a diagonal matrix $D$). With $r$ much smaller than $d$ it costs about $d\cdot r$ numbers instead of $d^2/2$. "Rank" is the number of independent directions in $WW^\top$.
It works when the posterior's correlations come from a few big trade-offs. It struggles when they come from many separate trade-offs (each needs its own direction) or from long chains where every parameter is tied to its neighbours.
Three ways to say it:
- Picture: an upright blob (the diagonal) with a few tilted sticks through it (the columns of $W$).
- Numbers: $d = 1\,000$, $r = 10$: $12\,000$ learned numbers instead of $501\,500$.
- Slogan: low-rank = mean-field + the $r$ most important correlation directions.
Rank 1 in three dimensions. $W = (2, 1, -1)^\top$ (one direction, $r = 1$) and $D = \text{diag}(1, 1, 1)$.
- $WW^\top$: multiply every pair of entries: $\begin{bmatrix} 4 & 2 & -2 \\ 2 & 1 & -1 \\ -2 & -1 & 1 \end{bmatrix}$.
- Add $D$ to the diagonal: $\Sigma_q = \begin{bmatrix} 5 & 2 & -2 \\ 2 & 2 & -1 \\ -2 & -1 & 2 \end{bmatrix}$.
- Correlations: $\rho_{12} = 2/\sqrt{5\cdot2} = 0.632$, $\rho_{13} = -2/\sqrt{10} = -0.632$, $\rho_{23} = -1/\sqrt{2\cdot2} = -0.5$. One direction created three correlations at once.
- Learned numbers: mean 3 + one column of $W$ 3 + diagonal 3 $= 9 = d(r+2) = 3\cdot3$.
- What rank 1 cannot do: every covariance it makes has $\Sigma_{12}\Sigma_{13}\Sigma_{23} = (w_1w_2)(w_1w_3)(w_2w_3) = (w_1w_2w_3)^2 \ge 0$. So a posterior where all three correlations are $-0.4$ (product negative) is out of reach for rank 1; it needs $r = 2$.
The low-rank-plus-diagonal Gaussian guide is
$$q(\mathbf{z}) = N\big(\mathbf{z};\ \boldsymbol\mu,\ WW^\top + D\big), \qquad \mathbf{z} = \boldsymbol\mu + W\boldsymbol\varepsilon_r + D^{1/2}\boldsymbol\varepsilon_d,$$with $W\in\mathbb{R}^{d\times r}$ (the covariance factor or loadings), $D$ diagonal and positive, $\boldsymbol\varepsilon_r\sim N(\mathbf 0, I_r)$ and $\boldsymbol\varepsilon_d\sim N(\mathbf 0, I_d)$ independent.
- NumPyro (
AutoLowRankMultivariateNormal(model, rank=r)):auto_loc$(d,)$,auto_cov_factor$(d, r)$,auto_scale$(d,)$, so $d(r+2)$ learned numbers. It builds the covariance as $\text{diag}(\mathbf{s})\,(FF^\top + I)\,\text{diag}(\mathbf{s})$, i.e. $W = \text{diag}(\mathbf{s})F$ and $D = \text{diag}(\mathbf{s}^2)$. Withrank=Noneit uses $r = \text{round}(\sqrt d)$. The factor starts at zero, so training starts with no correlations and learns them. - Cost: a sample costs $O(dr)$. The density uses the Woodbury identity and the matrix determinant lemma, which only need an $r\times r$ matrix $I + W^\top D^{-1}W$ (the "capacitance" matrix): $O(dr^2 + r^3)$. Memory $O(dr)$.
- Can capture: any correlation pattern made of at most $r$ directions; in practice the strongest shared directions of the posterior.
- Cannot capture: patterns that need more than $r$ directions: many separate pairwise trade-offs (one direction per pair), long chains of neighbour correlations, and sign patterns like the one in the example. Directions it leaves out are treated mean-field (too narrow).
- $r \ge d - 1$ could represent any covariance, but would learn $d(r+2) \gt d + d(d+1)/2$ numbers: never better than full-rank.
Why do we need it?
Full-rank memory grows like $d^2$ and becomes impossible for big models, while mean-field throws away every correlation. Low-rank keeps the important trade-offs at a cost that grows only like $d\cdot r$, so big models still get sensible uncertainty.
Where is it used?
NumPyro and Pyro AutoLowRankMultivariateNormal, factor analysis and probabilistic PCA (the same $WW^\top + D$ form), Gaussian-process approximations, large Bayesian regressions, and your forecasting model when its size makes full-rank too expensive.
How is it used?
Pass rank=r (or accept the default $\text{round}(\sqrt d)$). Choose $r$ from the number of large eigenvalues of the posterior correlation of a smaller fit, then check that a higher $r$ does not change the ELBO or the key intervals much.
"A low-rank guide is a cheaper version of mean-field."
It has more parameters than mean-field ($d(r+2)$ vs $2d$) and contains it: set $W = 0$ and you get mean-field. It sits between mean-field and full-rank.
"Rank $r$ captures the $r$ largest correlations."
It captures $r$ directions. One direction can create many correlations at once (a shared driver), but $r$ unrelated pairs need $r$ directions. And some sign patterns are impossible for small $r$.
"Choose $r$ to explain 90% of the variance, as in PCA."
The diagonal $D$ already carries each parameter's private variance. Count the eigenvalues that stand above the floor, then confirm by comparing ELBOs and key intervals at $r$ and at a larger $r$.
"Low-rank fits are as easy to optimize as full-rank."
The rank-$r$ objective can have several local optima (for example, which correlation sign to give up); different seeds can land in different ones. Compare runs.
In your forecasting model the big trade-offs are few: the base slope against the changepoint adjustments, a level against the yearly seasonality, a couple of overlapping regressors. A low-rank guide spends its $r$ directions on those. What it handles badly is a long chain: neighbouring changepoint adjustments $\delta_j$ and $\delta_{j+1}$ often trade off with each other (a slope change can happen a little earlier or a little later), and such local ties run all the way along the grid. If most of your $d$ comes from 25 or more changepoints, check how much the $\delta$ intervals change between rank $r$ and a full-rank fit on a shorter history.
"Low-rank is $O(dr)$, so it is just as good as full-rank but cheaper."
"Low-rank is cheaper because it stores less covariance: $r$ shared directions plus a diagonal. It is as good as full-rank only when the posterior's correlation really is dominated by about $r$ directions."
Model answer: "The low-rank guide is $N(\boldsymbol\mu, WW^\top + D)$ with $W$ of size $d\times r$. In NumPyro it learns $d(r+2)$ numbers: the mean, the $d\times r$ factor and a diagonal scale. Sampling costs $O(dr)$ and the density $O(dr^2 + r^3)$ through Woodbury. It keeps the strongest correlation directions and treats everything else as independent. I pick $r$ from the eigenvalue spectrum of a smaller full-rank fit and check that raising $r$ does not change the conclusions."
Low-rank: $\Sigma_q = WW^\top + D$, $W$: $d\times r$; NumPyro learns $d(r+2)$ (auto_loc, auto_cov_factor, auto_scale); default $r = \text{round}(\sqrt d)$.
Sample $O(dr)$, density $O(dr^2 + r^3)$ (Woodbury), memory $O(dr)$.
Captures $r$ shared directions; misses many separate pairs and chains. Choose $r$ = eigenvalues above the floor, not "90% of variance".
Quick check: a posterior has 8 parameter pairs, each pair strongly correlated, and the pairs are independent of each other. What rank do you need, and how many numbers does that guide learn ($d = 16$)?
One direction per pair: $r = 8$. Learned numbers $16\cdot(8+2) = 160$. Full-rank would learn $16 + 16\cdot17/2 = 152$, fewer! For this structure a full-rank guide (or a block-wise guide that handles each pair) is the better choice.
The three guides side by side on a posterior we know exactly core
To judge a guide you need the true posterior, and usually you do not have it. One model where you do: a linear regression with Normal noise of known size and a Normal prior. Its posterior is exactly Gaussian, with a covariance you can compute by hand. So we can fit all three guides with NumPyro and compare each one with the truth.
The regression is forecasting-flavoured: 120 days, an intercept, a trend, a temperature column, a "feels-like" column that is almost a copy of temperature, and a promotion flag. Temperature and feels-like trade off strongly; intercept and trend trade off moderately. These are two separate correlation directions.
Three ways to say it:
- Picture: three moulds pressed onto the same known shape; we measure the gaps.
- Numbers: true sd of the temperature effect 0.583; mean-field says 0.126, rank 1 says 0.564, full-rank 0.581.
- Slogan: each guide gets right exactly the directions it can store.
Read the real NumPyro results (code at the end of the chapter; 20 000 SVI steps per guide, then 20 000 guide draws).
| Guide | sd(temperature) | corr(temp, feels) | sd(temp + feels) | sd(intercept) |
|---|---|---|---|---|
| exact posterior | 0.583 | −0.974 | 0.133 | 0.193 |
| mean-field | 0.126 | −0.005 | 0.177 | 0.094 |
| low-rank, r = 1 | 0.564 | −0.975 | 0.127 | 0.089 |
| low-rank, r = 2 | 0.572 | −0.972 | 0.134 | 0.191 |
| full-rank | 0.581 | −0.975 | 0.131 | 0.196 |
- Mean-field reports the temperature effect as $0.583/0.126 = 4.6$ times more certain than it is, loses the −0.97 correlation entirely, and makes the combined effect $0.177/0.133 = 1.3$ times too wide (the short axis, as in the 2D example). Its sds match the theory $1/\sqrt{\Lambda_{ii}}$: 0.123 and 0.091 (the small gaps are SVI noise).
- Rank 1 spends its only direction on the strongest trade-off (temperature vs feels-like): those numbers are right. The intercept–trend trade-off is left mean-field: sd 0.089 instead of 0.193.
- Rank 2 has a second direction for intercept–trend: now everything is close.
- Full-rank matches everything (up to SVI noise in the last digit).
The three Gaussian guides in one table (latent dimension $d$, rank $r$):
| Mean-field | Low-rank | Full-rank | |
|---|---|---|---|
| NumPyro | AutoNormal, AutoDiagonalNormal | AutoLowRankMultivariateNormal | AutoMultivariateNormal |
| Covariance $\Sigma_q$ | $\text{diag}(\mathbf{s}^2)$ | $WW^\top + D$ | $LL^\top$ |
| Learned numbers | $2d$ | $d(r+2)$ | $d + d(d+1)/2$ |
| Work per sample | $O(d)$ | $O(dr)$; density $O(dr^2 + r^3)$ | $O(d^2)$ |
| Captures | each parameter's own spread | $r$ shared directions + own spreads | every pairwise correlation |
| Misses | all correlations | patterns needing more than $r$ directions | non-Gaussian shapes only |
All three share the same limits: they are Gaussian in unconstrained space and are fitted by minimizing $KL(q\,\|\,p)$, which prefers to fit inside the posterior.
Why do we need it?
The guide decides which uncertainties survive into your reported intervals. Comparing guides on a posterior you know exactly turns "it is an approximation" into numbers: which quantities are trustworthy, and by how much each guide is off.
Where is it used?
Validation studies for SVI (conjugate or linear-Gaussian test models), simulation-based calibration, the regression-like core of your forecasting model (trend, Fourier terms, holidays and regressors are all columns of one design matrix), and unit tests in an A/B framework (Beta-Binomial posteriors are known exactly).
How is it used?
Build a version of your model with a known posterior (Normal likelihood, known noise, Normal priors), fit each guide, and compare sds and correlations of the quantities you report. Keep this as a regression test when you change guides, ranks or optimizers.
"If the posterior means agree across guides, the guides agree."
In the table every guide gets the means right; they differ in the spreads and correlations. Always compare sds and the intervals of the quantities you report.
"Low-rank with a small $r$ is always at least as good as mean-field for every quantity."
It is never worse in KL, but for a quantity that lies in a direction it did not capture it can be about as wrong as mean-field (rank 1's intercept sd above).
The regression above has the same structure as the core of your forecasting model: trend, Fourier, holiday and regressor columns of one design matrix, with collinear columns producing strong posterior correlations. The lesson carries over directly: mean-field would make a holiday effect or a regressor coefficient look precise when it is not, and a low-rank guide is good exactly for the trade-off directions it has room for. A one-off check like this one (or against NUTS on a short history) is how you show that your guide choice does not change the conclusions.
Known-posterior test: linear-Gaussian regression → compare each guide's sds/correlations with $\Sigma = (X^\top X/\sigma^2 + I/\tau^2)^{-1}$.
Mean-field: means right, collinear coefficients 4–5× too certain. Rank $r$: right along its $r$ directions only. Full-rank: right (Gaussian case).
Trap: same means ≠ same uncertainty.
Quick check: why is the posterior of a linear regression with known noise sd and a Normal prior exactly Gaussian?
The log-likelihood $-\frac{1}{2\sigma^2}\|\mathbf{y} - X\boldsymbol\beta\|^2$ and the log-prior $-\frac{1}{2\tau^2}\|\boldsymbol\beta\|^2$ are both quadratic in $\boldsymbol\beta$, so their sum is a quadratic: the log of a Gaussian, with precision $X^\top X/\sigma^2 + I/\tau^2$. Every guide's error can be measured against it.
Model size and cost: parameter counts, memory and work per step core
Picture the numbers each guide must store as boxes. Mean-field is one row of boxes per parameter (a mean and a spread). Low-rank is a strip $r$ boxes wide. Full-rank is a triangle that is $d$ boxes on each side. Double the model size $d$: the row and the strip double, but the triangle's area quadruples. That is the whole story of Module 68 in one sentence: full covariance grows like $d^2$, rank-$r$ grows like $d\cdot r$.
Two things make the real bill bigger than the parameter count. The optimizer (Adam) keeps two extra running averages per learned number, and a gradient is computed for every one of them each step. And every step pushes each sample through the guide's matrix: through a $d\times d$ triangle for full-rank, a $d\times r$ strip for low-rank.
Three ways to say it:
- Picture: a row, a strip, a triangle.
- Numbers: at $d = 20\,000$ full-rank learns 200 million numbers (0.8 GB in float32 before Adam); rank 32 learns 680 thousand (2.7 MB).
- Slogan: full-rank memory is quadratic; that, not accuracy, is why big models switch to low-rank.
$d = 1\,000$ latent numbers, rank $r = 10$, float32 (4 bytes per number, JAX's default).
- Mean-field: $2d = 2\,000$ numbers → $8$ kB.
- Low-rank: $d(r+2) = 1\,000\times12 = 12\,000$ numbers → $48$ kB.
- Full-rank: $d + d(d+1)/2 = 1\,000 + 500\,500 = 501\,500$ numbers → $2.0$ MB.
- Adam keeps two moment estimates per learned number: multiply by 3 (parameters + 2 moments). Full-rank: $6.0$ MB. NumPyro also builds the constrained $d\times d$ matrix $L$ each step ($10^6$ cells, 4 MB), and the gradient has the size of the parameters.
- Now $d = 20\,000$: full-rank learns $20\,000 + 200\,010\,000 = 200\,030\,000$ numbers = $0.8$ GB; with Adam 2.4 GB, plus 1.6 GB for the $d\times d$ matrix and 0.8 GB of gradient: several GB on one device. Low-rank with $r = 32$: $20\,000 \times 34 = 680\,000$ numbers, 2.7 MB.
- Work per sample: full-rank multiplies by and solves with a $d\times d$ triangle, about $d^2 = 4\times10^8$ multiply-adds at $d = 20\,000$; low-rank about $d\cdot r$ for sampling plus $d\cdot r^2 \approx 2\times10^7$ for the density.
For latent dimension $d$, rank $r$, $b$ bytes per number (4 for float32, 8 for float64 with jax_enable_x64):
| Mean-field | Low-rank | Full-rank | |
|---|---|---|---|
| Learned numbers $n_\phi$ | $2d$ | $d(r+2)$ | $d + d(d+1)/2$ |
| Growth | $O(d)$ | $O(dr)$ | $O(d^2)$ |
| Training memory (rough) | $\approx b \cdot n_\phi \cdot (1 + 2 + 1)$: parameters, Adam's two moments, one gradient; full-rank adds $b\,d^2$ for the materialized $L$ | ||
| Work per sample | $\approx d$ | $\approx dr$ (sample) $+ dr^2 + r^3/3$ (density) | $\approx d^2$ (product + triangular solve) |
| Factorizing a dense $d\times d$ | $\approx d^3/3$: only if you build $L$ from a covariance, invert, or form a Laplace approximation, not in the SVI loop | ||
- With $S$ particles per step (
Trace_ELBO(num_particles=S)) the per-step work is multiplied by $S$. - The model's own cost (evaluating the likelihood over $T$ data points, often $O(Td)$ for a design matrix) is paid by every guide. For small $d$ it usually dominates; for large $d$ the full-rank guide's $d^2$ takes over.
- With NumPyro's default rank $r = \text{round}(\sqrt d)$, the density term $dr^2$ grows like $d^2$ as well, but it is a single dense matrix product, which hardware does very fast (see the timings).
Why do we need it?
Before you run a model you need to know whether the guide fits in memory and how long each step takes. These formulas turn "the model got bigger" into a concrete number of megabytes and milliseconds, and justify switching guides by size.
Where is it used?
Capacity planning for SVI jobs, the automatic full-rank/low-rank switch in your forecasting model, choosing $r$ under a memory budget, deciding when to use AutoGuideList blocks, and explaining JIT compile and run times (Chapter 6.17).
How is it used?
Count $d$, compute $n_\phi$ for each guide, multiply by bytes and by 4 for training state, and compare with the device memory. Time a few hundred jitted steps for the candidates at your real $d$ (timings vary by machine), then pick the richest guide that fits the budget.
"Memory = number of parameters × 4 bytes."
During training you also hold Adam's two moment estimates and a gradient (≈ 4× the parameters), plus whatever the guide materializes (full-rank: the $d\times d$ matrix $L$), plus the model's intermediate arrays.
"Full-rank costs $O(d^3)$ per step."
$O(d^2)$ per sample in NumPyro, because it learns $L$ directly. The measured times grew about 10× when $d$ tripled, close to $d^2$ (9×), not $d^3$ (27×).
"The guide is always the bottleneck."
For small $d$ and long series the likelihood over $T$ observations often costs more than any guide. Profile before blaming the guide.
This is the arithmetic behind your automatic guide choice. Your forecasting model's $d$ grows with the changepoint grid, the Fourier orders, the holidays and the regressors; the full-rank guide's memory grows like $d^2$, the low-rank guide's like $d\cdot r$. A size-based switch says: while $d + d(d+1)/2$ fits the budget, take the richer guide; beyond it, keep the $r$ strongest directions. Be ready to quote the formulas $2d$, $d(r+2)$ and $d + d(d+1)/2$ and to compute them for your model's typical $d$.
"Full covariance is $O(d^2)$, so it is fine for most models."
"Full covariance is $O(d^2)$ in storage and per-step work; that is fine up to a few hundred or a few thousand latent numbers and then becomes the dominant cost. Rank-$r$ is $O(dr)$ in storage."
Model answer: "For $d$ latent dimensions the full-rank guide learns $d + d(d+1)/2$ numbers, so storage, Adam state and per-sample work all grow quadratically; factorizing a dense covariance would be cubic. The low-rank guide learns $d(r+2)$, linear in $d$ for fixed $r$. At $d = 1\,000$, $r = 10$ that is $501\,500$ versus $12\,000$ numbers. In a benchmark (rerun it on your own machine before quoting numbers), tripling $d$ from 1 000 to 3 000 made a full-rank step about ten times slower and a rank-10 step about three times slower."
Learned numbers: MF $2d$ · LR $d(r+2)$ · FR $d + d(d+1)/2$. Memory ≈ bytes × numbers × 4 (params + Adam m, v + grad).
Per sample: MF $O(d)$ · LR $O(dr)$ + $O(dr^2 + r^3)$ density · FR $O(d^2)$; $O(d^3)$ only to factorize a dense matrix.
Trap: count $d$ in unconstrained numbers; the model's likelihood cost may dominate for small $d$.
Quick check: $d = 5\,000$, float32. Roughly how much memory do the full-rank parameters take, before Adam? And low-rank with $r = 20$?
Full-rank: $5\,000 + 5\,000\cdot5\,001/2 = 12\,507\,500$ numbers × 4 bytes ≈ 50 MB (≈ 150 MB with Adam's moments, plus 100 MB for the $5\,000\times5\,000$ matrix $L$). Low-rank: $5\,000\times22 = 110\,000$ numbers ≈ 0.44 MB.
Choosing the variational family, and a size-based selection rule core
Choosing a guide is like choosing the scale of a map. For a small town you can afford a map that shows every street (full-rank). For a whole country you keep only the motorways, the few roads that carry most of the traffic (low-rank), and you accept that small lanes are missing. A map with no roads at all, only town names (mean-field), is cheap but useless for planning a route that crosses many towns.
So the choice depends on four questions: How big is $d$? (cost) How are the parameters tied together? (a few shared directions, or many separate ties) What do you report? (point forecasts barely care; intervals of sums and probabilities of individual effects care a lot) Is a Gaussian shape good enough at all?
Three ways to say it:
- Picture: street map for a town, motorway map for a country.
- Numbers: $d = 61$ → full-rank costs 1 952 numbers, take it; $d = 1\,000$ → full-rank 501 500 vs rank 32 only 34 000.
- Slogan: the richest guide you can afford, validated on a version you can check.
Three model sizes, one decision each (timings from the table above; your machine will differ).
- $d = 61$ (the example forecasting model). Full-rank learns $1\,952$ numbers; one step is a fraction of a millisecond. Full-rank: no reason to give up any correlation.
- $d = 1\,000$. Full-rank: $501\,500$ numbers, about 4 ms per step, so 20 000 steps ≈ 80 s. Low-rank with $r = 32$: $34\,000$ numbers, about 0.2 ms per step, ≈ 4 s. Both fit in memory; the choice is about time and how many correlation directions the posterior really has.
- $d = 20\,000$. Full-rank: 200 million numbers, several GB of training state. Low-rank with $r$ chosen from the spectrum of a smaller fit, e.g. $r = 32$: $680\,000$ numbers.
- An illustrative rule (the budget is an example, not a recommended constant): "use full-rank while $d + d(d+1)/2 \le 50\,000$, otherwise low-rank with $r = \text{round}(\sqrt d)$". Then $d \le 314$ gives full-rank ($314 + 314\cdot315/2 = 49\,769$) and $d = 315$ switches ($315 + 315\cdot316/2 = 50\,085$) to rank $\text{round}(\sqrt{315}) = 18$.
A size-based selection rule has three parts:
- Measure the latent dimension $d$ (trace the model, add up the unconstrained sizes).
- Compare the full-rank cost $d + d(d+1)/2$ (or measured memory/time) with a budget $B$. If it fits, use
AutoMultivariateNormal. - Otherwise use
AutoLowRankMultivariateNormal(rank=r)with $r$ from the posterior's spectrum (eigenvalues above the floor), from a default such as $\text{round}(\sqrt d)$, or capped so that $d(r+2) \le B$.
def choose_guide(model, *args, budget=50_000, max_rank=None):
d = latent_dim(model, *args) # unconstrained latent numbers
if d + d * (d + 1) // 2 <= budget: # full-rank is affordable
return AutoMultivariateNormal(model), d, None
r = max(1, round(d ** 0.5)) # e.g. NumPyro's own default rank
if max_rank is not None:
r = min(r, max_rank)
return AutoLowRankMultivariateNormal(model, rank=r), d, r
Validation (what makes the rule defensible): on a smaller version of the model (shorter history, fewer changepoints) fit full-rank, low-rank and, if affordable, NUTS (Chapter 6.15); compare the ELBO and the intervals of the quantities you report. Raise $r$ until they stop changing.
Other families in NumPyro, for when a single Gaussian is the wrong tool:
AutoDelta: a point (MAP), no uncertainty at all.AutoLaplaceApproximation: a full Gaussian centred at the mode with covariance from the Hessian (a one-off $O(d^3)$ factorization).AutoIAFNormal,AutoBNAFNormal: normalizing flows; they can bend the Gaussian into skewed or curved shapes, at a higher cost.AutoGuideList: different guides for different blocks, e.g. full-rank for a tightly coupled block of sites andAutoNormalfor the rest (sites split withnumpyro.handlers.block).
Why do we need it?
No single guide is best for every model size. A fixed choice either wastes memory and time on small models' worth of correlations, or drops correlations a big model needs. An explicit rule makes the trade-off automatic, reproducible and explainable.
Where is it used?
Your forecasting model's automatic full-rank vs low-rank switch, production pipelines that fit many series of different sizes, Stan's ADVI choice between meanfield and fullrank, and block-wise guides (AutoGuideList) in large hierarchical A/B models.
How is it used?
Compute $d$ before building SVI, pick the guide with a rule like choose_guide, log the decision ($d$, guide, $r$) with the run for reproducibility, and keep a validation notebook that compares guides on a reduced model whenever the model structure changes.
"The richer guide is always better."
A richer guide is never worse at its optimum, but it has more numbers to learn: SVI is slower, noisier and may need more steps or a smaller learning rate. Under a fixed time budget a well-converged low-rank fit can beat a half-converged full-rank fit.
"The size threshold is a universal constant."
It depends on device memory, time per fit, the number of particles and how many series you fit. Any number you quote is a design choice; defend it with your measurements.
"SVI converged, so the guide family is adequate."
Convergence means the optimizer found the best member of the family. Whether the family is adequate is a separate question, answered by comparing with a richer guide or NUTS on a reduced model.
When an interviewer asks about your automatic guide selection, the structure of a good answer is: (1) what $d$ is in your model and what drives it (changepoints, Fourier orders, holidays, regressors); (2) the counts $d + d(d+1)/2$ vs $d(r+2)$ and why memory, not accuracy, forces the switch; (3) what low-rank keeps (the main trade-offs) and loses (many local ties); (4) how you checked it (a reduced model fitted with full-rank or NUTS). State the threshold your code uses as a design choice and how you arrived at it; do not present it as a law.
"We use low-rank because it is more accurate for large models."
"We use low-rank for large models because full-rank's $O(d^2)$ memory and per-step cost become prohibitive; low-rank is an approximation of the covariance that keeps the strongest directions."
Model answer: "Both guides are Gaussian in unconstrained space. Full-rank learns $d + d(d+1)/2$ numbers and captures every correlation, so I use it whenever that fits the budget. Beyond that, the quadratic growth in memory and per-step work makes it impractical, so I switch to a low-rank-plus-diagonal guide with $d(r+2)$ numbers, which keeps the $r$ strongest correlation directions. I validated the switch by fitting a smaller version with both guides and checking that the forecast intervals and key coefficients agree."
Choose by: size $d$ (cost), structure (few directions vs many ties), what you report (sums/effects need correlations), Gaussian enough?
Rule: full-rank while $d + d(d+1)/2 \le B$; else low-rank with $r$ from the spectrum (or $\text{round}(\sqrt d)$), capped by $d(r+2)\le B$.
Trap: thresholds are design choices; convergence ≠ adequate family; validate on a reduced model.
Quick check: budget $B = 20\,000$ learned numbers. What is the largest $d$ for which the rule picks full-rank?
Need $d + d(d+1)/2 \le 20\,000$. $d = 198$: $198 + 198\cdot199/2 = 198 + 19\,701 = 19\,899$ ✓. $d = 199$: $199 + 199\cdot200/2 = 199 + 19\,900 = 20\,099$ ✗. So $d \le 198$.
Recap, cheat sheet and practice
- A NumPyro autoguide flattens all latent sites into one unconstrained vector of length $d$ and puts a Gaussian on it. The guide type is the covariance structure of that Gaussian.
- Mean-field ($2d$ numbers): diagonal covariance. For a Gaussian posterior it gets the means right and gives each parameter its conditional variance $1/\Lambda_{ii} \le \Sigma_{ii}$: too narrow along correlated directions, too wide along the short axis.
- Full-rank ($d + d(d+1)/2$ numbers): $N(\boldsymbol\mu, LL^\top)$ with a learned Cholesky factor. Captures every correlation; $O(d^2)$ memory and per-sample work; $O(d^3)$ only to factorize a dense matrix.
- Low-rank ($d(r+2)$ numbers in NumPyro): $WW^\top + D$. Keeps $r$ shared directions plus private variances; $O(dr)$ memory and sampling, $O(dr^2 + r^3)$ density. Misses patterns needing more than $r$ directions (many pairs, chains, some sign patterns).
- Training memory ≈ bytes × learned numbers × 4 (parameters, Adam's two moments, gradient), plus the $d\times d$ matrix for full-rank. Measured: tripling $d$ made a full-rank step ~10× slower, a rank-10 step ~3× slower.
- Choose the richest guide the budget allows; pick $r$ from the eigenvalues above the floor; validate on a reduced model against full-rank or NUTS; treat any size threshold as a design choice.
Cheat sheet
| Idea | Formula / call | In words |
|---|---|---|
| Latent dimension | $d = \sum_{\text{latent sites}}$ unconstrained size | shape (25,) → 25; simplex of $K$ → $K-1$; σ → 1 |
| Mean-field | AutoNormal: $2d$; optimum $s_i^2 = 1/\Lambda_{ii}$ | independent bells; conditional variances |
| Full-rank | AutoMultivariateNormal: $d + d(d+1)/2$; $\Sigma = LL^\top$ | every correlation; stored as $d\times d$ auto_scale_tril |
| Low-rank | AutoLowRankMultivariateNormal(rank=r): $d(r+2)$; $\text{diag}(\mathbf{s})(FF^\top+I)\text{diag}(\mathbf{s})$ | $r$ shared directions + diagonal; default $r = \text{round}(\sqrt d)$ |
| Sampling | $\boldsymbol\mu + \mathbf{s}\odot\boldsymbol\varepsilon$ · $\boldsymbol\mu + W\boldsymbol\varepsilon_r + D^{1/2}\boldsymbol\varepsilon_d$ · $\boldsymbol\mu + L\boldsymbol\varepsilon$ | work $d$ · $dr$ · $d^2$ |
| Memory (training) | $\approx b \cdot n_\phi \cdot 4$ (+ $b\,d^2$ for full-rank) | $b$ = 4 bytes (float32) or 8 |
| Rank-1 limit | $\Sigma_{12}\Sigma_{13}\Sigma_{23} = (w_1w_2w_3)^2 \ge 0$ | some sign patterns need $r \ge 2$ |
| Choosing $r$ | count eigenvalues above $c\times$ floor | not "90% of variance": $D$ covers the floor |
| Selection rule (illustrative) | full-rank if $d + d(d+1)/2 \le B$, else low-rank | validate on a reduced model |
import numpy as np
import jax, jax.numpy as jnp
import numpyro, numpyro.distributions as dist
from numpyro.handlers import seed, trace
from numpyro.distributions import biject_to
from numpyro.infer import SVI, Trace_ELBO
from numpyro.infer.autoguide import (AutoNormal, AutoDiagonalNormal,
AutoMultivariateNormal, AutoLowRankMultivariateNormal)
from numpyro.optim import Adam
# 1) Count what each guide learns: look at the real parameter shapes (d = 10, rank 3)
def toy(d):
numpyro.sample("theta", dist.Normal(0, 1).expand([d]).to_event(1))
for G, kw in [(AutoNormal, {}), (AutoDiagonalNormal, {}),
(AutoMultivariateNormal, {}), (AutoLowRankMultivariateNormal, {"rank": 3})]:
opt = Adam(0.01)
svi = SVI(toy, G(toy, **kw), opt, Trace_ELBO())
state = svi.init(jax.random.PRNGKey(0), 10)
stored = {k: v.shape for k, v in svi.get_params(state).items()} # constrained arrays
learned = sum(v.size for v in opt.get_params(state.optim_state).values()) # what Adam updates
print(f"{G.__name__:30s} learned {learned:3d} stored {stored}")
# AutoNormal learned 20 stored {'theta_auto_loc': (10,), 'theta_auto_scale': (10,)}
# AutoDiagonalNormal learned 20 stored {'auto_loc': (10,), 'auto_scale': (10,)}
# AutoMultivariateNormal learned 65 stored {'auto_loc': (10,), 'auto_scale_tril': (10, 10)}
# AutoLowRankMultivariateNormal learned 50 stored {'auto_cov_factor': (10, 3), 'auto_loc': (10,), 'auto_scale': (10,)}
# 2) A posterior we know exactly: linear regression with known noise sd and a Normal prior
rng = np.random.default_rng(0)
T = 120
t = np.arange(T) / T
temp = np.sin(4 * np.pi * t) + 0.3 * rng.normal(size=T)
feels = temp + 0.15 * rng.normal(size=T) # "feels-like" temperature: almost the same column
promo = (rng.random(T) < 0.2).astype(float)
X = np.column_stack([np.ones(T), t, temp, feels, promo])
y = X @ np.array([10.0, 2.0, 1.0, 0.5, 1.5]) + rng.normal(size=T)
sigma, tau = 1.0, 5.0
Lam = X.T @ X / sigma**2 + np.eye(5) / tau**2 # posterior precision
Sig = np.linalg.inv(Lam) # exact posterior covariance
def model(X, y=None):
b = numpyro.sample("b", dist.Normal(0, tau).expand([X.shape[1]]).to_event(1))
numpyro.sample("y", dist.Normal(X @ b, sigma), obs=y)
def summary(C): # sd(temp), corr(temp, feels), sd(temp + feels), sd(intercept)
s = np.sqrt(np.diag(C)); a = np.array([0, 0, 1, 1, 0])
return s[2], C[2, 3] / (s[2] * s[3]), np.sqrt(a @ C @ a), s[0]
print("exact sd(temp) %.3f corr %.3f sd(temp+feels) %.3f sd(intercept) %.3f" % summary(Sig))
for name, guide in [("mean-field", AutoNormal(model)),
("low-rank r=1", AutoLowRankMultivariateNormal(model, rank=1)),
("low-rank r=2", AutoLowRankMultivariateNormal(model, rank=2)),
("full-rank", AutoMultivariateNormal(model))]:
svi = SVI(model, guide, Adam(0.005), Trace_ELBO(num_particles=4))
res = svi.run(jax.random.PRNGKey(1), 20_000, jnp.array(X), jnp.array(y), progress_bar=False)
b = guide.sample_posterior(jax.random.PRNGKey(2), res.params, sample_shape=(20_000,))["b"]
print(f"{name:14s} sd(temp) %.3f corr %.3f sd(temp+feels) %.3f sd(intercept) %.3f" % summary(np.cov(np.asarray(b).T)))
print("mean-field theory: sd = 1/sqrt(diag(precision)) =", (1 / np.sqrt(np.diag(Lam)))[[2, 0]].round(3))
# exact sd(temp) 0.583 corr -0.974 sd(temp+feels) 0.133 sd(intercept) 0.193
# mean-field sd(temp) 0.126 corr -0.005 sd(temp+feels) 0.177 sd(intercept) 0.094
# low-rank r=1 sd(temp) 0.564 corr -0.975 sd(temp+feels) 0.127 sd(intercept) 0.089
# low-rank r=2 sd(temp) 0.572 corr -0.972 sd(temp+feels) 0.134 sd(intercept) 0.191
# full-rank sd(temp) 0.581 corr -0.975 sd(temp+feels) 0.131 sd(intercept) 0.196
# mean-field theory: sd = 1/sqrt(diag(precision)) = [0.123 0.091]
# (SVI is stochastic: on another machine or version the last digit may differ.)
# 3) An illustrative size rule (the budget is an example, NOT a recommended constant)
def latent_dim(model, *args, **kwargs):
tr = trace(seed(model, 0)).get_trace(*args, **kwargs)
return sum(int(np.prod(biject_to(s["fn"].support).inverse_shape(jnp.shape(s["value"]))))
for s in tr.values() if s["type"] == "sample" and not s["is_observed"])
def choose_guide(model, *args, budget=50_000, max_rank=None):
d = latent_dim(model, *args)
if d + d * (d + 1) // 2 <= budget:
return AutoMultivariateNormal(model), d, None
r = max(1, round(d ** 0.5)) if max_rank is None else min(max_rank, max(1, round(d ** 0.5)))
return AutoLowRankMultivariateNormal(model, rank=r), d, r
for d in [61, 314, 315, 1000]:
g, dd, r = choose_guide(toy, d)
print(d, type(g).__name__, "rank", r)
# 61 AutoMultivariateNormal rank None
# 314 AutoMultivariateNormal rank None
# 315 AutoLowRankMultivariateNormal rank 18
# 1000 AutoLowRankMultivariateNormal rank 32
1. A model has $d = 100$ latent numbers. Which NumPyro guide learns exactly 1 000 numbers?
2. A Gaussian posterior has two parameters with sd 1 and correlation 0.8. What sd does the optimal mean-field guide report for each?
3. Which statement about AutoMultivariateNormal in NumPyro is correct?
4. A posterior over 12 parameters consists of 6 independent pairs, each pair strongly correlated. What is the smallest rank that captures every correlation?
5. A model has a Dirichlet over 6 categories, 8 segment effects, a population mean and a population sd. What is $d$?
6. A mean-field guide is fitted to a Gaussian posterior where two parameters are strongly positively correlated. Which statement is true?
Practice problems
A. For the example forecasting model ($d = 61$), compute the learned numbers of mean-field, low-rank with $r = 5$, NumPyro's default low-rank, and full-rank.
- Mean-field: $2\cdot61 = 122$.
- Low-rank $r = 5$: $61\cdot7 = 427$.
- Default rank: $\text{round}(\sqrt{61}) = \text{round}(7.81) = 8$, so $61\cdot10 = 610$ (NumPyro creates
auto_cov_factorof shape (61, 8)). - Full-rank: $61 + 61\cdot62/2 = 61 + 1\,891 = 1\,952$. All tiny: full-rank is the natural choice.
B. Derive the optimal mean-field sds for a 2D Gaussian posterior with sds $\sigma_1, \sigma_2$ and correlation $\rho$.
- $\Sigma = \begin{bmatrix} \sigma_1^2 & \rho\sigma_1\sigma_2 \\ \rho\sigma_1\sigma_2 & \sigma_2^2 \end{bmatrix}$, $\det\Sigma = \sigma_1^2\sigma_2^2(1-\rho^2)$.
- $\Lambda_{11} = \sigma_2^2/\det\Sigma = \dfrac{1}{\sigma_1^2(1-\rho^2)}$, and similarly $\Lambda_{22} = \dfrac{1}{\sigma_2^2(1-\rho^2)}$.
- Setting $\partial KL/\partial s_i^2 = \tfrac12(\Lambda_{ii} - 1/s_i^2) = 0$ gives $s_i^2 = 1/\Lambda_{ii}$, so $s_1 = \sigma_1\sqrt{1-\rho^2}$ and $s_2 = \sigma_2\sqrt{1-\rho^2}$.
- This is the conditional sd of each parameter given the other. With $\rho = 0.9$ it is $0.436\,\sigma$.
C. A full-rank guide has $\boldsymbol\mu = \mathbf{0}$ and $L = \begin{bmatrix} 1 & 0 & 0 \\ 0.5 & 1 & 0 \\ -1 & 0.5 & 2 \end{bmatrix}$. Compute $\Sigma$, the correlations, and the sample for $\boldsymbol\varepsilon = (1, 1, 1)$.
- $\Sigma_{ij}$ = (row $i$ of $L$)·(row $j$ of $L$): $\Sigma_{11} = 1$, $\Sigma_{12} = 0.5$, $\Sigma_{13} = -1$, $\Sigma_{22} = 0.25 + 1 = 1.25$, $\Sigma_{23} = -0.5 + 0.5 = 0$, $\Sigma_{33} = 1 + 0.25 + 4 = 5.25$.
- Correlations: $\rho_{12} = 0.5/\sqrt{1.25} = 0.447$; $\rho_{13} = -1/\sqrt{5.25} = -0.436$; $\rho_{23} = 0$.
- Sample: $\mathbf{z} = L\boldsymbol\varepsilon = (1,\ 0.5 + 1,\ -1 + 0.5 + 2) = (1, 1.5, 1.5)$.
- Marginal sds are the row norms of $L$: $1$, $\sqrt{1.25} = 1.118$, $\sqrt{5.25} = 2.291$; not the diagonal $1, 1, 2$.
D. $d = 10\,000$, float32. Compare the training memory of full-rank and of low-rank with $r = 50$.
- Full-rank learned numbers: $10\,000 + 10\,000\cdot10\,001/2 = 10\,000 + 50\,005\,000 = 50\,015\,000$, i.e. 200 MB.
- ×4 for parameters, Adam's two moments and the gradient: 800 MB, plus the $10\,000\times10\,000$ matrix $L$: $10^8\cdot4$ bytes = 400 MB. About 1.2 GB.
- Low-rank: $10\,000\cdot52 = 520\,000$ numbers = 2.08 MB; ×4 ≈ 8.3 MB.
- Ratio of learned numbers ≈ 96: memory, not accuracy, is what pushes a model this size to low-rank.
E. Show that a rank-1-plus-diagonal covariance can never have correlations $\rho_{12} = 0.5$, $\rho_{13} = 0.5$, $\rho_{23} = -0.4$. Is this a valid posterior correlation matrix, and what rank fixes it?
- With $W = (w_1, w_2, w_3)^\top$, the off-diagonal entries of $WW^\top + D$ are $w_iw_j$ ($D$ only touches the diagonal), and correlations have the same signs as covariances.
- Product of the three: $w_1w_2\cdot w_1w_3\cdot w_2w_3 = (w_1w_2w_3)^2 \ge 0$. But $0.5\cdot0.5\cdot(-0.4) = -0.1 \lt 0$. So rank 1 must give up one sign (the 3D widget shows the fit flipping one).
- Valid? The determinant is $1 + 2(0.5)(0.5)(-0.4) - (0.25 + 0.25 + 0.16) = 1 - 0.2 - 0.66 = 0.14 \gt 0$ (and the $2\times2$ minors are positive), so yes, it is positive definite.
- Rank $2 = d - 1$ can represent any $3\times3$ covariance, so $r = 2$ fits it exactly. (It learns $3\cdot4 = 12$ numbers, more than full-rank's $3 + 6 = 9$: at $d = 3$ just use full-rank.)
F. (Interview) "Your forecasting model automatically picks a full-rank or a low-rank Gaussian guide based on model size. Walk me through the design, what you lose, and how you would defend the threshold."
"The guide is a Gaussian over the flattened unconstrained latent vector. Its size $d$ grows with the changepoint grid, Fourier orders, holidays, regressors and the likelihood's extra parameters. A full-rank guide, $N(\boldsymbol\mu, LL^\top)$, learns $d + d(d+1)/2$ numbers; its memory, Adam state and per-step work grow like $d^2$, which is fine for small $d$ and the best Gaussian approximation, since it captures every trade-off such as slope vs changepoint adjustments. Beyond a size budget I switch to a low-rank-plus-diagonal guide, $WW^\top + D$, with $d(r+2)$ numbers: linear in $d$. It keeps the $r$ strongest correlation directions and treats the rest as independent, so I lose patterns that need many directions, such as long chains of neighbouring changepoints, and those parameters' intervals become too narrow. The threshold is a design choice from memory and time measurements, not a law. I defend it by fitting a reduced model with both guides (and NUTS where affordable) and checking that forecast intervals and key coefficients agree, and by checking that increasing $r$ does not change the ELBO or the intervals."
The custom SVI training loop: early stopping, patience, checkpointing
NumPyro's svi.run takes a fixed number of steps and hands back the last state. A custom loop decides for itself when to stop and which state to keep. It needs four ideas: a relative measure of improvement (so the same rule works on small and large ELBOs), patience (so one noisy reading does not end training), a minimum number of steps, and a checkpoint of the best state. This chapter builds each one with pictures and simulations, then runs the real NumPyro/JAX loop and reads its log line by line.
- Describe the SVI loop (initialize → stochastic ELBO → gradients → optimizer update → repeat) and how the learning rate, gradient noise and optimizer shape it
- Get the signs right:
svi.updatereturns the loss = −ELBO; define "improvement" correctly in either convention - Use relative early stopping $\frac{ELBO_t - ELBO_{best}}{|ELBO_{best}| + \epsilon}$, explain why absolute thresholds break across scales, and what $\epsilon$ really does
- Choose patience, minimum steps, the evaluation frequency and smoothing, and recognize premature stopping
- Keep the best state (not the last one), and know why "best" measured on a noisy number can be partly luck
- Write and explain a runnable NumPyro/JAX loop with a jitted update, relative early stopping, patience and checkpointing
What we need from earlier chapters: the ELBO and SVI (Chapter 6.12): the ELBO is a lower bound on $\log p(D)$ that is highest when the guide is closest to the posterior, and SVI climbs it using noisy Monte Carlo gradients; guides and their sizes (Chapter 6.13); Adam and learning rates (Optimization, Chapter 3.4) and the noise of stochastic gradients (Chapter 3.14). JIT compilation is only used here; how it works is taught in Chapter 6.17. Notation: $t$ = step number; $\mathcal{L}_t$ = the loss reported at step $t$ ($= -\widehat{ELBO}_t$, a noisy estimate); $k$ = evaluate every $k$ steps; $P$ = patience (in evaluations); $\tau$ = tolerance; "nat" = the unit of a natural-log probability.
The optimization loop: init → noisy ELBO → gradient → update → repeat core
Imagine tuning an old radio while the signal crackles. Each moment you listen briefly (one noisy reading of how good the station sounds), you guess which way to turn the dial (the gradient), you turn it a little (the learning rate), and you listen again. The crackle never fully goes away, so you never get a perfectly clean reading, but on average every turn brings you closer.
SVI is exactly that loop. In one sentence from Chapter 6.12: the ELBO measures how well the guide matches the posterior (higher is better), SVI estimates it with a few random draws from the guide, and an optimizer such as Adam nudges the guide's parameters $\phi$ uphill. A step is: draw $\boldsymbol\varepsilon$ → make $\mathbf z$ → evaluate the noisy ELBO → differentiate → let the optimizer update $\phi$.
Three ways to say it:
- Picture: a hiker climbing in fog, taking a small step uphill based on a quick, blurry look at the slope.
- Numbers: in the real run at the end of this chapter, one step's loss has a random wobble of about 1.2 nats on a loss of about 552.
- Slogan: many small noisy steps add up to a good fit; no single step can be trusted.
One SVI step by hand. Target posterior (unnormalized) $\tilde p(z) = e^{-(z-3)^2/2}$. Guide $q = N(\mu, s^2)$ with $s = e^{u}$; start at $\mu = 0$, $u = 0$ ($s = 1$). One particle, plain gradient step with learning rate 0.1.
- Draw $\varepsilon = 0.5$, so $z = \mu + s\varepsilon = 0.5$ (the reparameterization trick of 6.12).
- Loss estimate $= -\log\tilde p(z) + \log q(z) = \tfrac12(0.5-3)^2 + \big(-\tfrac12\varepsilon^2 - u - \tfrac12\log 2\pi\big) = 3.125 - 0.125 - 0 - 0.919 = 2.081$.
- Gradient in $\mu$: $-\frac{d}{dz}\log\tilde p(z) = z - 3 = -2.5$. Gradient in $u$: $(z-3)\,\varepsilon\, s - 1 = (-2.5)(0.5)(1) - 1 = -2.25$.
- Update: $\mu \leftarrow 0 - 0.1(-2.5) = 0.25$; $u \leftarrow 0 - 0.1(-2.25) = 0.225$, so $s = e^{0.225} = 1.25$.
- Compare with the exact values: exact loss at the start $= \tfrac12(s^2 + (\mu-3)^2) - \log s - \tfrac12 - \tfrac12\log 2\pi = 5 - 1.419 = 3.581$ (the one-sample estimate said 2.081); exact gradient in $u$ is $s^2 - 1 = 0$ (the estimate said −2.25). The step moved $\mu$ the right way and $s$ the wrong way, purely because of noise. Averaged over many steps, the noise cancels.
The SVI loop in NumPyro:
svi = SVI(model, guide, Adam(lr), Trace_ELBO(num_particles=S))
state = svi.init(rng_key, *data) # run model + guide once, create the initial φ₀
for t in range(1, max_steps + 1):
state, loss = svi.update(state, *data) # noisy ELBO, gradient, one optimizer step
- Initialization:
svi.inittraces the model and guide once with a PRNG key and creates the parameters. Autoguides place the initial means withinit_loc_fn(defaultinit_to_uniform: unconstrained values drawn uniformly in $(-2, 2)$; alternativesinit_to_median,init_to_value) and start every scale atinit_scale = 0.1. - Stochastic ELBO: each step averages $S$ draws (
num_particles); the noise sd falls like $1/\sqrt S$. - Gradient: automatic differentiation of the estimate (reparameterized).
- Optimizer: turns gradients into steps.
Adam(adaptive per-parameter step sizes; the usual default),ClippedAdam(also clips each gradient value to ±clip_norm, default 10),SGD, or any optax optimizer. - Learning rate: the step size; too small is slow, too large is jumpy or diverges.
- Convergence: with a constant learning rate the parameters never settle exactly; they jitter around the optimum and the reported loss keeps fluctuating. "Converged" means "no longer improving by a meaningful amount", which is what the stopping rule in the rest of this chapter measures.
Why do we need it?
The posterior of a real model has no formula; SVI turns finding it into an optimization that runs step by step. Knowing what each step does tells you why the loss is noisy, why the learning rate matters, and why "when do I stop?" needs a rule of its own.
Where is it used?
NumPyro and Pyro SVI, variational autoencoders, Bayesian neural networks, topic models, and both of your projects: the A/B framework fits its hierarchical models with SVI, and the forecasting model runs a custom SVI loop with jitted updates.
How is it used?
Build SVI(model, guide, optimizer, Trace_ELBO()), call svi.init once, then call a jitted svi.update in a Python loop, watching the loss every $k$ steps and deciding when to stop. Read the final parameters with svi.get_params(state).
svi.update. Every $k$ steps the outer loop evaluates progress, keeps a checkpoint of the best state, and stops after $P$ evaluations without a meaningful gain."The loss should go down to zero."
The loss is $-ELBO = -\log p(D) + KL(q\,\|\,p(\cdot\mid D))$. It goes down toward $-\log p(D)$, which can be any number, even negative. You can only compare losses of the same model on the same data.
"This step's loss is lower than the last one, so the parameters improved."
One step's loss is one noisy estimate (in the real run below its sd is about 1.2 nats). Compare averages over many steps, or use a low-noise evaluation.
"svi.run stops when the model has converged."
It runs exactly num_steps steps and returns the last state. NumPyro's documentation recommends init/update/evaluate for early stopping, which is why a custom loop exists.
Your forecasting model's training loop is this cycle with extra logic around it: svi.init once, a JIT-compiled update called repeatedly, and, at some interval of steps, a check of the relative ELBO improvement against the best value so far, a patience counter, and a checkpoint of the best state. In an A/B framework like yours, the same SVI cycle fits the hierarchical models; even with a simple fixed number of steps there, the noise and learning-rate lessons apply.
Loop: state = svi.init(key, *data); repeat state, loss = svi.update(state, *data); read svi.get_params(state).
Each step: draw ε → noisy ELBO (S particles, noise ∝ $1/\sqrt S$) → autodiff gradient → optimizer step. Autoguide init: init_to_uniform locs, scale 0.1.
Trap: one step's loss is noise; the loss tends to $-\log p(D)$, not 0; svi.run has no early stopping.
Quick check: why can the exact gradient for the guide's log-sd be 0 while the one-sample gradient is −2.25?
The one-sample gradient depends on the random $\varepsilon$. It is an unbiased estimate: averaged over many draws of $\varepsilon$ it equals the exact gradient (0 when the guide's sd already matches), but any single draw can point either way. That is the "gradient noise" of SVI.
Learning rate, gradient noise, and what "converged" means for SVI
Put a marble in a bowl on a table that keeps shaking. Gravity pulls it to the bottom; the shaking keeps knocking it around. It ends up rolling around inside a small region near the bottom, never resting exactly there. That region is the noise ball.
In SVI the shaking is the gradient noise (from the random draws, and from minibatches if you use them) and the learning rate decides how hard each knock moves the marble. A big learning rate gets to the bottom fast but bounces around in a wide ball; a small one creeps down slowly but settles closer; a too big one throws the marble out of the bowl.
Three ways to say it:
- Picture: a marble on a shaking table: fast and bouncy, or slow and calm.
- Numbers: on a simple bowl, cutting the learning rate 10× cuts the leftover error about 10× but makes the descent about 10× slower.
- Slogan: with a constant learning rate SVI never stops moving; "converged" means "stopped improving by a meaningful amount".
Plain SGD on the bowl $f(\theta) = \tfrac12\theta^2$ (curvature $a = 1$), with gradient noise of sd $\sigma_g = 1$: each step $\theta \leftarrow \theta - \text{lr}\,(\theta + \sigma_g\xi)$, $\xi \sim N(0, 1)$.
- Without noise, each step multiplies $\theta$ by $(1 - \text{lr})$. With lr = 0.1 the distance halves every $\ln 0.5/\ln 0.9 \approx 6.6$ steps; with lr = 0.01 every $\approx 69$ steps.
- With noise, $\theta$ settles into a spread with variance $\dfrac{\text{lr}\,\sigma_g^2}{a(2 - \text{lr}\,a)}$ (set the variance before and after a step equal and solve).
- lr = 0.1: variance $0.1/1.9 = 0.0526$, so the leftover loss $E[f] = \tfrac12\cdot0.0526 = 0.026$.
- lr = 0.01: variance $0.01/1.99 = 0.00503$, leftover loss $0.0025$: ten times smaller, ten times slower.
- lr = 2.5: each step multiplies $\theta$ by $1 - 2.5 = -1.5$: the size grows by half every step. Divergence starts when $\text{lr}\cdot a \gt 2$.
- Gradient noise: the random error of each step's gradient estimate. In SVI it comes from the Monte Carlo draws of $\mathbf z$ (always) and from minibatching (if used). Averaging $S$ particles divides its variance by $S$.
- Noise ball: with a constant learning rate the parameters end up fluctuating around the optimum; for SGD on a quadratic the spread has variance $\frac{\text{lr}\,\sigma_g^2}{a(2-\text{lr}\,a)}$, roughly proportional to the learning rate.
- Stability: for plain SGD the step diverges when $\text{lr} \gt 2/a$ along the steepest direction. Adam rescales each coordinate, so its steps are roughly lr in size whatever the gradient scale, but too large an lr still makes it jumpy.
- Ways to shrink the ball: a decaying learning-rate schedule, more particles, averaging the parameters over the last steps, or gradient clipping (
ClippedAdam) against rare huge gradients. - Convergence in practice: the smoothed loss stops improving by more than a tolerance over several evaluations. The parameters still jitter; the stopping rule decides when more steps are no longer worth it.
Why do we need it?
The learning rate and the noise decide how fast the loss falls, how low it can go and how much it wobbles at the end. Every stopping rule must work with that wobble, so you need to know its size and what controls it.
Where is it used?
Choosing Adam(lr) and num_particles in NumPyro, learning-rate schedules in deep learning (step decay, cosine), SGD theory (Optimization, Chapter 3.14), and setting the tolerance and patience of your loop.
How is it used?
Try a few learning rates (for example 0.001, 0.003, 0.01, 0.03) on a short run and look at the smoothed loss: pick the largest that does not jump around or blow up. If the final wobble is too large, lower lr late in training or add particles.
"A smaller learning rate is always the safe choice."
It is stabler but slower. With a tiny learning rate the run can hit the step limit long before it converges (try lr = 0.001 in the simulator later in this chapter), and slow, steady progress is easier for a stopping rule to mistake for a plateau.
"Adam makes the learning rate unimportant."
Adam adapts the scale per parameter, but lr still sets how far each parameter moves per step: the speed of descent and the size of the final jitter.
"Without minibatches the SVI loss is not noisy."
The ELBO is always estimated from random draws of $\mathbf z$ from the guide, so it is noisy even on the full dataset. Minibatches add a second source of noise.
Both settings of your loop's stopping rule depend on this: the size of the loss wobble (from the learning rate, the number of particles and the data) decides how large a tolerance must be to rise above noise, and how many evaluations of patience are needed before "no improvement" means something. If you change the learning rate or the particle count, re-check the tolerance and patience.
Noise ball (SGD, quadratic): Var $= \frac{\text{lr}\,\sigma_g^2}{a(2 - \text{lr}\,a)} \approx \text{lr}\,\sigma_g^2/(2a)$; diverges if $\text{lr}\,a \gt 2$.
Bigger lr: faster, higher and noisier floor. Fixes: lr decay, more particles ($\text{Var}/S$), clipping.
Trap: constant-lr SVI never "stops moving"; converged = no meaningful improvement.
Quick check: plain SGD on $\tfrac12\theta^2$ with $\sigma_g = 2$ and lr = 0.05. What is the leftover loss $E[f]$?
Variance $= \frac{0.05\cdot4}{1\cdot(2 - 0.05)} = \frac{0.2}{1.95} = 0.1026$, so $E[f] = \tfrac12\cdot0.1026 = 0.051$.
Getting the signs right: svi.update returns the loss, which is −ELBO core
Golf and basketball both have scoreboards, but in golf lower is better and in basketball higher is better. The ELBO is a basketball score: higher means the guide fits the posterior better. Optimizers are built to minimize, so NumPyro reports the golf version: the loss, which is minus the ELBO estimate. Same information, opposite direction.
Every stopping rule asks one question: "did we improve?" In ELBO terms that means "went up"; in loss terms, "went down". Mixing the two, or dropping the sign with an absolute value, turns a worsening into an "improvement", and the loop then waits forever while the model gets worse.
Three ways to say it:
- Picture: one number line, read left-to-right as ELBO, right-to-left as loss.
- Numbers: loss 552.40 means ELBO −552.40; loss falling to 552.39 is the ELBO rising to −552.39.
- Slogan: improvement = $ELBO_t - ELBO_{best}$ = $\mathcal{L}_{best} - \mathcal{L}_t$; keep the sign.
Two evaluations from the real run in this chapter's code (smoothed losses; the best so far is $\mathcal{L}_{best} = 552.4030$ from step 1800).
- Step 2000: $\mathcal{L}_t = 552.6436$. In ELBO terms $ELBO_{best} = -552.4030$, $ELBO_t = -552.6436$.
- Improvement $= ELBO_t - ELBO_{best} = -552.6436 - (-552.4030) = -0.2406$. In loss terms $\mathcal{L}_{best} - \mathcal{L}_t = 552.4030 - 552.6436 = -0.2406$. Same number: it got worse.
- Relative improvement $= -0.2406 / (552.4030 + \epsilon) = -4.36\times10^{-4}$: below any positive tolerance, so the patience counter goes up.
- Step 2200: $\mathcal{L}_t = 552.3972$. Improvement $= 552.4030 - 552.3972 = 0.0058$; relative $= 0.0058/552.4030 \approx 1.05\times10^{-5}$. It is a new best (so it is checkpointed) but smaller than a tolerance of $10^{-4}$ (so the patience counter still goes up).
- The trap: $|ELBO_t - ELBO_{best}|/|ELBO_{best}|$ at step 2000 gives $+4.36\times10^{-4} \gt 10^{-4}$. The absolute value calls a worsening a "meaningful change" and resets patience.
Let $\mathcal{L}_t = -\widehat{ELBO}_t$ be the (smoothed) loss at evaluation $t$ and $\mathcal{L}_{best}$ the lowest loss so far. The signed improvement and the relative improvement are
$$\Delta_t = ELBO_t - ELBO_{best} = \mathcal{L}_{best} - \mathcal{L}_t, \qquad \text{rel}_t = \frac{\Delta_t}{|ELBO_{best}| + \epsilon} = \frac{\mathcal{L}_{best} - \mathcal{L}_t}{|\mathcal{L}_{best}| + \epsilon}.$$- $|ELBO_{best}| = |\mathcal{L}_{best}|$, so both conventions give exactly the same number. Positive = better.
- The syllabus writes the size of the change, $\frac{|ELBO_t - ELBO_{best}|}{|ELBO_{best}| + \epsilon}$. That is fine for describing how much things moved; for the decision "did we improve?" use the signed $\Delta_t$, or equivalently count only changes where the ELBO went up.
- The loss can be negative (the ELBO can be positive, because densities of continuous data can exceed 1). Always divide by $|\cdot|$, never by the signed value, or the sign of the test flips.
Why do we need it?
A sign error in the stopping rule does not crash anything: the loop simply stops too early or never stops, and nobody notices until results look odd. Defining improvement once, carefully, avoids that silent failure.
Where is it used?
Every custom NumPyro/Pyro loop (svi.update and svi.evaluate return the loss), PyTorch and JAX training loops (losses are minimized), Keras-style EarlyStopping(mode="min"), and your forecasting model's relative-ELBO rule.
How is it used?
Work in one convention throughout. With NumPyro, keep losses: rel = (best_loss - cur) / (abs(best_loss) + eps), improvement if rel > tol. If you log ELBOs for humans, convert with a single minus sign at print time.
"Relative improvement is $|ELBO_t - ELBO_{best}|/|ELBO_{best}|$, so a big number means progress."
That measures how much the ELBO moved, in either direction. A big drop gives a big number too. For stopping decisions use the signed improvement.
"With NumPyro I compute (cur - best_loss) / best_loss and continue while it is above the tolerance."
For losses, lower is better: the improvement is best_loss - cur, and the denominator must be abs(best_loss) because losses can be negative.
"Our loop stops when the ELBO stops changing."
"Our loop stops when the ELBO stops improving by more than a relative tolerance: the signed gain over the best value so far, divided by its magnitude, stays below τ for P evaluations."
Model answer: "NumPyro's svi.update returns the loss, the negative ELBO estimate, so lower is better. I define improvement as best loss minus current loss, divide by the absolute best loss plus a small ε, and count it only if it exceeds the tolerance. That is identical to the ELBO formulation, and the sign matters: an absolute value would reset patience when the ELBO gets worse."
loss = −ELBO. Improvement $\Delta_t = ELBO_t - ELBO_{best} = \mathcal{L}_{best} - \mathcal{L}_t$ (positive = better).
$\text{rel}_t = \Delta_t / (|ELBO_{best}| + \epsilon)$, same in both conventions.
Trap: absolute value or flipped sign turns a worsening into "progress"; divide by $|\cdot|$ because losses can be negative.
Quick check: a model's best loss so far is −250 (ELBO +250). The current loss is −252. Did it improve, and what is the relative improvement?
Yes: the loss went down (−252 < −250), the ELBO went up from 250 to 252. $\text{rel} = (\mathcal{L}_{best} - \mathcal{L}_t)/|\mathcal{L}_{best}| = (-250 - (-252))/250 = 2/250 = 0.008$. Dividing by the signed −250 instead would give −0.008 and wrongly call it a worsening.
Relative ELBO early stopping: why a fixed number of nats is the wrong yardstick core
"Stop when the ELBO improves by less than 1 per check." Is 1 a lot? It depends completely on the size of the ELBO. For a tiny model whose ELBO is about −50, gaining 1 is a 2% jump: you are far from done. For a model on a million data points whose ELBO is about −500 000, gaining 1 is two parts in a million: nothing. The ELBO is a sum over data points, so its size grows with the amount of data (and changes with the units of $y$). A fixed absolute threshold means something different for every model and dataset.
A relative threshold asks instead: "did the ELBO improve by more than a small fraction of its own size?" That question has the same meaning at every scale, which is why one rule can serve many series and model sizes.
Three ways to say it:
- Picture: judge a salary raise as a percentage, not in dollars.
- Numbers: +0.5 nats is $10^{-2}$ of 50, $2.5\times10^{-4}$ of 2 000, and $10^{-6}$ of 500 000.
- Slogan: measure progress in fractions of the ELBO, not in nats.
The same +0.5 nat gain at three scales, tolerance $\tau = 10^{-4}$, $\epsilon = 10^{-8}$.
- $ELBO_{best} = -50$, $ELBO_t = -49.5$: rel $= 0.5/50 = 10^{-2} \gt 10^{-4}$: meaningful.
- $ELBO_{best} = -2\,000$, $ELBO_t = -1\,999.5$: rel $= 0.5/2\,000 = 2.5\times10^{-4} \gt 10^{-4}$: meaningful.
- $ELBO_{best} = -500\,000$, $ELBO_t = -499\,999.5$: rel $= 0.5/500\,000 = 10^{-6} \lt 10^{-4}$: not meaningful.
- Equivalently, the rule needs a gain of at least $\tau\,|ELBO_{best}|$ nats: $0.005$, $0.2$ and $50$ nats in the three cases. The threshold scales with the problem.
- In the real run of this chapter, $\tau|ELBO| = 10^{-4}\times552 = 0.055$ nats per evaluation, while the smoothed loss (an average of 200 steps) still wobbles by about $1.2/\sqrt{200} \approx 0.085$ nats. The threshold is close to the noise, so a single evaluation cannot be trusted: that is the job of patience (next section).
Relative early stopping. At each evaluation compute
$$\text{rel}_t = \frac{ELBO_t - ELBO_{best}}{|ELBO_{best}| + \epsilon}.$$The evaluation is a meaningful improvement if $\text{rel}_t \gt \tau$ (the tolerance). Training stops when there has been no meaningful improvement for $P$ evaluations in a row (patience), after at least a minimum number of steps.
- Scale-free: multiply every log-density by $c \gt 0$ (more data with the same fit per point, or a change of units that rescales the objective): numerator and denominator both scale by $c$, so $\text{rel}_t$ is unchanged (as long as $\epsilon$ is negligible). An absolute rule $ELBO_t - ELBO_{best} \gt \delta$ is not.
- What ε really does: it prevents division by zero, and it decides what happens near zero. When $|ELBO_{best}| \gg \epsilon$ the rule is relative; when $|ELBO_{best}| \ll \epsilon$ it becomes an absolute rule with threshold $\tau\epsilon$ nats. A tiny ε (like $10^{-8}$) only protects against a literal zero.
- Not shift-free: adding a constant to the objective (for example, rescaling $y$ shifts every Normal log-density by a constant) changes $|ELBO|$ and so the effective threshold. An ELBO that happens to sit near 0 makes the relative rule very strict, which is where a larger ε helps.
- Choosing τ: a rule of thumb for SVI is $10^{-4}$ to $10^{-5}$ per evaluation. Translate it into nats, $\tau|ELBO|$, and compare it with the noise of your evaluation: far below the noise means the decision is driven by noise (so you need more patience or smoothing).
Why do we need it?
One training loop often fits many models: different series, history lengths, numbers of regressors, guide sizes. Their ELBOs differ by orders of magnitude. A relative rule gives the same meaning of "done" for all of them without retuning a threshold per model.
Where is it used?
Your forecasting model's relative-ELBO stopping; SciPy's L-BFGS-B (ftol: stop when $(f^k - f^{k+1})/\max(|f^k|, |f^{k+1}|, 1) \le$ ftol, a relative rule whose "1" plays the role of ε); convergence checks in EM algorithms; Stan's ADVI (tol_rel_obj, default 0.01, on a running relative change).
How is it used?
Keep the best smoothed loss; every $k$ steps compute rel = (best_loss - cur) / (abs(best_loss) + eps); reset the patience counter if rel > tol, otherwise increase it. Start with tol $10^{-4}$, check $\tau|ELBO|$ against the noise, and pick ε on the scale of a meaningful absolute change if your ELBO can be near zero.
"The relative rule is scale-free, so it needs no tuning."
It is free of multiplicative scale, not of additive shifts, and it ignores noise. You still choose τ (and check $\tau|ELBO|$ against the evaluation noise), ε, patience and the minimum number of steps.
"ε is just there to avoid dividing by zero; any tiny number works."
If your ELBO can be close to zero, a tiny ε makes the rule extremely strict there (every noise blip looks like a big relative change). Pick ε on the scale of the smallest change in nats you care about; then the rule is absolute near zero and relative far from it.
"Early stopping in SVI prevents overfitting, as in deep learning."
In deep learning, early stopping watches a validation loss to avoid overfitting. Here it watches the training ELBO to stop when optimization has converged, to save compute. The posterior does not overfit by training longer; the guide just gets closer to its best fit.
This is why your forecasting loop uses a relative ELBO criterion: the ELBO of a two-year daily series and of a ten-year one differ by roughly the ratio of their lengths, and so do models with and without many regressors. One relative tolerance gives them all the same meaning of "converged", where any fixed number of nats would be too strict for some and too loose for others. If the same loop ever trains on standardized and raw targets, remember that rescaling $y$ shifts the ELBO, which moves the effective threshold.
"We stop when the ELBO changes by less than 0.01."
"We stop when the relative improvement $\frac{ELBO_t - ELBO_{best}}{|ELBO_{best}| + \epsilon}$ stays below τ for P evaluations, because the ELBO's magnitude grows with the data size."
Model answer: "The ELBO is a sum over observations, so its size depends on how much data and which model I fit. An absolute threshold would stop small problems too early and large ones too late. Dividing the improvement by the magnitude of the best ELBO makes the rule invariant to rescaling. ε guards the case where the ELBO is near zero; there the rule becomes absolute with threshold τε."
$\text{rel}_t = \dfrac{ELBO_t - ELBO_{best}}{|ELBO_{best}| + \epsilon}$; meaningful if $\text{rel}_t \gt \tau$ ⇔ gain $\gt \tau(|ELBO_{best}| + \epsilon)$ nats.
Invariant to multiplying the ELBO by $c \gt 0$; ε makes it absolute near zero. Rule of thumb τ = $10^{-4}$–$10^{-5}$; compare $\tau|ELBO|$ with the evaluation noise.
Trap: SVI early stopping = convergence, not overfitting control.
Quick check: with τ = $10^{-4}$ and ε = 1, what gain in nats is needed when $ELBO_{best} = -3\,000$? And when $ELBO_{best} = 0.2$?
$10^{-4}\times(3\,000 + 1) = 0.30$ nats. For $0.2$: $10^{-4}\times(0.2 + 1) = 0.00012$ nats; there ε dominates and the rule acts like an absolute threshold of about $\tau\epsilon = 10^{-4}$ nats.
Patience and minimum steps: do not quit after one bad reading core
One cloudy day in July does not mean summer is over. You wait for several cloudy days in a row before you pack away the fan. The ELBO estimate is noisy in the same way: even while the model is still improving, a single evaluation can come out worse than the best so far, just by bad luck. If the loop stopped at the first such reading, it would stop far too early.
Patience $P$ says: stop only after $P$ evaluations in a row without a meaningful improvement. The chance that noise alone produces $P$ bad readings in a row shrinks very fast as $P$ grows, while the extra cost grows only in a straight line: $P$ more evaluations of $k$ steps each. A minimum number of steps adds a second guard: never stop during the early phase, when the loss is still dropping erratically, Adam is still warming up its averages, or the guide's scales are still growing from their small starting value.
Three ways to say it:
- Picture: wait for several cloudy days, not one.
- Numbers: if an unlucky reading has a 40% chance, five unlucky readings in a row have a $0.4^5 = 1\%$ chance.
- Slogan: small patience = fast but premature; large patience = safe but slow.
The real log (rule: τ = $10^{-4}$, patience 5, minimum 2 000 steps, evaluate every 200 steps).
- Steps 200–1 800: every evaluation improves by more than τ, so the counter stays 0.
- Step 2 000: worse than the best (rel $-4.36\times10^{-4}$): counter 1. Patience 1 would stop here.
- Step 2 200: a new best by a hair (rel $1.04\times10^{-5} \lt \tau$): checkpointed, but counter 2.
- Steps 2 400, 2 600, 2 800: worse again: counter 3, 4, 5. At 2 800 the counter reaches 5 and $2\,800 \ge 2\,000$, so the loop stops.
- Cost of patience: the last meaningful improvement was at step 1 800; the loop spent $5\times200 = 1\,000$ more steps making sure. Why it is worth it: if, say, 40% of evaluations look bad by chance while the model is still improving, then patience 1 stops prematurely with probability 0.4 at each such evaluation, patience 3 with $0.4^3 = 0.064$, and patience 5 with $0.4^5 \approx 0.01$ (treating readings as independent, a simplification).
- Patience $P$: the number of consecutive evaluations without a meaningful improvement ($\text{rel}_t \le \tau$) that triggers a stop. The counter resets to 0 at every meaningful improvement. Measured in evaluations, so the waiting time is $P\cdot k$ steps.
- Minimum steps $t_{min}$: no stop before step $t_{min}$, whatever the counter says.
- Maximum steps: a hard cap, so a run that never meets the rule still ends (and you get a warning to investigate).
- Premature stopping: stopping while the model was still improving meaningfully, because noise or a temporary plateau hid the progress. Its cost is a worse guide (lower ELBO, and spreads or correlations that have not finished settling, in either direction); it is invisible in the loss log unless you look for it.
- Trade-off: the probability of a noise-caused stop falls roughly geometrically with $P$; the extra compute grows linearly, about $P\cdot k$ steps past the true convergence point.
Why do we need it?
The ELBO estimate is noisy and real training curves have flat stretches. Without patience and a minimum, a relative rule fires on the first unlucky or flat evaluation and returns a half-trained guide, with no error message.
Where is it used?
Keras EarlyStopping(patience=...) and PyTorch Lightning's early-stopping callback, learning-rate schedulers such as ReduceLROnPlateau(patience=...), gradient-boosting libraries (early_stopping_rounds in XGBoost and LightGBM), and your forecasting model's SVI loop.
How is it used?
Measure the noise of your evaluation, pick $k$ and τ, then choose $P$ so that a false stop is unlikely (often 3–10 evaluations), and $t_{min}$ beyond the early transient (look at a few full runs). Log the step at which the rule fired and the step of the best state.
"Patience 5 means waiting 5 steps."
Patience counts evaluations. With an evaluation every 200 steps, patience 5 waits 1 000 steps. Change $k$ and you change the waiting time.
"Every new best value resets the patience counter."
In the rule of this chapter only a meaningful improvement (rel > τ) resets it; a tiny new record is still checkpointed but does not buy more time. Other rules reset on any new best; Keras's EarlyStopping exposes this choice as min_delta (default 0: any improvement counts). Know which one your code does.
"With enough patience, the minimum number of steps is unnecessary."
They protect different things. Patience protects against noise; the minimum protects the early phase, where the curve can be flat for reasons that have nothing to do with convergence (Adam's warm-up, scales growing from 0.1, a low-rank factor starting at zero).
If your loop compares against the best ELBO so far with patience, the relevant numbers are: how noisy one evaluation is (set by the learning rate, particles, data and $k$), how big τ·|ELBO| is compared with that noise, and how long the flat stretches in your training curves last. Changepoint-heavy models can show exactly the plateau-then-drop shape when a group of $\delta_j$ starts to move late; a minimum number of steps and a patience window longer than the typical plateau guard against it.
"Patience is there in case the model gets stuck."
"Patience is there because the ELBO estimate is noisy: one non-improving evaluation is weak evidence, several in a row is strong evidence."
Model answer: "Each evaluation of the stochastic ELBO can look worse than the best by chance even while training is progressing. Requiring P consecutive evaluations without a relative improvement above τ makes false stops roughly geometrically unlikely in P, at a linear cost of about P·k extra steps. A minimum number of steps protects the early transient. Too small a patience stops prematurely; too large wastes compute."
Stop if (no rel > τ for $P$ consecutive evaluations) AND step ≥ $t_{min}$; cap at max steps.
False-stop chance ≈ $b^P$ (falls fast); extra cost ≈ $P\cdot k$ steps (grows linearly).
Trap: patience is in evaluations, not steps; plateaus longer than $P\cdot k$ cause premature stops.
Quick check: you evaluate every 50 steps with patience 8. A known plateau in your training curves lasts about 600 steps. Is that enough?
The rule waits $8\times50 = 400$ steps without gain, shorter than the 600-step plateau, so it can stop on the plateau. Use patience of 13 or more at $k = 50$ (or $k = 100$ with patience of 7 or more, which also halves each evaluation's noise variance), or a minimum number of steps beyond the plateau.
How often to evaluate, and what to evaluate: smoothing and low-noise checks
To judge whether a crowd is getting louder you would not listen for one second; you would listen for a while and average. The same holds for the loss. The loop does not need to judge progress at every step: it evaluates every $k$ steps, and it judges a smoothed number (for example the mean of the last $k$ losses) instead of one noisy value.
Smoothing has a price: an average of the last 200 steps describes where the parameters were on average during those steps, about 100 steps ago. That lag is harmless when progress is slow, which is exactly when stopping decisions are made. A second option is a dedicated low-noise evaluation: every $k$ steps, estimate the ELBO with many particles and the same random draws each time, so that differences between evaluations reflect the parameters, not the luck of the draws.
Three ways to say it:
- Picture: listen for a while before judging the noise level.
- Numbers: per-step wobble 1.2 nats; the mean of 200 steps wobbles by $1.2/\sqrt{200} \approx 0.085$.
- Slogan: judge averages, not single steps; average over more steps when the threshold is close to the noise.
- Per-step noise sd $\sigma = 1.2$ nats (measured in the real run of this chapter: the sd of single-particle losses at the end was 1.18).
- Window mean of $k = 50$: sd $1.2/\sqrt{50} = 0.17$. Of $k = 200$: $1.2/\sqrt{200} = 0.085$. Of $k = 800$: $0.042$. (Assuming the steps' noise is independent; it nearly is, since each step draws fresh random numbers.)
- The rule's threshold at τ = $10^{-4}$ and $|ELBO| \approx 552$ is $0.055$ nats: below the noise at $k = 50$, comparable at $k = 200$.
- Lag: a window of $k$ steps is on average $(k-1)/2$ steps old: about 100 steps for $k = 200$.
- An exponential moving average $m_t = \beta m_{t-1} + (1-\beta)\mathcal{L}_t$ with $\beta = 0.99$ has sd $\sigma\sqrt{(1-\beta)/(1+\beta)} = 1.2\sqrt{0.01/1.99} = 0.085$ and lag $\beta/(1-\beta) = 99$ steps: the same trade-off as a window of 200.
- Low-noise evaluation: 64 particles at a fixed key every 200 steps. Across 10 seeds of the real model, the checkpoints it picked were on average 0.29 nats better (judged by a precise 4 000-particle ELBO) than those picked by the 200-step training-loss average.
- Evaluation frequency $k$: the loop computes the stopping statistic every $k$ steps. Patience is counted in evaluations.
- Window mean: $\bar{\mathcal L} = \frac1k\sum_{i=t-k+1}^{t}\mathcal{L}_i$; noise sd ≈ $\sigma/\sqrt k$; lag ≈ $(k-1)/2$ steps. This is free: the losses are returned by
svi.updateanyway. - Exponential moving average (EMA): $m_t = \beta m_{t-1} + (1-\beta)\mathcal{L}_t$; noise sd ≈ $\sigma\sqrt{(1-\beta)/(1+\beta)}$; lag ≈ $\beta/(1-\beta)$ steps.
- Low-noise evaluation:
Trace_ELBO(num_particles=M).loss(key_eval, params, model, guide, *data)with a fixedkey_eval. Its noise sd is $\sigma/\sqrt M$, and reusing the key (common random numbers) makes differences between evaluations far less noisy still. (svi.evaluate(state, *data)also exists, but uses the training ELBO's particle count.) - Cost of evaluating: in a jitted loop, reading a loss into Python (
float(loss)) makes the host wait for the device. Keep losses as device arrays and convert once per evaluation, not once per step (details of this in Chapter 6.17).
Why do we need it?
The relative threshold is often close to the per-step noise. Without smoothing, the stopping rule mostly reacts to noise; with too much smoothing or too rare evaluations, it reacts late and wastes steps. These knobs set that balance.
Where is it used?
Every training dashboard (TensorBoard's smoothing slider is an EMA), Stan's ADVI (it evaluates the ELBO every eval_elbo = 100 iterations), NumPyro loops that average svi.update losses per chunk, and checkpoint selection with a held-fixed evaluation key.
How is it used?
Measure σ from a short run, pick $k$ so that $\sigma/\sqrt k$ is below or near $\tau|ELBO|$, accumulate losses on the device and average every $k$ steps. If you checkpoint based on the evaluation, consider a fixed-key, many-particle ELBO every $k$ steps; it costs about $M/k$ extra steps' worth of work.
"Evaluating every step gives the most information."
Each single-step loss is mostly noise, and pulling it into Python every step forces the device to synchronize, which slows a jitted loop. Evaluate every $k$ steps on an average.
"A lower smoothed loss means the latest parameters are better."
A window mean describes the parameters of the last $k$ steps on average, not the latest ones. When you checkpoint on it, you store the parameters at the end of the window, which is a small but real mismatch; a fixed-key evaluation of the current parameters avoids it.
In your forecasting loop, the choice of how often to check the relative ELBO improvement and on which number (a single loss, an average of the last few hundred, or a separate multi-particle estimate) sets how noisy each check is. The noisier the check, the more patience you need and the more the "best" checkpoint is partly luck (next section). Averaging the losses svi.update already returns is free; a fixed-key evaluation costs a little more and picks better checkpoints.
Window mean of $k$: sd $\sigma/\sqrt k$, lag $(k-1)/2$. EMA(β): sd $\sigma\sqrt{(1-\beta)/(1+\beta)}$, lag $\beta/(1-\beta)$.
Low-noise check: Trace_ELBO(num_particles=M).loss(fixed_key, params, model, guide, *data) every $k$ steps.
Trap: per-step evaluation = mostly noise + device syncs; smoothed loss lags the current parameters.
Quick check: per-step noise sd 2 nats, $|ELBO| \approx 10\,000$, τ = $10^{-5}$. Roughly how many steps should each evaluation average so its noise is at most the threshold?
Threshold $= 10^{-5}\times10\,000 = 0.1$ nats. Need $2/\sqrt k \le 0.1$, so $\sqrt k \ge 20$, $k \ge 400$ steps.
Best-state checkpointing: return the best state, not the last one core
Video games teach this: save before the boss fight. If the fight goes badly you reload the save instead of living with the mess. In SVI the "mess" can be a late blow-up (a learning rate that is a bit too high finally throws the parameters out of the bowl, a NaN from an extreme draw) or simply the jitter of the noise ball, which means the last state is a random point near the optimum rather than the best one seen. A checkpoint keeps a copy of the parameters at the best evaluation so far; at the end the loop returns that copy.
One subtlety. "Best" is judged with a noisy measuring stick. If you pick the fastest of ten runners from one shaky stopwatch reading each, the winner is partly the one whose reading was luckiest. The checkpoint with the lowest noisy loss is likewise partly lucky: its true loss is usually a bit worse than its measured one. This is the winner's curse. A less noisy evaluation shrinks it.
Three ways to say it:
- Picture: save the game before the boss; but the save button is pressed by a slightly shaky hand.
- Numbers: in the real run, the checkpoint's measured loss was 552.40 but a precise re-measurement gave 552.67, while the final state measured 552.35.
- Slogan: keep the best state to survive late failures; measure "best" with low noise to avoid keeping a lucky one.
What the checkpoint bought in the real run (precise judge: a 4 000-particle ELBO with one fixed key, applied to the stored parameters).
- Smoothed-training-loss rule: checkpoint from step 2 200 (measured 552.397), stopped at step 2 800. Precise loss: checkpoint 552.67, final state 552.35. The checkpoint won its noisy comparison by luck; the final state was actually better here.
- Same loop with a 64-particle fixed-key evaluation: checkpoint from step 2 000, stopped at 3 000. Precise loss: checkpoint 551.97, final 552.63. Now the checkpoint is genuinely better.
- Over 10 seeds: with the smoothed loss, checkpoints beat final states by 0.33 nats on average; with the low-noise evaluation by 0.38, and its checkpoints were 0.29 nats better than the smoothed-loss ones.
- These gaps are small (about $5\times10^{-4}$ of the loss) because this model trains smoothly. The checkpoint's real value shows when training goes wrong late: the final state can then be hundreds of nats worse, and the checkpoint is untouched (try "blows up late" in the simulator of the next section).
- Current state vs best state: the current state is whatever the optimizer holds now; the best state is a copy of the parameters at the evaluation with the lowest loss so far. Return the best state when the loop ends, for whatever reason (patience, max steps, NaN).
- What to store: for prediction,
svi.get_params(state)(the guide's constrained parameters) plus the step and loss. To resume training, the wholesvi_state(it also holds Adam's moments and the PRNG key). - Checkpoint frequency: only at evaluations (every $k$ steps). Each stored copy costs the guide's size in memory (for full-rank $\approx d^2/2$ numbers), and writing it to disk costs time, so do not checkpoint every step.
- JAX detail: JAX arrays are immutable, so keeping a reference to
svi.get_params(state)is already a true snapshot. Exception: if you jit the update withdonate_argnums, old buffers are deleted, and a stored reference raises "Array has been deleted"; then copy withjax.tree.map(jnp.copy, params). - Off by one step:
svi.update(state)returns the loss of the parameters it was given, together with the state after the step. The difference is one tiny step; with smoothed or separate evaluations it does not matter, but be precise if asked. - Winner's curse: choosing the minimum of noisy evaluations biases the chosen state's measured loss downward; its true loss is higher. More evaluations compared and noisier evaluations make it worse. Remedies: smoothing, more particles, a fixed evaluation key, re-evaluating the top candidates.
- Reproducibility: with the same seed, data, code, library versions and hardware, the run (and so the checkpoint) is repeated exactly. Record the seed, the stopping settings, the best step and loss, $d$, the guide type and rank, and the versions with the saved parameters.
Why do we need it?
Stochastic optimization does not improve monotonically, and it can fail late. Returning the last state hands you a random point of the noise ball, or a broken state after a blow-up; returning the best state gives the best fit you actually reached.
Where is it used?
Keras ModelCheckpoint(save_best_only=True) and EarlyStopping(restore_best_weights=True), PyTorch Lightning checkpoint callbacks, XGBoost's best_iteration, Orbax checkpoints in JAX, and your forecasting model's best-state checkpointing.
How is it used?
At each evaluation, if the loss is the lowest so far, store best_params = svi.get_params(state), the step and the loss. When the loop ends, use best_params with guide.sample_posterior or Predictive. Log the best step next to the stop step; a best step far before the stop is worth a look.
"The checkpoint is the truly best state the run ever visited."
It is the state with the best measured loss among the evaluated ones. Measurements are noisy, so it is partly chosen by luck, and its measured loss is optimistic. Report a fresh, precise evaluation if you need its loss.
"In JAX I must deep-copy the parameters before storing them."
Not normally: JAX arrays are immutable, so a stored reference never changes. You must copy only if you donate buffers to the jitted update (donate_argnums), because donated arrays are deleted.
"To resume training later, saving the best parameters is enough."
Resuming exactly also needs the optimizer state (Adam's moments and step count) and the PRNG key, i.e. the whole svi_state. Restarting Adam from scratch at the best parameters changes the trajectory.
Your forecasting model's best-state checkpointing returns the parameters of the best evaluation instead of the final ones. Two things are worth being able to say about it: it is mainly insurance against late degradation (a blow-up or NaN after many good steps), and on smooth runs its advantage is small because "best" is measured with noise. If your loop compares against the best ELBO so far, the same number serves both the stopping rule and the checkpoint, so its noise level matters twice.
"We checkpoint the best state, so we return the optimal parameters."
"We return the parameters with the best evaluated objective, which protects against late degradation; with a noisy objective, that choice is itself slightly optimistic."
Model answer: "SVI's objective is a noisy estimate and the parameters jitter with a constant learning rate, so the last state is not the best one and training can degrade late. At each evaluation I store the parameters if the smoothed loss is the lowest so far, and return them at the end. Because the minimum of noisy values is biased, I evaluate on a smoothed loss or a fixed-key multi-particle ELBO, and I record the step, loss, seed and settings with the checkpoint so the run is reproducible."
At each evaluation: if loss < best → best_params = svi.get_params(state), store step & loss; return best_params at the end.
Store the whole svi_state to resume (Adam moments + key). JAX arrays immutable → reference = snapshot (copy only with donate_argnums).
Trap: the checkpoint is chosen by a noisy number (winner's curse); its value is insurance against late failures more than a big gain on smooth runs.
Quick check: a run stops by patience at step 9 000; the checkpoint is from step 3 000. What might this tell you?
After step 3 000 no evaluation beat the best one, so the patience counter grew at every evaluation; stopping only at 9 000 means the waiting window $P\cdot k$ (or the minimum number of steps) was about 6 000 steps or more. Either the rule waited far longer than needed (wasted compute), or the step-3 000 reading was an unusually lucky one that later, honest evaluations could not beat. Re-evaluate the checkpoint and a late state precisely before trusting either.
The whole stopping rule, simulated: six knobs, one decision core
Each knob so far was studied alone. In a real loop they act together. Two knobs shape the curve: the learning rate (how fast the loss falls and how much the parameters jitter at the end) and the noise (how much each reported loss wobbles). Four knobs shape the decision: the tolerance τ, the patience $P$, the minimum number of steps, and the evaluation interval $k$. A good setting stops soon after the curve has really flattened, never on a temporary plateau, and always returns a good state even if training goes wrong late.
The simulator below is a stand-in for SVI (not NumPyro itself): a loss curve with the same features as real runs (a fast drop, a noisy floor whose height grows with the learning rate, optional plateaus and late blow-ups) and exactly the stopping rule of this chapter's code. Because it is a simulation, it can show what a real run never shows: the true loss behind the noisy numbers.
Three ways to say it:
- Picture: a referee (the rule) watching a blurry replay (the noisy loss) and deciding when the game is over.
- Numbers: defaults stop at step 3 300 with the checkpoint 0.12 nats above the optimum; on the plateau scenario the same rule stops 12 nats short.
- Slogan: tune the rule on runs where you can see the truth, then trust it on runs where you cannot.
Three scenarios with the default rule (lr 0.01, noise 1.2 nats per step, τ = $10^{-4}$, $P = 5$, minimum 1 000 steps, $k = 100$, default noise draw).
- Smooth: stops at step 3 300 by patience; checkpoint from step 2 800, 0.12 nats above the optimum. Good.
- Plateau then improve: also stops at 3 300, but the checkpoint is 12.1 nats above the optimum: it stopped on the plateau. With patience 20 it stops at 6 600 and the checkpoint is 0.19 above: fixed, at the price of 3 300 more steps.
- Blows up late: stops at 2 100; the final state is 167 nats above the optimum, the checkpoint (step 1 600) 9.8 nats above. Without the checkpoint you would ship the broken state.
- Too small a learning rate (0.001, smooth): never satisfies the rule and hits the 8 000-step cap, still 91 nats above the optimum. Too large (0.05): stops early at 1 600 with a checkpoint 2 nats above the optimum, because the parameters jitter in a wide noise ball.
The complete rule (exactly what the code in the next section does):
- Every step:
state, loss = update(state, *data); keep the loss. - Every $k$ steps: $\text{cur}$ = mean of the last $k$ losses (or a low-noise evaluation). If it is not finite, stop.
- $\text{rel} = (\text{best} - \text{cur})/(|\text{best}| + \epsilon)$ (∞ at the first evaluation).
- If $\text{cur} \lt \text{best}$: checkpoint the parameters, $\text{best} = \text{cur}$.
- If $\text{rel} \gt \tau$: counter = 0, else counter + 1.
- If step ≥ $t_{min}$ and counter ≥ $P$: stop. Always stop at the maximum number of steps.
- Return the checkpoint, the best step, the stop step and the reason.
How to read a finished run when you cannot see the truth: compare the stop step with the best step (a gap of about $P\cdot k$ is normal); look at the smoothed curve for plateaus; re-run once with a stricter rule (smaller τ or larger $P$) and check whether the ELBO or the reported intervals change meaningfully.
Why do we need it?
Settings that work alone can fail together: a tolerance that is fine at one learning rate is too strict at another; a patience that handles noise can be too short for a plateau. Seeing them interact on a run with a known truth builds the judgement to set them on real models.
Where is it used?
Tuning the stopping rule of your forecasting loop, choosing defaults for a training library that many models share, and explaining in an interview why your loop stopped where it did (or why it would have failed with other settings).
How is it used?
Start from τ = $10^{-4}$, $P$ = 5, $k$ = 100–200, a minimum around the end of the early drop. Check on a few full-length runs that the stop comes after the curve flattens; if your models show plateaus, lengthen $P\cdot k$ or the minimum; keep the checkpoint always on.
"If the loop stopped by patience, the model has converged."
It means the rule saw no meaningful improvement for $P$ evaluations. A plateau, a too-small learning rate or unlucky noise produce the same signal. Check the curve and re-run with a stricter rule once.
"Hitting the maximum number of steps is fine; it trained longer."
It means the rule never fired: the model may still be improving (learning rate too small) or the tolerance is below what the noise allows. Treat it as a warning to investigate, not as success.
Use this as a checklist for your forecasting loop's settings: is the evaluation noise at your learning rate small compared with τ·|ELBO|? Is $P\cdot k$ longer than the plateaus you have seen in full-length runs? Does the minimum number of steps cover the early drop? Is the checkpoint always returned, including after a NaN? Being able to say how your rule behaves in each of the three scenarios is a strong interview answer.
Curve knobs: lr (speed vs noise-ball height), noise. Decision knobs: τ, $P$, $t_{min}$, $k$.
Failure modes: premature (plateau / lucky reading), never fires (tiny lr, τ below noise), late blow-up (rescued by the checkpoint).
Trap: "stopped by patience" ≠ "converged"; "hit max steps" ≠ success.
Quick check: in the simulator, why does a larger learning rate (0.05) give a checkpoint about 2 nats above the optimum even on the smooth scenario?
A larger learning rate makes the parameters jitter in a wider noise ball around the optimum, so every state the loop can checkpoint is, on average, further from the optimum (the floor of the true loss rises with the learning rate). The stopping rule cannot fix that; a smaller learning rate late in training (a schedule) can.
The loop in NumPyro and JAX, line by line core
Everything in this chapter fits in about thirty lines of Python. The heavy work, one SVI step, is a single call to a JIT-compiled function: JAX traces svi.update once, compiles it with XLA, and reuses the compiled program for every later step (how and why is Chapter 6.17). Around it sits ordinary Python: a counter, a running list of losses, a few comparisons and a stored copy of the best parameters. The split matters: the stopping decisions need concrete numbers, so they live outside the compiled step.
Three ways to say it:
- Picture: a fast engine (the jitted update) with a simple dashboard and a driver (the Python loop) deciding when to park.
- Numbers: the real run below: 14 evaluations, stop at step 2 800, checkpoint from step 2 200; the whole script (three fits) runs in about 5 seconds on a laptop.
- Slogan: compile the step, keep the decisions in Python, convert to floats only when you evaluate.
The real log of the code below (two years of daily data, level + trend + weekly Fourier terms, full-rank guide with $d = 9$; rule τ = $10^{-4}$, $P = 5$, minimum 2 000 steps, $k = 200$):
step 200 loss 2122.360 ELBO -2122.360 rel inf bad 0
step 400 loss 1239.132 ELBO -1239.132 rel 4.16e-01 bad 0
step 600 loss 638.937 ELBO -638.937 rel 4.84e-01 bad 0
step 800 loss 554.256 ELBO -554.256 rel 1.33e-01 bad 0
step 1000 loss 553.041 ELBO -553.041 rel 2.19e-03 bad 0
step 1200 loss 552.891 ELBO -552.891 rel 2.70e-04 bad 0
step 1400 loss 552.675 ELBO -552.675 rel 3.92e-04 bad 0
step 1600 loss 552.567 ELBO -552.567 rel 1.94e-04 bad 0
step 1800 loss 552.403 ELBO -552.403 rel 2.98e-04 bad 0
step 2000 loss 552.644 ELBO -552.644 rel -4.36e-04 bad 1
step 2200 loss 552.397 ELBO -552.397 rel 1.04e-05 bad 2
step 2400 loss 552.580 ELBO -552.580 rel -3.31e-04 bad 3
step 2600 loss 552.694 ELBO -552.694 rel -5.38e-04 bad 4
step 2800 loss 552.737 ELBO -552.737 rel -6.15e-04 bad 5
stopped at step 2800 | checkpoint from step 2200
- Steps 200–800: the loss falls from 2 122 to 554; relative improvements of 13–48% per evaluation.
- Steps 1 000–1 800: still improving by more than $10^{-4}$ each time (0.1–1.2 nats), so
badstays 0. - Step 2 000: worse than the best (552.403):
bad= 1. - Step 2 200: a new best by 0.006 nats (rel $1.04\times10^{-5}$): checkpointed, but below τ, so
bad= 2. (This tiny number is a difference of two nearly equal float32 values; on repeated runs its last digit varied between 1.03 and 1.05.) - Steps 2 400–2 800: worse each time; at 2 800
bad= 5 and step ≥ 2 000: stop. Return the parameters from step 2 200.
The loop (from the runnable script at the end of the chapter):
def fit_svi(model, guide, *args, lr=0.01, max_steps=20_000, eval_every=200,
rel_tol=1e-4, patience=5, min_steps=2_000, eval_particles=0,
seed=0, verbose=False):
svi = SVI(model, guide, Adam(lr), Trace_ELBO())
state = svi.init(jax.random.PRNGKey(seed), *args) # runs model + guide once, creates the params
update = jax.jit(svi.update) # traced + compiled on the first call
if eval_particles: # optional low-noise evaluation, fixed key
ev, key_ev = Trace_ELBO(num_particles=eval_particles), jax.random.PRNGKey(seed + 1)
eval_loss = jax.jit(lambda p: ev.loss(key_ev, p, model, guide, *args))
best_loss, best_params, best_step = np.inf, svi.get_params(state), 0
bad, window, log = 0, [], []
for step in range(1, max_steps + 1):
state, loss = update(state, *args) # loss = -ELBO estimate (lower is better)
window.append(loss)
if step % eval_every: # evaluate only every k steps
continue
if eval_particles:
cur = float(eval_loss(svi.get_params(state)))
else:
cur = float(jnp.mean(jnp.stack(window))) # smoothed: mean loss of the last k steps
window = []
if not np.isfinite(cur): # NaN or inf: stop and keep the best state
break
rel = (best_loss - cur) / (abs(best_loss) + 1e-8) if np.isfinite(best_loss) else np.inf
if cur < best_loss: # best so far -> checkpoint it
best_loss, best_params, best_step = cur, svi.get_params(state), step
bad = 0 if rel > rel_tol else bad + 1 # only a MEANINGFUL gain resets patience
log.append((step, cur, rel, bad))
if verbose:
print(f"step {step:5d} loss {cur:8.3f} ELBO {-cur:8.3f} rel {rel:9.2e} bad {bad}")
if step >= min_steps and bad >= patience:
break
return best_params, dict(best_loss=best_loss, best_step=best_step, stop_step=step,
final_params=svi.get_params(state), log=log)
jax.jit(svi.update): the first call traces and compiles (a one-off cost); later calls with arguments of the same shapes and dtypes reuse the compiled program. Changing a data shape triggers a recompile (Chapter 6.17). An equivalent design callssvi.updateinside your own jitted function, for example one that runs $k$ steps withjax.lax.scanand returns their losses.window.append(loss)keeps device arrays;float(...)happens once per evaluation, so the host waits for the device only every $k$ steps.- The first evaluation has no previous best, so
relis set to ∞ (a meaningful improvement by definition) instead of computing ∞/∞. - Two separate decisions: is it the best so far? (checkpoint) and was it a meaningful gain? (patience). A tiny new record is stored but does not reset patience.
- On NaN or ∞ the loop breaks and still returns the checkpoint. The stopping logic is plain Python because it needs concrete numbers, which a traced (jitted) function does not have.
Why do we need it?
svi.run has no early stopping and returns the final state. A custom loop adds a principled stopping rule, a best-state checkpoint, NaN handling and logging while keeping each step fully compiled and fast.
Where is it used?
Your forecasting model's training loop; NumPyro's documentation recommends init/update/evaluate for early stopping; the same pattern appears in Flax/Optax training loops and in Pyro with PyTorch.
How is it used?
Build the model and guide, call fit_svi(model, guide, *data), log the returned best_step, stop_step and best_loss with the run, then sample from the guide with guide.sample_posterior(key, best_params, sample_shape=(S,)).
"jax.jit(svi.update) recompiles at every call."
It compiles on the first call and reuses the program while the arguments keep the same shapes and dtypes. Passing data of a new shape (a longer series, a different batch size) triggers a new compile.
"Put the early-stopping if inside the jitted function to make it faster."
Inside jit values are abstract tracers, so a Python if on them fails. The decisions belong in the Python loop (or in lax.cond/lax.while_loop if you really need them compiled).
"Calling float(loss) every step costs nothing."
It forces the host to wait for the device at every step. Keep losses as arrays and convert once per evaluation.
This is a minimal version of the loop described on your résumé: JIT-compiled updates, relative-ELBO early stopping with patience and a minimum number of steps, and best-state checkpointing. In an interview, walk through it in this order: svi.init; the jitted update and why it is compiled once; the evaluation every $k$ steps on a smoothed loss; the signed relative improvement against the best; the two separate decisions (checkpoint vs patience); the stop condition; returning the checkpoint. Then mention what you would add: NaN handling, logging the best step, and a low-noise evaluation if checkpoint quality matters.
"The loop runs until the ELBO converges, then returns the parameters."
"The loop runs jitted SVI steps, evaluates a smoothed loss every k steps, stops after P evaluations without a relative gain above τ (and not before a minimum), and returns the best checkpointed parameters, not the last ones."
Model answer: "I initialize SVI, jit svi.update so each step runs as one compiled XLA program, and call it in a Python loop. Every k steps I average the returned losses (loss = −ELBO), compute the signed improvement over the best loss divided by its magnitude, checkpoint the parameters if it is a new best, and increase a patience counter unless the gain exceeds the tolerance. After a minimum number of steps, P evaluations without meaningful gain end training, and I return the checkpoint. NaNs stop the loop early and also return the checkpoint."
update = jax.jit(svi.update); loop: state, loss = update(state, *data); every k: cur = float(mean(window)).
rel = (best - cur)/(abs(best) + eps); if cur < best: checkpoint; bad = 0 if rel > tol else bad + 1; stop if step >= min_steps and bad >= patience.
Trap: decisions outside jit; convert to float only at evaluations; return the checkpoint (also after NaN).
Quick check: in the real log, why did step 2 200 count as "bad" although it set a new best?
Its relative gain was $1.04\times10^{-5}$, below τ = $10^{-4}$. The loop separates the two questions: any new best is checkpointed, but only a gain larger than τ resets the patience counter.
Recap, cheat sheet and practice
- The SVI loop:
svi.initonce, then a jittedsvi.updaterepeatedly (noisy ELBO → gradient → optimizer step). With a constant learning rate the parameters jitter in a noise ball whose size grows with the learning rate; "converged" means "no meaningful improvement any more". - Signs:
svi.updatereturns the loss = −ELBO estimate. Improvement $= ELBO_t - ELBO_{best} = \mathcal{L}_{best} - \mathcal{L}_t$; keep the sign and divide by the absolute value. - Relative stopping: $\text{rel}_t = \frac{ELBO_t - ELBO_{best}}{|ELBO_{best}| + \epsilon} \gt \tau$ counts as meaningful. Scale-free under multiplication; ε turns it into an absolute rule near zero. Rule of thumb τ = $10^{-4}$–$10^{-5}$; compare $\tau|ELBO|$ with the evaluation noise.
- Patience $P$ (in evaluations) protects against noise: false stops fall roughly like $b^P$, cost grows like $P\cdot k$ steps. Minimum steps protect the early phase. Plateaus longer than $P\cdot k$ still cause premature stops.
- Evaluate every $k$ steps on a smoothed loss (sd $\sigma/\sqrt k$, lag $(k-1)/2$) or a fixed-key multi-particle ELBO; convert to Python floats only then.
- Checkpoint the best state and return it (also after NaN). It mainly insures against late failures; "best" chosen on noisy numbers is partly luck (winner's curse), less so with low-noise evaluation.
Cheat sheet
| Idea | Formula / code | In words |
|---|---|---|
| One step | state, loss = jax.jit(svi.update)(state, *data) | loss = −ELBO estimate at the parameters given |
| Improvement | $\Delta_t = ELBO_t - ELBO_{best} = \mathcal{L}_{best} - \mathcal{L}_t$ | positive = better; keep the sign |
| Relative improvement | $\text{rel}_t = \Delta_t / (|ELBO_{best}| + \epsilon)$ | meaningful if $\gt \tau$; threshold $\tau(|ELBO| + \epsilon)$ nats |
| Patience | stop if $P$ evaluations in a row with $\text{rel} \le \tau$ and step $\ge t_{min}$ | waits $P\cdot k$ steps; false stops ≈ $b^P$ |
| Smoothing | window: sd $\sigma/\sqrt k$, lag $(k-1)/2$; EMA: sd $\sigma\sqrt{\tfrac{1-\beta}{1+\beta}}$, lag $\tfrac{\beta}{1-\beta}$ | judge averages, not single steps |
| Noise ball (SGD) | Var $= \frac{\text{lr}\,\sigma_g^2}{a(2-\text{lr}\,a)}$, diverges if $\text{lr}\,a \gt 2$ | bigger lr: faster but noisier floor |
| Checkpoint | if cur < best: best_params = svi.get_params(state) | JAX arrays immutable; copy only with donate_argnums |
| Low-noise check | Trace_ELBO(num_particles=M).loss(fixed_key, params, model, guide, *data) | common random numbers → better checkpoints |
import numpy as np
import jax, jax.numpy as jnp
import numpyro, numpyro.distributions as dist
from numpyro.infer import SVI, Trace_ELBO
from numpyro.infer.autoguide import AutoMultivariateNormal
from numpyro.optim import Adam
# 1) Two years of daily data: level + trend + weekly seasonality (Fourier order 3) + noise
rng = np.random.default_rng(42)
T = 730
t = np.arange(T) / T # time scaled to [0, 1)
def fourier(day, period, order):
x = 2 * np.pi * day[:, None] * np.arange(1, order + 1) / period
return np.concatenate([np.sin(x), np.cos(x)], axis=1)
F = fourier(np.arange(T), 7.0, 3) # 6 weekly columns
y = 5 + 2 * t + F @ np.array([0.8, -0.3, 0.1, 0.5, 0.2, -0.1]) + rng.normal(0, 0.5, T)
t, F, y = jnp.array(t), jnp.array(F), jnp.array(y)
def model(t, F, y=None):
m = numpyro.sample("m", dist.Normal(0, 10))
k = numpyro.sample("k", dist.Normal(0, 10))
beta = numpyro.sample("beta", dist.Normal(0, 1).expand([F.shape[1]]).to_event(1))
sigma = numpyro.sample("sigma", dist.HalfNormal(1.0))
numpyro.sample("y", dist.Normal(m + k * t + F @ beta, sigma), obs=y)
# 2) The custom loop
def fit_svi(model, guide, *args, lr=0.01, max_steps=20_000, eval_every=200,
rel_tol=1e-4, patience=5, min_steps=2_000, eval_particles=0,
seed=0, verbose=False):
svi = SVI(model, guide, Adam(lr), Trace_ELBO())
state = svi.init(jax.random.PRNGKey(seed), *args) # runs model + guide once, creates the params
update = jax.jit(svi.update) # traced + compiled on the first call
if eval_particles: # optional low-noise evaluation, fixed key
ev, key_ev = Trace_ELBO(num_particles=eval_particles), jax.random.PRNGKey(seed + 1)
eval_loss = jax.jit(lambda p: ev.loss(key_ev, p, model, guide, *args))
best_loss, best_params, best_step = np.inf, svi.get_params(state), 0
bad, window, log = 0, [], []
for step in range(1, max_steps + 1):
state, loss = update(state, *args) # loss = -ELBO estimate (lower is better)
window.append(loss)
if step % eval_every: # evaluate only every k steps
continue
if eval_particles:
cur = float(eval_loss(svi.get_params(state)))
else:
cur = float(jnp.mean(jnp.stack(window))) # smoothed: mean loss of the last k steps
window = []
if not np.isfinite(cur): # NaN or inf: stop and keep the best state
break
rel = (best_loss - cur) / (abs(best_loss) + 1e-8) if np.isfinite(best_loss) else np.inf
if cur < best_loss: # best so far -> checkpoint it
best_loss, best_params, best_step = cur, svi.get_params(state), step
bad = 0 if rel > rel_tol else bad + 1 # only a MEANINGFUL gain resets patience
log.append((step, cur, rel, bad))
if verbose:
print(f"step {step:5d} loss {cur:8.3f} ELBO {-cur:8.3f} rel {rel:9.2e} bad {bad}")
if step >= min_steps and bad >= patience:
break
return best_params, dict(best_loss=best_loss, best_step=best_step, stop_step=step,
final_params=svi.get_params(state), log=log)
# 3) Run it (d = 9 latent numbers, so a full-rank guide is cheap)
guide = AutoMultivariateNormal(model)
params, info = fit_svi(model, guide, t, F, y, verbose=True)
print("stopped at step", info["stop_step"], "| checkpoint from step", info["best_step"])
# step 200 loss 2122.360 ELBO -2122.360 rel inf bad 0
# ... (the full log is printed in the section "The loop in NumPyro and JAX")
# step 2800 loss 552.737 ELBO -552.737 rel -6.15e-04 bad 5
# stopped at step 2800 | checkpoint from step 2200
# 4) Is the checkpoint better than the final state? Judge both with a precise ELBO estimate
judge = Trace_ELBO(num_particles=4000)
def precise_loss(p):
return float(judge.loss(jax.random.PRNGKey(123), p, model, guide, t, F, y))
print("precise loss checkpoint %.2f final %.2f" % (precise_loss(params), precise_loss(info["final_params"])))
# precise loss checkpoint 552.67 final 552.35 (here the checkpoint won its noisy comparison by luck)
# 5) Same loop, but evaluate with 64 particles at a fixed key instead of the smoothed training loss
p64, i64 = fit_svi(model, guide, t, F, y, eval_particles=64)
print("low-noise eval: stop", i64["stop_step"], "checkpoint", i64["best_step"],
"| precise loss checkpoint %.2f final %.2f" % (precise_loss(p64), precise_loss(i64["final_params"])))
# low-noise eval: stop 3000 checkpoint 2000 | precise loss checkpoint 551.97 final 552.63
# 6) An impatient rule (patience 1, no minimum, tolerance 1e-3)
p1, i1 = fit_svi(model, guide, t, F, y, patience=1, min_steps=0, rel_tol=1e-3)
print("impatient: stop", i1["stop_step"], "| precise loss %.2f" % precise_loss(p1))
# impatient: stop 1200 | precise loss 552.35 (this easy model was already close at step 1200)
# 7) Use the checkpoint: posterior means (truth: m = 5, k = 2, sigma = 0.5)
post = guide.sample_posterior(jax.random.PRNGKey(1), params, sample_shape=(2000,))
print({name: round(float(post[name].mean()), 2) for name in ["m", "k", "sigma"]})
# {'m': 5.01, 'k': 1.92, 'sigma': 0.49}
# (float32 sums on CPU can differ in the last digit between runs, e.g. rel at step 2200: 1.03e-05 to 1.05e-05.)
1. The best loss so far is 552.40 and the current smoothed loss is 552.64 (NumPyro losses). What is the signed relative improvement?
2. Why does the loop compare the improvement with a fraction of $|ELBO_{best}|$ instead of a fixed number of nats?
3. You evaluate every $k = 250$ steps with patience $P = 4$. How many steps without a meaningful improvement does the loop wait before stopping (after the minimum)?
4. Which statement about best-state checkpointing is true?
5. τ = $10^{-4}$, ε = 10, and $ELBO_{best} = 0.5$. Roughly what gain in nats counts as meaningful?
6. Each step's loss has noise sd 1.6 nats. What is the noise sd of the mean of a window of 64 steps?
Practice problems
A. Apply the rule by hand. Smoothed losses at evaluations 1–5: 1 000.00, 990.00, 989.95, 989.97, 989.90. τ = $10^{-4}$, patience 3, no minimum. When does it stop, and which state is returned?
- Eval 1: first evaluation, rel = ∞ → checkpoint (1 000.00), bad = 0.
- Eval 2: rel $= (1\,000 - 990)/1\,000 = 0.01 \gt 10^{-4}$ → checkpoint (990.00), bad = 0.
- Eval 3: rel $= (990.00 - 989.95)/990.00 = 5.05\times10^{-5} \lt 10^{-4}$ → new best, checkpoint (989.95), but bad = 1.
- Eval 4: rel $= (989.95 - 989.97)/989.95 = -2.0\times10^{-5}$ → not a new best, bad = 2.
- Eval 5: rel $= (989.95 - 989.90)/989.95 = 5.05\times10^{-5} \lt 10^{-4}$ → new best, checkpoint (989.90), bad = 3 = patience → stop. Return the parameters from evaluation 5.
B. Two models: ELBO ≈ −120 and ELBO ≈ −1 200 000. Compare an absolute rule (stop below 1 nat per evaluation) with a relative rule (τ = $10^{-4}$).
Relative thresholds: $10^{-4}\times120 = 0.012$ nats and $10^{-4}\times1\,200\,000 = 120$ nats. The absolute rule's 1 nat is $8.3\times10^{-3}$ of the small ELBO (it would stop the small model while it still gains almost 1% per evaluation: premature) and $8.3\times10^{-7}$ of the big one (it would keep the big model running long after gains of 100 nats, i.e. $10^{-4}$, stopped mattering, until noise alone ends it). One relative tolerance gives both the same meaning.
C. Plain SGD on $\tfrac12\theta^2$ with gradient noise sd 1: compare lr = 0.2 and lr = 0.02 (leftover loss and speed).
- lr = 0.2: Var $= 0.2/(1\cdot(2 - 0.2)) = 0.111$, leftover $E[f] = 0.056$; distance halves every $\ln 0.5/\ln 0.8 \approx 3.1$ steps.
- lr = 0.02: Var $= 0.02/1.98 = 0.0101$, leftover $0.0051$; halving every $\approx 34$ steps.
- Ten times smaller lr: about ten times lower floor, ten times slower. A schedule (large lr first, small later) gets both.
D. While the model is still improving, each evaluation has a 30% chance of looking non-improving by noise alone. What patience keeps the chance of a false stop at a given point below 1%?
Treating evaluations as independent, $P$ bad readings in a row have probability $0.3^P$: $0.3^3 = 0.027$, $0.3^4 = 0.0081$. So $P = 4$. The cost: $4k$ steps of waiting after the real convergence.
E. A colleague writes rel = abs(cur - best_loss) / abs(best_loss) and resets patience when rel > tol. What goes wrong, and how do you fix it?
The absolute value counts a worsening as a meaningful change. If the loss jumps up (a noisy bad evaluation, or real degradation), patience is reset and the loop keeps running, possibly forever while the model gets worse. Fix: rel = (best_loss - cur) / (abs(best_loss) + eps) (positive = better, in NumPyro's loss convention), and reset only when rel > tol. Keep abs only in the denominator, because losses can be negative.
F. (Interview) "Walk me through your custom SVI training loop and justify each design choice. What would you change if the model became ten times bigger?"
"I call svi.init once, then a jitted svi.update in a Python loop, so each step is one compiled XLA program and the loop logic stays in Python, where it has concrete numbers. svi.update returns the loss, the negative ELBO estimate. Every k steps I average the recent losses to beat the per-step noise and compute the signed improvement over the best loss, divided by its magnitude: relative, because the ELBO's size depends on the dataset and model, so one tolerance works everywhere. A new best is checkpointed; only a gain above the tolerance resets a patience counter, because single evaluations are noisy; after a minimum number of steps, P evaluations without meaningful gain stop training. I return the checkpoint rather than the last state, which protects against late blow-ups and NaNs. If the model were ten times bigger: d grows, so I would re-check the guide (full-rank memory grows like d², so probably low-rank), expect longer compiles and steps, re-measure the loss noise and the ELBO scale (the relative tolerance carries over, but patience and k may need to change), and consider a fixed-key multi-particle evaluation so checkpoints are less driven by luck."
SVI vs NUTS
You now know both engines. NUTS walks around the posterior and records where it goes (Chapters 6.9–6.10). SVI picks a simpler distribution and tunes it until it overlaps the posterior as well as it can (Chapters 6.11–6.14). This chapter puts them side by side on the same posteriors: how fast each one is, how each one grows with data and parameters, what kind of error each one makes, how good its uncertainty is, which alarms each one needs, and when each one is the right tool. Your syllabus marks this as a top-priority (P0) interview topic, so every section ends with words you can say out loud.
- Explain the core difference: SVI optimizes a stand-in distribution; NUTS samples the posterior itself
- Count the cost of each in gradient evaluations, and say why NUTS pays per draw and SVI pays per step
- Explain scalability: why SVI can use minibatches (scaled by $N/B$) and plain NUTS cannot, and what grows with the number of parameters
- Separate the two kinds of error: approximation error (SVI: a bias that more steps cannot remove) and Monte Carlo error (NUTS: a wobble that shrinks like $1/\sqrt{\text{ESS}}$)
- Describe uncertainty quality: why mean-field SVI is under-dispersed on correlated posteriors, and why that can push a decision quantity like $P(\theta_B \gt \theta_A \mid D)$ either way
- Say why SVI is guide-dependent, and what NUTS depends on instead
- List the diagnostics each method needs: ELBO traces and comparisons for SVI; $\hat R$, ESS and divergences for NUTS
- Choose the right tool for a situation, and design a hybrid workflow that validates SVI against NUTS on a subset or a smaller model
- Never say "NUTS gives the exact posterior". Say asymptotically exact, and explain Monte Carlo error and convergence
What we need from earlier chapters: MCMC, warmup and autocorrelation (Chapter 6.9); HMC, NUTS, $\hat R$, ESS and divergences (Chapter 6.10); variational inference and KL divergence (Chapter 6.11); the ELBO and SVI (Chapter 6.12); mean-field, full-rank and low-rank guides (Chapter 6.13); your training loop (Chapter 6.14); hierarchical models and the funnel (Chapter 6.5, Chapter 6.7); bias and variance of an estimator (Chapter 5.1); the multivariate Normal and its covariance ellipse (Chapter 5.15). Notation: θ is the vector of the model's unknown (latent) parameters, with $d$ entries; $D$ is the data, with $N$ rows; $q_\phi(\theta)$ is the guide with variational parameters $\phi$; $S$ is the number of posterior draws; ESS is the effective sample size. Colours in this chapter: green = the true posterior, blue = NUTS or HMC draws, orange = mean-field SVI, teal = full-rank SVI.
Two roads to a posterior: optimize a stand-in, or sample the real thing core
Imagine you must describe where people live in a city. There are two ways to do it.
The template way (SVI). You pick a ready-made shape from a catalogue, say an oval, and you slide and stretch it until it covers the busy parts of the city as well as an oval can. At the end you hand over a short formula: "an oval, centred here, this wide, tilted like this". It is quick. But if the city is shaped like a horseshoe, your answer is still an oval.
The walker way (NUTS). You send a walker who wanders the streets and spends more time where more people live. Every minute you write down where the walker is. After many hours your notebook of positions is the map, whatever shape the city has. It is slow, and if you stop too early, or the walker never finds one neighbourhood, the map has holes.
The posterior $p(\theta\mid D)$ is the city. SVI optimizes a simple stand-in distribution $q_\phi(\theta)$, called the guide. NUTS samples: it produces a long list of plausible values of θ, called draws.
Three ways to say it:
- Picture: SVI stretches a ready-made template over the posterior; NUTS walks around inside it and records the footprints.
- Numbers: SVI with an AutoNormal guide on 10 parameters answers with 20 numbers (a centre and a width each); NUTS answers with a table of, say, 4 000 draws × 10 parameters = 40 000 numbers.
- Slogan: SVI optimizes a stand-in; NUTS samples the real thing, slowly.
One posterior, two kinds of answer. In Chapter 6.3 a conversion rate had the exact posterior Beta(8, 52). Its exact 90% interval is [0.0693, 0.2113]. Pretend we did not know the formula and ran both engines.
- SVI with a Normal guide on the log-odds $\eta = \log\frac{\theta}{1-\theta}$ (this is what AutoNormal does with a parameter that lives between 0 and 1) ends with two numbers: centre $m = -1.925$ and width $s = 0.383$.
- To get a 90% interval from the guide, go 1.645 widths each way and turn log-odds back into rates: $m \pm 1.645\,s = -1.925 \pm 0.630$, which gives $-2.555$ and $-1.295$; then $\theta = 1/(1 + e^{-\eta})$ gives [0.072, 0.215].
- NUTS-style draws: suppose we have 4 000 draws of θ (here, for simplicity, independent draws; real NUTS draws are correlated). Sort them and read the 5% and 95% points: [0.0695, 0.2117]. Their average is 0.1330 (exact mean 0.1333).
- Compare with the exact [0.0693, 0.2113]. SVI is off by about 0.003 at each end, and that offset stays however long we optimize: it is the best a logit-Normal can do. The draws are off by about 0.0004, a random wobble that changes with every new set of draws and shrinks as we draw more.
Same question, two kinds of answer: a formula with a small built-in bias, or a pile of draws with a small random error.
SVI (stochastic variational inference). Choose a family of distributions $\{q_\phi\}$ (the guide). Find
$$\phi^\star = \arg\max_\phi \ \text{ELBO}(\phi), \qquad \text{ELBO}(\phi) = E_{q_\phi}\big[\log p(D, \theta) - \log q_\phi(\theta)\big],$$with stochastic gradient steps (Chapter 6.12). The output is the fitted distribution $q_{\phi^\star}$: a formula you can draw from as often as you like, cheaply.
NUTS (the No-U-Turn Sampler). Build a Markov chain whose long-run (stationary) distribution is $p(\theta\mid D)$. Each step simulates a frictionless puck on the landscape $-\log p(D, \theta)$ with leapfrog steps and stops the trajectory when it starts to turn back (Chapter 6.10). After a warmup phase (which tunes the step size and a scale for each parameter), it records $S$ draws $\theta^{(1)}, \dots, \theta^{(S)}$. The output is the draws; any summary is an average over them:
$$E[f(\theta)\mid D] \approx \frac{1}{S}\sum_{s=1}^{S} f\big(\theta^{(s)}\big).$$- Optimization means "change numbers to make one score as large as possible". Sampling means "produce random values with the right frequencies".
- Both engines need only the unnormalized log joint $\log p(D, \theta) = \log p(D\mid\theta) + \log p(\theta)$ and its gradient. Neither needs the evidence $p(D)$.
- A gradient evaluation is one computation of $\nabla_\theta \log p(D, \theta)$ at one value of θ. It is the main unit of cost for both.
Why do we need it?
Most real posteriors have no formula (no conjugacy, Chapter 6.3), so we must compute them. Knowing that one engine gives a fitted stand-in and the other gives draws of the real thing tells you what each answer can and cannot be trusted for.
Where is it used?
NumPyro's SVI with AutoNormal or AutoMultivariateNormal guides, and its MCMC(NUTS(...)); Stan (NUTS by default, plus ADVI as its variational option); PyMC (NUTS, plus ADVI); Pyro. Both of your projects fit their models with SVI.
How is it used?
Write the model once. For SVI, add a guide, run svi.run, then draw from the guide with guide.sample_posterior. For NUTS, run MCMC(NUTS(model), num_warmup=..., num_samples=...) and use mcmc.get_samples(). Downstream code (intervals, $P(\theta_B \gt \theta_A)$, forecasts with Predictive) works on draws from either.
"SVI also gives samples, so it is a kind of sampler."
SVI's samples come from the fitted guide $q$, not from the posterior. Drawing a million of them describes $q$ perfectly and the posterior no better than before.
"NUTS climbs to the best value of θ."
NUTS does not climb anywhere. It wanders, spending time in each region in proportion to its posterior probability. The only "optimizing" in NUTS is warmup tuning its own step size, not finding a best θ.
"MAP is a quick version of NUTS."
MAP (NumPyro's AutoDelta guide) is SVI with the most extreme template: a single point. It has no uncertainty at all.
Both of your projects use SVI, so their answers are draws from a fitted guide. In an A/B framework like yours, $P(\theta_B \gt \theta_A \mid D)$ is the fraction of guide draws where B beats A. In your forecasting model, forecast intervals come from guide draws pushed through the model (for example with NumPyro's Predictive). The useful consequence: the same downstream code runs unchanged on NUTS draws, so swapping the engine is an easy way to check the guide.
"SVI is just a faster NUTS."
They answer the same question in different ways: SVI optimizes a chosen family of distributions to be close to the posterior; NUTS draws from the posterior itself with a Markov chain.
Model answer: "SVI turns inference into optimization: I pick a guide family and maximize the ELBO, and I get back a fitted distribution. NUTS is MCMC: it simulates Hamiltonian dynamics to produce correlated draws whose distribution approaches the posterior. So SVI's error is mainly the gap between the family and the posterior, while NUTS's error is mainly Monte Carlo noise from a finite number of draws, plus any convergence problems."
SVI: $\phi^\star = \arg\max_\phi \text{ELBO}(\phi)$ → a formula $q_{\phi^\star}$. NUTS: a Markov chain → draws $\theta^{(1..S)}$; summaries are averages over draws.
Both need only $\log p(D, \theta)$ and its gradient; cost is counted in gradient evaluations.
Trap: draws from the guide describe the guide, not the posterior.
Quick check: after SVI finishes, you draw 1 000 000 samples from the guide. Is your posterior summary now exact?
No. A million draws describe the guide $q_{\phi^\star}$ almost perfectly, but $q_{\phi^\star}$ itself differs from $p(\theta\mid D)$ whenever the posterior is not in the guide's family. In the example above, even infinitely many guide draws give the interval [0.072, 0.215], not the exact [0.0693, 0.2113].
Speed: NUTS pays per draw, SVI pays per step core
Both engines spend almost all of their time on one job: computing the gradient $\nabla_\theta \log p(D, \theta)$, the direction in which the log posterior rises fastest. So "how fast is it?" really means "how many gradients does it need, and how much data does each one touch?"
NUTS needs a whole trajectory of leapfrog steps (small physics moves of the puck) for every single draw, and every leapfrog step needs a fresh gradient over all the data. On top of that you want thousands of draws, a warmup phase that costs the same per iteration, and several chains.
SVI needs one gradient per step per particle (a particle is one random draw from the guide used to estimate the ELBO; NumPyro's default is 1), and it can use a small random slice of the data (a minibatch, next section). It stops once the ELBO stops improving.
Three ways to say it:
- Picture: NUTS takes many long walks, and on every step of every walk it re-reads the whole book; SVI takes many short nudges, each after reading a few pages.
- Numbers: 4 chains × 2 000 iterations × 31 leapfrog steps = 248 000 full-data gradients, against 20 000 SVI steps × 1 particle = 20 000.
- Slogan: cost = (gradients needed) × (data per gradient), and the posterior's shape sets both bills.
A back-of-the-envelope bill. Assume one gradient over the full data takes 2 milliseconds (an illustration; real numbers depend on the model and the machine).
- NUTS: 4 chains, each with 1 000 warmup + 1 000 kept iterations = 2 000 iterations per chain. Suppose trajectories reach tree depth 5, so up to $2^5 - 1 = 31$ leapfrog steps per iteration.
- NUTS gradients: $4 \times 2\,000 \times 31 = 248\,000$. Time: $248\,000 \times 2\text{ ms} = 496$ s, about 8.3 minutes.
- SVI: 20 000 steps × 1 particle = 20 000 full-data gradients. Time: $20\,000 \times 2\text{ ms} = 40$ s. That is $248\,000 / 20\,000 = 12.4$ times fewer gradients.
- SVI with minibatches of 1% of the rows: each step costs about 1/100 of a full gradient, so 20 000 steps ≈ 200 full-data gradients ≈ 0.4 s of gradient work: about $248\,000 / 200 = 1\,240$ times less than NUTS.
- Reality check from this chapter's code: on a regression with 3 unknowns and 200 rows, NUTS took under 2 s and SVI under 1 s on our machine, both including JIT compilation (your times will differ). For small models the difference hardly matters.
Count cost in row-gradients: the gradient contribution of one data row. Then, roughly,
$$\text{cost}_{\text{NUTS}} \approx C \times (W + S) \times \bar L \times N, \qquad \text{cost}_{\text{SVI}} \approx T \times K \times B.$$- $C$ = number of chains; $W$ and $S$ = warmup and kept iterations per chain (warmup iterations cost the same as kept ones).
- $\bar L$ = average number of leapfrog steps per iteration. A NUTS tree of depth $j$ has up to $2^j - 1$ steps; NumPyro's default
max_tree_depth=10caps it at 1 023. NumPyro reports it per draw withextra_fields=("num_steps",). - $T$ = SVI steps until the ELBO stops improving; $K$ = particles per step (
Trace_ELBO(num_particles=1)by default); $B$ = rows per minibatch ($B = N$ without minibatching). - What makes $\bar L$ large: a badly scaled or strongly correlated posterior forces a small step size, so trajectories need many steps. What makes $T$ large: the same geometry (elongated posteriors slow Adam down too), a small learning rate, and noisy gradients.
- Both also pay a one-time JIT compilation cost (Chapter 6.17). For NUTS, the fair speed measure is ESS per second, not draws per second.
Why do we need it?
"SVI is faster" is only true for some models. Counting gradients tells you when the speed gap is a factor of 2 (small models: just use NUTS) and when it is a factor of 1 000 (huge data, many refits: SVI is the only practical option).
Where is it used?
Planning fits in NumPyro, Stan and PyMC; reading NUTS's tree depth (num_steps) and SVI's step count; budgeting nightly refits of many forecasting series; deciding whether a validation run with NUTS fits into a CI job.
How is it used?
Time a short run of each engine after compilation. For NUTS, look at average num_steps and ESS per second; for SVI, at how many steps the loss needs to flatten. Multiply out to the full job, and remember which parts scale with $N$.
"SVI is always faster than NUTS."
For small models (tens of parameters, thousands of rows) NUTS often finishes in seconds and the gap does not matter. SVI can also be slow: an elongated posterior or a small learning rate can need tens of thousands of steps.
"NUTS costs the same per draw on every model."
The number of leapfrog steps per draw is set by the posterior's shape. Trajectories that keep hitting max_tree_depth (1 023 steps) are an efficiency warning: reparameterize or rescale before buying more compute.
"Compare speed with the wall-clock time of the first run."
The first call includes JIT compilation. Compare like with like, and for NUTS compare effective draws per second (ESS/s), since 1 000 highly correlated draws may be worth only 100 independent ones.
In your forecasting model, the custom SVI loop runs a JIT-compiled update: after compiling once, each step is one cheap gradient (per particle) of the log joint. NUTS on the same model would need a full trajectory of gradients for every draw, and the trend, changepoint, seasonality and holiday parameters are correlated, which tends to raise the tree depth. If the model is refit often or for many series, that multiplies quickly. In an A/B framework like yours, a single experiment's model is often small enough that NUTS takes seconds, which makes it a practical reference for checking the SVI answers.
NUTS cost ≈ chains × (warmup + draws) × leapfrog steps × N rows. SVI cost ≈ steps × particles × B rows.
Leapfrog steps per draw ≤ $2^{\text{depth}} - 1$ (NumPyro cap: depth 10 → 1 023). Bad geometry raises both bills.
Trap: "SVI is faster" is a statement about large models; measure ESS/s for NUTS and steps-to-flat for SVI.
Quick check: NUTS runs 4 chains of 2 000 iterations and the trajectories reach tree depth 6. At most how many full-data gradients is that?
Depth 6 means up to $2^6 - 1 = 63$ leapfrog steps per iteration, so at most $4 \times 2\,000 \times 63 = 504\,000$ full-data gradients. Each one touches every data row.
Scalability: more rows, more parameters, and minibatches core
A model can grow in two directions. It can get more rows of data ($N$: users, days, orders), or more unknowns ($d$: segments, changepoints, Fourier coefficients, regressors). The two engines feel these differently.
More rows. SVI's objective is an average, and an average can be estimated from a random handful of rows: read a random batch of $B$ rows, add up their log-likelihoods, and multiply by $N/B$ to stand in for all $N$. The estimate is noisy but right on average, which is all a stochastic optimizer needs. This is a minibatch. Plain NUTS cannot do this: it relies on exact energies (HMC's accept/reject step compares the energy at the start and the end of a trajectory, and NUTS weights the points of its trajectory by their energies), and a noisy estimate breaks that.
More unknowns. NUTS stores $S \times d$ numbers and each gradient costs more, but its trajectories lengthen only slowly with $d$ on well-behaved posteriors. SVI's guide grows with $d$ in a way that depends on the guide: $2d$ numbers for mean-field, about $d^2/2$ for full-rank, $d(r+2)$ for low-rank (Chapter 6.13).
Three ways to say it:
- Picture: NUTS must read the whole book before every footstep; SVI may read one random page and multiply by the number of pages.
- Numbers: with $N = 1\,000\,000$ rows and batches of $B = 1\,000$, each SVI step reads 0.1% of the data and multiplies the batch's log-likelihood by 1 000.
- Slogan: minibatches are SVI's superpower; plain NUTS has no minibatch mode.
The N/B factor, and what grows with d.
- Rows: $N = 1\,000\,000$, batch $B = 1\,000$, so the scale factor is $N/B = 1\,000$.
- The batch's log-likelihoods add up to $-3\,250.4$. The estimate of the full-data log-likelihood is $1\,000 \times (-3\,250.4) = -3\,250\,400$.
- Averaged over all possible random batches, this estimate equals the true full-data sum exactly. This chapter's code checks it on 200 rows with batches of 20: full-data log density $-289.9$; the average of 4 000 minibatch estimates is $-289.5$ (their spread is 33.5, so that small difference is just noise).
- Parameters: take $d = 500$ latent parameters. AutoNormal stores $2d = 1\,000$ numbers; a rank-10 low-rank guide stores $d(r+2) = 500 \times 12 = 6\,000$; a full-rank guide stores $d + d(d+1)/2 = 500 + 125\,250 = 125\,750$.
- NUTS with 4 000 kept draws stores $4\,000 \times 500 = 2\,000\,000$ numbers, plus its mass matrix: 500 numbers if diagonal (NumPyro's default,
dense_mass=False), 250 000 if dense.
Minibatch ELBO. If the likelihood is a sum over rows that are independent given θ, $\log p(D\mid\theta) = \sum_{i=1}^{N} \log p(y_i\mid\theta)$, then for a uniformly random batch $\mathcal{B}$ of $B$ rows and $\theta \sim q_\phi$,
$$\widehat{\text{ELBO}} = \frac{N}{B}\sum_{i \in \mathcal{B}} \log p(y_i\mid\theta) + \log p(\theta) - \log q_\phi(\theta)$$is an unbiased estimate of the ELBO: its average over random batches and draws is the true ELBO. In NumPyro, numpyro.plate("data", N, subsample_size=B) draws the batch (without replacement) and applies the factor $N/B$ for you.
- Assumption: rows are conditionally independent given θ. Models where rows are linked (autoregressive terms, latent state chains) need special care.
- Price: extra gradient noise, so SVI may need a smaller learning rate or more steps.
- Why not NUTS: NUTS needs the exact log density for the accept step and exact gradients to keep the energy nearly constant along a trajectory. Subsampling MCMC methods exist (stochastic-gradient Langevin dynamics; NumPyro's
HMCECS), but they are different algorithms with their own approximations, not plain NUTS. - Growth with d: NUTS: per-gradient work and memory grow with $d$; trajectory length grows slowly for well-behaved targets (theory for simple targets: about $d^{1/4}$). SVI: mean-field $O(d)$, low-rank $O(dr)$, full-rank $O(d^2)$ memory, with per-step work growing at least as fast.
Why do we need it?
With millions of rows, one full-data gradient can take seconds, and NUTS needs hundreds of thousands of them. Minibatches make each SVI step cost a tiny fraction of that, and knowing how guides grow with $d$ tells you which guide fits in memory.
Where is it used?
Minibatch SVI in NumPyro and Pyro (plate(..., subsample_size=B)), variational autoencoders, Bayesian neural networks, topic models (the original "stochastic variational inference" paper was about topic models on millions of documents), and large hierarchical models in industry.
How is it used?
Wrap the likelihood in a plate with subsample_size, index the data with the plate's indices, and pick $B$ by trying a few sizes: the smallest batch whose ELBO trace is still smooth enough. Check that the scaling is applied (forgetting it makes the prior far too strong).
"Minibatching changes the answer SVI converges to."
With the $N/B$ factor the minibatch ELBO has the same average as the full ELBO; it only adds noise, which more steps or a smaller learning rate averages away. Without the factor, the data are down-weighted by $B/N$ and the prior dominates.
"The smallest batch is the fastest way to train."
Tiny batches give very noisy gradients, so you need more steps and a smaller learning rate. On a GPU, a batch of a few thousand rows often takes about as long as a batch of ten. Pick $B$ by trying a few sizes.
"You can run NUTS on minibatches the same way."
Plain NUTS needs exact log densities and gradients. Subsampling MCMC methods exist, but they are different algorithms with their own errors.
"Any time series can be minibatched by days."
Only if, given the parameters, the days' likelihood terms are independent (as with a trend + seasonality + independent noise model). Autoregressive terms or latent state chains link the days.
If your forecasting history is a few years of daily data, that is around a thousand rows, so full-batch SVI is cheap and avoids extra gradient noise; minibatches would matter only for very long, high-frequency, or many-series data fitted jointly. Given the parameters, a trend + seasonality + holidays + regressors + independent-noise model's likelihood is a sum over days, so minibatching days would be valid. In an A/B framework like yours, conversion data can be summarized as counts per variant and segment ($k$ conversions of $n$ users): a Binomial likelihood on those counts gives exactly the same posterior for θ as one Bernoulli term per user, so the data size that matters is the number of cells, not the number of users.
Minibatch ELBO: $\frac{N}{B}\sum_{i\in\mathcal{B}} \log p(y_i\mid\theta) + \log p(\theta) - \log q_\phi(\theta)$, unbiased if rows are conditionally independent. NumPyro: plate("data", N, subsample_size=B).
NUTS: every leapfrog step reads all N rows; no plain minibatch mode. Guides: mean-field $2d$, low-rank $d(r+2)$, full-rank $d + d(d+1)/2$.
Trap: forgetting the $N/B$ factor makes the prior $N/B$ times too strong.
Quick check: N = 2 000 000 rows, batches of B = 500. What factor multiplies the batch log-likelihood, and what happens if you forget it?
$N/B = 4\,000$. Forgetting it makes the data term 4 000 times too small compared with the prior, so the posterior is pulled toward the prior and is far too wide: the model behaves as if it had seen only 500 rows.
Uncertainty quality: under-dispersion, correlations and shape core
SVI minimizes $KL(q\,\|\,p)$ (Chapter 6.11). That direction of KL punishes $q$ very hard for putting probability where the posterior has almost none, and barely punishes it for missing parts of the posterior. So $q$ plays safe: it sits in the heart of the posterior and keeps its tails short. Statisticians call this under-dispersed: spread out less than it should be.
A mean-field guide makes this much worse when parameters are correlated. It can only draw upright ovals (each parameter on its own). The posterior of two correlated parameters is a tilted cigar. The best upright oval for a tilted cigar is a small, round blob in the middle: much too short along the cigar, and a little too fat across it. A full-rank guide can tilt, so it fixes the correlation, but it is still an oval: it cannot bend into a banana or narrow into a funnel. NUTS draws land wherever the posterior has mass, whatever its shape, given enough draws and healthy chains.
Three ways to say it:
- Picture: an upright oval squeezed onto a tilted cigar is a small round blob: it misses the cigar's length.
- Numbers: with correlation ρ = 0.9, the true sd is 1 but the mean-field sd is $\sqrt{1 - 0.81} = 0.44$: its 90% interval is less than half as wide as it should be.
- Slogan: mean-field is confidently narrow; full-rank fixes the tilt, not the shape; NUTS sees the shape but needs time.
Two correlated segment effects. In a hierarchical A/B model, two segment effects θA and θB both lean on the same population mean, so their posterior is correlated. Say both have sd 1, correlation ρ = 0.9, and posterior means 0 and 0.4.
- Mean-field sd of each effect: $\sqrt{1 - \rho^2} = \sqrt{1 - 0.81} = \sqrt{0.19} = 0.436$ (instead of 1).
- 90% interval half-width: true $1.645 \times 1 = 1.645$; mean-field $1.645 \times 0.436 = 0.717$. Far too confident about each effect.
- The difference θB − θA: true variance $1 + 1 - 2(0.9)(1)(1) = 0.2$, sd $0.447$. Mean-field ignores the covariance: variance $0.19 + 0.19 = 0.38$, sd $0.616$. Here mean-field is too wide.
- $P(\theta_B \gt \theta_A \mid D)$: true $\Phi(0.4/0.447) = \Phi(0.894) = 0.814$; mean-field $\Phi(0.4/0.616) = \Phi(0.649) = 0.742$. ($\Phi$ is the standard Normal CDF.)
- The sum θA + θB: true sd $\sqrt{2 + 1.8} = 1.95$; mean-field $\sqrt{0.38} = 0.62$, three times too narrow.
So mean-field is not just "too narrow everywhere": it is too narrow for each effect and for the sum, too wide for the difference, and the decision probability moved from 0.81 to 0.74.
Let the posterior be Gaussian, $N(\mu, \Sigma)$, with precision matrix $\Lambda = \Sigma^{-1}$. The mean-field guide that minimizes $KL(q\,\|\,p)$ is
$$q^\star(\theta) = \prod_{i=1}^{d} N\!\left(\theta_i;\ \mu_i,\ \frac{1}{\Lambda_{ii}}\right).$$- $1/\Lambda_{ii}$ is the conditional variance of $\theta_i$ when all the other parameters are held fixed. It is never larger than the marginal variance $\Sigma_{ii}$, and equal only when $\theta_i$ is uncorrelated with the rest. In two dimensions: $1/\Lambda_{ii} = \sigma_i^2 (1 - \rho^2)$.
- So mean-field gets the means right here (Gaussian case) but its marginal variances are too small and its correlations are zero. Derived quantities such as differences can come out too narrow or too wide.
- A full-rank Gaussian guide $N(m, LL^\top)$ recovers a Gaussian posterior exactly. For non-Gaussian posteriors (skewed, banana-shaped, funnel-shaped, heavy-tailed, several modes) it still misses the shape and usually under-covers the tails.
- NUTS draws approach the true distribution, shape included. Its uncertainty is as good as its ESS and its convergence (next sections).
Why do we need it?
Decisions use uncertainty: an interval, a probability that B beats A, a forecast quantile. If the method squeezes or distorts the uncertainty, the decision changes, even when the point estimates look fine.
Where is it used?
Choosing between AutoNormal, AutoLowRankMultivariateNormal and AutoMultivariateNormal in NumPyro; hierarchical A/B models (correlated segment effects); regressions with correlated predictors; trend and changepoint parameters in Prophet-style models; Bayesian neural networks, where mean-field is common and known to be overconfident.
How is it used?
Before trusting SVI intervals, look at the posterior correlations (from a full-rank fit or a NUTS run on a subset). If important parameters are strongly correlated, do not use mean-field for decisions about them; compare the decision quantities across guides and against NUTS.
"Mean-field overstates certainty about everything."
It makes each parameter's marginal spread too small, but a combination along the posterior's narrow direction (like a difference of correlated effects) can come out too wide. So $P(\theta_B \gt \theta_A)$ can be wrong in either direction.
"A full-rank guide is exact."
Only if the posterior is Gaussian on the guide's (unconstrained) scale. Skew, bananas, funnels, heavy tails and several modes all remain.
"Correlation 0 in the posterior means mean-field is fine."
The banana has correlation 0 and strong dependence: θ₂'s spread depends on θ₁. Mean-field (and full-rank) still fail on it.
"NUTS uncertainty is automatically right."
Only with enough effective draws and healthy chains. In the funnel, HMC misses the neck; its divergences are the warning.
In an A/B framework like yours with hierarchical partial pooling, the segment effects share the population mean, so their posterior is correlated, and $P(\theta_B \gt \theta_A \mid D)$ or a lift interval depends on that correlation. A mean-field guide can push it either way, so a full-rank guide or a NUTS check is worth it for the decision quantities. In your forecasting model, the trend slope and the changepoint adjustments δj, the Fourier coefficients and holiday effects, and the intercept and regressor coefficients compete to explain the same bumps (identifiability, Chapter 6.8), so the posterior is correlated. That is exactly why a full-rank or low-rank guide is used rather than mean-field: forecast intervals add up many correlated pieces.
"Variational inference underestimates uncertainty."
"Reverse-KL VI, especially mean-field, tends to underestimate marginal posterior variances, because it ignores correlations and avoids putting mass where the posterior has little. Derived quantities can be off in either direction."
Model answer: "Minimizing KL(q‖p) makes q mode-seeking and under-dispersed. With a mean-field guide on a correlated Gaussian posterior, each marginal variance shrinks to the conditional variance, σ²(1 − ρ²). A full-rank guide fixes that for Gaussian posteriors but not for skewed or curved ones. NUTS has no family restriction; its uncertainty is limited by the effective sample size and convergence instead."
Mean-field optimum on $N(\mu, \Sigma)$: $q_i = N(\mu_i, 1/\Lambda_{ii})$, $\Lambda = \Sigma^{-1}$; in 2D, sd $= \sigma\sqrt{1 - \rho^2}$ (ρ = 0.9 → 0.44 σ).
Full-rank fixes correlations, not shape. NUTS sees the shape, limited by ESS and convergence.
Trap: mean-field is too narrow for each parameter but can be too wide for a contrast; decision probabilities can move either way.
Quick check: two parameters with posterior sds 2 and correlation 0.6. What sd does the best mean-field guide give each one?
$2\sqrt{1 - 0.6^2} = 2\sqrt{0.64} = 2 \times 0.8 = 1.6$. The 90% interval half-width shrinks from $1.645 \times 2 = 3.29$ to $1.645 \times 1.6 = 2.63$: about 20% too narrow.
Guide dependence: in SVI, the template is part of the answer core
Run SVI twice on the same model and data with two different guides, and you get two different posteriors. Run it twice with the same guide from two different starting points, and on some models you again get two different answers. The SVI result is a joint product of the model and the guide family and the optimizer settings (starting point, learning rate, number of steps, particles, random seed).
NUTS has no template. Its result depends on the model (and how you write it: centered or non-centered, Chapter 6.7) and on its tuning settings (warmup length, target acceptance, maximum tree depth), but if the chains converge, every setting aims at the same target.
The clearest case is a posterior with two modes (two separate peaks). A single Gaussian guide cannot cover both, so it settles on whichever peak is closer to where it started, and reports it with full confidence. NUTS chains can get stuck too, but if you start several chains in different places, they disagree, and $\hat R$ shouts.
Three ways to say it:
- Picture: SVI is a cookie cutter; the cutter's shape is printed on every cookie.
- Numbers: two equal peaks at ±2.5: SVI started on the right reports "mean +2.5, sd 0.8, P(θ₁ > 0) = 0.999"; the truth is "mean 0, sd 2.62, P(θ₁ > 0) = 0.5".
- Slogan: change the guide, change the answer; in SVI the guide is part of your uncertainty model.
A two-peaked posterior. The posterior is an equal mix of $N(-2.5, 0.8^2)$ and $N(+2.5, 0.8^2)$ for θ₁.
- Truth: mean $0$ (the peaks balance); variance $= 0.8^2 + 2.5^2 = 0.64 + 6.25 = 6.89$, so sd $= 2.62$; $P(\theta_1 \gt 0) = 0.5$.
- A Gaussian guide started at θ₁ = +1 slides to the right peak: mean about $+2.5$, sd about $0.8$.
- Its probability: $P_q(\theta_1 \gt 0) = \Phi(2.5/0.8) = \Phi(3.125) = 0.9991$.
- Started at θ₁ = −1, the same guide reports mean $-2.5$ and $P_q(\theta_1 \gt 0) = 0.0009$. Same model, same data, opposite conclusions.
- Guide sizes for $d = 10$ parameters (Chapter 6.13): AutoDelta 10 numbers (a point), AutoNormal 20, low-rank with $r = 2$: $10 \times 4 = 40$, AutoMultivariateNormal $10 + 55 = 65$. Each is a different template, and each gives a different answer on a non-Gaussian posterior.
SVI returns $\hat q = q_{\hat\phi}$, where $\hat\phi$ is where a stochastic optimizer, started at $\phi_0$, stopped while trying to solve $\max_{\phi} \text{ELBO}(\phi)$ over the family $\mathcal{Q}$. So
$$\hat q = \hat q\,(\text{model},\ \mathcal{Q},\ \phi_0,\ \text{learning rate},\ \text{steps},\ \text{particles},\ \text{seed}).$$- Family dependence: the best member $q^\star \in \mathcal{Q}$ has a gap $KL(q^\star\,\|\,p) \ge 0$ that only a richer family can reduce.
- Start dependence: when the ELBO has several local optima (several modes, label switching in mixtures), different $\phi_0$ give different $\hat q$.
- NumPyro's ready-made guides, from least to most flexible:
AutoDelta(a point: MAP),AutoNormal(independent Normals),AutoLowRankMultivariateNormal,AutoMultivariateNormal; plusAutoLaplaceApproximation(a Gaussian at the MAP, from the curvature) and normalizing-flow guides (for exampleAutoBNAFNormal,AutoIAFNormal), and hand-written guides. - NUTS depends on the model's parameterization and its tuning settings; if it converges, it targets $p(\theta\mid D)$ itself. Multimodality is hard for it too, but it is detectable with several chains and $\hat R$.
Why do we need it?
If you report an SVI posterior without saying which guide produced it, nobody can judge it. Knowing that the guide and the start shape the answer tells you what to vary when you check robustness.
Where is it used?
Choosing among NumPyro's AutoNormal, AutoLowRankMultivariateNormal and AutoMultivariateNormal; mixture models and latent class models (several modes by construction); Bayesian neural networks (many modes); any report or model card that must state its inference settings.
How is it used?
Record the guide type, rank, learning rate, steps, particles and seed with every fit. Refit with two or three seeds and starting points and with a richer guide; if the decision quantities move, the guide is part of the problem. On a subset, compare with several NUTS chains started far apart.
"AutoNormal is the default, so it is a neutral choice."
It is a strong assumption: every parameter independent and Normal on the unconstrained scale. "Default" means convenient, not assumption-free.
"Two SVI runs agree, so the answer is right."
Two runs with the same guide share the same family gap. Agreement shows the optimization is stable, not that the guide matches the posterior.
"NUTS handles several modes automatically."
It does not. Chains can stay in one peak for the whole run. The difference is that several chains started far apart reveal the problem through $\hat R$; a single SVI fit does not.
Your forecasting model chooses a full-rank or a low-rank Gaussian guide automatically from the model size. That makes the guide part of the forecast: for a large model, the rank $r$ decides how many directions of posterior correlation the guide can represent. A natural robustness check is to refit a smaller version of the model with both guides (and with NUTS) and see whether the forecast intervals and the trend and changepoint summaries change. In the A/B framework, recording the guide type with every reported $P(\theta_B \gt \theta_A \mid D)$ makes results comparable across experiments.
"My SVI posterior is the posterior of the model."
"My SVI posterior is the best approximation to the model's posterior within the guide family I chose, as found by my optimizer."
Model answer: "SVI results are guide-dependent. The family sets a floor on the error: the KL gap of the best member. The optimizer's start and settings can add more, especially with several modes. NUTS has no family; it depends on the parameterization and tuning, and its failures are visible through R̂ and divergences. So when the decision is important, I compare guides and check against NUTS on a smaller problem."
SVI answer = f(model, guide family, start, learning rate, steps, particles, seed). Family gap: $KL(q^\star\,\|\,p)$, fixed by the family.
Two peaks: a Gaussian guide picks one, silently. Several NUTS chains disagree, loudly ($\hat R \gg 1.01$).
Trap: agreement between SVI runs with the same guide is not evidence of accuracy.
Quick check: you switch from AutoNormal to AutoMultivariateNormal and your 90% interval for a key effect widens by 60%. What does that tell you?
That the posterior has strong correlations involving that effect, which AutoNormal could not represent, so its interval was much too narrow. The full-rank interval is more credible, but if the posterior is also non-Gaussian it may still be too narrow: a NUTS run on the same model (or a smaller version) settles it.
Two kinds of error: approximation error vs Monte Carlo error core
Think of two dart players aiming at the bullseye (the true posterior answer).
The SVI player throws a tight group, but in the wrong place. The darts land close to each other (little randomness), and the whole group sits off-centre because the guide family cannot reach the bullseye. Throwing longer (more SVI steps) tightens the group around the wrong spot. This is approximation error, a kind of bias (a systematic offset).
The NUTS player aims at the bullseye, but each dart wobbles. With few effective draws the darts scatter widely; with more they cluster closer to the centre, and the scatter shrinks like $1/\sqrt{\text{ESS}}$. This is Monte Carlo error, a kind of variance (random noise). There is one catch: if the chain has not converged (it is stuck, or never visits part of the posterior), the NUTS player is also aiming at the wrong spot.
Three ways to say it:
- Picture: SVI's darts are grouped tightly in the wrong place; NUTS's darts are centred but spread out, and the spread shrinks as you throw more.
- Numbers: on a ρ = 0.9 posterior, mean-field SVI's sd is 56% too small after 1 000 or 1 000 000 steps; NUTS's estimate of a posterior mean with sd 0.5 and ESS 400 wobbles by about $0.5/\sqrt{400} = 0.025$.
- Slogan: SVI's error is bias that more steps cannot fix; NUTS's error is noise that more draws can fix, if the chain is healthy.
Putting numbers on both errors.
- NUTS, posterior mean: the posterior sd of a parameter is 0.5 and the chains give ESS = 400. Monte Carlo standard error (MCSE) $= 0.5/\sqrt{400} = 0.5/20 = 0.025$. Report "mean ± 0.025 (Monte Carlo)".
- To halve it you need $\sqrt{\text{ESS}}$ twice as big, so ESS = 1 600: four times the draws.
- NUTS, a probability: $\hat p = 0.30$ with ESS = 400. MCSE $= \sqrt{0.3 \times 0.7/400} = \sqrt{0.000525} = 0.023$.
- SVI: mean-field on a ρ = 0.9 posterior gives sd $\sqrt{1 - 0.81} = 0.436$ instead of 1: an error of $-56\%$. After 1 000 steps, 10 000 steps or a million steps, it is still $-56\%$, because that is the best member of the family.
- So the questions differ: for NUTS, "is ESS big enough, and did the chains converge?"; for SVI, "is the family rich enough, and did the optimizer finish?"
SVI's total error splits into three parts:
$$\underbrace{p \to q^\star}_{\text{family gap}} \ +\ \underbrace{q^\star \to \hat q}_{\text{optimization error}} \ +\ \underbrace{\hat q \to \text{summary from draws of } \hat q}_{\text{small Monte Carlo error}}$$- Family gap: $q^\star$, the best member of the family, still differs from the posterior ($KL(q^\star\,\|\,p) \gt 0$ unless the posterior is in the family). More steps cannot remove it.
- Optimization error: the optimizer stopped before $q^\star$ (too few steps, too large a learning rate, a local optimum). Your training loop's stopping rule controls this part only.
- The Monte Carlo error from summarizing $\hat q$ is tiny and cheap to shrink, because drawing from a guide is cheap.
NUTS's error for an estimate $\bar f = \frac{1}{S}\sum_s f(\theta^{(s)})$ of $E[f(\theta)\mid D]$:
$$\text{MCSE}(\bar f) \approx \frac{\text{sd}(f(\theta)\mid D)}{\sqrt{\text{ESS}}} \quad\text{(if the chains have converged)}, \qquad \text{plus any bias from non-convergence.}$$- ESS (effective sample size) is the number of independent draws that would carry the same information as your correlated chain (Chapter 6.10).
- The MCSE formula assumes the chain explores the whole posterior. A stuck chain, or one that cannot enter a region (divergences), can show a small MCSE and a large error.
- Bias–variance language (Chapter 5.1): SVI trades a little variance for a fixed bias; healthy NUTS has a vanishing bias and a variance that shrinks with ESS.
Why do we need it?
The two errors need opposite remedies. Monte Carlo error is fixed by more draws; approximation error is fixed only by a richer guide (or by switching to NUTS). Mixing them up wastes compute: running SVI ten times longer does not fix a too-narrow mean-field posterior.
Where is it used?
Reporting NUTS results with MCSE (ArviZ and NumPyro summaries print ESS; MCSE = sd/√ESS); deciding how many draws to keep; choosing between AutoNormal, low-rank and full-rank guides; comparing an SVI result with a NUTS reference "up to Monte Carlo error".
How is it used?
For NUTS: compute sd/√ESS for each reported number and keep only the digits it supports. For SVI: refit with a richer guide; if a decision quantity moves by more than you care about, the family gap is too big. When comparing the two, ignore differences smaller than about 2 MCSE of the NUTS reference.
"More SVI steps bring you closer to the posterior."
More steps bring you closer to the best member of the guide family, $q^\star$. The family gap between $q^\star$ and the posterior does not move.
"A small MCSE means the NUTS answer is accurate."
MCSE measures random wobble assuming the chain explored the whole posterior. A chain that never visits a region can report a tiny MCSE and still be wrong (next section).
"Thinning (keeping every 10th draw) reduces Monte Carlo error."
Thinning throws away information; it never raises ESS. It only saves memory (Chapter 6.9).
When your training loop stops (relative ELBO change below a tolerance for "patience" evaluations, then restore the best state), it is controlling the optimization error only. The family gap was decided earlier, when the model size picked a full-rank or a low-rank guide. So "my loop converged" and "my posterior is accurate" are different claims; the second needs a comparison with a richer guide or with NUTS.
"SVI has sampling error, just like MCMC."
"SVI's main error is approximation error: a bias from the guide family that does not shrink with more steps. MCMC's main error is Monte Carlo error, which shrinks like 1/√ESS, plus any convergence problems."
Model answer: "I think of it as bias versus variance. SVI gives a low-variance answer with a fixed bias set by the guide family, plus some optimization error. Healthy NUTS gives an answer whose bias vanishes as the chain runs, with a Monte Carlo standard error of about sd/√ESS that I can report. That is why I report MCSE for NUTS, and why for SVI I compare guides instead of running longer."
SVI error = family gap ($p \to q^\star$) + optimization error ($q^\star \to \hat q$) + tiny draw noise.
NUTS error ≈ MCSE = sd/√ESS (×4 draws → ÷2 error), plus bias if not converged.
Trap: more steps never fix the family gap; a small MCSE never proves convergence.
Quick check: NUTS gives a posterior mean of 2.31 with posterior sd 0.8 and ESS 256. How many decimal places can you honestly report?
MCSE $= 0.8/\sqrt{256} = 0.8/16 = 0.05$. A 2-MCSE band is about ±0.1, so "2.3" is honest and the "1" in 2.31 is Monte Carlo noise. To trust the second decimal you would need MCSE ≈ 0.005, an ESS of about $(0.8/0.005)^2 = 25\,600$.
"Asymptotically exact": why you never say NUTS gives the exact posterior core
Asymptotically means "in the limit, as something grows without end". NUTS is asymptotically exact: if you could run its chains forever, the averages over its draws would equal the true posterior averages. That is a wonderful guarantee, and it is a guarantee about a run you will never do.
Every real run stops after a finite number of draws. So every NUTS number carries Monte Carlo error. Worse, a finite chain can fail in ways that its own error bars do not show: it may not have finished warming up, it may be stuck in one of several peaks, or it may never enter a narrow region (the neck of a funnel), which is exactly what divergences warn about. And even a perfect run gives the posterior of your model; if the model is wrong, the posterior is precisely computed and still wrong about the world.
Three ways to say it:
- Picture: NUTS is a perfectly honest surveyor who would need infinite time; you always stop the survey early.
- Numbers: 3 000 healthy draws estimate a tail probability as 0.067 ± 0.005; a chain that never reaches the funnel's neck reports 0.000 ± 0.000 when the truth is 0.067.
- Slogan: asymptotically exact, practically approximate: check the diagnostics, report the MCSE.
One true probability, two chains. Both questions below have the same true answer, $P = \Phi(-1.5) = 0.0668$.
- Healthy case: correlated Gaussian, event θ₁ > 1.5. After 3 000 kept draws, the widget's first chain gives $\hat p = 0.0667$, and the ESS of the event's yes/no indicator is about 2 480. MCSE $= \sqrt{0.0667 \times 0.9333 / 2\,480} = 0.0050$. Interval $\hat p \pm 2\,\text{MCSE} = [0.057, 0.077]$: it contains the truth.
- Unhealthy case: Neal's funnel, event $v \lt -4.5$ (the narrow neck), where $v \sim N(0, 3^2)$, so the truth is again $\Phi(-4.5/3) = \Phi(-1.5) = 0.0668$.
- After 1 000 kept draws, the same kind of chain had never gone below $v = -4.5$: $\hat p = 0$. The MCSE formula gives $\sqrt{0 \times 1/1\,000} = 0$. "Zero, plus or minus zero". Confidently wrong.
- After 3 000 draws it had dipped in briefly: $\hat p = 0.013$, but the ESS of the indicator is only about 78, so MCSE $\approx 0.013$ and the interval is about [−0.013, 0.039]: it still misses 0.0668. (Other seeds: some chains never enter the neck at all; others dip in and report something like 0.09 with an ESS below 20, which is useless.)
- The warning: 115 divergences (trajectories whose energy blew up) in 3 100 iterations. Divergences say "there is a region this sampler cannot enter safely", exactly the region the question is about.
A valid MCMC algorithm whose chain has the posterior as its stationary distribution and can reach every region (ergodic) satisfies, for any function $f$ with a finite posterior mean,
$$\frac{1}{S}\sum_{s=1}^{S} f\big(\theta^{(s)}\big) \ \longrightarrow\ E[f(\theta)\mid D] \qquad \text{as } S \to \infty.$$That is what asymptotically exact means. For finite $S$:
- Monte Carlo error remains, about $\text{sd}(f)/\sqrt{\text{ESS}}$.
- Convergence concerns remain: incomplete warmup, chains stuck in different modes ($\hat R \gt 1.01$), regions the integrator cannot enter (divergences), very long autocorrelation (low ESS).
- "Exact" also hides two more things: floating-point arithmetic, and model misspecification: NUTS targets $p(\theta\mid D)$ under your model, not the truth about the world.
- SVI, by contrast, is not asymptotically exact in steps: infinitely many steps converge (at best) to $q^\star$, not to the posterior.
Why do we need it?
Overclaiming is the classic interview mistake on this topic. The precise words ("asymptotically exact; finite runs have Monte Carlo error and convergence risk") show that you know what the guarantee is and what it is not.
Where is it used?
Every MCMC report: Stan and NumPyro print ESS and R̂ for exactly this reason; ArviZ prints MCSE; validation studies that use NUTS as a reference for SVI must state the reference's own Monte Carlo error.
How is it used?
Call NUTS "the reference" or "asymptotically exact", never "the truth". Report MCSE next to NUTS estimates, check R̂, ESS and divergences before using the draws, and judge SVI–NUTS differences relative to the NUTS MCSE.
"NUTS gives the exact posterior; SVI gives an approximation."
Both give approximations in practice. SVI is approximate by design (the family gap, which does not shrink). NUTS is asymptotically exact: its error shrinks with more draws, and it can be measured (MCSE) and checked (R̂, ESS, divergences).
"If NUTS ran without errors, its answer is the truth."
"Ran without errors" is not "converged". And a converged run is the posterior of your model, which is only as right as the model's likelihood and priors.
"A small MCSE proves the estimate is right."
The funnel chain had MCSE 0 and was wrong. MCSE is meaningful only after the convergence diagnostics pass.
"NUTS gives the exact posterior."
"NUTS is asymptotically exact: as the number of draws grows, its estimates converge to the true posterior expectations. Any finite run has Monte Carlo error, about sd/√ESS, and possible convergence problems, which is why we check R̂, ESS and divergences."
"SVI is approximate and NUTS is exact."
"SVI is approximate by design, with a bias set by the guide family; NUTS is approximate in practice, with an error we can measure and shrink."
Model answer: "I'd never call NUTS exact. It targets the posterior: run long enough with a healthy chain and the averages converge to the posterior expectations. With a finite run I have Monte Carlo error, so I report MCSE, and I can have convergence failures, like stuck chains or divergences in a funnel, so I check R̂, ESS and divergences first. And it's the posterior of my model, so I still need posterior predictive checks. That's still a stronger guarantee than SVI's, whose error is a fixed family gap that more steps can't remove."
If you validate your SVI fits against NUTS (on a subset or a smaller model), call the NUTS run the reference, not the truth, and quote its MCSE: an SVI–NUTS difference smaller than about two NUTS MCSEs is not evidence of a problem. For the hierarchical A/B model with few segments or little data per segment, the funnel is a real risk; check that the NUTS reference itself has zero divergences before trusting it.
Asymptotically exact: $\frac{1}{S}\sum f(\theta^{(s)}) \to E[f(\theta)\mid D]$ as $S \to \infty$ (valid, ergodic chain).
Finite run: Monte Carlo error ≈ sd/√ESS + convergence risk (R̂, ESS, divergences) + it is the posterior of your model.
Say: "asymptotically exact; finite chains have Monte Carlo error and convergence concerns". Never: "exact".
Quick check: a colleague says "We used NUTS, so the credible interval is exact." Give a one-sentence correction.
"NUTS is asymptotically exact, but this run had a finite number of draws, so the interval endpoints carry Monte Carlo error (check their MCSE), and they are only trustworthy if R̂ ≈ 1, the ESS is large enough and there are no divergences; and the interval is exact only for our model, not for reality."
Diagnostics: what each engine's warning lights can and cannot see core
Each engine fails in its own way, so each needs its own warning lights.
SVI's main light is the ELBO trace: the ELBO (or NumPyro's loss, which is minus the ELBO) plotted against the step number. When it flattens, the optimizer has stopped improving. That is like a car's fuel gauge reading "arrived at the best spot this road reaches". It cannot tell you how far that spot is from where you wanted to go, because the distance, $KL(q\,\|\,p) = \log p(D) - \text{ELBO}$, involves $\log p(D)$, which you never know. What you can do is compare: two guides on the same model and data share the same $\log p(D)$, so the one with the higher ELBO is closer to the posterior.
NUTS has several lights: $\hat R$ (do chains started in different places agree?), ESS (how many independent draws is this worth?), divergences (did the simulation blow up somewhere, meaning a region it cannot enter?), tree depth (is it working too hard?), and trace plots. They are designed to catch exactly the failures that make a finite chain wrong.
Three ways to say it:
- Picture: SVI's dashboard says "the engine has stopped"; NUTS's dashboard says "the survey teams agree, covered enough ground, and found no cliffs".
- Numbers: NUTS is "green" with $\hat R \le 1.01$, ESS ≥ 400 and 0 divergences (rules of thumb); SVI is "done" when the smoothed ELBO's relative change stays below your tolerance for your patience window.
- Slogan: a flat ELBO means "done optimizing", not "correct".
Reading the gap between two flat ELBO curves. Correlated Gaussian posterior, ρ = 0.9. (In this toy problem we happen to know $\log p(D)$; in real problems you do not.)
- Here $\log p(D) = \log\big(2\pi\sqrt{1 - \rho^2}\big) = \log(2\pi \times 0.436) = \log 2.739 = 1.008$.
- Full-rank guide: it can match the posterior exactly, so its ELBO flattens at $1.008$ (KL = 0).
- Mean-field guide: its KL to a correlated Gaussian is $-\tfrac12\log(1 - \rho^2) = -\tfrac12 \log 0.19 = 0.830$, so its ELBO flattens at $1.008 - 0.830 = 0.177$.
- Both curves are flat, so both loops would stop. In real life you see only 0.177 and 1.008. Their difference, 0.83, is $KL_{\text{MF}} - KL_{\text{FR}}$: the full-rank guide is closer. You cannot tell from the ELBO alone that it is exactly right.
- For comparison, NUTS on this chapter's hierarchical model (the Code-it block): $\hat R$ max 1.002, ESS min 1 180, 0 divergences: all lights green, and each light is a check you can actually read.
| Engine | Diagnostic | What it checks | Rule of thumb |
|---|---|---|---|
| SVI | ELBO (loss) trace, smoothed | the optimizer has stopped improving | relative change below your tolerance for your patience window |
| SVI | several seeds and starting points | optimization stability, local optima | decision quantities agree across runs |
| SVI | ELBO comparison across guides | which guide is closer in KL (same model and data) | a higher average ELBO (beyond seed-to-seed noise) wins |
| SVI | posterior predictive checks | does the fitted model reproduce the data? (Chapter 6.8) | no systematic misfit |
| SVI | importance-sampling check (Pareto $\hat k$) | how far $q$ is from the posterior, from weights $p/q$ | $\hat k \lt 0.7$ usable (a rule of thumb from the literature) |
| SVI | comparison with NUTS on a subset | the family gap, directly | differences small compared with the posterior sd |
| NUTS | split $\hat R$ | chains started apart agree | $\le 1.01$ (older advice: 1.1) |
| NUTS | bulk and tail ESS, MCSE | enough independent information | ESS ≥ 400 in total (about 100 per chain) |
| NUTS | divergences | regions the integrator cannot enter | 0; any divergence needs investigating |
| NUTS | tree depth, E-BFMI, trace plots | efficiency, energy exploration, visual mixing | not stuck at max depth; E-BFMI not below about 0.3 |
Key identity: $\log p(D) = \text{ELBO}(q) + KL(q\,\|\,p(\theta\mid D))$, so for the same model and data, $\text{ELBO}(q_1) - \text{ELBO}(q_2) = KL(q_2\,\|\,p) - KL(q_1\,\|\,p)$. In NumPyro: svi.run(...).losses (loss = −ELBO); mcmc.print_summary() prints n_eff and r_hat and the number of divergences; mcmc.get_extra_fields()["diverging"] gives them per draw.
Why do we need it?
Neither engine announces failure on its own. SVI happily returns a too-narrow guide with a beautifully flat ELBO; NUTS happily returns draws from chains that never mixed. Diagnostics are the only way to know which answers to trust.
Where is it used?
Every NumPyro, Stan and PyMC workflow: print_summary, ArviZ's az.summary (R̂, ESS, MCSE) and az.plot_trace; ELBO traces in training loops with early stopping; the "Yes, but did it work?" Pareto-$\hat k$ check for variational fits.
How is it used?
For SVI: plot the smoothed loss, refit with 2–3 seeds and a richer guide, compare average ELBOs and decision quantities, run posterior predictive checks. For NUTS: run at least 4 chains from dispersed starts, require R̂ ≤ 1.01, enough ESS and zero divergences before reading any number.
"The ELBO has flattened, so the posterior approximation is good."
A flat ELBO means the optimizer has stopped improving within the family. The remaining gap, $KL(q\,\|\,p)$, is invisible in the trace.
"A higher ELBO always means a better posterior approximation."
Only for the same model and the same data. Across different models the ELBO is a (lower bound on the) evidence, which answers a different question: which model explains the data better.
"R̂ = 1.00, so NUTS has converged."
R̂ near 1 is necessary, not sufficient. All chains can miss the same region (the funnel's neck) and still agree. Read divergences and ESS too.
"A few divergences are fine; the estimates look reasonable."
Divergences mark places the sampler cannot explore, so the draws can be biased exactly there. Raise target_accept_prob (to 0.9–0.99) or reparameterize, and re-check.
Your forecasting loop's relative-ELBO early stopping with patience and best-state checkpointing is an excellent optimization diagnostic: it tells you when the guide has stopped improving and keeps the best state despite noise. It cannot measure the family gap. The cheap additions are: refit with two or three seeds; compare the average final ELBO of the full-rank and low-rank guides on the same data (higher is closer in KL); posterior predictive checks of the forecasts; and, on a smaller version of the model, a NUTS run whose $\hat R$, ESS and divergences are clean.
"How do you know SVI converged?" "The ELBO stopped improving, so the posterior is right."
"The ELBO plateau tells me the optimizer converged. Whether the approximation is good is a separate question, which I check by comparing guides, seeds and NUTS."
Model answer: "For SVI I monitor a smoothed ELBO and stop on relative improvement with patience; that shows optimization convergence. Because log p(D) = ELBO + KL and log p(D) is unknown, the plateau says nothing about the KL gap, so I compare guides by their ELBO on the same data, check stability across seeds, run posterior predictive checks, and validate against NUTS on a smaller problem. For NUTS I require R̂ ≤ 1.01, enough bulk and tail ESS, and zero divergences before I read any estimate."
$\log p(D) = \text{ELBO} + KL(q\,\|\,p)$: flat ELBO = optimizer done; KL invisible. Same model and data: higher ELBO = closer guide.
NUTS: R̂ ≤ 1.01, ESS ≥ 400 total, 0 divergences, not stuck at max tree depth (rules of thumb).
Trap: R̂ ≈ 1 is necessary, not sufficient; a flat ELBO is not accuracy.
Quick check: on the same model and data, AutoNormal ends at an average ELBO of −1 520.4 and AutoMultivariateNormal at −1 512.1. What can you conclude, and what can't you?
The difference, 8.3 nats, equals $KL_{\text{AutoNormal}} - KL_{\text{AutoMVN}}$: the full-rank guide is closer to the posterior by 8.3 nats of KL. You cannot conclude that AutoMVN is close in absolute terms, because $\log p(D)$ is unknown; its own KL could still be large. (Also check that the gap is larger than the noise between seeds.)
When each is the right tool core
Think of SVI as a fast satellite photo and NUTS as a detailed survey on foot. You would not survey a whole country on foot every morning; you also would not build a bridge from a satellite photo alone. The good choice depends on four things:
- Size: how many rows each gradient must touch ($N$) and how many unknowns ($d$).
- Stakes: does a decision hang on the tails or on a correlation (a 95% interval, $P(\theta_B \gt \theta_A)$ near a threshold, a forecast's upper quantile)?
- Shape: is the posterior roughly Gaussian, strongly correlated, funnel-shaped, or multi-peaked?
- Schedule: one analysis, daily refits of many series, or answers within a latency budget?
Three ways to say it:
- Picture: satellite photo for coverage and speed, survey on foot for the details that matter, and the survey to check the photo.
- Numbers: a 10-parameter A/B model on aggregated counts runs NUTS in seconds; a model with a few hundred latents refit daily for 500 series needs SVI.
- Slogan: NUTS when you can afford it and the details matter; SVI for scale and speed; both when it counts.
Two jobs, worked through. (Illustrations, not descriptions of your actual code.)
- Job 1: one A/B experiment. 8 segments, a population mean and spread: $d = 10$ latents. Data aggregated to 8 (conversions, users) pairs, so each gradient touches 8 cells. This chapter's code ran NUTS on exactly this model in about 1–2 s with $\hat R$ 1.002 and 0 divergences.
- Verdict for job 1: NUTS is affordable and gives the reference posterior; use it (or use SVI and check it against NUTS once).
- Job 2: daily forecasts for 500 series. Suppose each model has 25 changepoint adjustments + 20 yearly Fourier coefficients (order 10) + 6 weekly (order 3) + 10 holidays + 5 regressors + 4 others = 70 latents and about 1 100 daily rows.
- NUTS with 4 chains × 2 000 iterations × (say) 63 leapfrog steps = 504 000 gradients per series; × 500 series = 252 million gradients every day. SVI with 5 000 steps = 5 000 gradients per series, 2.5 million per day: about 100 times fewer.
- Verdict for job 2: SVI with a full-rank guide (70 latents: $70 + 70 \times 71/2 = 2\,555$ guide numbers, cheap), validated against NUTS on a few series.
| Situation | Usually the right tool | Why |
|---|---|---|
| Small or medium model, full-data gradients affordable | NUTS | no family gap; diagnostics tell you when it fails |
| Decision depends on tails, intervals or correlations | NUTS, or SVI validated against NUTS | mean-field under-dispersion moves tail probabilities |
| New model, still being developed | NUTS (as the reference) | separates model problems from inference problems |
| Huge N (millions of rows) | SVI with minibatches | plain NUTS reads all rows per leapfrog step |
| Large d (thousands of latents) | SVI with a low-rank or mean-field guide | full-rank guides grow like $d^2$ |
| Many refits (many series, many experiments, daily) | SVI, with periodic NUTS checks | cost per fit matters more than the last decimal |
| Latency budget, amortized inference | SVI | a fitted guide is instant to sample |
| Funnel geometry | reparameterize first, then either | both engines struggle in the centered funnel |
| Several modes | neither alone: several NUTS chains + R̂, several SVI starts, rethink the model | a Gaussian guide picks one peak; chains can stick |
Why do we need it?
Picking the engine by habit gives either wasted compute (NUTS where SVI would do) or silently wrong decisions (mean-field SVI where tails matter). A short checklist turns this into a defensible choice you can explain.
Where is it used?
Design reviews of Bayesian systems; choosing inference for an experimentation platform versus a forecasting pipeline; Stan's and PyMC's documentation (NUTS by default, ADVI as the fast option); interviews ("why SVI and not NUTS?").
How is it used?
Answer the four questions (size, stakes, shape, schedule). If NUTS is affordable for the decision that matters, use it; otherwise use the richest guide you can afford and validate it against NUTS on something smaller. Write the reason down with the results.
"NUTS for small data, SVI for big data."
Size is one factor. A small model refit thousands of times a day may need SVI; a large model whose single decision depends on a tail probability may justify an expensive NUTS run (or a NUTS check on part of it).
"NUTS is too slow on my model, so I will switch to SVI."
First ask why NUTS is slow. Deep trees and divergences usually mean bad geometry (a funnel, wildly different scales). That geometry hurts SVI too, just silently. Reparameterize or rescale, then decide.
"Once we choose SVI, NUTS is irrelevant."
NUTS remains the natural reference for checking the guide, on a subset or a smaller model.
Both of your projects use SVI, and both can have good reasons: if your forecasting model is refit as data arrive and has tens to hundreds of correlated latents (check your own refit schedule and $d$), SVI's cost per fit is far lower, and your loop is already JIT-compiled with early stopping. In an A/B framework like yours, one engine for many metrics and experiments keeps the platform simple and fast. The honest addition, in an interview, is how the SVI answers are (or would be) checked: NUTS on individual experiments or on a smaller version of the forecasting model.
"We used SVI because NUTS is too slow." (and nothing else)
Give the trade-off: size, refit schedule, what the guide captures, and how you know it is good enough.
Model answer (swap in your real refit schedule and latent count): "SVI fit our constraints: the model is refit regularly, it has dozens to hundreds of correlated latent parameters, and NUTS would need hundreds of thousands of full-data gradients per fit. SVI with a full-rank guide, or low-rank for big models, captures the main posterior correlations at a fraction of the cost. The price is a guide-dependent approximation that tends to under-state uncertainty, so I would validate it against NUTS on a smaller version of the model and compare the decision quantities, such as forecast intervals or P(B > A)."
Four questions: size ($N$, $d$), stakes (tails, correlations), shape (Gaussian, correlated, funnel, modes), schedule (once, daily, real-time).
NUTS when affordable and details matter; SVI for scale, speed and many refits; hybrid when SVI is used and tails matter.
Trap: slow NUTS usually signals bad geometry, which also hurts SVI; reparameterize before switching.
Quick check: an experiment platform runs 300 small A/B tests a week, each with about 10 latent parameters, and reports P(B > A). Which engine?
Each model is small, so NUTS takes seconds per test: 300 × a few seconds is affordable, and it gives reference-quality tails. If the platform already uses SVI for speed and uniformity, a reasonable compromise is SVI with a full-rank guide plus routine NUTS checks on a sample of tests, comparing P(B > A).
Hybrid workflows: validate SVI with NUTS, and let SVI help NUTS core
You do not have to pick one engine for ever. The two can work as a team.
NUTS checks SVI. Before a ship sails fast through a channel, a small pilot boat goes first and checks the depth. Run NUTS on something small enough to afford (a subset of the data, a few series, a few experiments, a simplified model), run SVI with your candidate guides on the same small problem, and compare the numbers your decisions use. If the guide passes there, run SVI at full scale with more confidence; re-check from time to time.
SVI helps NUTS. A fitted guide is a rough map of the posterior: where it is (its centre) and how it is stretched and tilted (its covariance). NUTS can start from that centre instead of a random point, and it can run in coordinates where the guide's ellipse becomes a circle (preconditioning), so its steps fit the posterior's shape. NumPyro's NeuTra reparameterization does a curved version of this with a flexible guide.
Three ways to say it:
- Picture: a pilot boat (NUTS) checks the channel before the big ship (SVI) sails at speed; and the ship's chart (the guide) helps the pilot steer.
- Numbers: in this chapter's code, the full-rank guide's μ was within 0.21 posterior sds of NUTS with 92% of its spread; AutoNormal's τ had only 63% of NUTS's spread.
- Slogan: validate small, deploy big, re-check occasionally.
A validation report from this chapter's code (the 8-segment hierarchical model; numbers from the Code-it block at the end).
- Reference (NUTS, R̂ max 1.002, ESS min 1 180, 0 divergences): μ = −2.157 ± 0.141, τ = 0.155 ± 0.132. Its own Monte Carlo error for the mean of μ is at most $0.141/\sqrt{1\,180} = 0.004$, so differences much bigger than 0.01 are real.
- Choose a tolerance before looking (your choice, not a standard): means within 0.25 posterior sd of NUTS, sds within ±25% of NUTS's.
- AutoMultivariateNormal: μ = −2.187 ± 0.130. Mean gap $0.030/0.141 = 0.21$ sd ✓; sd ratio $0.130/0.141 = 0.92$ ✓. τ = 0.131 ± 0.108: gap $0.024/0.132 = 0.18$ sd ✓; sd ratio $0.108/0.132 = 0.82$ ✓ (only just).
- AutoNormal: μ = −2.163 ± 0.120 ✓ (gap 0.04 sd, ratio 0.85). τ = 0.111 ± 0.083: gap $0.044/0.132 = 0.33$ sd ✗; ratio $0.083/0.132 = 0.63$ ✗.
- Decision: use the full-rank guide for this model, and remember that even it under-states τ's spread by about 18%, because τ's posterior is skewed against zero and a Gaussian (on log τ) cannot match it exactly.
A validation workflow (one common recipe):
- Build the model and run prior predictive checks (Chapter 6.2).
- Reference: run NUTS on a problem small enough to afford: a random subset of rows, a few groups or series, a shorter history, or a simplified model (fewer changepoints). Require clean diagnostics; a NUTS run with divergences is not a reference.
- Candidates: fit SVI with each candidate guide on the same problem.
- Compare decision quantities (posterior means and sds, 90% intervals, $P(\theta_B \gt \theta_A)$, forecast quantiles), relative to the posterior sd, and ignore gaps smaller than about 2 NUTS MCSE.
- Pick the cheapest guide that passes; run SVI at full scale; re-validate when the model, the data or the guide change.
- Caveat: a subset posterior is wider and often less Gaussian than the full-data one, so it is usually a harder test for a Gaussian guide; but some problems appear only at full scale (many groups, many changepoints, more correlated latents). That is why a smaller version of the full model is a valuable second check.
- SVI → NUTS: start NUTS at the guide's median,
NUTS(model, init_strategy=init_to_value(values=guide.median(params))); precondition with the guide's covariance (NUTS's owndense_mass=Truelearns a full mass matrix during warmup); or useNeuTraReparam(guide, params), which needs a guide of the AutoContinuous kind (e.g. AutoMultivariateNormal, AutoLowRankMultivariateNormal, AutoIAFNormal, AutoBNAFNormal; not AutoNormal).
Why do we need it?
SVI's error cannot be read off its own output, and full-scale NUTS may be unaffordable. A small, affordable NUTS run is the cheapest direct measurement of SVI's family gap on your actual model; and a fitted guide can make NUTS cheaper and more reliable.
Where is it used?
Validating production variational fits in industry (experimentation platforms, demand forecasting); NumPyro's init_to_value and NeuTraReparam; Stan users comparing ADVI with NUTS; simulation-based calibration studies that test inference code on simulated data.
How is it used?
Write one script that fits NUTS and each candidate guide on the same subset and prints a side-by-side table of the decision quantities with a pass/fail column; keep it in the repository and rerun it when the model changes. Use the guide's median as NUTS's starting point when chains start badly.
"NUTS on a subset gives the truth for the full data."
It gives a reference for the subset posterior. Compare SVI and NUTS on the same subset; that measures the guide's gap on a posterior of the same model, which is what you want.
"Validated once, validated for ever."
A new regressor, more segments, a different likelihood, a longer history or a different guide rank can all change the posterior's shape. Re-run the comparison when they change.
"If SVI and NUTS disagree, SVI is wrong."
Usually, but check the reference first: a NUTS run with divergences, low ESS or $\hat R \gt 1.01$ is not a reference.
"Starting NUTS at the SVI answer biases it toward SVI."
A good starting point only shortens warmup; a valid chain forgets where it started. Preconditioning with the guide changes coordinates, not the target.
For your forecasting model, a natural reference is NUTS on a smaller version: one or a few series, a shorter history, perhaps fewer candidate changepoints, using the same likelihood and priors. Compare the full-rank and low-rank guides' forecast intervals, trend and changepoint summaries against it. For the A/B framework, run NUTS on a sample of past experiments and compare the reported $P(\theta_B \gt \theta_A \mid D)$ and lift intervals. In both cases the fitted guide's median is a good NUTS starting point (init_to_value).
"How do you know your SVI posterior is good enough?" "Because the ELBO converged."
"Because I compared it with a NUTS reference on a problem small enough for NUTS, using the quantities my decisions depend on."
Model answer: "SVI's error is invisible in its own output, so I validate it. I run NUTS on a subset or a smaller version of the model, check its R̂, ESS and divergences, fit my candidate guides on the same problem, and compare means, sds, intervals and the decision quantities like P(B > A) or forecast quantiles, against a tolerance I set in advance and the NUTS Monte Carlo error. I pick the cheapest guide that passes, run it at scale, and re-validate when the model changes. The guide can also help NUTS: I can initialize from its median or precondition with its covariance."
Validate: NUTS reference (clean diagnostics) on a subset or smaller model → SVI guides on the same problem → compare decision quantities (in posterior-sd units, beyond 2 MCSE) → cheapest passing guide at scale → re-validate on change.
SVI → NUTS: init_to_value(values=guide.median(params)), dense mass / preconditioning, NeuTraReparam (AutoContinuous guides).
Trap: a subset is usually a harder test, but some problems appear only at full scale.
Quick check: NUTS on a subset says P(θ_B > θ_A) = 0.91 with MCSE 0.006; your SVI guide says 0.86. Is the difference real, and does it matter?
The gap is 0.05, about 8 MCSE, so it is real, not Monte Carlo noise. Whether it matters depends on the decision: if you ship at P > 0.9, SVI says "don't ship" and NUTS says "ship", so yes. Try a richer guide (full-rank or low-rank instead of mean-field) and see whether it closes the gap; if not, use NUTS for this decision.
Recap, cheat sheet and practice
- SVI optimizes a guide $q_\phi$ to maximize the ELBO and returns a formula; NUTS samples with a Markov chain and returns draws. Both need only $\log p(D, \theta)$ and its gradient.
- Speed: NUTS pays chains × iterations × leapfrog steps (up to $2^{\text{depth}} - 1$) × all $N$ rows; SVI pays steps × particles × $B$ rows. The posterior's shape drives both bills; for small models both take seconds.
- Scale: SVI can use minibatches scaled by $N/B$ (unbiased, noisier); plain NUTS cannot. Guides grow like $2d$, $d(r+2)$ or $d + d(d+1)/2$.
- Uncertainty: reverse KL makes SVI under-dispersed; mean-field on a correlated posterior gives marginal sd $\sigma\sqrt{1-\rho^2}$ and zero correlation, so decision probabilities can move either way. Full-rank fixes correlations, not shape.
- Guide dependence: SVI's answer depends on the family, the start and the optimizer settings; with several modes it silently picks one. NUTS depends on parameterization and tuning, and its failures show in $\hat R$ and divergences.
- Errors: SVI = family gap + optimization error (a bias that more steps cannot fix). NUTS = Monte Carlo error sd/√ESS + any non-convergence bias.
- Never say "NUTS gives the exact posterior". Say: asymptotically exact; finite chains have Monte Carlo error and convergence concerns; and it is the posterior of your model.
- Diagnostics: a flat ELBO means "optimizer done", not "correct"; compare guides' ELBOs on the same data, seeds, PPCs, NUTS. For NUTS: $\hat R \le 1.01$, ESS ≥ 400, 0 divergences.
- Choose by size, stakes, shape and schedule; combine them: validate SVI against NUTS on a subset or a smaller model, and use the guide to start or precondition NUTS.
Cheat sheet
| Topic | SVI | NUTS |
|---|---|---|
| What it does | $\max_\phi \text{ELBO}(\phi)$ → formula $q_{\phi^\star}$ | Markov chain → draws $\theta^{(1..S)}$ |
| Cost (row-gradients) | $T \times K \times B$ | $C \times (W + S) \times \bar L \times N$ |
| Many rows | minibatches, scale by $N/B$ | every leapfrog step reads all rows |
| Many parameters | $2d$ · $d(r+2)$ · $d + d(d+1)/2$ | $S \times d$ draws; trajectories lengthen slowly |
| Main error | family gap + optimization error (bias) | MCSE $= \text{sd}/\sqrt{\text{ESS}}$ + non-convergence bias |
| Infinite effort gives | $q^\star$, the best member of the family | the posterior (asymptotically exact) |
| Uncertainty | under-dispersed; mean-field sd $= \sigma\sqrt{1-\rho^2}$ | shape included, if converged |
| Depends on | guide family, start, learning rate, steps, seed | parameterization, warmup, target acceptance |
| Diagnostics | ELBO trace, seeds, ELBO across guides, PPC, Pareto $\hat k$, NUTS check | $\hat R \le 1.01$, ESS ≥ 400, 0 divergences, tree depth, traces |
| Right tool when | huge $N$ or $d$, many refits, latency | affordable, tails matter, new model, reference |
| Together | NUTS on a subset validates SVI; the guide's median and covariance start and precondition NUTS | |
import time
import numpy as np
import jax, jax.numpy as jnp
import numpyro
import numpyro.distributions as dist
from numpyro import handlers
from numpyro.infer import MCMC, NUTS, SVI, Trace_ELBO
from numpyro.infer.autoguide import AutoNormal, AutoMultivariateNormal
from numpyro.infer.util import log_density
from numpyro.diagnostics import summary
# Wall-clock times below are from one laptop CPU and include JIT compilation:
# yours WILL differ. The posterior numbers are seeded and should match closely.
def fit_nuts(model, *args):
t0 = time.perf_counter()
mcmc = MCMC(NUTS(model), num_warmup=500, num_samples=1000, num_chains=2,
chain_method="sequential", progress_bar=False)
mcmc.run(jax.random.PRNGKey(0), *args, extra_fields=("diverging",))
draws = mcmc.get_samples(group_by_chain=True)
jax.block_until_ready(draws)
secs = time.perf_counter() - t0 # includes JIT compilation
diag = summary(draws) # split R-hat and ESS per site
rhat = max(float(np.max(v["r_hat"])) for v in diag.values())
ess = min(float(np.min(v["n_eff"])) for v in diag.values())
ndiv = int(mcmc.get_extra_fields()["diverging"].sum())
flat = {k: np.asarray(v).reshape((-1,) + v.shape[2:]) for k, v in draws.items()}
return flat, secs, f"R-hat max {rhat:.3f} | ESS min {ess:.0f} | divergences {ndiv}"
def fit_svi(model, Guide, *args, steps=20000, lr=0.01):
t0 = time.perf_counter()
guide = Guide(model)
svi = SVI(model, guide, numpyro.optim.Adam(lr), Trace_ELBO())
res = svi.run(jax.random.PRNGKey(0), steps, *args, progress_bar=False)
draws = guide.sample_posterior(jax.random.PRNGKey(1), res.params, sample_shape=(4000,))
jax.block_until_ready(draws)
secs = time.perf_counter() - t0 # includes JIT compilation
a, b = res.losses[-4000:-2000].mean(), res.losses[-2000:].mean()
rel = abs(float(b - a)) / abs(float(a)) # did the loss (= -ELBO) flatten?
return {k: np.asarray(v) for k, v in draws.items()}, secs, f"loss {float(b):.2f} | relative change {rel:.0e}"
def report(name, d, secs, diag, sites):
cells = " ".join(f"{s} {d[s].mean():.3f} ± {d[s].std():.3f}" for s in sites)
print(f" {name:10s} {cells} [{secs:4.1f} s] {diag}")
# ---------- 1. A correlated posterior: two nearly identical predictors ----------
rng = np.random.default_rng(0)
x1 = rng.normal(size=200)
x2 = x1 + 0.3 * rng.normal(size=200) # x2 is almost a copy of x1
y = 1.0 * x1 + 1.0 * x2 + rng.normal(size=200)
X, y = jnp.array(np.c_[x1, x2]), jnp.array(y)
def reg(X, y=None):
b = numpyro.sample("b", dist.Normal(0.0, 5.0).expand([2]).to_event(1))
sigma = numpyro.sample("sigma", dist.HalfNormal(2.0))
numpyro.sample("y", dist.Normal(X @ b, sigma), obs=y) # Normal(loc, scale = sd)
print("1. Correlated regression (b1 and b2 share one signal)")
for name, fit in [("NUTS", lambda: fit_nuts(reg, X, y)),
("AutoNormal", lambda: fit_svi(reg, AutoNormal, X, y)),
("AutoMVN", lambda: fit_svi(reg, AutoMultivariateNormal, X, y))]:
d, secs, diag = fit()
b = d["b"]
print(f" {name:10s} b1 {b[:,0].mean():.3f} ± {b[:,0].std():.3f} b2 {b[:,1].mean():.3f} ± {b[:,1].std():.3f}"
f" corr {np.corrcoef(b.T)[0,1]:+.2f} sd(b1+b2) {(b[:,0] + b[:,1]).std():.3f} [{secs:.1f} s]")
# NUTS b1 0.871 ± 0.230 b2 1.194 ± 0.225 corr -0.95 sd(b1+b2) 0.074 [1.7 s]
# AutoNormal b1 0.879 ± 0.077 b2 1.205 ± 0.077 corr +0.00 sd(b1+b2) 0.109 [0.7 s]
# AutoMVN b1 0.883 ± 0.237 b2 1.211 ± 0.219 corr -0.95 sd(b1+b2) 0.075 [0.9 s]
# -> mean-field: sds 3x too small, no correlation, and sd(b1+b2) too WIDE.
# ---------- 2. A small hierarchical model: 8 segments, partial pooling ----------
n = jnp.array([100, 60, 240, 20, 30, 180, 15, 80]) # visitors per segment
k = jnp.array([8, 7, 24, 3, 2, 20, 1, 11]) # conversions per segment
def segments(n, k=None):
mu = numpyro.sample("mu", dist.Normal(-2.2, 1.0)) # population log-odds
tau = numpyro.sample("tau", dist.HalfNormal(0.5)) # between-segment spread
with numpyro.plate("seg", n.shape[0]):
z = numpyro.sample("z", dist.Normal(0.0, 1.0)) # non-centered
numpyro.sample("k", dist.Binomial(total_count=n, logits=mu + tau * z), obs=k)
print("2. Hierarchical segments (10 latent dimensions: mu, tau, z1..z8)")
for name, fit in [("NUTS", lambda: fit_nuts(segments, n, k)),
("AutoNormal", lambda: fit_svi(segments, AutoNormal, n, k)),
("AutoMVN", lambda: fit_svi(segments, AutoMultivariateNormal, n, k))]:
d, secs, diag = fit()
report(name, d, secs, diag, ["mu", "tau"])
# NUTS mu -2.157 ± 0.141 tau 0.155 ± 0.132 [ 1.2 s] R-hat max 1.002 | ESS min 1180 | divergences 0
# AutoNormal mu -2.163 ± 0.120 tau 0.111 ± 0.083 [ 0.7 s] loss 19.38 | relative change 3e-03
# AutoMVN mu -2.187 ± 0.130 tau 0.131 ± 0.108 [ 1.3 s] loss 19.44 | relative change 3e-04
# -> both loss traces are flat; AutoNormal's tau spread is only 63% of NUTS's.
# ---------- 3. Minibatching: the N/B-scaled minibatch log density is unbiased ----------
def reg_mb(X, y, batch=None):
b = numpyro.sample("b", dist.Normal(0.0, 5.0).expand([2]).to_event(1))
sigma = numpyro.sample("sigma", dist.HalfNormal(2.0))
with numpyro.plate("data", X.shape[0], subsample_size=batch) as idx: # scales by N/B
numpyro.sample("y", dist.Normal(X[idx] @ b, sigma), obs=y[idx])
point = {"b": jnp.array([1.0, 1.0]), "sigma": 1.0}
full = log_density(reg_mb, (X, y), {}, point)[0]
one = lambda key: log_density(handlers.seed(reg_mb, key), (X, y), {"batch": 20}, point)[0]
mb = jax.jit(jax.vmap(one))(jax.random.split(jax.random.PRNGKey(0), 4000))
print(f"3. log density: full data {float(full):.1f} | 4000 minibatches of 20: "
f"mean {float(mb.mean()):.1f}, sd {float(mb.std()):.1f}")
# 3. log density: full data -289.9 | 4000 minibatches of 20: mean -289.5, sd 33.5
# (whole script: about 7-10 s on our machine)
1. Which sentence about NUTS would you say in an interview?
2. A Gaussian posterior has sd 2 for each of two parameters and correlation 0.8. What marginal sd does the best mean-field guide give?
3. You run your mean-field SVI loop ten times longer and the posterior sds do not change at all. What is the most likely reason?
4. N = 500 000 rows and minibatches of B = 250. By what factor must the batch log-likelihood be multiplied?
plate(..., subsample_size=B) applies it automatically.5. Which check tells you most directly how far an SVI guide is from the true posterior?
6. A NUTS estimate has MCSE 0.05 with ESS 100. What ESS do you need for MCSE 0.025?
Practice problems
A. Interview: "Explain the difference between SVI and NUTS in about a minute."
"Both approximate the posterior of the same model. SVI turns inference into optimization: I choose a family of distributions, the guide, and maximize the ELBO with stochastic gradients, so I get back a fitted distribution. It is fast, scales with minibatches and low-rank guides, but it is approximate by design: its error is the gap between the family and the posterior, it depends on the guide, and reverse KL makes it under-dispersed, badly so for mean-field with correlated parameters. NUTS is MCMC: it uses Hamiltonian dynamics to produce draws whose distribution converges to the posterior. It is asymptotically exact and captures shape and correlations, but it is expensive, needs full-data gradients, and a finite run has Monte Carlo error and convergence risks, so I check R̂, ESS and divergences. In practice I use SVI at scale and validate it against NUTS on a smaller problem."
B. Two segment effects have posterior sds 1, correlation 0.8, and means 0 (A) and 0.5 (B). Compute the mean-field sds, the sd of θ_B − θ_A under both, and P(θ_B > θ_A) under both.
Mean-field sd $= \sqrt{1 - 0.64} = 0.6$ each. True $Var(\theta_B - \theta_A) = 1 + 1 - 2(0.8) = 0.4$, sd $0.632$. Mean-field ignores the covariance: $0.36 + 0.36 = 0.72$, sd $0.849$ (too wide). True $P = \Phi(0.5/0.632) = \Phi(0.791) = 0.786$; mean-field $P = \Phi(0.5/0.849) = \Phi(0.589) = 0.722$. Mean-field is too narrow for each effect but too wide for their difference, so it understates the evidence that B beats A.
C. N = 2 000 000 rows. NUTS: 4 chains × 1 500 iterations, about 15 leapfrog steps each. SVI: 10 000 steps, 1 particle, minibatches of 2 000 rows. Compare the row-gradient bills.
NUTS: $4 \times 1\,500 \times 15 = 90\,000$ full-data gradients × 2 000 000 rows $= 1.8 \times 10^{11}$ row-gradients. SVI: $10\,000 \times 1 \times 2\,000 = 2 \times 10^{7}$. Ratio: $1.8 \times 10^{11} / 2 \times 10^{7} = 9\,000$. NUTS needs about 9 000 times more gradient work here, before counting any extra SVI steps the minibatch noise may require.
D. Interview: "Your NUTS reference has 37 divergences in 4 000 draws, and your SVI answer disagrees with it. What do you do?"
"First I fix the reference, because a NUTS run with divergences is not trustworthy: the divergences mark regions it cannot explore, often a funnel in a hierarchical model. I would try the non-centered parameterization, raise target_accept_prob to 0.95 or 0.99, and rerun until there are no divergences and R̂ and ESS are fine. Then I compare SVI with the clean reference. If they still disagree on the decision quantities, I try a richer guide, and if none passes I use NUTS for that decision. Note that the same geometry that caused the divergences probably also hurts SVI, so the reparameterization may close the gap on its own."
E. NUTS reports P(θ_B > θ_A) = 0.94 with ESS 250 for the indicator. Your shipping rule is P > 0.95. What do you report, and how many effective draws would make the answer clear?
MCSE $= \sqrt{0.94 \times 0.06 / 250} = \sqrt{0.0002256} = 0.015$. A ±2 MCSE band is about [0.91, 0.97], which straddles 0.95, so the run cannot decide. For MCSE 0.005 you need ESS $= 0.94 \times 0.06 / 0.005^2 = 0.0564 / 0.000025 = 2\,256$: about 9 times more effective draws. (And remember the threshold rule itself, not just its Monte Carlo error, is a decision choice.)
F. Interview: "Why does your forecasting model use SVI rather than NUTS, and what would make you switch?"
(Illustrative numbers: replace them with your real d, refit schedule and step counts.) "The model has dozens to hundreds of correlated latent parameters (trend, changepoint adjustments, Fourier coefficients, holidays, regressors) and is refit as new data arrive, possibly for many series. NUTS would need hundreds of thousands of full-data gradients per fit; a JIT-compiled SVI loop with relative-ELBO early stopping needs a few thousand. A full-rank guide, or low-rank for large models, keeps the main correlations that drive forecast intervals, which mean-field would lose. I would switch to NUTS, or add it, when a decision depends on tail quantiles (capacity planning at the 99th percentile), when validation on a smaller version shows the guide's intervals are too narrow, or when the model is small enough that NUTS is cheap. I'd also never treat NUTS as exact: it is the reference, with its own Monte Carlo error and diagnostics."
JAX fundamentals: arrays, pure functions, grad, vmap, scan, PRNG keys
NumPyro is built on JAX. Every model you write, every SVI step and every NUTS trajectory is JAX code. JAX looks like NumPy, but it plays by a few strict rules: arrays never change, functions must be pure, and randomness comes from keys you hold in your hand. Those rules are exactly what let JAX hand you gradients, batching and fast loops for free. This chapter teaches the rules and the tools, one at a time, and shows where each one lives inside NumPyro.
- Work with JAX arrays: immutable, typed, float32 by default, updated with
.at[...] - Say what a pure function is and predict what goes wrong when a function reads a global, prints, or uses NumPy's random numbers under a transformation
- Write code in the functional style: state goes in, new state comes out (the SVI state, the parameters, the key)
- Use the transformations
grad(exact derivatives, checked against finite differences) andvmap(batching without loops), and stack them - Write step-by-step recursions (an AR process, a cumulative trend) with
lax.scanand its carried state - Handle randomness with PRNG keys: split, never reuse, and explain why reproducible Bayesian inference needs explicit keys
- See each idea inside NumPyro: models are traced functions,
Predictivemaps over posterior draws,svi.updatetakes gradients and threads a key
What we need from earlier chapters: derivatives, the chain rule and automatic differentiation (Calculus Chapter 2.9); arrays, shapes, broadcasting and vectorisation (Linear Algebra Chapter 1.16); the likelihood of a Normal mean (Chapter 6.1); the posterior predictive (Chapter 6.1); SVI and the ELBO (Chapter 6.12). Words used everywhere below. An array is a grid of numbers; its shape lists how many numbers it has along each axis (shape (3, 5) = 3 rows, 5 columns); its dtype ("data type") says what kind of number each entry is (float32 = a decimal number stored in 32 bits, about 7 significant digits; float64 = 64 bits, about 16 digits; int32 = a whole number). A device is the hardware that does the arithmetic: the CPU, a GPU or a TPU. In code, import jax.numpy as jnp is JAX's NumPy-like library and from jax import lax gives the lower-level building blocks. All timings in this chapter were measured on one laptop CPU (JAX 0.11, NumPyro 0.22); timings vary by machine, so read them as orders of magnitude.
JAX arrays: like NumPy, but they never change core
A NumPy array is like a whiteboard: you can rub out one number and write a new one in the same place (a[0] = 10). A JAX array is like a printed page: you cannot change it. To "change" one number you print a new page that is the same except for that number, and the old page stays exactly as it was.
Why would anyone want that? Because JAX's tools (compiling, differentiating, batching) read your code once and then reason about it. If any array could be secretly rewritten at any moment, that reasoning would break. Arrays that never change make the code easy to analyse, safe to reorder and safe to run on a GPU.
Three ways to say it:
- Picture: NumPy = whiteboard you edit; JAX = printed pages, and an "edit" prints a new page.
- Numbers:
x = [1, 2, 3];y = x.at[0].set(10)givesy = [10, 2, 3]whilexis still[1, 2, 3]. - Slogan: in JAX you never change an array; you make a new one.
Five lines in the JAX console (real outputs, JAX 0.11).
x = jnp.array([1.0, 2.0, 3.0])gives an array of shape (3,) and dtype float32. (NumPy would have used float64: JAX uses 32-bit numbers unless you switch on 64-bit mode.)x[0] = 10.0raisesTypeError: JAX arrays are immutable and do not support in-place item assignment. Instead of x[idx] = y, use x = x.at[idx].set(y).y = x.at[0].set(10.0)returns[10. 2. 3.]. Printingxstill shows[1. 2. 3.].- Other updates work the same way:
x.at[1].add(5.0)gives[1. 7. 3.];x.at[2].multiply(2.0)gives[1. 2. 6.]. - A surprise:
x[10]does not raise an error. It returns3.0(the index is clamped to the last element), andx.at[10].set(99.0)silently returns an unchanged[1. 2. 3.].
A JAX array (type jax.Array, made with jnp.array, jnp.zeros, jnp.arange…) is an n-dimensional block of numbers with a fixed shape, a fixed dtype, and a home device. It is immutable: no operation changes it in place.
- Functional updates:
x.at[idx].set(v),.add(v),.multiply(v),.min(v),.max(v)each return a new array. "Functional" means "returns a new value instead of changing the old one". - Default precision: float32 (about 7 significant digits). For 64-bit numbers call
jax.config.update("jax_enable_x64", True)at the very start of the program. - Indexing out of range does not raise: reads are clamped to the nearest valid index, and out-of-range updates are dropped. Check your indices yourself.
- The API copies NumPy (
jnp.sum,jnp.exp, broadcasting, slicing), so most NumPy code needs onlynp→jnp, except for in-place assignment.
Why do we need it?
JAX's transformations (compile, differentiate, batch) need to know that nothing changes behind their back. Immutable arrays guarantee that a value, once made, means the same thing everywhere in the program, so JAX can safely reorder, fuse and differentiate the operations.
Where is it used?
Every NumPyro model and guide: the data arrays you pass in, the design matrix of the forecasting model (trend, Fourier, holiday and regressor columns), the guide's parameters, the optimizer state, and every posterior sample array that mcmc.get_samples() returns.
How is it used?
Build arrays with jnp, compute with whole-array operations, and replace every a[i] = v with a = a.at[i].set(v). Keep data preparation (cleaning, joining, selecting) in NumPy or pandas before handing arrays to the model.
y and the original x keeps its values, so any code still holding x sees exactly what it saw before.".at[i].set(v) copies the whole array every time, so JAX code must be slow."
Outside a compiled function it does make a copy. Inside a jitted function (Chapter 6.17) the compiler can usually reuse the memory when the old array is not needed any more, so the update can happen in place behind the scenes.
"Indexing past the end raises an IndexError, like NumPy."
JAX clamps reads (x[10] returned 3.0) and drops out-of-range updates without a word. An off-by-one bug produces wrong numbers, not a crash.
"float32 is always precise enough."
float32 keeps about 7 significant digits. That is usually fine for SVI and NUTS, but a finite-difference gradient check or a sum of millions of log-probabilities can lose digits. Turn on 64-bit mode when you need it, and expect results to differ slightly from float64 NumPy.
In an A/B framework like yours, the per-user (or per-group) conversion arrays, the group index array and the segment labels become JAX arrays once they reach the model. In your forecasting model, the time index, the changepoint features, the Fourier columns, the holiday indicators and the regressor matrix $X_t$ are JAX arrays (check how your code builds them). Anything that needs editing (filling gaps, selecting rows, building columns) is easiest to do in NumPy or pandas before the arrays enter the model; Chapter 6.17 shows why selecting rows inside traced code causes trouble.
JAX arrays are immutable: y = x.at[i].set(v) (also .add, .multiply, .min, .max) returns a new array; x is unchanged.
Default dtype float32 (≈ 7 digits); 64-bit needs jax_enable_x64.
Trap: out-of-range indices are clamped (reads) or dropped (writes), never an error.
Quick check: x = jnp.zeros(3); y = x.at[1].set(5.0); z = y.at[1].add(1.0). What are x, y and z?
x = [0, 0, 0] (never changed), y = [0, 5, 0], z = [0, 6, 0]. Each .at call builds a new array from the one it was called on.
Pure functions: why JAX insists on them core
A pure function is like a vending machine: the same coins and the same button always give the same snack, and the machine does nothing else in the shop. An impure function is like a cook who glances out of the window and changes the recipe when it rains, or who writes in a diary every time they cook.
Why does JAX care? Its transformations work by running your Python function once with stand-in values called tracers (placeholders that know only the shape and dtype, not the numbers), writing down every array operation, and then reusing that recording. During the recording, anything that is not an input (a global variable, a random number from NumPy) gets frozen into the recording as a constant, and any side effect (printing, appending to a list) happens only that one time.
Three ways to say it:
- Picture: JAX copies your recipe onto a card once, then cooks from the card; whatever was on the counter at copying time is now written on the card.
- Numbers: with a global
scale = 2the jitted function returned 10 for input 5; after settingscale = 100it still returned 10. - Slogan: the output depends only on the inputs, and nothing else goes in or out.
An impure function under jax.jit (jit compiles a function; the full story is Chapter 6.17). Real run:
- Define
scale = 2.0anddef f(x): print("python ran, x =", x); return x * scale; thenjf = jax.jit(f). jf(jnp.array(1.0))printspython ran, x = JitTracer(~float32[])and returns2.0. The print shows a tracer, not the number 1.0.jf(jnp.array(5.0))prints nothing and returns10.0: the recording was reused; the Python body did not run.- Set
scale = 100.0, calljf(jnp.array(5.0)): it returns10.0, not 500. The old value 2.0 was frozen into the recording. - Call with a new shape,
jf(jnp.array([1.0, 2.0])): a new recording is needed, the print appears again, and the result is[100. 200.], because this recording froze the new scale. Same code, different answers, depending on the call history. - NumPy randomness:
jax.jit(lambda x: x + np.random.normal())returned-1.043316on two calls in a row. The "random" number was drawn once, during recording, and then frozen.
A function is pure when (1) its output depends only on its explicit inputs, and (2) it has no side effects: it does not print, write files, change global variables, change its inputs, or draw from a hidden random-number state.
- A tracer is the placeholder JAX passes in while recording; it carries an abstract value: shape and dtype (for example
float32[2]), but no numbers. - JAX's transformations (
jit,grad,vmap, and the bodies oflax.scan,lax.cond…) assume purity. Impure code does not always crash: it often runs and quietly gives stale or frozen results. - Pure replacements: pass changing values as arguments; return new state instead of modifying it (next section); use
jax.debug.printto print real values inside compiled code; use PRNG keys for randomness (section 8).
Why do we need it?
Tracing records a function once and reuses the recording thousands of times. That is only correct if the function's result is fully decided by its inputs. Purity is the contract that makes compiling, differentiating and batching give the same answers as plain Python.
Where is it used?
NumPyro models and guides (traced inside svi.update, MCMC and Predictive), loss functions passed to jax.grad, step functions passed to lax.scan, and the update function of a custom SVI loop that you wrap in jax.jit.
How is it used?
Before you transform a function, scan it for globals it reads, objects it mutates, prints and NumPy random calls. Turn globals into arguments, mutation into "return the new value", prints into jax.debug.print, and random calls into jax.random with a key argument.
"I put a print inside my jitted function, so I can see the values."
A Python print runs only while tracing and shows a tracer like JitTracer(float32[2]). Use jax.debug.print("x = {}", x), which prints the real value on every call (verified: it printed 1.0, then 2.0).
"It gave the right answer the first time, so it is fine."
Frozen globals only show up later, when the global changes and the recording is reused. Impure code fails quietly, which is worse than a crash.
"np.random inside a model gives fresh noise each step."
Under a transformation it is drawn once and frozen. In NumPyro, all randomness goes through numpyro.sample (which uses keys); in plain JAX, through jax.random with a key.
When your custom loop calls a jitted svi.update, your NumPyro model is traced. We measured this: in 100 jitted updates of a small Beta-Binomial model, the Python body of the model ran once, and inside it theta was a tracer, not a number. So inside your models (A/B and forecasting alike), anything that should change between runs (prior scales, the likelihood choice, the data) must be an argument of the model, not a global you edit later, and any Python-side logging inside the model only fires at trace time. Record values with numpyro.deterministic instead.
"JAX cannot handle side effects."
JAX transformations assume pure functions. Side effects are not forbidden, they just run at trace time (once per compilation) instead of on every call, and hidden inputs get frozen. That is why the result can be silently wrong rather than an error.
Model answer: "jit, grad and vmap trace the function with abstract values and reuse the recorded program. A pure function gives the same output for the same inputs, so the recording is valid forever. Globals are baked in as constants and prints run once, so I pass everything as arguments and use jax.debug.print and explicit PRNG keys."
Pure = output depends only on inputs + no side effects. JAX traces once (with tracers: shape + dtype, no numbers) and replays.
Globals → frozen constants; print → runs at trace time only; np.random → frozen number.
Fixes: pass as arguments, return new state, jax.debug.print, jax.random keys.
Quick check: a jitted function reads a global list weights that you append to between calls. What happens?
Nothing visible at first: the first trace converted the list's contents into constants, so later calls with the same input shapes keep using the old weights. Only a re-trace (new shape, cleared cache) picks up the change. Fix: pass the weights in as an array argument.
Functional style: state goes in, new state comes out
If nothing can be changed in place, how does anything evolve, like parameters during training? Answer: you pass the current state into a function, and it hands back a new state. You keep the new one and carry on, like a relay runner passing a baton forward. This style, "functions plus values that never change", is called functional programming.
A bonus falls out for free: every old state is still there, untouched. Want to keep the best state seen so far? Just keep a reference to it; nothing will ever overwrite it.
Three ways to say it:
- Picture: a chain of snapshots; training moves along the chain, and you can hold on to any snapshot.
- Numbers:
state₀ → update → state₁ → update → state₂ …; the best snapshot at step 37 is still exactly the step-37 values when you reach step 200. - Slogan: never modify, always return.
The NumPyro SVI loop in functional style.
svi_state = svi.init(key, data)builds the first state. It is a small named tupleSVIState(optim_state, mutable_state, rng_key): the optimizer's numbers, any mutable model state, and a PRNG key.svi_state, loss = svi.update(svi_state, data)returns a new state and the loss (the negative ELBO estimate). The old state object is unchanged.params = svi.get_params(svi_state)gives a dictionary such as{"theta_auto_loc": …, "theta_auto_scale": …}. A nested container of arrays like this is called a pytree.jax.tree.map(lambda p: p * 0.5, params)applies a function to every array ("leaf") in the pytree and returns a new pytree with the same structure.- Checkpointing:
if loss < best_loss: best_loss, best_state = loss, svi_state. No copy is needed:best_statepoints at a snapshot that can never change.
Functional style: all changing quantities are explicit values that flow through functions, $\text{state}_{t+1} = \text{update}(\text{state}_t, \text{data})$. Nothing is modified in place.
- Pytree: any nesting of tuples, lists, dicts (and named tuples) whose leaves are arrays. Parameters, optimizer states and posterior-sample dictionaries are pytrees. JAX transformations accept and return pytrees.
- Tools:
jax.tree.map(f, tree)(apply to every leaf),jax.tree.leaves(tree)(list the leaves),jax.tree.map(f, t1, t2)(combine two trees with the same structure). - Because arrays are immutable, "keep the best state" is just "keep a reference". In libraries where tensors are updated in place (for example PyTorch's optimizer steps), the same line would keep a reference to values that keep changing, so you must copy explicitly.
Why do we need it?
Pure functions cannot change anything, so the only way to make progress is to return the new state. Explicit state also makes every run reproducible and every intermediate state inspectable and savable.
Where is it used?
NumPyro's SVI.init / SVI.update, optax and numpyro.optim optimizers, MCMC sampler states (HMCState), the carry of lax.scan, and best-state checkpointing in a custom training loop.
How is it used?
Write loops as state = update(state, batch), keep any snapshot you care about in a variable, and use jax.tree.map to transform all parameters at once (for example to compute parameter changes between two states).
"Returning a new state every step must waste a lot of memory."
Old states are freed as soon as nothing refers to them, and inside compiled code the compiler reuses buffers. You only pay for the snapshots you deliberately keep (such as best_state).
"best_state = svi_state makes a copy."
It stores a reference. That is safe in JAX only because the arrays inside can never change. Convert to NumPy (jax.device_get) when you want to save the checkpoint to disk.
Your forecasting model's custom SVI loop with best-state checkpointing (Chapter 6.14) is this pattern: if the current loss beats the best so far, store the current state (or svi.get_params(svi_state)) as the best, and at the end return the best rather than the last. In JAX that store is a single assignment. The early-stopping rule (relative ELBO improvement with patience) only needs the loss values, which svi.update returns next to the new state.
$\text{state}_{t+1} = \text{update}(\text{state}_t, \text{data})$: svi_state, loss = svi.update(svi_state, data).
Pytree = nested tuples/lists/dicts of arrays; jax.tree.map applies a function to every leaf.
Best-state checkpoint = keep a reference (safe because immutable). Trap: in-place libraries need a deep copy.
Quick check: why does svi.update return the state instead of updating svi internally?
Because the update function must be pure to be jitted and differentiated. All information that changes (optimizer numbers, the PRNG key) therefore travels in and out as an explicit value. A side benefit: you can keep, compare or restore any state.
Transformations: functions in, functions out
JAX's famous tools do not compute numbers directly. Each one takes a function and gives you back a new function. Think of kitchen gadgets that take a recipe and return a changed recipe: jit = "the same recipe, but done by a fast machine"; grad = "a recipe for how the result changes when you nudge the input"; vmap = "the same recipe for a whole tray of inputs at once".
Because they return ordinary functions, you can stack them: jax.jit(jax.vmap(jax.grad(f))) is "a compiled function that gives the slope of f at many points at once".
Three ways to say it:
- Picture: machines that wrap other machines; each wrapper changes what goes in or what comes out.
- Numbers:
loglik(95)= −17.89;grad(loglik)(95)= 0.25;vmap(grad(loglik))([90, 95, 100])= [0.5, 0.25, 0]. - Slogan: transform the function first, then call it.
Stacking transformations on one log-likelihood. Data: five days of orders 96, 104, 110, 90, 100; model $y_i \sim N(\mu, 10^2)$ (the example of Chapter 6.1). loglik(mu) returns $\sum_i \log N(y_i\mid\mu, 10^2)$.
loglik(95.0): distances $1, 9, 15, -5, 5$; squares add to $1 + 81 + 225 + 25 + 25 = 357$; result $-357/200 - 5\log(10\sqrt{2\pi}) = -1.785 - 16.108 = -17.893$.g = jax.grad(loglik)is a function.g(95.0)$= \sum_i (y_i-\mu)/\sigma^2 = 25/100 = 0.25$.gv = jax.vmap(g)takes a vector of μ values:gv(jnp.array([90., 95., 100.]))$= [50/100, 25/100, 0/100] = [0.5, 0.25, 0]$ (JAX printed[0.49999997 0.24999999 -0.]).jax.jit(gv)gives the same numbers from a compiled program (first call slower, later calls faster: Chapter 6.17).- The other order fails:
jax.grad(jax.vmap(loglik))raisesTypeError: Gradient only defined for scalar-output functions. Output had shape: (2,), becausevmap(loglik)returns a vector.
A transformation is a function that takes a Python function made of JAX operations and returns a new function. The main ones:
| Transformation | Returns a function that… | Shape rule (for f: one μ → one number) |
|---|---|---|
jax.jit(f) | computes the same thing with a compiled program | same shapes as f |
jax.grad(f) | returns $\partial f/\partial(\text{first argument})$ | output shaped like the argument; f must return a scalar |
jax.vmap(f) | applies f to every slice along a new leading axis | input gains an axis of size B; output gains the same axis |
jax.value_and_grad(f) | returns (f, ∂f) together, sharing the work | (scalar, shaped like the argument) |
- They compose: the output of one is a valid input to another.
- They work by tracing (previous section), so they all need pure functions made of JAX operations.
lax.scan,lax.fori_loopandlax.condare not transformations; they are control-flow building blocks that can be traced, differentiated and batched (section 7 and Chapter 6.17).
Why do we need it?
Inference needs gradients (SVI, NUTS), many evaluations at once (posterior draws, chains) and speed (thousands of steps). Writing each of these by hand for every model would be slow and error-prone; transformations derive them from the one function you wrote.
Where is it used?
NumPyro builds the log density of your model and then applies value_and_grad (SVI and HMC), vmap (vectorized chains, Predictive(parallel=True)) and jit (MCMC compiles its sampler; you jit svi.update). Flax and optax training loops are built the same way.
How is it used?
Write the plainest version of the computation (one data set, one parameter value), check it on small numbers, then wrap it: grad for derivatives, vmap for batches, jit last, around the whole step you will call many times.
loglik maps one μ to one number; grad turns that into "one μ → one slope"; vmap into "B values of μ → B slopes"; jit compiles the whole thing."jax.grad(f) gives me the gradient."
It gives you a function that computes the gradient. You still have to call it: jax.grad(f)(x).
"The order of the wrappers does not matter."
vmap(grad(f)) works (a slope per μ); grad(vmap(f)) fails, because grad needs a scalar output. Put jit outermost, around the biggest piece of work you call repeatedly.
Transformations take a function and return a function: jit (compile), grad (derivative; scalar output), vmap (batch), value_and_grad.
They compose: jax.jit(jax.vmap(jax.grad(f))). Example: slopes [0.5, 0.25, 0] at μ = 90, 95, 100.
Trap: grad(vmap(f)) fails (vector output).
Quick check: f(w) returns a scalar loss for one parameter vector w of shape (4,). What shape does jax.vmap(jax.grad(f)) return for a batch W of shape (10, 4)?
(10, 4): grad(f) maps a (4,) vector to a (4,) gradient, and vmap adds the leading batch axis of size 10.
grad: exact derivatives by automatic differentiation core
A gradient says how fast a function's output changes when you nudge each input (the slope, in one dimension). There are three ways to get one. You could nudge and measure (finite differences: approximate). You could do the algebra by hand (symbolic: exact, but laborious and easy to get wrong). Or you could let the computer apply the chain rule to every small operation your code performs: automatic differentiation (exact up to rounding, and automatic). jax.grad is the third way. It was built step by step in Calculus Chapter 2.9.
Three ways to say it:
- Picture: the slope of the tangent line that just touches the log-likelihood curve.
- Numbers: at μ = 95 the slope of the five-day log-likelihood is exactly 25/100 = 0.25;
jax.gradreturns 0.24999999 (float32). - Slogan: write the function, get its gradient for free.
Two slopes of the five-day log-likelihood, checked three ways. Data 96, 104, 110, 90, 100; μ = 95; σ = 10 written as σ = exp(s) with s = log σ (positive parameters are usually optimized on the log scale).
- By hand, in μ: $\frac{\partial}{\partial\mu}\sum_i\big(-\frac{(y_i-\mu)^2}{2\sigma^2}\big) = \sum_i\frac{y_i-\mu}{\sigma^2} = \frac{1+9+15-5+5}{100} = \frac{25}{100} = 0.25$.
- By hand, in s = log σ: $\log N(y\mid\mu,\sigma^2) = -\frac{(y-\mu)^2}{2e^{2s}} - s - \tfrac12\log 2\pi$, so $\frac{\partial}{\partial s}\sum_i = \sum_i\frac{(y_i-\mu)^2}{\sigma^2} - n = \frac{357}{100} - 5 = -1.43$.
- Autodiff:
jax.grad(loglik, argnums=(0, 1))(95.0, jnp.log(10.0))returned(0.24999999, -1.4300001). - Finite differences $\frac{f(x+h)-f(x-h)}{2h}$ in float32 for μ: h = 0.1 gave 0.249996; h = 0.001 gave 0.248909; h = 0.0001 gave 0.247955 (the exact digits vary a little between runs and machines). Smaller h was worse: the two function values agree in almost all of their 7 digits, so their difference is mostly rounding noise.
- The same check in float64 (
jax_enable_x64): h = 0.0001 gave 0.25000000000829914. With 16 digits the rounding noise is tiny.
jax.grad(f, argnums=0) returns a function computing $\nabla f$ with respect to the chosen argument(s), by reverse-mode automatic differentiation (the chain rule applied backwards through the recorded operations).
fmust return a real scalar; inputs must be floating point (jax.grad(lambda x: x**2)(3)raises a TypeError about int32;(3.0)gives 6.0). The gradient has the same shape (the same pytree) as the argument.jax.value_and_grad(f)returns $(f(x), \nabla f(x))$ in one pass. Reverse mode costs a small constant multiple of one evaluation of $f$, however many parameters there are (cost of autodiff).- Central finite difference: $\frac{f(x+h)-f(x-h)}{2h} = f'(x) + \underbrace{\tfrac{h^2}{6}f'''(x)}_{\text{truncation}} + \underbrace{O(\varepsilon\,|f|/h)}_{\text{rounding}}$, where ε ≈ 1.2 × 10⁻⁷ (float32) or 2.2 × 10⁻¹⁶ (float64) is the machine epsilon. Truncation shrinks as h shrinks; rounding grows. For a parabola $f''' = 0$, so only rounding remains.
- Gradient checks:
from jax.test_util import check_grads, thencheck_grads(f, (x,), order=1), preferably in float64.
Why do we need it?
SVI climbs the ELBO and NUTS steers its trajectories with gradients. Models have dozens to thousands of parameters; deriving and coding every partial derivative by hand is impossible to keep correct, and finite differences are slow (one or two evaluations per parameter) and inexact.
Where is it used?
svi.update (NumPyro's optimizer calls jax.value_and_grad on the loss), the leapfrog steps of HMC and NUTS (gradient of the log joint density), MAP fitting with AutoDelta, Laplace approximations (Hessians via jax.hessian), and every neural-network trainer written in JAX.
How is it used?
Write the loss or log density as a pure function of a parameter pytree, call jax.value_and_grad(loss)(params), feed the gradient to an optimizer. When you write a custom log density, verify it once against central differences in float64 at a few random points.
"jax.grad works by finite differences, so it is approximate."
It is automatic differentiation: the chain rule applied to each recorded operation. It is exact up to floating-point rounding and needs no step size.
"For a gradient check, use the smallest h you can."
Rounding error grows like ε/h. In float32 a tiny h makes the check fail even for a correct gradient (h = 10⁻⁴ gave 0.248 instead of 0.25). Use float64 and h around 10⁻⁵ to 10⁻⁶ for central differences, or check_grads from jax.test_util.
"I can take grad of a function that returns a vector of per-day log-likelihoods."
grad needs a scalar: sum them first (that is what a log-likelihood is), or use jax.jacrev for the full Jacobian.
Every SVI step in your two projects is a grad call: NumPyro computes the Monte Carlo ELBO estimate and its gradient with respect to the guide's parameters (locations, scales, Cholesky or low-rank factors) using jax.value_and_grad, then hands the gradient to the optimizer. If you write a custom likelihood (for example your own Negative Binomial parameterization or a Student-t with a learned ν), a one-time float64 gradient check at a few parameter values catches sign and parameterization bugs before they silently slow down training.
"Autodiff is just symbolic differentiation done by the computer."
Symbolic differentiation manipulates formulas and can blow up in size; autodiff never builds a formula. It applies the chain rule numerically to the actual operations executed for the given input, giving the derivative's value at that point, exact up to rounding.
Model answer: "Finite differences approximate the slope with a secant and suffer from truncation and rounding error; symbolic differentiation produces an expression; reverse-mode autodiff, which jax.grad uses, propagates derivatives backwards through the recorded computation and costs a few times one function evaluation, regardless of the number of parameters."
jax.grad(f, argnums): reverse-mode autodiff; f returns a real scalar; float inputs; gradient has the argument's shape.
Five days, μ = 95, σ = 10: ∂/∂μ = Σ(yᵢ − μ)/σ² = 0.25; ∂/∂log σ = Σ(yᵢ − μ)²/σ² − n = −1.43.
Central difference error ≈ h²f‴/6 + O(ε|f|/h): check gradients in float64, not with the tiniest h.
Quick check: you get a gradient-check mismatch of 2 × 10⁻³ in float32 with h = 10⁻⁴. Is your gradient wrong?
Not necessarily. In float32, ε|f|/h ≈ 1.2 × 10⁻⁷ × 18 / 10⁻⁴ ≈ 2 × 10⁻², so an error of that size can be pure rounding. Redo the check in float64 (or with h ≈ 10⁻²) before suspecting the code.
vmap: write it for one, run it for many core
A posterior is not one parameter value; it is thousands of draws (samples). To forecast, you push every draw through the model. The natural way to write that is a for loop, but loops in Python are slow, and rewriting the model by hand to work on a whole matrix of draws is fiddly. vmap ("vectorizing map") does the rewriting for you: write the function for one draw, and vmap turns it into a function for a whole stack of draws, computed with whole-array operations.
Think of a cookie cutter: you could press it once per cookie, or press a sheet cutter with the same shape over the whole tray at once. Same cookies, one press.
Three ways to say it:
- Picture: one function, a stack of inputs, a stack of outputs, one call.
- Numbers: 3 draws of a trend (slope k, intercept m) over 5 days give an output of shape (3, 5). For 4 000 draws × 365 days,
jit(vmap(trend))took about 0.12 ms; a Python loop of 4 000 jitted calls took about 28 ms on one machine. - Slogan: write for one, vmap for many.
A trend forecast over posterior draws. def trend(params, t): k, m = params; return k * t + m handles one draw $(k, m)$ and a vector of days t = [0, 1, 2, 3, 4].
- One draw (k = 1, m = 10):
trend([1., 10.], t)= [10, 11, 12, 13, 14]. - Three draws in a (3, 2) array: [[1, 10], [2, 9], [0.5, 11]].
jax.vmap(trend, in_axes=(0, None))(draws, t):in_axes=(0, None)means "slice the first argument along axis 0, share the second argument unchanged".- Result, shape (3, 5): row 1 = [10, 11, 12, 13, 14]; row 2 = 9 + 2t = [9, 11, 13, 15, 17]; row 3 = 11 + 0.5t = [11, 11.5, 12, 12.5, 13]. It matches stacking three separate calls (checked:
jnp.allclose→ True). - Posterior predictive: add noise with one key per draw:
keys = jax.random.split(key, 3)andjax.vmap(predict_one)(draws, keys), shape (3, 5).
jax.vmap(f, in_axes=0, out_axes=0) returns a function that applies f to every slice of its inputs along the mapped axis and stacks the results along out_axes:
in_axesgives one entry per argument: an axis number to map over, orNoneto share the argument across the batch.- It is not a loop: JAX rewrites each operation into its batched version. The recorded program for the example is just broadcast, multiply, add on (3, 5) arrays (one
muland oneadd, no loop). - Memory: the whole batch is materialized at once. For very large batches, map in chunks (
lax.map, or NumPyro'sPredictivewith its default sequential mode).
Why do we need it?
Bayesian outputs are averages over many draws: forecasts, intervals, P(θ_B > θ_A). Python loops over draws are slow; hand-batched code is bug-prone. vmap gives batched speed while you keep writing and testing the simple one-draw version.
Where is it used?
NumPyro's Predictive (with parallel=True it vmaps the model over posterior draws), MCMC(chain_method="vectorized") (chains run side by side), per-example gradients, ensembles, and evaluating a log-likelihood at a grid of parameter values.
How is it used?
Write f for a single draw, test it, then call jax.vmap(f, in_axes=(0, None))(draws, shared). Check the output shape (batch axis first). If memory runs out, process the draws in chunks.
vmap builds one batched program that takes all S draws as a single (S, 2) array and returns the (S, T) matrix of forecasts."vmap runs my function in parallel threads, one per draw."
It rewrites the computation into whole-array operations on a batched array. The hardware may then parallelize those operations, but there is no per-draw thread and no Python loop.
"in_axes=0 for every argument."
Arguments shared by every draw (the time grid, the design matrix, the data) need None; otherwise JAX tries to slice them too and complains about mismatched sizes.
"vmap is always better than a loop."
It holds the whole batch in memory. 10 000 draws × 365 days × 100 regressor columns is 365 million numbers; chunk it.
In your forecasting model, the posterior predictive for T future days is an (S, T) array: one row per posterior (or guide) draw. NumPyro's Predictive builds it by splitting one key into S keys and mapping the model over the draws: with parallel=True that map is one vmap; the default (parallel=False) runs the draws one at a time with lax.map to keep memory flat (both read from the NumPyro 0.22 source). In the A/B framework, any per-draw decision quantity (lift, expected loss, P(θ_B − θ_A > δ)) is a vectorized computation over the draws, and running the same analysis across many segments is a natural vmap.
jax.vmap(f, in_axes=(0, None))(draws, t): map over axis 0 of draws, share t; output gets a leading batch axis.
3 draws × 5 days → (3, 5); it is batched array code, not a loop. Measured: 4 000 × 365 trend in ≈ 0.12 ms vs ≈ 28 ms for a loop of jitted calls (one machine).
Trap: shared arguments need None; huge batches need chunking.
Quick check: loglik(theta, y) returns a scalar for one parameter vector θ of shape (3,) and data y of shape (500,). How do you evaluate it for 1 000 posterior draws stored in an array of shape (1000, 3), and what shape comes out?
jax.vmap(loglik, in_axes=(0, None))(draws, y), output shape (1000,): one log-likelihood per draw. The data are shared, so their axis entry is None.
lax.scan: loops that carry a state from step to step
Some calculations are naturally step by step: today's value depends on yesterday's. An autoregressive process (AR, Chapter 7.3) says $y_t = \phi\,y_{t-1} + \varepsilon_t$; a trend whose slope changes at changepoints says "today's level = yesterday's level + today's slope". You cannot compute day 10 before day 9.
lax.scan is JAX's loop for this. Picture a conveyor belt with a backpack (the carry) handed from one step to the next. At each step the worker reads the backpack and today's input, writes one output for today, and hands an updated backpack forward. JAX records the worker's job once, however many steps there are, so a 1 000-step loop is still a small program.
Three ways to say it:
- Picture: a backpack passed down a line of workers, each leaving one output on the table.
- Numbers: φ = 0.5, shocks ε = [1, 0, 2, −1], start 0 → outputs [1, 0.5, 2.25, 0.125]; final carry 0.125.
- Slogan: state in, state out, one output per step.
An AR(1) recursion with scan. def step(carry, e): y = 0.5 * carry + e; return y, y returns (new carry, output). Then last, ys = lax.scan(step, 0.0, jnp.array([1., 0., 2., -1.])).
- Start: carry = 0.
- Step 1: $0.5 \times 0 + 1 = 1$. Output 1, carry 1.
- Step 2: $0.5 \times 1 + 0 = 0.5$. Output 0.5, carry 0.5.
- Step 3: $0.5 \times 0.5 + 2 = 2.25$. Output 2.25, carry 2.25.
- Step 4: $0.5 \times 2.25 - 1 = 0.125$. Output 0.125, carry 0.125. JAX printed
0.125 [1. 0.5 2.25 0.125].
A trend with slope changes. Carry = (slope k, level). Start k = 1, level = 10; slope changes δ = [0, 0, 0.5, 0, −1, 0]. Each step: k ← k + δ_t, level ← level + k. Slopes 1, 1, 1.5, 1.5, 0.5, 0.5; levels 11, 12, 13.5, 15, 15.5, 16, the same as the closed form 10 + cumsum(1 + cumsum(δ)). (For a purely linear recursion like this one, the vectorized cumsum is simpler; scan earns its keep when each step depends non-linearly on the previous one, as in AR or state-space models.)
Gradients flow through scan. For the AR loop, the loss $\sum_t y_t^2$ has $\partial/\partial\phi = 6.1875$ at φ = 0.5 (by hand: $2y_2\cdot1 + 2y_3\cdot 2\phi + 2y_4\,(3\phi^2+2) = 1 + 4.5 + 0.6875$); jax.grad through lax.scan returned 6.1875, finite differences 6.187.
lax.scan(f, init, xs) with f(carry, x) → (new_carry, y) computes, in Python terms,
carry = init; ys = []
for x in xs:
carry, y = f(carry, x)
ys.append(y)
return carry, stack(ys)
- The carry must keep the same shape and dtype (the same pytree structure) at every step;
xsis sliced along its first axis; the outputs are stacked along a new first axis. - The body
fis traced once; the jaxpr contains a singlescanwithlength=4, not four copies. - It is differentiable (forward and reverse mode) and can be jitted and vmapped. Relatives:
lax.fori_loop(carry only, no stacked outputs) andlax.while_loop(unknown number of steps; not reverse-mode differentiable). - The steps run in order; scan does not parallelize over time. To run many independent series,
vmapthe scan.
Why do we need it?
A Python for loop inside a traced function is unrolled: 1 000 steps become 1 000 copies of the body, and compile time explodes (Chapter 6.17 measures 188 ms for 1 000 unrolled steps against 12 ms for a compiled loop). scan keeps one copy of the body and still lets gradients flow.
Where is it used?
AR and ARMA likelihoods, local-level and state-space trends, Kalman filters, RNNs, NumPyro's svi.run (with progress_bar=False it runs all optimization steps inside one lax.scan), and NumPyro's numpyro.contrib.control_flow.scan for time-series models with latent states.
How is it used?
Write the one-step function (carry, x_t) → (new_carry, y_t), pick the initial carry, pass the per-step inputs as an array with time first, and call lax.scan. Check the first few steps by hand, as in the example.
f, applied step after step. The purple carry is the only thing that passes between steps; JAX compiles f once and loops over it."A Python for loop inside a jitted function is the same as scan."
It gives the same numbers, but it is unrolled at trace time: n steps become n copies of the body (a 5 000-step loop became 10 000 equations and took 1.8 s to compile on one machine). Use scan or fori_loop for long loops.
"The carry can grow, for example by appending to a list."
The carry must keep exactly the same shape and dtype. Things that accumulate go in the stacked outputs ys, or in a fixed-size array updated with .at[t].set.
"scan runs the time steps in parallel."
A recursion is sequential by nature. Parallelism comes from vmap-ing the scan over independent series (segments, posterior draws).
In your forecasting model, the piecewise-linear trend can be computed without a loop (slope = k + changepoint matrix × δ, as in Chapter 7.8), so scan is not needed there. It becomes the right tool if you extend the model with something truly recursive: an AR term on the residuals, a local-level (random-walk) trend or a state-space component, the fixes for autocorrelated residuals listed in Chapter 7.17. NumPyro's own svi.run(..., progress_bar=False) already runs its whole optimization loop as one lax.scan; a custom loop with early stopping cannot, because the stopping decision is made in Python between steps (Chapter 6.14).
carry, ys = lax.scan(f, init, xs) with f(carry, x) → (new_carry, y); body traced once; carry shape fixed.
AR(1), φ = 0.5, ε = [1, 0, 2, −1] → ys = [1, 0.5, 2.25, 0.125]; differentiable (∂/∂φ Σy² = 6.1875).
Trap: Python loops unroll (slow compile); scan is sequential, vmap it across series.
Quick check: you want the running maximum of a series, m_t = max(m_{t-1}, x_t). What are the carry, the input and the output of the scan step?
Carry = the running maximum so far (start at −∞ or at x[0]); input = x_t; step: m = jnp.maximum(carry, x_t); return m, m. The stacked outputs are the running maxima. (For this special case lax.cummax exists too.)
PRNG keys: randomness you can replay core
NumPy's random numbers come from a hidden global dice-roller: every call quietly changes a hidden state, so the next call gives new numbers. JAX cannot have hidden state (functions must be pure). Instead, you hold the randomness in your hand: a small array called a key. A PRNG ("pseudo-random number generator") turns a key into numbers that look random but are completely determined by the key.
Same key in, same numbers out, every time. When you need new numbers, you split the key into fresh child keys and use each child once. Think of a family tree of raffle tickets: each ticket can be scratched once; to get more tickets, you split one into two new ones.
Three ways to say it:
- Picture: a tree of keys, each leaf used exactly once.
- Numbers:
jax.random.normal(PRNGKey(0), (3,))= [1.6226, 2.0253, −0.4336] every single time; its two children give [1.0040, −0.9063, −0.7482] and [−2.4425, −2.0357, 0.2055]. - Slogan: split, don't reuse.
Keys in the JAX console and in NumPyro (real outputs).
key = jax.random.PRNGKey(0)is the array[0 0](two uint32 numbers).jax.random.normal(key, (3,))→[1.6226422, 2.0252647, -0.43359444]. Call it again with the same key: identical.k1, k2 = jax.random.split(key);normal(k1, (3,))→[1.0040143, -0.9063372, -0.7481722];normal(k2, (3,))→[-2.4424558, -2.0356805, 0.20554423].- The reuse trap:
a = normal(key, (5,))andb = normal(key, (5,))are identical; their correlation is exactly 1.0, although you meant them to be independent. - NumPyro:
mcmc.run(jax.random.PRNGKey(0), …)twice on the Beta-Binomial model gave the same posterior mean 0.12404 both times;PRNGKey(1)gave 0.12760 (the exact posterior mean is 0.125; the gaps are Monte Carlo error).
A PRNG key is a small array that fully determines the output of a JAX random function: jax.random.normal(key, shape), uniform, bernoulli, categorical, gamma….
- Create:
jax.random.PRNGKey(seed)(classic, a uint32[2] array, used in most NumPyro examples) orjax.random.key(seed)(newer "typed" key; gives the same numbers here). Both work with NumPyro 0.22. jax.random.split(key, n)returns n new, independent-looking keys;jax.random.fold_in(key, i)derives a key from a key and an integer (handy inside loops: one key per step i).- Rules: (1) same key + same call → same numbers; (2) never use one key for two things that should be independent; (3) once you split a key, treat the parent as used up; (4) pass keys as function arguments.
- JAX's default generator (threefry) is counter-based: numbers are computed from the key, not from a running state. That is what makes it pure, splittable and safe to use inside
jit,vmapand parallel chains.
Why do we need it?
A global seed breaks under JAX's rules: inside a traced function it would be frozen (NumPy's random number above), and with chains or draws run side by side, a shared stream would make each chain's numbers depend on the execution order. Explicit keys make every random number a pure function of a key you can record.
Where is it used?
mcmc.run(rng_key, …), svi.init(rng_key, …) (the SVI state then carries a key and splits it every step), Predictive(…)(rng_key, …) (split into one key per draw), numpyro.handlers.seed(model, rng_seed=0), data simulation, dropout and initialization in Flax.
How is it used?
Make one root key from a recorded seed, split it into named purposes (init, training, prediction, simulation), split further inside loops or vmap, and never reuse a key. Save the seed with the run so results can be reproduced.
"Set one seed at the top of the script, like np.random.seed(0), and you are reproducible."
JAX has no global seed. Make a root key from a recorded seed and pass (split) keys explicitly to everything random: mcmc.run, svi.init, Predictive, your simulations.
"Reusing a key is harmless; the numbers are random anyway."
The same key gives the same numbers. Reusing it creates fake perfect correlations: identical noise for two variants, identical initial points for "different" chains.
"Splitting keys makes the results unpredictable."
split is deterministic: the same parent always gives the same children. The whole tree is reproducible from the root seed.
Both projects need reproducible inference for interviews, audits and debugging. In the A/B framework, a simulated A/A or power study should split one key per simulated experiment (and per variant inside it), or the "independent" experiments are copies of each other. In your forecasting loop, use one key for svi.init, let the SVI state carry and split its own key during training, and use a separate key for Predictive: then changing the number of forecast draws cannot change the fitted parameters. Record the root seed alongside the guide type, optimizer and iteration count (the reproducibility checklist of Chapter 7.19).
"JAX uses keys because it is a quirky library."
Keys follow from purity: a random function must get its randomness as an input, otherwise tracing would freeze it and parallel code would depend on execution order.
Model answer: "JAX functions are pure, so there is no hidden global random state. A PRNG key is an explicit input; the same key always gives the same numbers, and split creates new independent keys. That makes results reproducible under jit, vmap and parallel chains. The main rule is never to reuse a key for two things that should be independent."
key = jax.random.PRNGKey(seed); k1, k2 = jax.random.split(key); jax.random.normal(k1, shape).
Same key → same numbers (PRNGKey(0) → [1.6226, 2.0253, −0.4336]). Split, never reuse; a split key is used up.
NumPyro: keys for mcmc.run, svi.init, Predictive; separate branches per purpose; record the seed.
Quick check: inside a loop you write noise = jax.random.normal(key, (T,)) for each of 100 simulated datasets, using the same key. What do you get, and how do you fix it?
100 identical noise vectors, so 100 identical datasets. Fix: keys = jax.random.split(key, 100) and use keys[i] for dataset i (or jax.random.fold_in(key, i)), or vmap the simulation over the split keys.
Putting it together: a NumPyro model is a traced JAX function
Everything in this chapter meets inside NumPyro. Your model is an ordinary Python function made of numpyro.sample statements. NumPyro's effect handlers (seed, trace, substitute, condition, mask) are themselves function transformations: they wrap your model and intercept each sample statement to supply a key, record a value, plug in a parameter, or switch off a likelihood term. From the wrapped model NumPyro builds one pure function, "parameters → log density", and then uses grad, vmap and jit on it.
So your model's Python body is mostly run as a recording, not as a calculation. Once the recording exists, the compiled program does all the work.
Three ways to say it:
- Picture: your model is a recipe card; NumPyro and JAX read it once and build a machine from it.
- Numbers:
svi.initran the model body 3 times; the following 100 jittedsvi.updatecalls ran it once in total, with a tracer instead of a number. - Slogan: your model is a function that JAX transforms.
Counting how often the model's Python body runs (a counter appended to a list inside the model; real run, Beta-Binomial model with AutoNormal guide and Adam).
svi.init(PRNGKey(0), 20, 3): the body ran 3 times (types seen: a concrete array, a tracer, a concrete array). NumPyro runs the model to discover its sample sites and initial values.upd = jax.jit(svi.update, static_argnums=1), then 100 callssvi_state, loss = upd(svi_state, 20, 3): the body ran once in total (a tracer), during the first call's tracing. The other 99 steps ran only the compiled program. Final loss 1.85.- One extra update without jit: the body ran once more. Un-jitted, every step re-runs the Python body (and is much slower: Chapter 6.17).
handlers.trace(handlers.seed(model, rng_seed=0)).get_trace(20, 3)records the sites:theta(a sample, value 0.2575 for this key) andk(observed, value 3). Callingmodel(20, 3)with no seed handler fails: there is no key to samplethetawith.
Where each JAX idea lives in NumPyro:
| JAX idea | In NumPyro |
|---|---|
| Immutable arrays | data, parameters, posterior samples (mcmc.get_samples() returns a dict of arrays) |
| Pure functions | model and guide must be pure apart from numpyro.sample/param/deterministic statements, which the handlers manage |
| Functional state | SVIState(optim_state, mutable_state, rng_key); svi.update returns a new state; MCMC sampler states |
grad | SVI: value_and_grad of the ELBO loss with respect to guide parameters; HMC/NUTS: gradient of the potential energy (minus the log joint density) in each leapfrog step |
vmap | Predictive(parallel=True), MCMC(chain_method="vectorized"); numpyro.plate vectorizes by broadcasting, the same "whole array at once" idea |
scan | svi.run(…, progress_bar=False) runs its steps in a lax.scan; numpyro.contrib.control_flow.scan for recursive models |
| PRNG keys | handlers.seed; mcmc.run(key), svi.init(key), Predictive(…)(key) |
jit | MCMC compiles its sampler; you jit svi.update in a custom loop (Chapter 6.17) |
Why do we need it?
Knowing that the model is traced explains most NumPyro surprises: prints that appear once, Python if statements on data that fail, slow first steps, recompilation when data sizes change, and why randomness must go through numpyro.sample.
Where is it used?
Every NumPyro inference call: SVI with any autoguide, MCMC with NUTS or HMC, Predictive for prior and posterior predictive checks, and log_density utilities for model comparison and debugging.
How is it used?
Write the model with all data and settings as arguments, keep data-dependent logic in array operations, use numpyro.deterministic to record derived quantities, debug with handlers.trace(handlers.seed(model, 0)).get_trace(...) and jax.debug.print.
jit compiles the whole update once; the training loop then reuses that compiled update."My model function runs once per SVI step, so I can put Python logic on the data in it."
Under jit it runs once per compilation, with tracers. A Python if on a data value, a NumPy call on a sampled value, or a boolean mask that changes an array's length will fail or freeze (Chapter 6.17).
"numpyro.sample is like calling a random number generator."
It is a named site that the handlers interpret: under seed it draws with a key, under substitute it returns the guide's value, with obs= it scores the data. The same model line means different things to different inference algorithms.
Both of your projects are NumPyro models driven by SVI, so both are traced JAX functions. Practical consequences: pass the data arrays, the scaler's statistics, prior scales and switches (likelihood family, guide type) as arguments; keep the shapes fixed across the steps of a training run; use numpyro.deterministic to expose quantities such as the trend component or the lift $\theta_B - \theta_A$; and expect the first SVI step to be slow (compilation) and every later step fast. The boolean-masking problem you met in the A/B framework is the clearest example of "Python logic on traced values"; Chapter 6.17 takes it apart.
A NumPyro model is a Python function that NumPyro's handlers (seed, trace, substitute, condition, mask) turn into a pure log-density / loss function; JAX then applies grad, vmap, jit.
Measured: init ran the body 3 times; 100 jitted updates ran it once (a tracer).
Trap: Python logic on data inside the model runs at trace time only; use arrays, jax.debug.print, numpyro.deterministic.
Quick check: you add print(mu.shape, mu) inside your forecasting model and train with a jitted update for 2 000 steps. How many lines do you expect, and what do they show?
A handful from svi.init (some with concrete arrays) and then one from the first jitted update, showing a tracer with the shape (for example float32[T]) but no numbers. Nothing for the remaining steps. Use jax.debug.print to see values every step.
Recap, cheat sheet and practice
- JAX arrays are immutable, typed (float32 by default) and live on a device. Update with
x.at[i].set(v), which returns a new array. Out-of-range indices are clamped or dropped, not errors. - Pure functions: output depends only on inputs; no side effects. Transformations trace a function once with tracers (shape + dtype, no numbers) and reuse the recording, so globals freeze, prints run once and NumPy randomness is frozen.
- Functional style:
state = update(state, data). Old states never change, so keeping the best state is just keeping a reference. Pytrees (nested dicts/tuples of arrays) hold parameters and states. - Transformations take a function and return a function, and they compose:
jit(vmap(grad(f))). gradis reverse-mode autodiff: exact up to rounding, scalar output, float inputs. Check custom gradients against central differences in float64; small h in float32 drowns in rounding error.vmapturns a one-draw function into a batched one: (S draws) → (S, T) forecasts in one call, within_axes=Nonefor shared arguments.lax.scanruns a recursion with a fixed-shape carry, traced once, differentiable; Python loops unroll.- PRNG keys: same key → same numbers; split, never reuse; separate branches for init, training and prediction make NumPyro runs reproducible.
- A NumPyro model is a traced function: handlers turn it into a pure loss, JAX differentiates and compiles it, and the Python body runs at trace time.
Cheat sheet
| Idea | Code | Remember |
|---|---|---|
| Array update | y = x.at[i].set(v) (.add, .multiply) | x unchanged; float32 default; out-of-range silently clamped/dropped |
| Purity | pass everything as arguments | globals frozen at trace; print once; use jax.debug.print |
| Functional state | svi_state, loss = svi.update(svi_state, data) | best checkpoint = keep a reference |
| Pytrees | jax.tree.map(f, params) | dicts/tuples of arrays; transformations accept them |
| Gradient | jax.grad(f, argnums), jax.value_and_grad(f) | scalar output; float inputs; Σ(yᵢ − μ)/σ² = 0.25 in the example |
| Gradient check | $[f(x+h)-f(x-h)]/2h$ | error ≈ h²f‴/6 + ε|f|/h; use float64, h ≈ 10⁻⁵ |
| Batching | jax.vmap(f, in_axes=(0, None)) | write for one, run for many; output gets a leading batch axis |
| Recursion | carry, ys = lax.scan(step, init, xs) | step(carry, x) → (carry, y); carry shape fixed |
| Randomness | k1, k2 = jax.random.split(key) | same key → same numbers; never reuse; record the seed |
| NumPyro | handlers.seed, mcmc.run(key), Predictive(...)(key) | model body runs at trace time; data and settings as arguments |
import jax
import jax.numpy as jnp
from jax import lax
from jax.scipy.stats import norm
import numpyro
import numpyro.distributions as dist
from numpyro.infer import MCMC, NUTS, Predictive
# 1) Arrays are immutable: update with .at[], which returns a NEW array
x = jnp.array([1.0, 2.0, 3.0])
y = x.at[0].set(10.0)
print(x, y, x.dtype) # [1. 2. 3.] [10. 2. 3.] float32
print(x[10], x.at[10].set(99.0)) # 3.0 [1. 2. 3.] (out of range: clamped / dropped, no error!)
# 2) Impure functions break under jit: the global is baked in at trace time
scale = 2.0
def f(x):
return x * scale
jf = jax.jit(f)
print(jf(5.0)) # 10.0
scale = 100.0
print(jf(5.0)) # 10.0 (stale! the traced program still uses 2.0)
# 3) grad: exact derivatives of a log-likelihood, checked by finite differences
data = jnp.array([96.0, 104.0, 110.0, 90.0, 100.0])
def loglik(mu, log_sigma):
return jnp.sum(norm.logpdf(data, mu, jnp.exp(log_sigma)))
g_mu, g_ls = jax.grad(loglik, argnums=(0, 1))(95.0, jnp.log(10.0))
print(g_mu, g_ls) # 0.24999999 -1.4300001 (by hand: 25/100 and 357/100 - 5)
h = 1e-2
fd = (loglik(95.0, jnp.log(10.0) + h) - loglik(95.0, jnp.log(10.0) - h)) / (2 * h)
print(fd) # -1.4298439 (truncation + float32 rounding error)
# 4) vmap: write the trend for ONE posterior draw, apply it to all draws at once
t = jnp.arange(5.0)
def trend(params, t):
k, m = params
return k * t + m
draws = jnp.array([[1.0, 10.0], [2.0, 9.0], [0.5, 11.0]])
print(jax.vmap(trend, in_axes=(0, None))(draws, t)) # shape (3, 5): one row per draw
print(jax.vmap(jax.grad(loglik))(jnp.array([90.0, 95.0, 100.0]), jnp.full(3, jnp.log(10.0))))
# [ 0.49999997 0.24999999 -0. ]
# 5) scan: an AR(1) recursion y_t = 0.5 * y_(t-1) + eps_t with a carried state
def ar_step(carry, e):
y_new = 0.5 * carry + e
return y_new, y_new # (new carry, output of this step)
last, ys = lax.scan(ar_step, 0.0, jnp.array([1.0, 0.0, 2.0, -1.0]))
print(last, ys) # 0.125 [1. 0.5 2.25 0.125]
# 6) PRNG keys: same key -> same numbers; split for new numbers
key = jax.random.PRNGKey(0)
print(jax.random.normal(key, (3,))) # [ 1.6226422 2.0252647 -0.43359444]
print(jax.random.normal(key, (3,))) # identical: the key fully determines the output
k1, k2 = jax.random.split(key)
print(jax.random.normal(k1, (3,))) # [ 1.0040143 -0.9063372 -0.7481722]
# 7) NumPyro: explicit keys make inference reproducible
def model(n, k=None):
theta = numpyro.sample("theta", dist.Beta(2.0, 18.0))
numpyro.sample("k", dist.Binomial(n, theta), obs=k)
means = []
for seed in [0, 0, 1]:
mcmc = MCMC(NUTS(model), num_warmup=300, num_samples=1000, progress_bar=False)
mcmc.run(jax.random.PRNGKey(seed), 20, 3)
means.append(round(float(mcmc.get_samples()["theta"].mean()), 4))
print(means) # [0.124, 0.124, 0.1276]: same key, same draws (exact mean 0.125)
pred = Predictive(model, posterior_samples=mcmc.get_samples())(jax.random.PRNGKey(2), 50)
print(pred["k"].shape) # (1000,): one simulated count per posterior draw (one key each)
1. x = jnp.array([1., 2., 3.]); y = x.at[1].set(0.). What is x now?
.at[1].set(0.) returns a new array (y = [1, 0, 3]); JAX arrays are immutable, so x keeps its values. Only the NumPy-style x[1] = 0. raises a TypeError.2. A jitted function reads a global prior_scale = 1.0. You set prior_scale = 5.0 and call it again with inputs of the same shape. What does it use?
jit the Python body ran once, during tracing, and the global's value became a constant in the program. Same shapes → the cached program is reused. Pass the scale as an argument instead.3. f(theta) returns the vector of 365 daily log-likelihoods. What does jax.grad(f)(theta) do?
grad is defined for scalar-output functions ("Gradient only defined for scalar-output functions"). The total log-likelihood is the sum; for per-day gradients use jax.jacrev or vmap(grad(...)) over days.4. draws has shape (1000, 2) and t has shape (30,). What is the shape of jax.vmap(trend, in_axes=(0, None))(draws, t) if trend(params, t) returns shape (30,)?
draws (1 000 draws), shares t, and stacks the (30,) outputs along a new leading axis: (1000, 30).5. What must the step function of lax.scan(step, init, xs) return?
step(carry, x) → (carry, y). scan stacks the y's; the carry must have a fixed shape and dtype so the body can be traced once.6. A simulation draws noise for variant A and variant B with jax.random.normal(key, (n,)) using the same key. What is the right fix?
Practice problems
A. Make this pure and jit-safe: counts = np.zeros(3); def add_day(day_counts): counts[:] += day_counts; return counts.sum().
- Problems: it reads and mutates a global array (in place), so under
jitthe global would be frozen and the mutation would not happen. - Pure version: pass the state in and return the new state:
def add_day(counts, day_counts): new = counts + day_counts; return new, new.sum(). - Use:
counts = jnp.zeros(3); thencounts, total = jax.jit(add_day)(counts, jnp.array([1., 0., 2.]))each day. The caller keeps the state; old states stay valid.
B. Data y = 2, 4, 9; model $y_i \sim N(\mu, \sigma^2)$ with σ = 2. Find ∂/∂μ and ∂/∂log σ of the log-likelihood at μ = 4, and say what a float32 central difference with h = 10⁻⁴ is likely to show.
- Residuals $y_i - \mu$: −2, 0, 5. ∂/∂μ $= \sum(y_i-\mu)/\sigma^2 = 3/4 = 0.75$.
- ∂/∂log σ $= \sum(y_i-\mu)^2/\sigma^2 - n = (4 + 0 + 25)/4 - 3 = 7.25 - 3 = 4.25$.
jax.gradreturned exactly(0.75, 4.25). - The log-likelihood is about −8.46, so the float32 rounding error of a central difference is roughly $1.2\times10^{-7} \times 8.5 / 10^{-4} \approx 10^{-2}$. Indeed h = 10⁻⁴ gave 0.7534, while h = 10⁻² gave 0.75002. The gradient is fine; the check is too fine-grained for float32.
C. Posterior draws of a regression: a, b, sigma each of shape (S,); covariate x of shape (T,). Write a vmapped posterior predictive and give its output shape.
def predict_one(a, b, s, key): mu = a + b * x; return mu + s * jax.random.normal(key, mu.shape) (with x passed in or closed over as a constant). Then keys = jax.random.split(key, S) and jax.vmap(predict_one)(a, b, sigma, keys) (all four mapped along axis 0). Output shape (S, T): one simulated future path per draw, each with its own noise key. Quantiles over axis 0 give the predictive band.
D. A trend starts with slope 2 and level 0; slope changes δ = [0, −1, 0, 1.5]. Write the scan step and compute the outputs.
def step(carry, d): k, level = carry; k = k + d; level = level + k; return (k, level), level; calllax.scan(step, (2.0, 0.0), deltas).- Slopes: 2, 1, 1, 2.5. Levels: 0 + 2 = 2; 2 + 1 = 3; 3 + 1 = 4; 4 + 2.5 = 6.5.
- Outputs [2, 3, 4, 6.5]; final carry (2.5, 6.5). JAX printed exactly these.
E. (Interview) "Why does NumPyro ask for a PRNG key in mcmc.run, and what would go wrong with a global seed?"
"JAX functions are pure, so randomness has to be an input: the key. The sampler splits it into keys for each chain and step, so the whole run is a deterministic function of the key: the same key reproduced the same posterior mean (0.12404) exactly. A global seed doesn't fit: inside a compiled function the hidden state would be frozen at trace time, and with chains run side by side (vmap or parallel devices) a shared stream would make each chain's numbers depend on execution order. Separate keys per chain, per purpose (init, training, prediction) give reproducible, independent randomness."
F. (Interview) "Walk me through what happens the first time and the hundredth time you call a jitted svi.update."
"First call: JAX traces svi.update with abstract values. Inside, NumPyro runs the guide and the model under its handlers (so my model's Python body runs once, with tracers), builds the Monte Carlo ELBO loss, applies value_and_grad, and records the Adam update; the recording is compiled by XLA for these exact shapes and dtypes. That takes a fraction of a second. Hundredth call: same shapes, so JAX looks up the cached executable and runs it directly; no Python model code runs; a step takes on the order of a tenth of a millisecond for a small model. If I changed the data's shape, it would trace and compile again (Chapter 6.17)."
JIT compilation and model size
Your custom SVI loop calls a JIT-compiled update thousands of times. The first call is slow, every later call is fast, and some innocent-looking lines (a Python if, a boolean mask, a new data length) either crash or quietly trigger recompilation. This chapter opens the compiler's box: how Python becomes a compiled program, why shapes must be fixed, how to replace control flow and boolean masks, what moving data costs, and how the size of a model (its latent dimension and its guide) drives memory and compute.
- Trace the path Python function → tracing → jaxpr → XLA → compiled executable, and read a real
jax.make_jaxprprintout - Explain compile once, run fast, time JAX code honestly (
block_until_ready), and compute when compiling pays off - Predict recompilation (new shapes, dtypes, static arguments) and avoid it with padding and bucketing
- Replace Python control flow on traced values with
jnp.where,lax.cond,lax.fori_loopandlax.scan, including the NaN-gradient trap ofwhere - Explain and fix the boolean-masking problem:
x[mask]has a data-dependent shape; use masked sums,jnp.whereor NumPyro'smask - Account for compile overhead and device transfer in a training loop
- Count a model's latent dimension and guide parameters, and reason with O(d), O(d²), O(d³) and O(dr) costs, backed by real measurements
What we need from earlier chapters: pure functions, tracers, grad, vmap, scan and PRNG keys (Chapter 6.16); the three Gaussian guides and their parameter counts (Chapter 6.13); the custom SVI loop with early stopping and checkpointing (Chapter 6.14); computational complexity and big-O (Linear Algebra, computational complexity); the Cholesky factorization (Linear Algebra, Cholesky). Words used below. Compile: translate a program into fast machine instructions before running it. Host: the Python process on the CPU that runs your script. Device: the hardware that runs the compiled array programs (here also the CPU; often a GPU). Latency: the time one call takes. All timings below were measured on one laptop CPU (Apple M5, JAX 0.11, NumPyro 0.22). Timings vary by machine and even between runs; where they did, we give the range we saw. Real JAX cannot run in this page, so the interactives simulate JAX's rules using these measured numbers.
From Python to a compiled program: tracing, jaxpr, XLA core
Plain Python runs your code line by line, and every jnp operation is sent to the device separately. jax.jit ("just-in-time compilation") does something different. The first time you call the function, JAX runs your Python once with tracers (stand-ins that know only shape and dtype, Chapter 6.16) and writes down every array operation. That written-down program is called a jaxpr. It is then translated for XLA (Google's "accelerated linear algebra" compiler), which optimizes it for your CPU or GPU, for example by fusing several small operations into one loop over memory. The result is a compiled executable, stored in a cache. Later calls with the same input shapes skip all of that and run the executable directly.
It is like a court stenographer and a factory. The stenographer records exactly what you did, in a simple shorthand (the jaxpr); the factory builds a machine that repeats that sequence very fast (the executable). The machine is built for one size of input.
Three ways to say it:
- Picture: record once, build a machine, reuse the machine.
- Numbers: the function
0.5 * sum((y − μ)²)on 5 numbers is recorded as 4 operations (subtract, square, sum, multiply), each tagged with its shape, such asf32[5]. - Slogan: trace once, compile once, run many times.
Reading a real jaxpr. def nll(mu, y): return 0.5 * jnp.sum((y - mu) ** 2), called with μ = 95 (a float32 scalar) and the five days y = [96, 104, 110, 90, 100]. jax.make_jaxpr(nll)(jnp.float32(95.0), y) printed:
{ lambda ; a:f32[] b:f32[5]. let
c:f32[5] = sub b a
d:f32[5] = integer_pow[y=2] c
e:f32[] = reduce_sum[axes=(0,) out_sharding=None] d
f:f32[] = mul 0.5:f32[] e
in (f,) }
a:f32[]is μ: a float32 with shape[](a scalar).b:f32[5]is y: five float32 numbers. No values appear anywhere: only types and shapes.c = sub b a: the residuals y − μ (the scalar is broadcast), shape [5].d = integer_pow[y=2] c: square each residual.e = reduce_sum d: add them up, giving a scalar.f = mul 0.5 e: halve it. Outputf.- Running it: residuals 1, 9, 15, −5, 5; squares sum to 357; 0.5 × 357 = 178.5 (JAX printed 178.5).
- Called with y of length 1 000, the same four lines appear with
f32[1000]instead off32[5]: a different program, compiled separately.
The next stage is the lowered program that XLA receives (jax.jit(nll).lower(mu, y).as_text(), first lines):
func.func public @main(%arg0: tensor<f32>, %arg1: tensor<5xf32>) -> (tensor<f32>) {
%1 = stablehlo.broadcast_in_dim %0, dims = [] : (tensor<f32>) -> tensor<5xf32>
%2 = stablehlo.subtract %arg1, %1 : tensor<5xf32>
%3 = stablehlo.multiply %2, %2 : tensor<5xf32>
%4 = stablehlo.reduce(%3 init: %cst) applies stablehlo.add across dimensions = [0] ...
JIT compilation of a function f (jax.jit(f), or the @jax.jit decorator) happens on the first call for each new signature of the inputs:
- Tracing: run the Python body with tracers that carry an abstract value (shape, dtype, and whether the type is "weak", i.e. came from a Python number). Python-level code runs; JAX operations are recorded.
- jaxpr: the recorded program, a list of primitive operations (
sub,mul,dot_general,scan…) on typed variables. See it withjax.make_jaxpr(f)(*args). - Lowering: translation to StableHLO, the input language of XLA (
jax.jit(f).lower(*args)). - XLA compilation: optimization (operation fusion, memory planning) into an executable for the target device (
.compile()). - Caching and execution: the executable is stored under a key made of the function and the inputs' abstract signature (plus any static arguments). Later calls with a matching key run it directly.
Only array operations are recorded. Python arithmetic on concrete things (Python ints, NumPy constants, configuration flags) is done during tracing and its result is baked in as a constant.
Why do we need it?
Without compilation every small operation is dispatched from Python one at a time, with overhead each time. An SVI step is hundreds of small operations, so the overhead dominates. Compiling the whole step into one program removes that overhead and lets XLA fuse operations.
Where is it used?
The jitted svi.update in your custom loop, NumPyro's MCMC (which compiles its NUTS sampler), svi.run, Flax and optax training steps, and any function you decorate with @jax.jit. TensorFlow (XLA) and PyTorch 2 (torch.compile) use the same trace-then-compile idea.
How is it used?
Wrap the function you call many times (the whole update step, not tiny pieces). Inspect it with jax.make_jaxpr when something is surprising, and switch on jax.config.update("jax_log_compiles", True) to see a log line every time something is compiled.
"jit translates my Python into C."
It never compiles your Python. It records the array operations your Python performed on tracers and compiles that record. Python logic (loops, ifs on settings, helper functions) runs once during tracing and disappears from the compiled program.
"The jaxpr stores my data."
It stores operations and shapes, not values (unless a value was a constant at trace time, such as a global NumPy array, which becomes a baked-in constant).
When you write update = jax.jit(svi.update) in your custom loop, the "function" being traced is the whole SVI step: run the guide, run your model under NumPyro's handlers (your forecasting model's trend, Fourier, holiday and regressor terms, or your A/B framework's Beta-Binomial and hierarchical terms), compute the ELBO estimate, take its gradient, and apply the Adam update. All of that becomes one compiled program per input signature. The signature includes the shapes of the data arrays, so the history length, the number of regressors or the number of users fixes which program you get.
jit pipeline: Python → tracing (abstract values: shape, dtype) → jaxpr → StableHLO → XLA → executable, cached by (function, shapes, dtypes, static args).
jax.make_jaxpr(f)(*args) shows the recorded program: operations + shapes, no values.
Trap: Python code runs only during tracing; the compiled program contains only array operations.
Quick check: f(x) = jnp.sin(x) * 2.0 + jnp.sum(x) is jitted and called with an array of shape (3,). Roughly what does its jaxpr contain, and what changes if you call it with shape (4,)?
Something like sin (f32[3]), mul by 2.0 (f32[3]), reduce_sum (f32[]), and an add that broadcasts the scalar sum back to f32[3]. With shape (4,), the same operations appear with f32[4], and JAX traces and compiles a second program.
Compile once, run fast (and how to time it honestly) core
Compiling is a fixed setup cost, like setting up a production line before making the first item. Once the line exists, each item is cheap. So jit pays off when you will call the same function many times with the same input shapes, which is exactly what an SVI training loop does (thousands of steps), and it does not pay off for a function you call once.
A second surprise: JAX is asynchronous. When you call a jitted function, Python gets back a "future" array immediately and moves on while the device computes. If you time a call without waiting for the result, you only time the sending of the order, not the work.
Three ways to say it:
- Picture: a tall first bar (compile), then a long row of tiny bars (runs).
- Numbers (one laptop CPU): a jitted
svi.updatefor a small hierarchical model took 0.43–0.75 s on its first call and 0.1–0.2 ms per call after that; without jit, each step took 40–70 ms. - Slogan: pay once, reuse many times, as long as the shapes do not change.
When does compiling pay off? Use round versions of the measured numbers: compile cost $C = 0.5$ s, jitted step $r = 0.15$ ms, un-jitted ("eager") step $e = 50$ ms.
- Total time for $N$ steps with jit: $C + N r = 0.5 + 0.00015\,N$ seconds. Without jit: $N e = 0.05\,N$ seconds.
- Break-even: $C + N r = N e \Rightarrow N^* = C/(e - r) = 0.5/(0.05 - 0.00015) = 0.5/0.04985 \approx 10$ steps.
- For 5 000 steps: jit $0.5 + 5000 \times 0.00015 = 0.5 + 0.75 = 1.25$ s; eager $5000 \times 0.05 = 250$ s, about 4 minutes.
- Smaller example, the
nllabove on 100 000 numbers: first call 13–35 ms across runs (trace + compile), later calls 0.03–0.09 ms. - Honest timing:
t0 = time.perf_counter(); f(x).block_until_ready(); t1 = time.perf_counter(). Without.block_until_ready()the measured time can be far too small, because the work has not finished.
For a jitted function called $N$ times with one input signature:
$$T_{\text{total}} \approx T_{\text{trace}} + T_{\text{compile}} + N\,(T_{\text{dispatch}} + T_{\text{run}}),$$and every new signature adds another $T_{\text{trace}} + T_{\text{compile}}$.
- Asynchronous dispatch: a JAX call returns as soon as the work is queued. Anything that needs the actual value waits:
x.block_until_ready(),float(x),print(x),np.asarray(x), a Pythonifonx. - Warm-up: call once before timing, so the measurement excludes compilation. Report medians of many calls.
- Ahead-of-time:
jax.jit(f).lower(*args).compile()compiles without running. A persistent cache (jax.config.update("jax_compilation_cache_dir", path)) can reuse compiled programs across Python sessions.
Why do we need it?
To know when jit helps (many repeated calls with the same shapes) and when it hurts (one-off calls, constantly changing shapes), and to measure speed correctly instead of measuring how fast Python can queue work.
Where is it used?
SVI training loops, MCMC (NumPyro compiles the sampler once, so the first chain starts slowly), model serving where a "warm-up" request is sent at startup, and every benchmark of JAX code.
How is it used?
Jit the whole repeated step. Warm it up once. Time with block_until_ready and take medians. Expect the first step of every training run (and every run with new data shapes) to be slow, and do not mistake it for a performance bug.
"jit makes everything faster."
It removes per-operation overhead for functions called many times with the same shapes. For a one-off call, compile time dominates. For one huge operation (a single big matrix multiply), the operation already runs a compiled kernel, so jit adds little.
"My jitted function takes 0.01 ms. Amazing!"
Did you call .block_until_ready()? Without it you timed the dispatch, not the computation.
"The first call is slow, so my model is slow."
The first call includes tracing and compilation. Time the second and later calls (and report medians) to judge the step speed.
Your custom SVI loop uses JIT-compiled updates: expect the first update call of every run to take a noticeable fraction of a second (more for a larger model, since tracing and compiling scale with the program's size), and every later step to be fast. When you report "training takes X seconds", separate compile time from step time. If you refit the forecasting model every day with one more day of history, the data shape changes, so each refit compiles once; that is fine. What hurts is recompiling inside a run (next section).
$T \approx T_{\text{compile}} + N\,T_{\text{step}}$ per input signature; break-even $N^* = C/(e - r)$ (≈ 10 steps with C = 0.5 s, e = 50 ms, r = 0.15 ms).
Measured (one CPU): jitted svi.update 0.43–0.75 s first call, 0.1–0.2 ms later; eager 40–70 ms per step.
Trap: async dispatch: time with .block_until_ready(), after a warm-up call.
Quick check: compile 2 s, jitted step 1 ms, eager step 20 ms. Is jit worth it for a 50-step job? For a 5 000-step job?
Break-even $N^* = 2000/(20 - 1) \approx 105$ steps. 50 steps: jit 2 + 0.05 = 2.05 s vs eager 1 s, so eager wins. 5 000 steps: jit 2 + 5 = 7 s vs eager 100 s, so jit wins by about 14×.
Static shapes and recompilation; padding and bucketing core
A compiled program is built for exact shapes and dtypes. Every array size in it is a fixed number. Call it with a different shape and JAX has to trace and compile a new program. It is like a cookie-cutter machine built for 5 cm cookies: a 6 cm cookie needs a new machine.
If your inputs come in many sizes (experiments with different numbers of users, segments of different sizes), you can end up compiling over and over. The fix is to pad each input up to one of a few standard sizes (buckets, such as 128, 256, 512) and pass a mask that says which entries are real.
Three ways to say it:
- Picture: one machine per cookie size; padding makes most cookies the same size.
- Numbers: calls with n = 100, 100, 101, 102, 100, 103 compiled 4 programs (about 35 ms each); padding sizes 100, 101, 102, 130, 200, 260 into buckets 128/256/512 compiled only 3.
- Slogan: same shape, same program.
Counting compilations with a Python counter inside the jitted function (it only runs while tracing). Real run:
- n = 100: compile, 33 ms. n = 100 again: cache hit, 0.02 ms. n = 101: compile, 38 ms. n = 102: compile, 37 ms. n = 100: cache hit. n = 103: compile, 36 ms. Total 4 traces for 6 calls.
- Same shape (100,) but dtype int32 instead of float32: one more compile (5 in total). A NumPy float32 array of shape (100,): no new compile (same abstract signature). A new value of the scalar μ: no new compile (values are not part of the key).
- Padding: bucket(n) = the smallest of 64, 128, 256, 512, … that is ≥ n. Sizes 100, 101, 102 → 128; 130, 200 → 256; 260 → 512. Three compiles for six sizes, and the masked mean gave the correct value every time.
- The price of padding: n = 100 padded to 128 computes 28 extra entries, $28/128 \approx 22\%$ of the work is padding. With power-of-two buckets, padding is less than half of each padded array (as long as $n$ is above half of the smallest bucket).
- Static arguments: a Fourier-features function with
ordermarked static (static_argnames="order") compiled once for order 3 (shape (14, 6)), reused it for a different period (period is traced), and compiled again for order 10 (shape (14, 20)). Without marking it static,jnp.arange(order)failed withConcretizationTypeError: an array's length cannot be a traced value.
The jit cache key contains: the function, the pytree structure of the arguments, each array argument's shape and dtype (and weak-type flag), the values of static arguments, and the device. A call whose key is new is retraced and recompiled; values of ordinary (traced) arguments never trigger recompilation.
- Static arguments (
jax.jit(f, static_argnums=…)orstatic_argnames=…) are treated as Python constants during tracing: they may decide shapes and Python control flow, but each new value compiles a new program, and they must be hashable (passing a JAX array raisedValueError: Non-hashable static arguments are not supported). - Shapes must be known from shapes, never from traced values:
jnp.zeros(n),jnp.arange(n),x[:n]need a concreten. - Padding + masking: pad arrays to a bucket size $B \ge n$, carry a boolean mask $m$ (True for the $n$ real entries), and write every reduction with the mask, e.g. mean $= \sum_i m_i x_i / \sum_i m_i$.
Why do we need it?
Compiling takes milliseconds to seconds. A loop that recompiles at every call can be slower than not compiling at all, and the slowdown is easy to miss because the results are still correct.
Where is it used?
Batched model serving (bucketed batch sizes), NLP models (sequences padded to bucket lengths with attention masks), ragged groups in hierarchical models (padded to a rectangle with a mask), and configuration values such as a Fourier order or a number of changepoints passed as static arguments.
How is it used?
Turn on jax.config.update("jax_log_compiles", True) to see every compilation. Fix shapes for a whole training run. If inputs vary in size, pad to a few buckets and pass a mask. Make only small configuration values static, never data arrays.
"Calling with new values recompiles."
Values never do. New shapes, new dtypes, new static-argument values and new pytree structures do. Passing int32 data once and float32 later silently doubles the compiles.
"Mark everything static, to be safe."
Every new static value is a new compile, static values must be hashable, and arrays should never be static. Make only small configuration values static (an order, a flag, a likelihood name).
"Padding changes the answer."
Only if you forget the mask in a reduction: the mean of [1, 1, 0, 0] (two real, two padding) is 0.5, the masked mean is 1. Every sum, mean and log-likelihood must use the mask.
In an A/B framework like yours, experiments and segments arrive with different user counts. Jitting a function on the exact per-segment arrays compiles once per distinct size; with many segments that adds seconds of compile time and hides in "the framework is slow". Padding to a few buckets plus a mask keeps it to a handful of compiles. In your forecasting model, the history length T, the horizon, the number of changepoints and the Fourier orders set the shapes: they are fixed within a training run (one compile per run). Fourier order and changepoint count are natural static values, because they decide array sizes.
Cache key = function + shapes + dtypes + static args (+ pytree structure). New key → retrace + recompile (≈ 35 ms for a tiny function, ≈ 0.5 s for an SVI step on one CPU).
Fix varying sizes with padding to buckets (64, 128, 256, …) + a mask; wasted work < 50% per call (for sizes above half of the smallest bucket).
Traps: dtype changes recompile; static args must be hashable; shapes cannot depend on traced values.
Quick check: a jitted function is called with arrays of length 300, 310, 900, 1 000 and 1 020. How many compiles with exact shapes, and with buckets 64, 128, 256, …?
Exact: 5 compiles (five different lengths). Buckets: 300 and 310 → 512; 900, 1 000 and 1 020 → 1024. So 2 compiles. Wasted work for 300 → 512: 212/512 ≈ 41%.
Python control flow under tracing: where, cond, fori_loop, scan
During tracing, data values are unknown: a tracer knows its shape, not its numbers. So a Python if x > 0: cannot decide which way to go, and JAX raises an error. Python for loops do run during tracing, but they are unrolled: a loop of 1 000 steps is recorded as 1 000 copies of its body, and compiling that long program is slow.
JAX gives array-friendly replacements. jnp.where: compute both branches for every element and pick per element. lax.cond: put a switch into the compiled program and choose one branch when it runs. lax.fori_loop, lax.scan, lax.while_loop: put a real loop into the compiled program, with the body recorded once.
Three ways to say it:
- Picture: at a fork whose sign is still blank, either walk both roads and keep the right result (
where) or build a railway switch that flips later (cond). - Numbers: a 5 000-step Python loop became 10 000 recorded equations and took 1.8 s to compile;
lax.fori_loopwas 1 equation and about 16 ms. - Slogan: branches on data →
where/cond; long loops →fori_loop/scan; branches on settings → plain Python.
A Huber-style loss ($\tfrac12 r^2$ for small residuals, $|r| - \tfrac12$ for large ones; Huber's robust loss with threshold 1). Real runs:
@jax.jit def f(r): if r > 1.0: return r - 0.5 …→TracerBoolConversionError: Attempted boolean conversion of traced array with shape bool[].jnp.where(jnp.abs(r) > 1.0, jnp.abs(r) - 0.5, 0.5 * r**2)on r = [−3, 0.5, 2]: for −3, $|{-3}| - 0.5 = 2.5$; for 0.5, $0.5 \times 0.25 = 0.125$; for 2, $2 - 0.5 = 1.5$. Output[2.5 0.125 1.5].lax.cond(jnp.abs(r) > 1.0, lambda r: jnp.abs(r) - 0.5, lambda r: 0.5 * r**2, r): r = 2.0 → 1.5; r = 0.5 → 0.125. (For a scalar condition.)- A Python
ifon a setting is fine:if likelihood == "normal": … else: …withlikelihoodmarked static compiled one program per choice (values −1.5 and −3.0 on the test data). - Loop unrolling, measured compile times: Python loop of 100 / 1 000 / 5 000 steps: 32 ms / 188 ms / 1 795 ms (200 / 2 000 / 10 000 equations);
lax.fori_loop: 13 / 12 / 16 ms (1 equation each). Same result, 2.0.
| Situation | Tool | What it does |
|---|---|---|
| Condition on a setting known before tracing | plain Python if (mark the setting static) | chooses at trace time; one program per setting |
| Elementwise choice on data | jnp.where(cond, a, b) | computes both a and b everywhere, selects per element |
| One scalar choice between expensive branches | lax.cond(pred, f_true, f_false, x) (lax.switch for several) | runs one branch at run time (under vmap with a batched predicate it becomes a select: both run) |
| Loop with a fixed number of steps | lax.fori_loop(lo, hi, body, init) | compiled loop, body recorded once |
| Loop producing an output per step | lax.scan (Chapter 6.16) | compiled loop with stacked outputs |
| Loop until a condition holds | lax.while_loop(cond, body, init) | compiled loop; not reverse-mode differentiable |
The NaN-gradient trap of where. Because both branches are computed, an "unsafe" branch can poison the gradient: jnp.where(x > 0, jnp.log(x), 0.0) at x = 0 gives value 0 but gradient nan (the backward pass multiplies 0 by the infinite slope of log at 0). Fix with a "double where": make the unsafe branch's input safe first, jnp.where(x > 0, jnp.log(jnp.where(x > 0, x, 1.0)), 0.0), which gives gradient 0.
Why do we need it?
Models need data-dependent choices (robust losses, clipping, piecewise functions, censoring) and long loops (time steps, iterations). Plain Python versions either crash under tracing or explode the program size and compile time.
Where is it used?
Piecewise trends and clipping in forecasting models, robust losses, safe logs and divisions, NumPyro's NUTS (its tree building uses while_loop), svi.run (a scan over steps), and RNN or state-space models (scan over time).
How is it used?
Ask: does the condition depend on traced data? If not, use Python and mark the setting static. If yes and it is elementwise, use jnp.where (with the double-where trick for unsafe branches). Replace long Python loops by fori_loop or scan.
where, cond or a compiled loop."jnp.where skips the branch it does not pick."
It computes both. Errors in the unused branch (log of 0, division by 0, overflow in exp) can still produce NaN or infinite gradients. Make each branch's input safe first.
"lax.cond is always faster than where."
For cheap elementwise branches, where is usually as fast or faster. And under vmap with a batched condition, cond is turned into a select that runs both branches anyway (we checked the jaxpr).
"Python loops are forbidden in JAX."
They are fine when short or over settings (three seasonal blocks, a list of likelihood terms): they unroll into a few copies. Long loops over data or time steps should be fori_loop or scan.
In both projects, switches like "Normal, Student-t or Negative Binomial likelihood", "include holidays or not", or "full-rank or low-rank guide" are settings, decided before tracing: plain Python if is right, and each choice compiles its own program. Choices that depend on data values (clip a forecast at zero, skip a segment with no conversions, use a different formula for large residuals) must be written with jnp.where, jnp.maximum, masks or lax.cond. A softplus or jnp.maximum(mu, 1e-6) on a count mean is the elementwise version of "if mu < 0".
Python if on traced data → TracerBoolConversionError. Use jnp.where (elementwise, both branches computed) or lax.cond (scalar, one branch at run time).
Long Python loops unroll (5 000 steps: 10 000 equations, 1.8 s compile); lax.fori_loop/scan/while_loop keep one body (≈ 12–16 ms).
Trap: where(x > 0, log x, 0) has a nan gradient at 0; use a double where.
Quick check: inside a jitted model you write if sigma < 1e-3: sigma = 1e-3, where sigma is a sampled value. What happens, and what is the fix?
Tracing fails with TracerBoolConversionError, because sigma < 1e-3 is a traced boolean. Fix: sigma = jnp.maximum(sigma, 1e-3) (or better, give sigma a prior and transform that keep it positive).
The boolean-masking problem (your A/B framework) core
In NumPy, x[mask] keeps only the entries where the mask is True. "Conversions of the users in group B" is conv[group == 1]. Natural, and fine in NumPy. But look at the length of the result: it is the number of True entries, which depends on the data. Under jit, every shape must be known from the input shapes alone, before any values exist. So JAX refuses: NonConcreteBooleanIndexError.
The fix is a change of habit: do not cut rows out; grey them out. Keep the full array and switch off the unwanted entries: multiply by the mask, or jnp.where(mask, x, 0), then divide by the number of True entries. The shape never changes, so one compiled program serves every experiment.
Three ways to say it:
- Picture: instead of scissors that make the table shorter, a highlighter that marks which rows count.
- Numbers: 6 users, group B is users 2, 4, 5 with conversions 0, 1, 0:
x[mask]has length 3 (another experiment: length 4); the masked version always has length 6 and gives the same rate 1/3. - Slogan: don't cut, mask.
Conversion rate of variant B, conversions conv = [1, 0, 1, 1, 0, 1], groups group = [0, 1, 0, 1, 1, 0] (0 = A, 1 = B). Real outputs:
- Eager (no jit):
conv[group == 1]= [0, 1, 0], shape (3,); mean 0.3333. - Jitted:
jax.jit(lambda c, g: c[g == 1].mean())(conv, group)raisesNonConcreteBooleanIndexError: Array boolean indices must be concrete; got bool[6]. - Masked sum: mask m = [0, 1, 0, 1, 1, 0]; $\sum_i m_i c_i = 0 + 0 + 0 + 1 + 0 + 0 = 1$; $\sum_i m_i = 3$; rate $= 1/3$. Code:
jnp.sum(jnp.where(g == 1, c, 0.0)) / jnp.sum(g == 1)→0.33333334under jit, shapes fixed at (6,). - Other fixed-shape options: multiply,
jnp.sum(m * c) / jnp.sum(m)(same value); or fixed-size indices,jnp.nonzero(g == 1, size=6, fill_value=-1)→[1 3 4 -1 -1 -1]. - Selecting before tracing also works: a mask computed in NumPy from concrete data and then used as a constant gave 0.3333 under jit. But each different selection size is a different input shape, so each one compiles again.
- In-place masked updates fail the same way:
y.at[mask].set(0.0)with a traced boolean mask also raisedNonConcreteBooleanIndexError. Usejnp.where(mask, 0.0, y).
An operation has a data-dependent shape when the size of its output depends on the values of its inputs (boolean indexing x[mask], jnp.nonzero(x) without size=, jnp.unique without size=). Such operations need concrete values and cannot be traced inside jit.
Fixed-shape replacements, with mask $m_i \in \{0, 1\}$:
$$\text{masked sum} = \sum_i m_i x_i, \qquad \text{masked mean} = \frac{\sum_i m_i x_i}{\max\big(\sum_i m_i,\ 1\big)}, \qquad \text{masked log-likelihood} = \sum_i m_i \log p(y_i\mid\theta).$$- The
max(·, 1)guards the case "no entries selected" (an empty group); decide explicitly what the answer should be then. - Masked-out entries are still computed, so their values must be finite: $0 \times \infty$ and $0 \times \text{NaN}$ are NaN. Replace them first (
jnp.where(m, x, safe_value)). - NumPyro:
dist.X(...).mask(m)orwith numpyro.handlers.mask(mask=m):makes masked-out observations contribute 0 to the log density. (obs_mask=is different: it treats masked observations as missing values to impute.)
Why do we need it?
Selecting subsets (a variant, a segment, the valid days, non-missing values) is everywhere in analysis code, and the NumPy habit x[mask] crashes under jit. Masked arithmetic keeps one compiled program for all subsets and all experiments.
Where is it used?
Per-variant and per-segment metrics in A/B analysis, ragged groups padded into a rectangle for hierarchical models, missing days in time series, attention masks in Transformers, padded batches in serving, and NumPyro likelihoods over padded data.
How is it used?
Pass the mask in as an array argument, keep the data full-length (or padded), write every reduction as a masked sum divided by the mask count, and in NumPyro wrap the likelihood with .mask(m) or handlers.mask. Check once that the masked result equals the eager x[mask] result.
"Filtering rows inside the model with x[mask] is the cleanest way."
It is the NumPy way, and it fails under jit because the output length depends on the data. Inside traced code, mask arithmetically.
"Multiplying by the mask always works."
Not if the masked-out entries are infinite or NaN (padding with 0 inside a log, a Gamma or log-normal likelihood at 0): $0 \times (-\infty)$ is NaN. Pad with safe values, or use jnp.where to replace them before the arithmetic, and let NumPyro's mask zero the log-probabilities.
"The mask handler removes the padded data."
It keeps all shapes and only sets the masked log-probability terms to 0. The padded entries are still computed. In the NumPyro version checked here (0.22), masked observations are swapped for a harmless in-support value before their log-probability is evaluated, so padded observations are safe there. But anything else computed from padded entries (for example the log of a padded covariate that feeds the mean) still flows through and can give NaN gradients, so give those safe values too, and test the gradient once on a padded input.
This is the problem you hit in your experimentation framework with boolean masking before tracing. Picking a variant's or a segment's rows with a boolean mask produces an array whose length depends on the data; inside a function that JAX traces (a jitted step, an SVI update, a NumPyro model) that is not allowed, so JAX raises NonConcreteBooleanIndexError (or, if the mask was computed earlier from concrete data, every new subset size quietly forces a recompile). There are two clean designs, and you should be able to explain whichever your code uses: (a) do all boolean selection in NumPy or pandas outside the traced code, on concrete data, and accept one compile per distinct shape (fine for a handful of segments); or (b) keep fixed shapes, pass a group index and a mask into the model, and use masked sums, jnp.where, or NumPyro's .mask(m) / handlers.mask, so one compiled program serves every experiment. We checked that NumPyro's masked log density on padded data, −3.875885, equals the log density of the unpadded data, $4\log 0.6 + 2\log 0.4 = -3.875884$.
"Boolean indexing is a JAX bug."
It is a consequence of compiling for shapes: the output length of x[mask] is only known after looking at the data, and a compiled program must know every shape in advance.
Model answer: "jit traces the function with abstract values that have shapes but no values, and compiles a program for those shapes. x[mask] returns as many elements as there are True entries, so its shape depends on the data, and JAX raises NonConcreteBooleanIndexError. I keep the shape fixed and apply the mask arithmetically: jnp.where(mask, x, 0).sum() / mask.sum(), or in NumPyro I wrap the likelihood with a mask so masked observations contribute zero log-probability. If I really need a subset, I select it in NumPy before tracing and accept a recompile per shape, or pad to a fixed size."
x[mask] → data-dependent shape → NonConcreteBooleanIndexError under jit (also .at[mask].set).
Fix: masked sum Σ mᵢxᵢ, masked mean Σ mᵢxᵢ / max(Σ mᵢ, 1), jnp.where(mask, x, 0); NumPyro dist.X(...).mask(m) or handlers.mask(mask=m); or select in NumPy before tracing (one compile per shape).
Trap: masked-out entries must be finite in hand-written arithmetic (0 × inf = NaN); NumPyro's mask protects padded observations, not padded covariates.
Quick check: rewrite revenue[converted == 1].mean() so it runs under jit, and say what it returns when nobody converted.
m = (converted == 1); jnp.sum(jnp.where(m, revenue, 0.0)) / jnp.maximum(jnp.sum(m), 1). With nobody converted the numerator is 0 and the denominator is max(0, 1) = 1, so it returns 0 instead of a NaN; if "no data" should be reported differently, return a flag as well (jnp.sum(m) == 0).
Compile overhead, asynchronous dispatch and device transfer
Think of a manager (your Python program, the host) and a workshop (the device that runs compiled programs). The manager sends work orders and immediately moves on: that is asynchronous dispatch. As long as the manager only sends orders, the workshop stays busy. But whenever the manager asks "what is the number?" (printing the loss, float(loss), converting to NumPy, an if on a value), it has to wait for the workshop to finish and for the result to be shipped back. Asking at every step keeps interrupting the flow.
Data also has to travel: NumPy arrays live in host memory, and each time you pass one to a jitted function, JAX copies it to the device. On a CPU that is a quick memory copy; on a GPU it crosses a bus and costs much more.
Three ways to say it:
- Picture: a courier between office and workshop; fewer trips, faster work.
- Numbers (one laptop CPU): 2 000 jitted SVI steps took 235–240 ms when the loss was read every step and 141 ms when it was read every 100 steps; passing NumPy arrays every step instead of device arrays took 157 ms.
- Slogan: keep data on the device; look at it rarely.
Where the time goes in the loop (2 000 steps of a jitted svi.update, small regression model, 2 000 data points; measured twice).
- Reading
float(loss)every 100 steps: 141 ms in total, so about $141/2000 \approx 0.07$ ms per step of real work and dispatch. - Reading it every step: 235–240 ms, about $0.119$ ms per step.
- Extra cost per read: $0.119 - 0.0705 \approx 0.048$ ms. With 2 000 reads that is $2000 \times 0.048 \approx 96$ ms, which explains the gap ($141 + 96 \approx 237$ ms).
- Passing the NumPy arrays (instead of JAX arrays already on the device) at every step: 157 ms, about 8 µs extra per step for converting and copying 2 × 2 000 numbers.
- An explicit
jax.device_putof a 1 000 × 1 000 float32 array (4 MB) took about 0.04 ms here: on a CPU the "device" is the same memory. On a GPU every one of these numbers is larger; measure on your own hardware.
- Asynchronous dispatch: a JAX call queues work and returns an array "future" at once. Blocking operations wait for the result:
x.block_until_ready(),float(x),x.item(),print(x),np.asarray(x),jax.device_get(x), Python control flow onx. - Device transfer: host → device with
jax.device_put(x)(or implicitly when a NumPy array is passed to a JAX function); device → host withjax.device_get(x)or any blocking conversion. - A training loop costs about $N\,t_{\text{step}} + (N/k)\,t_{\text{sync}}$ when the host reads results every $k$ steps, plus compile time.
- NumPyro's own
svi.runwith a progress bar keeps the losses on the device and fetches them in batches; its source comments that this avoids "blocking on the device at every step".
Why do we need it?
A fast compiled step can be slowed down by the loop around it: a sync every step, a NumPy conversion every step, or data copied in every step. On a GPU these overheads can dominate the actual computation.
Where is it used?
Custom SVI loops with early stopping (the stopping rule needs loss values on the host), logging and progress bars, data pipelines that feed batches, and model serving (keep weights on the device, move only requests and answers).
How is it used?
Convert data to JAX arrays once before the loop (jnp.asarray or jax.device_put), keep the state on the device, collect losses in a device array or list, and read them on the host every k steps (averaging them, which also smooths ELBO noise).
"Printing the loss every step is free."
Each print forces the host to wait for the device and copy the value back. On the measured CPU it made the loop about 1.7 times slower (about 40% of the time was spent waiting); on a GPU it can be worse.
"Passing NumPy arrays to the jitted step is the same as passing JAX arrays."
The result is the same, but each call converts and copies the arrays. Convert once, before the loop.
"The first step's time tells me the transfer cost."
The first step also contains tracing and compilation. Measure transfer and sync costs on warm, later steps.
Your forecasting loop's relative-ELBO early stopping with patience (Chapter 6.14) needs loss values in Python, so it must sync at some point. Evaluating every k steps (and using the average of those k losses) cuts both the sync overhead and the ELBO noise in the stopping rule; the best-state checkpoint can be taken at the same moments. Move the design matrix and the targets to the device once, before the loop, and keep the SVI state on the device between steps.
JAX calls are asynchronous; float(x), print(x), np.asarray(x), .block_until_ready() wait for the device.
Loop cost ≈ N·t_step + (N/k)·t_sync. Measured (CPU, 2 000 steps): read every step 235–240 ms, every 100 steps 141 ms.
Put data on the device once; read losses every k steps.
Quick check: a GPU loop does 0.02 ms of work per step and each loss read costs 0.15 ms. What fraction of the time is spent waiting if you read every step? Every 50 steps?
Every step: $0.15/(0.02 + 0.15) \approx 88\%$ waiting. Every 50 steps: per 50 steps, work 1 ms and one read 0.15 ms, so $0.15/1.15 \approx 13\%$. (These GPU numbers are illustrative; measure your own.)
Model size: latent dimension and parameter counts core
"How big is my model?" has two different answers, and mixing them up is a classic mistake. The first is the latent dimension $d$: how many unknown numbers the model has (slopes, changepoint adjustments, Fourier coefficients, segment effects, noise scales). The second is how many numbers the optimizer must learn to describe the approximate posterior: the guide's variational parameters. That second number depends on the guide family: a mean-field guide needs 2 numbers per latent value, a full-rank guide needs a whole triangle of covariance numbers, a low-rank guide needs a thin rectangle.
The number of data points is neither of these. A year of daily data and ten years of daily data give the same $d$ if the model has the same components.
Three ways to say it:
- Picture: $d$ dials on the model; the guide needs one position and one wobble per dial, plus (for full rank) a link between every pair of dials.
- Numbers: d = 50: mean-field 100 parameters, full-rank 1 325, low-rank (r = 5) 350, checked in NumPyro 0.22. At d = 2 000: 4 000, 2 003 000 and 24 000 (r = 10).
- Slogan: full-rank grows like d², low-rank like d·r.
Counting d for an illustrative Prophet-style model (not your exact code; plug in your own component sizes).
- Trend: base slope k and offset m: 2.
- Changepoints: 25 slope adjustments $\delta_j$: 25.
- Yearly seasonality, Fourier order 10: a sine and a cosine coefficient per order: $2 \times 10 = 20$. Weekly, order 3: $2 \times 3 = 6$.
- 10 holiday effects: 10. 5 exogenous regressors: 5. Noise scale σ (Normal likelihood): 1. (Student-t would add ν; Negative Binomial has a concentration α instead of σ.)
- Total: $d = 2 + 25 + 20 + 6 + 10 + 5 + 1 = 69$.
- Guide parameters: mean-field $2d = 138$; full-rank $d + d(d+1)/2 = 69 + 69 \cdot 70/2 = 69 + 2415 = 2484$; low-rank with r = 10: $d(r + 2) = 69 \times 12 = 828$.
- Latent dimension $d$: the total number of scalar latent values after flattening every sample site (vectors and matrices count all their entries). NumPyro's autoguides work in the unconstrained space: a positive σ is represented by log σ, a probability by its logit, but the count is the same.
- Guide parameter counts (verified with
svi.initin NumPyro 0.22):
| Guide | Parameters | Stored as | Order |
|---|---|---|---|
AutoNormal (mean-field) | $2d$ | loc (d), scale (d) | O(d) |
AutoMultivariateNormal (full-rank) | $d + d(d+1)/2$ | loc (d), lower-triangular Cholesky factor L (d(d+1)/2 free numbers, shown as a d×d matrix) | O(d²) |
AutoLowRankMultivariateNormal (rank r) | $d(r + 2)$ | loc (d), cov_factor (d×r), scale (d) | O(dr) |
- The full-rank count is $d(d+1)/2$, not $d^2$, because a covariance matrix is symmetric (or, equivalently, its Cholesky factor is triangular).
- Memory of the parameters: count × 4 bytes in float32. Adam keeps two more arrays of the same size (first and second moments), so the optimizer state is about 3 × the parameters (we counted 3 976 = 3 × 1 325 + 1 step counter for d = 50).
Why do we need it?
The guide you can afford depends on $d$: a full-rank guide for d = 69 is trivial (2 484 numbers), for d = 5 000 it is 12.5 million numbers and slow. You cannot choose the family, or explain an automatic choice, without counting $d$ first.
Where is it used?
Choosing between AutoNormal, AutoLowRankMultivariateNormal and AutoMultivariateNormal; deciding how many changepoints, Fourier terms, holidays and segments a model can carry; estimating memory before a run; your forecasting model's automatic full-rank vs low-rank switch.
How is it used?
Add up the sizes of all latent sites (or read them off svi.get_params / the guide's loc), compute the three parameter counts, multiply by 4 bytes and by 3 for Adam, and compare with your memory and time budget.
"d is the number of data points."
d counts unknowns, not observations. More days of history do not change d (unless the model has a latent value per day, as a local-level trend would).
"The full-rank guide has d² parameters."
It has $d + d(d+1)/2$: the covariance is symmetric, so only one triangle is free. Still O(d²): doubling d roughly quadruples it.
"Parameter count = memory used."
Multiply by 4 bytes (float32) and by about 3 for Adam's moment estimates, and add room for gradients and intermediate arrays during the step.
Your forecasting model chooses automatically between a full-rank and a low-rank Gaussian guide based on model size. This section is the reason: every changepoint, Fourier order, holiday and regressor adds to d, and the full-rank parameter count grows like d²/2 while the low-rank one grows like d·(r + 2). A rule of the form "full-rank if d is at most some threshold, otherwise low-rank with rank r" keeps the cost bounded. Check the exact threshold and rank in your code and be ready to state them with the resulting counts; the guides themselves, and what each can and cannot capture, are in Chapter 6.13.
"The full-rank guide costs O(d³) memory."
Memory and parameters are O(d²) (a triangle of about d²/2 numbers). O(d³) is the time to factorize a d×d covariance matrix (Cholesky); NumPyro's full-rank guide learns the factor L directly, so each SVI step costs O(d²) work, not a fresh O(d³) factorization.
Model answer: "With d latent variables, mean-field has 2d parameters, full-rank has d + d(d+1)/2, which is O(d²) memory and O(d²) work per step, and a rank-r low-rank guide has d(r + 2), linear in d. That is why the model switches to low-rank as d grows: it keeps the strongest posterior correlations at a cost close to O(dr)."
d = number of scalar latents (not data points). Example: 2 + 25 + 20 + 6 + 10 + 5 + 1 = 69.
Mean-field 2d · full-rank d + d(d+1)/2 · low-rank d(r + 2). d = 69, r = 10: 138 · 2 484 · 828.
Memory ≈ count × 4 bytes × 3 (Adam). Trap: full-rank memory is O(d²), not O(d³).
Quick check: an A/B model has 2 variants × 30 segments, non-centered: for each variant a mean μ, a spread τ and 30 standardized effects z. What is d, and how many parameters does a full-rank guide need?
Per variant 1 + 1 + 30 = 32, so d = 64. Full-rank: 64 + 64 · 65/2 = 64 + 2 080 = 2 144. Mean-field 128; low-rank with r = 5: 64 × 7 = 448.
Memory and compute: O(d²), O(d³) and O(dr), measured core
Big-O notation says how a cost grows when the size doubles. O(d): doubles. O(d²): four times. O(d³): eight times. (Constants are ignored: O(d²) means "about some constant × d²", not "d² seconds".)
A full covariance matrix is a d×d square: storing it is O(d²). Factorizing it (the Cholesky decomposition, which turns a covariance Σ into a triangular L with $\Sigma = LL^\top$) takes about $d^3/3$ multiply-adds: O(d³). A low-rank-plus-diagonal covariance $\Sigma \approx DD^\top + \text{diag}$ is a thin d×r rectangle plus a vector: O(dr) to store, and its log-determinant and solves need only an r×r matrix (the Woodbury and matrix-determinant identities), so the work stays close to linear in d.
Three ways to say it:
- Picture: a square that grows in both directions vs a thin strip that only grows in length.
- Numbers (one laptop CPU): d = 4 000: the full covariance is 16 million numbers, 64 MB in float32, and its Cholesky took 33–180 ms across runs; a rank-10 factor is 40 000 numbers (0.16 MB) and its log-determinant took under 0.15 ms.
- Slogan: squares grow fast; thin strips stay cheap.
Arithmetic, then measurements.
- d = 1 000, full covariance: $d^2 = 10^6$ numbers × 4 bytes = 4 MB. The full-rank guide stores $d + d(d+1)/2 = 501\,500$ numbers (NumPyro reported exactly 501 500): 2.0 MB, about 6 MB with Adam.
- d = 10 000: $d + d(d+1)/2 = 50\,015\,000$ numbers: 200 MB, about 600 MB with Adam. A Cholesky of a 10 000 × 10 000 matrix needs about $d^3/3 \approx 3.3 \times 10^{11}$ multiply-adds.
- Low-rank, d = 10 000, r = 10: $d(r+2) = 120\,000$ numbers: 0.48 MB, about 1.4 MB with Adam.
- Measured SVI step times (median of 30 warm steps, two runs averaged; a linear-regression model with d coefficients and 200 observations):
| d | mean-field | full-rank | low-rank (r = 10) |
|---|---|---|---|
| 100 | 0.06 ms (200 params) | 0.27 ms (5 150) | 0.15 ms (1 200) |
| 500 | 0.13 ms (1 000) | 2.9 ms (125 750) | 0.27 ms (6 000) |
| 1 000 | 0.24 ms (2 000) | 8.0 ms (501 500) | 0.36 ms (12 000) |
| 2 000 | 0.30 ms (4 000) | 38.5 ms (2 003 000) | 0.70 ms (24 000) |
From d = 1 000 to 2 000 the full-rank step got 4.8× slower (close to the ×4 of O(d²)); the low-rank step 1.9× (close to ×2). First calls (compile) took 0.19–0.6 s for all three. Timings vary by machine.
| Gaussian family over d latents | Storage | Per SVI step (per draw) | Notes |
|---|---|---|---|
| Mean-field (diagonal) | O(d) | O(d): $\theta = \mu + \sigma \odot \varepsilon$ | no correlations |
| Full-rank, learning L directly | O(d²) | O(d²): $\theta = \mu + L\varepsilon$, a triangular solve for $\log q$ | all correlations |
| Full covariance Σ given, needs its factor | O(d²) | O(d³) Cholesky, once per new Σ | e.g. Laplace approximations, MultivariateNormal(covariance_matrix=Σ) |
| Low-rank + diagonal, $DD^\top + \text{diag}(s^2)$ | O(dr) | O(dr) to sample; O(dr² + r³) for $\log q$ via an r×r "capacitance" matrix | the r strongest correlation directions |
Here ⊙ means elementwise multiplication and ε is a vector of independent standard Normal draws (the reparameterization of Chapter 6.12). For fixed small r, the low-rank costs grow linearly in d.
Why do we need it?
To predict whether a model will fit in memory and finish in time before running it, to explain why a guide choice depends on model size, and to know which part of a step dominates when the model grows.
Where is it used?
Choosing full-rank vs low-rank guides, Gaussian processes (O(n³) kernel Cholesky, with low-rank and inducing-point approximations), Kalman filters, Laplace approximations, attention cost in Transformers (O(n²)), and any memory budget for GPU training.
How is it used?
Write the cost of each piece as a power of d, plug in your d, compare with measured timings at two sizes (the ratio tells you the exponent), and switch to a cheaper family (or reduce d) when the dominant term gets too large.
"Low-rank is always better: cheaper and almost as good."
It captures only r directions of correlation. If the posterior has many strongly correlated groups of parameters (say, changepoints that trade off with regressors in many ways), a small r misses some, and the uncertainty can be too small. Check against full-rank (or NUTS) on a smaller version of the model (Chapter 6.15).
"O(d²) means it is slow."
Big-O is about growth, not speed at your size. At d = 100 the full-rank step took 0.27 ms; the d² term only bites when d is in the thousands.
"Every full-rank SVI step does an O(d³) Cholesky."
Not in NumPyro's AutoMultivariateNormal, which learns the Cholesky factor directly: steps are O(d²). O(d³) appears when you start from a covariance matrix (a Laplace approximation's inverse Hessian, a Gaussian-process kernel, a covariance_matrix= argument).
For your forecasting model, d grows with every changepoint, Fourier order, holiday and regressor, so the guide's cost grows with your modelling choices: full-rank memory and per-step work like d², low-rank like d·r. That is the engineering reason behind the automatic full-rank vs low-rank switch, and also a reason to keep the candidate changepoint grid and Fourier orders no larger than needed (they cost compute and, as Chapter 7.18 shows, variance). For the A/B framework, d grows with segments × variants × metrics; a hierarchical model over many segments is a natural place for mean-field or low-rank guides. Compile time grows with the size of the traced program as well, which is one more reason to keep the model's array code vectorized.
Storage: mean-field O(d), full-rank O(d²) (d + d(d+1)/2), low-rank O(dr) (d(r + 2)).
Work: full-rank step O(d²) (L learned directly); Cholesky of a given Σ O(d³) ≈ d³/3; low-rank log q O(dr² + r³).
Measured (CPU): d 1 000 → 2 000: full-rank step 8.0 → 38.5 ms (×4.8), low-rank 0.36 → 0.70 ms (×1.9).
Quick check: a full-rank step takes 8 ms at d = 1 000. Roughly how long at d = 4 000 if the O(d²) part dominates, and how much memory does the optimizer state need?
×16 (4² = 16): about 130 ms per step (in practice somewhat more, as the 1 000 → 2 000 ratio of 4.8 suggests). Parameters: 4 000 + 4 000 · 4 001/2 = 8 006 000; × 4 bytes × 3 ≈ 96 MB.
Recap, cheat sheet and practice
- jit pipeline: Python → tracing with abstract values (shape, dtype) → jaxpr → StableHLO → XLA → executable, cached by function + shapes + dtypes + static arguments. Python code runs only at trace time.
- Compile once, run fast: the first call per signature pays trace + compile (≈ 0.4–0.75 s for an SVI step on one CPU), later calls are fast (≈ 0.1–0.2 ms vs 40–70 ms eager). Break-even $N^* = C/(e - r)$. Time with
block_until_readyafter a warm-up. - Static shapes: new shapes, dtypes or static values recompile; values do not. Pad varying sizes to a few buckets and carry a mask.
- Control flow: Python
ifon traced data fails; usejnp.where(both branches computed; beware NaN gradients, use a double where) orlax.cond. Long Python loops unroll; usefori_loop,scan,while_loop. Pythonifon settings is fine. - Boolean masking:
x[mask]has a data-dependent shape →NonConcreteBooleanIndexError. Don't cut, mask: masked sums,jnp.where, NumPyro.mask(m)/handlers.mask; or select before tracing and accept a compile per shape. - Async dispatch and transfer: reading values blocks; copying NumPy data costs time. Put data on the device once; read losses every k steps.
- Model size: d = number of scalar latents. Guide parameters: 2d, d + d(d+1)/2, d(r + 2). Memory ≈ × 4 bytes × 3 with Adam.
- Growth: full-rank O(d²) storage and per-step work, O(d³) to factorize a given covariance; low-rank O(dr) storage, about linear work. Measured: full-rank step ×4.8 from d = 1 000 to 2 000, low-rank ×1.9.
Cheat sheet
| Topic | Tool / formula | Remember |
|---|---|---|
| See the recording | jax.make_jaxpr(f)(*args); jax.jit(f).lower(*args).as_text() | operations + shapes, no values |
| Find recompiles | jax.config.update("jax_log_compiles", True) | one log line per compilation |
| Timing | f(x).block_until_ready(), warm-up, median | first call = trace + compile |
| Break-even | $N^* = C/(e - r)$ | C = 0.5 s, e = 50 ms, r = 0.15 ms → ≈ 10 steps |
| Static settings | jax.jit(f, static_argnames="order") | new value → new compile; must be hashable |
| Varying sizes | pad to buckets 64, 128, 256, … + mask | waste < 50% per call; few compiles |
| Data-dependent branch | jnp.where / lax.cond | where computes both branches; double-where for log/div |
| Long loops | lax.fori_loop, lax.scan, lax.while_loop | body traced once; Python loops unroll |
| Subsets | $\sum m_i x_i / \max(\sum m_i, 1)$; NumPyro .mask(m) | never x[mask] inside traced code |
| Loop overhead | $N t_{\text{step}} + (N/k)\,t_{\text{sync}}$ | read losses every k steps; data on device once |
| Guide sizes | 2d · d + d(d+1)/2 · d(r + 2) | d = 69, r = 10: 138 · 2 484 · 828 |
| Growth | O(d) · O(d²) · O(d³) Cholesky · O(dr) | doubling d: ×2 · ×4 · ×8 · ×2 |
import time
import numpy as np
import jax
import jax.numpy as jnp
from jax import lax
import numpyro
import numpyro.distributions as dist
from numpyro import handlers
from numpyro.infer import SVI, Trace_ELBO
from numpyro.infer.util import log_density
from numpyro.infer.autoguide import AutoNormal, AutoMultivariateNormal, AutoLowRankMultivariateNormal
from numpyro.optim import Adam
# 1) Tracing: the jaxpr records operations and SHAPES, not values
def nll(mu, y):
return 0.5 * jnp.sum((y - mu) ** 2)
y5 = jnp.array([96.0, 104.0, 110.0, 90.0, 100.0])
print(jax.make_jaxpr(nll)(jnp.float32(95.0), y5)) # sub, integer_pow, reduce_sum, mul on f32[5]
print(nll(95.0, y5)) # 178.5
# 2) Compile once, run fast. Time with block_until_ready (JAX dispatches asynchronously)
y = jax.random.normal(jax.random.PRNGKey(0), (100_000,))
jnll = jax.jit(nll)
t0 = time.perf_counter(); jnll(0.1, y).block_until_ready(); t1 = time.perf_counter()
for _ in range(100):
jnll(0.1, y).block_until_ready()
t2 = time.perf_counter()
print(f"first call {1e3*(t1-t0):.1f} ms, later calls {1e3*(t2-t1)/100:.3f} ms") # we saw 13-35 ms vs 0.03-0.09 ms
# 3) Every new input shape is a new trace + compile; padding to buckets avoids most of them
traces = 0
@jax.jit
def masked_mean(x, mask):
global traces; traces += 1 # Python side effect: runs only while tracing
return jnp.sum(jnp.where(mask, x, 0.0)) / jnp.maximum(jnp.sum(mask), 1)
def bucket(n):
b = 64
while b < n:
b *= 2
return b
for n in [100, 101, 102, 130, 200, 260]:
B = bucket(n)
xp = np.zeros(B, np.float32); xp[:n] = 1.0
m = np.zeros(B, bool); m[:n] = True
masked_mean(xp, m)
print("6 sizes ->", traces, "compiles") # 6 sizes -> 3 compiles (buckets 128, 256, 512)
# 4) Python control flow on traced values fails; use jnp.where / lax.cond / lax.fori_loop
def huber_bad(r):
return r - 0.5 if r > 1.0 else 0.5 * r ** 2
try:
jax.jit(huber_bad)(2.0)
except Exception as e:
print(type(e).__name__) # TracerBoolConversionError
huber = jax.jit(lambda r: jnp.where(jnp.abs(r) > 1.0, jnp.abs(r) - 0.5, 0.5 * r ** 2))
print(huber(jnp.array([-3.0, 0.5, 2.0]))) # [2.5 0.125 1.5 ]
print(jax.jit(lambda n: lax.fori_loop(0, n, lambda i, v: 0.5 * v + 1.0, 0.0))(1000)) # 2.0
# 5) The boolean-masking problem: x[mask] has a data-dependent shape
conv = jnp.array([1.0, 0.0, 1.0, 1.0, 0.0, 1.0])
group = jnp.array([0, 1, 0, 1, 1, 0])
try:
jax.jit(lambda c, g: c[g == 1].mean())(conv, group)
except Exception as e:
print(type(e).__name__) # NonConcreteBooleanIndexError
rate_b = jax.jit(lambda c, g: jnp.sum(jnp.where(g == 1, c, 0.0)) / jnp.sum(g == 1))
print(rate_b(conv, group)) # 0.33333334: same answer, fixed shapes
# NumPyro: padded observations switched off with a mask handler (contribute 0 log-probability)
def model(y, valid):
p = numpyro.sample("p", dist.Beta(1, 1))
with numpyro.plate("n", y.shape[0]), handlers.mask(mask=valid):
numpyro.sample("y", dist.Bernoulli(p), obs=y)
y_pad = jnp.array([1., 0., 1., 1., 0., 1., 0., 0.])
valid = jnp.array([True] * 6 + [False] * 2)
lp, _ = log_density(model, (y_pad, valid), {}, {"p": jnp.array(0.6)})
print(lp, 4 * np.log(0.6) + 2 * np.log(0.4)) # -3.875885 -3.87588...: only the 6 real rows count
# 6) Model size: parameters of the three Gaussian guides for d latent values
d, r = 50, 5
def big(): numpyro.sample("theta", dist.Normal(jnp.zeros(d), 1.0).to_event(1))
for G in [AutoNormal(big), AutoMultivariateNormal(big), AutoLowRankMultivariateNormal(big, rank=r)]:
st = SVI(big, G, Adam(0.01), Trace_ELBO()).init(jax.random.PRNGKey(0))
leaves = jax.tree.leaves(st.optim_state) # Adam keeps params + 2 moment arrays (+ a step counter)
n_par = sum(l.size for l in leaves[1:]) // 3
print(type(G).__name__, n_par) # 100 (2d), 1325 (d + d(d+1)/2), 350 (d(r+2))
# 7) O(d^3): time a Cholesky factorization as d doubles (timings vary by machine and run)
chol = jax.jit(jnp.linalg.cholesky)
for dd in [1000, 2000, 4000]:
A = np.random.default_rng(0).normal(size=(dd, dd)).astype(np.float32)
S = jnp.asarray(A @ A.T / dd + np.eye(dd, dtype=np.float32))
chol(S).block_until_ready()
t0 = time.perf_counter(); chol(S).block_until_ready()
print(dd, f"{1e3*(time.perf_counter()-t0):.1f} ms", f"{dd*dd*4/1e6:.0f} MB")
# runs varied: d=1000 0.6-4.9 ms, 2000 3.5-27 ms, 4000 33-180 ms; memory 4, 16, 64 MB
1. The first call of your jitted svi.update takes 0.6 s and later calls take 0.15 ms. Why?
2. Which change makes a jitted function recompile?
3. Inside a jitted step you write if loss > best_loss: …, where loss is computed in the step. What happens?
loss > best_loss is a traced boolean with no value during tracing, so Python cannot branch on it. Early-stopping logic belongs in the Python loop around the jitted update.4. How do you compute variant B's conversion rate inside traced code?
5. How many variational parameters does AutoMultivariateNormal have for d = 100 latent values?
6. You double the latent dimension d. Roughly how does the full-rank guide's parameter storage change?
Practice problems
A. Compile cost 0.8 s, jitted step 0.3 ms, eager step 30 ms. Find the break-even number of steps and compare totals for a 3 000-step run.
- $N^* = C/(e - r) = 800\text{ ms}/(30 - 0.3)\text{ ms} = 800/29.7 \approx 27$ steps.
- 3 000 steps with jit: $0.8 + 3000 \times 0.0003 = 0.8 + 0.9 = 1.7$ s.
- Eager: $3000 \times 0.03 = 90$ s. jit is about 53× faster for this run.
B. Your framework analyses 40 segments with between 50 and 900 users each, one jitted call per segment. How many compiles with exact shapes (worst case) and with power-of-two buckets starting at 64? What is the worst-case padding waste?
- Exact shapes: up to 40 compiles (one per distinct size). At about 0.5 s each for a model-sized function, that is up to 20 s of compiling.
- Buckets: sizes 50–900 fall into 64, 128, 256, 512 or 1024: at most 5 compiles.
- Worst waste: a size just above a bucket boundary, e.g. 257 → 512: $255/512 \approx 50\%$; on average much less. Alternatively, pad all segments into one (40, 1024) rectangle with a mask and compile once.
C. Make this jit-safe: def avg_rev(rev, conv, seg, s): r = rev[(conv == 1) & (seg == s)]; r = r[r > 0]; return r.mean() if len(r) > 0 else 0.0.
- Problems: two boolean selections (data-dependent shapes) and a Python
ifon a data-dependent length. - One mask for all conditions:
m = (conv == 1) & (seg == s) & (rev > 0). - Masked mean with a guard:
n = jnp.sum(m); total = jnp.sum(jnp.where(m, rev, 0.0)); return jnp.where(n > 0, total / jnp.maximum(n, 1), 0.0). Shapes never change;scan stay a traced argument, so one compile serves every segment.
D. An A/B model: 2 variants × 30 segments, non-centered hierarchical (μ, τ and 30 z's per variant). Compute d and the three guide parameter counts (low-rank r = 5), and the full-rank memory with Adam.
- Per variant: 1 + 1 + 30 = 32 latents; d = 64.
- Mean-field: 128. Full-rank: $64 + 64 \cdot 65/2 = 64 + 2080 = 2144$. Low-rank: $64 \times (5 + 2) = 448$.
- Full-rank memory with Adam: $2144 \times 4 \times 3 \approx 25.7$ kB. Tiny: at this size, full-rank is affordable and captures the μ–τ–z correlations.
E. d = 5 000, float32, Adam. Compare the optimizer-state memory of a full-rank guide and a rank-20 low-rank guide.
- Full-rank parameters: $5000 + 5000 \cdot 5001/2 = 5000 + 12\,502\,500 = 12\,507\,500$.
- × 4 bytes = 50.0 MB; × 3 for Adam ≈ 150 MB (plus gradients and temporaries during the step).
- Low-rank: $5000 \times 22 = 110\,000$ parameters; × 4 × 3 ≈ 1.3 MB, over 100× smaller. Per-step work: O(d²) = 25 million vs O(dr²) = 2 million operations (rough counts).
F. (Interview) "Your forecasting model switches automatically between a full-rank and a low-rank Gaussian guide based on model size. Why, and what is the trade-off?"
"The model's latent dimension d grows with the changepoints, Fourier orders, holidays and regressors. A full-rank Gaussian guide stores a d-vector plus a Cholesky factor with d(d+1)/2 entries, so memory and per-step work grow like d²; Adam triples the memory. That is cheap at a few hundred latents but becomes slow and memory-hungry in the thousands: in a benchmark of a regression with d coefficients (rerun it on your own machine before quoting), a full-rank step was about 4.8× slower going from d = 1 000 to 2 000. A low-rank guide uses a d×r factor plus a diagonal, d(r + 2) parameters, so its cost grows linearly in d. The trade-off is fidelity: full-rank captures every pairwise posterior correlation, low-rank only the r strongest directions plus independent variances, so it can understate uncertainty when many parameters trade off. So the rule is: full-rank while it is affordable, low-rank above a size threshold, and I validate the low-rank fit against full-rank or NUTS on a smaller version of the model."
Glossary
Every important word of this guide in one place, explained in plain English (218 terms). Type in the box to filter: it searches the terms and their explanations. The small numbers after each entry link to the section that teaches it (hover for its title).
All terms, A to Z
No term matches. Try a shorter word, or press / to search the whole guide.
- Acceptance rate
- The fraction of proposed moves the sampler actually takes. Too low means steps are too big; very high often means steps are too small and the chain crawls. 6.7 6.9
- Adam memory (about three times the parameters)
- Parameters take count $\times$ 4 bytes in float32, and Adam keeps two more arrays of the same size (first and second moments), so the optimizer state is about 3 times the parameters. 6.17 6.13
- Approximation vs optimization error
- Approximation error is $\min_\phi KL(q_\phi\|p)$: the best member of the family is still not the posterior. Optimization error comes from stopping early, noisy gradients or a local optimum. What you get is the sum. 6.11 6.11
- Asymptotically exact
- NUTS draws converge to the posterior as the run grows, but any finite run has Monte Carlo error and may not have converged. Passing the diagnostics means “no evidence of sampling problems”, never “exact posterior”. 6.10 6.15
- Asynchronous dispatch
- A JAX call returns as soon as the work is queued; anything needing the value waits (
block_until_ready(),float(x),print(x),np.asarray(x), a Pythonifonx). A loop costs about $N\,t_{\text{step}}+(N/k)\,t_{\text{sync}}$ if the host reads results every $k$ steps. 6.17 6.17 - Autocorrelation (of a chain)
- $\rho_k=\text{Corr}(\theta_t,\theta_{t+k})$ along the chain. It makes MCMC less efficient, not wrong: averages still converge, they just need more draws. 6.9
- Autoguide initialization
- Autoguides place initial means with
init_loc_fn(defaultinit_to_uniform: unconstrained values drawn uniformly in $(-2,2)$) and start every scale atinit_scale=0.1. 6.14 - Bayes factor
- The ratio of two models' evidences, $p(D\mid M_1)/p(D\mid M_2)$. It compares models but swings strongly with how wide each prior is. 6.1
- Bayesian model
- A two-step story of how the data were made: nature first picks the unknown parameter $\theta$ from the prior $p(\theta)$, then the data come from the likelihood $p(D\mid\theta)$. Prior plus likelihood is a complete model. 6.1
- Bernstein–von Mises theorem
- With enough data per parameter (and regularity conditions such as a fixed number of parameters) the posterior is approximately Normal with sd close to the usual standard error, and the prior stops mattering. 6.2 6.9
- Best-state checkpointing
- Keep a copy of the parameters at the evaluation with the lowest loss so far and return it when the loop ends, for whatever reason (patience, maximum steps, NaN). To resume training store the whole
svi_state(Adam moments and PRNG key); for predictionsvi.get_params(state)is enough. Mainly insurance against late degradation; because the minimum of noisy values is optimistic, it is partly luck. 6.14 - Beta-Binomial (model vs distribution)
- The model: Beta prior plus Binomial likelihood, giving the posterior Beta($\alpha+k,\beta+n-k$). The distribution: the marginal or predictive count when $\theta$ is integrated out. Say which one you mean. NumPyro writes the Beta as
Beta(concentration1=α, concentration0=β), withconcentration1counting successes. 6.3 6.3 - Beta-binomial predictive
- The distribution of the number of conversions among the next $m$ users when the rate has a Beta($\alpha',\beta'$) posterior: the same mean as Binomial($m,\bar p$) but a variance inflated by $\frac{\alpha'+\beta'+m}{\alpha'+\beta'+1}\ge 1$. 6.3
- Between-group $\tau$ vs within-group $\sigma$
- $\tau$ is how much the groups' true values differ; $\sigma$ is the noise inside a group. $Var(y)=\sigma^2+\tau^2$, and a group average has variance $\tau^2+\sigma^2/n_g$. 6.5 6.5
- Boolean-masking problem
- Selecting rows with
x[mask]inside traced code raisesNonConcreteBooleanIndexError, because a compiled program must know every shape in advance. Don't cut, mask: masked sums,jnp.where(mask, x, 0), NumPyro.mask(m); or select in NumPy before tracing and accept one compile per shape. 6.17 - Borrowing strength
- Letting small groups use information from all groups through the shared population distribution, so a 40-user segment's estimate is steadied by the others. 6.6 6.2
- Burn-in
- The first iterations of a chain, discarded because they still remember the starting point. Convergence is checked (traces, several chains, $\hat R$), never proven; burn-in cannot fix slow mixing or a missed mode. 6.9
- Centered parameterization
- Write group effects directly as $\theta_g\sim N(\mu,\tau^2)$, so the sampler explores $(\mu,\tau,\theta_1,\dots,\theta_G)$. Best when each group has strong data ($SE_g\ll\tau$). 6.7 6.7
- Cholesky factor
- A lower-triangular $L$ with positive diagonal such that $\Sigma=LL^\top$. Learning $L$ rather than $\Sigma$ guarantees a valid covariance, makes sampling a matrix-vector product, and gives $\log|\Sigma|=2\sum\log L_{ii}$ in $O(d)$. 6.13
- Common random numbers (fixed-key evaluation)
- Evaluating the ELBO with the same PRNG key each time makes differences between evaluations far less noisy than the evaluations themselves, which gives better checkpoints. 6.14
- Compile once, run fast
- $T\approx T_{\text{trace}}+T_{\text{compile}}+N(T_{\text{dispatch}}+T_{\text{run}})$ per input signature; every new signature adds another trace and compile. Report compile time and step time separately, and time after a warm-up call with
.block_until_ready(). 6.17 - Complete pooling
- One shared parameter for every group, $\hat\theta_g=\bar y$ (the hierarchical model with $\tau=0$). Ignores real differences between groups. 6.6 6.6
- Conditional variance $1/\Lambda_{ii}$
- The variance of $\theta_i$ given all the other parameters, which never exceeds its marginal variance $\Sigma_{ii}$. It is what a mean-field guide converges to (correlation 0.9 gives 0.19 instead of 1). 6.13
- Confidence interval
- A frequentist interval from a procedure that captures the fixed true value in 95% of repeated samples. The 95% belongs to the recipe, not to this one interval. 6.4
- Confounded components (credit assignment)
- Components whose design-matrix columns move together (holiday and promotion, trend and yearly season over a short history, a changepoint at a holiday) are only identified as a sum; the priors write the split. Report the sum or add data where they vary separately. 6.8
- Conjugacy
- A property of a (prior family, likelihood) pair: the posterior stays in the prior's family for every dataset, so updating just changes hyperparameters, in closed form (an exact formula: no grid, sampling or optimization). Beta is conjugate to the Binomial but not to a Normal likelihood. 6.3
- Conjugate prior
- A prior family that the likelihood maps back into itself: prior in the family gives posterior in the family, for any data. A convenience, not evidence the prior is right. 6.2 6.3
- Control flow under tracing
- A plain Python
ifon traced data fails (TracerBoolConversionError) but is fine on settings known before tracing. Usejnp.where(elementwise, both branches computed),lax.cond(one scalar choice) andlax.fori_loop,lax.scan,lax.while_loop(loops traced once; Python loops unroll). 6.17 - Converged (in SVI)
- “No longer improving by a meaningful amount”, judged on a smoothed loss. The parameters still jitter, and “stopped by patience” is not the same as “converged”. 6.14 6.14
- Cost growth of guides: $O(d)$, $O(d^2)$, $O(d^3)$, $O(dr)$
- Mean-field storage and work $O(d)$; full-rank learning $L$ directly $O(d^2)$ per step; a Cholesky factorization of a given covariance $O(d^3)$ (about $d^3/3$); low-rank $O(dr)$ to sample and $O(dr^2+r^3)$ for the density. Doubling $d$: $\times2$, $\times4$, $\times8$, $\times2$. 6.17 6.17
- Cost in row-gradients
- A bookkeeping unit for speed: NUTS $\approx C\times(W+S)\times\bar L\times N$ (chains, warmup plus kept draws, leapfrog steps, rows), SVI $\approx T\times K\times B$ (steps, particles, rows per batch). For NUTS the fair measure is ESS per second, not draws per second. 6.15
- Credible interval
- An interval $[L,U]$ with $P(L\le\theta\le U\mid D)=0.95$: a direct probability statement about $\theta$ given the model, prior and data. Many intervals have that mass; from draws, take
np.quantile(draws, [0.025, 0.975]). 6.4 - Credible vs confidence interval
- Confidence: $\theta$ fixed, data random, 95% is the procedure's long-run hit rate. Credible: data fixed, $\theta$ uncertain, 95% is the posterior probability, which depends on the prior. With weak priors and lots of data the numbers are close. 6.4
- Data-dependent shape
- An operation whose output size depends on input values (
x[mask],jnp.nonzero,jnp.uniquewithoutsize=). It needs concrete values and cannot be traced inside jit. 6.17 - De Finetti's theorem (informal)
- If you would treat any number of groups as exchangeable, your beliefs can be written as “iid given some unknown parameters $\phi$, with a prior on $\phi$”. With $\phi=(\mu,\tau)$ that is the hierarchical model. 6.5
- Decision rule
- A function from posterior to action (ship, stop, continue), fixed in advance together with $\delta$, the probability threshold, $\varepsilon$ and a maximum sample size. How often it errs depends on how it is run, so simulate it. 6.4
- Default rank
- With
rank=NoneNumPyro uses $r=\text{round}(\sqrt d)$. Choose $r$ from the eigenvalues of the posterior correlation above the floor, not from “90% of variance”, and check that raising $r$ does not change the conclusions. 6.13 6.13 - Detailed balance
- $\pi_iP_{ij}=\pi_jP_{ji}$: the flow from $i$ to $j$ equals the flow back, which makes $\pi$ stationary. Metropolis is built to satisfy it; it is sufficient, not necessary. 6.9 6.9
- Device transfer
- Copying data between host and device (
jax.device_put,jax.device_get, or implicitly when NumPy arrays are passed to JAX). Put the design matrix and targets on the device once, keep the SVI state there, and read losses every $k$ steps. 6.17 - Diffuse (vague) prior
- A very wide prior such as $N(0, 1000^2)$ that tries to say “anything goes”. Nearly flat where the likelihood lives, but on transformed quantities it can be wildly informative. 6.2 6.2
- Diffusive behaviour (random-walk slowness)
- A random walk with step $\varepsilon$ travels about $\varepsilon\sqrt n$ after $n$ steps, so crossing a distance $D$ costs about $(D/\varepsilon)^2$ steps. Momentum makes it about $D/\varepsilon$. 6.10
- Dirichlet-Multinomial
- The categorical analogue of the Beta-Binomial: Dirichlet($\boldsymbol\alpha$) prior plus Multinomial counts $\mathbf c$ gives Dirichlet($\boldsymbol\alpha+\mathbf c$). Each single share is a Beta; whole-mix questions need joint draws because the shares sum to 1. 6.3
- Discounting (power prior)
- Multiplying a historical prior's pseudo-counts by a factor $a_0$ between 0 and 1: Beta($a_0\alpha, a_0\beta$) keeps the mean and is worth $a_0(\alpha+\beta)$. 6.2
- Divergence (divergent transition)
- An HMC/NUTS trajectory whose simulated energy error exploded (over 1000 in NumPyro by default) or became NaN, because the step was too large for the local curvature. It marks a region the sampler could not explore, so estimates are biased, not just noisy. The cause is curvature too high for the step: in the funnel's neck the curvature is $1/\tau^2$, and leapfrog is stable only if $\varepsilon\lt2\tau$. 6.7 6.10
- Double where (NaN-gradient trap)
- Because
jnp.wherecomputes both branches,jnp.where(x>0, jnp.log(x), 0.0)has a NaN gradient at $x=0$. Make the unsafe branch's input safe first:jnp.where(x>0, jnp.log(jnp.where(x>0, x, 1.0)), 0.0). 6.17 - Dual averaging and target_accept_prob
- The warmup rule that tunes the step size toward a target average acceptance,
target_accept_prob(default 0.8). A higher target (0.95) gives smaller steps: more accurate, slower trajectories. 6.10 - E-BFMI
- The estimated Bayesian fraction of missing information: an HMC energy diagnostic that compares successive changes of the energy with the spread of the energy itself. Low values mean the sampler explores energy levels slowly. Chapter 6.15 gives the rule of thumb of not below about 0.3 (together with tree depth and trace plots). 6.15
- Effective number of parameters
- How many independent numbers the pooled estimates really use, $\sum_g\partial\hat\theta_g/\partial\bar y_g$: from 1 (complete pooling) to $G$ (no pooling). The same idea as the degrees of freedom of ridge regression. 6.6
- Effective sample size (ESS)
- How many independent draws the correlated MCMC draws are worth, $mn/\hat\tau$. ArviZ splits it into bulk-ESS (centre: mean, median) and tail-ESS (interval endpoints). Rule of thumb: at least about 400 in total (100 per chain). 6.10 6.9
- Eight schools
- The standard small hierarchical example, with eight groups and noisy estimates. It shows pooled vs separate estimates (school A: 11.4 $\pm$ 8.3 under full Bayes) and the funnel-shaped posterior that makes such models hard to sample. 6.6 6.7
- ELBO comparison across guides
- For the same model and data, $\text{ELBO}(q_1)-\text{ELBO}(q_2)=KL(q_2\|p)-KL(q_1\|p)$, so a higher average ELBO (beyond seed-to-seed noise) means a guide closer in KL. A flat ELBO only says the optimizer is done. 6.15
- ELBO (evidence lower bound)
- $\text{ELBO}(q)=E_q[\log p(D,\theta)-\log q(\theta)]$, computable from the joint and draws from $q$. It satisfies $\log p(D)=\text{ELBO}(q)+KL(q\|p(\theta\mid D))$, so ELBO $\le\log p(D)$ with equality only when $q$ is the posterior, and maximizing it minimizes the KL. 6.12
- Empirical Bayes (type-II maximum likelihood)
- Estimate $\tau$ by maximizing the marginal likelihood $p(y\mid\tau)$ and plug it in. It drops the uncertainty about $\tau$, so intervals are too narrow with few groups, badly so when $\hat\tau$ sits at 0. 6.6
- Energy error $\Delta H$
- The change in $H$ over a simulated trajectory. Acceptance is $e^{-\Delta H}$ (0.905 at $\Delta H=0.1$, 0.135 at 2); it measures simulation accuracy, not exploration. 6.10
- Entropy $H[q]$
- A measure of how spread out $q$ is (for a Normal, $\frac12\log(2\pi e\sigma^2)$). In the ELBO it stops the guide from collapsing to a single point. 6.12
- Epsilon ($\epsilon$) in the relative rule
- A small constant that prevents division by zero. When $|ELBO_{best}|\gg\epsilon$ the rule is relative; near zero it becomes an absolute rule with threshold $\tau\epsilon$ nats. 6.14
- Equal-tailed interval (ETI)
- The credible interval cut at the $\gamma/2$ and $1-\gamma/2$ quantiles. Its ends transform with any increasing change of scale, which makes it safe to convert units afterwards. 6.4
- Evaluation frequency $k$ and smoothing
- The stopping statistic is computed every $k$ steps on a window mean (sd $\sigma/\sqrt k$, lag $(k-1)/2$) or an exponential moving average, or on a fixed-key multi-particle ELBO. Convert to Python floats only at evaluations to avoid waiting on the device every step. 6.14
- Evidence (marginal likelihood) $p(D)$
- $p(D) = \int p(D\mid\theta)\,p(\theta)\,d\theta$: the likelihood averaged over the prior, one number per model. It normalizes the posterior; it is not the likelihood at the best $\theta$. 6.1
- Exchangeability
- The joint distribution of the groups does not change when you reorder them: they are interchangeable before the data. It is not “identical” and not “independent”; it is the assumption that justifies $\theta_g\sim N(\mu,\tau^2)$. Groups that differ in a known way $x_g$ are exchangeable given $x_g$ (use $\theta_g\sim N(\mu+\beta x_g,\tau^2)$); control and treatment are never exchangeable groups. 6.5
- Expected loss (of a decision)
- The posterior average cost of choosing wrongly: $E[\max(\theta_A-\theta_B,0)\mid D]$ for choosing B, in the metric's own units so it can be turned into money. 6.4
- Factorizes (independent posteriors)
- In the two-variant Beta-Binomial model the joint posterior is a product of two Betas, so $\theta_A$ and $\theta_B$ are independent given the data; the gap $\theta_B-\theta_A$ is then not a Beta and needs draws. 6.3
- Fake-data (parameter-recovery) check
- Fix parameters, simulate a dataset from the model, fit the same model, compare the posterior with the truth, repeat. It tests the inference code and what the data can tell you, not whether real data follow the model. Drawing the truth from the prior each time is simulation-based calibration: a 90% interval then contains the truth in exactly 90% of runs. 6.5
- Family gap
- The part of SVI's error that remains even at the best member of the family: $KL(q^\star\|p)>0$ unless the posterior is in the family. More steps cannot remove it; only a richer family can. 6.15 6.15
- Flat prior and improper prior
- A flat prior has constant density over the allowed range (Beta(1, 1) on a rate). A constant density over an unbounded range does not integrate to 1: it is called improper and is acceptable only if the posterior still integrates to 1. 6.2
- Forward KL (mass-covering)
- $KL(p\|q)$ averages over $p$, so it punishes $q$ for missing any mass of $p$. It stretches over all modes (too wide between them); with a Normal $q$ it is moment matching (same mean and variance); reverse KL has no such shortcut and can have several local minima. 6.11
- Full Bayes (over the hyperparameters)
- Average the group estimates over the whole posterior of $\tau$ (and $\mu$). The variance gains $Var(E[\theta_j\mid\tau,y])$, the extra uncertainty from not knowing $\tau$. 6.6
- Full-rank guide (AutoMultivariateNormal)
- $N(\boldsymbol\mu,LL^\top)$ with a learned lower-triangular Cholesky factor $L$: $d+d(d+1)/2$ learned numbers (
auto_loc,auto_scale_tril). Captures every correlation; $O(d^2)$ memory and work per sample; $O(d^3)$ only to factorize a dense matrix. Still Gaussian, so full-rank is not exact. 6.13 - Functional state ($\text{state}_{t+1}=\text{update}(\text{state}_t,\text{data})$)
- All changing quantities flow through functions as explicit values, as in
svi_state, loss = svi.update(svi_state, data). Nothing is modified in place. 6.16 - Functional update
- “Returns a new value instead of changing the old one”. Because JAX arrays are immutable, keeping a reference to the best parameters is already a true snapshot. 6.16 6.16
- Funnel (Neal's funnel)
- The shape of a centered hierarchical posterior with weak data: the room for $\theta_g$ is proportional to $\tau$, so the posterior is a narrow neck at small $\log\tau$ and a wide mouth at large $\log\tau$. See it in a pairs plot of $\theta_g$ against $\log\tau$. 6.7
- Funnel under variational inference
- A Gaussian guide has the same width at every height, so in centered coordinates it cannot follow a funnel and silently understates the uncertainty in $\tau$; there is no divergence alarm. The form with the higher final ELBO has the smaller KL. 6.7
- Gamma-Poisson
- The count-rate conjugate pair: Gamma($a$, rate $b$) prior on $\lambda$ plus Poisson counts with exposures $t_i$ gives Gamma($a+\sum y_i,\ b+\sum t_i$); the predictive is a Negative Binomial. SciPy's gamma uses scale $=1/b$. 6.3
- Global scaler
- One standardization (one mean and one sd) applied to all groups, so group differences survive and a single prior scale means the same thing everywhere. Scaling each group by its own mean and sd would make every group average 0 and erase $\tau$. 6.5 6.2
- grad (reverse-mode autodiff)
jax.gradapplies the chain rule backwards through the recorded operations: exact up to rounding, scalar output, float inputs, cost a small multiple of one evaluation however many parameters. Every SVI step is avalue_and_gradof the ELBO loss. 6.16- Gradient check (finite differences)
- Compare autodiff with the central difference $[f(x+h)-f(x-h)]/2h$, whose error is about $h^2f'''/6+O(\varepsilon|f|/h)$. Use float64 and a moderate $h$ (about $10^{-5}$); the tiniest $h$ drowns in rounding error. 6.16
- Gradient noise and the noise ball
- With a constant learning rate the parameters never settle exactly; they jitter around the optimum in a noise ball whose size grows roughly in proportion to the learning rate (SGD on a quadratic: $\frac{\text{lr}\,\sigma_g^2}{a(2-\text{lr}\,a)}$). Fixes: a decaying schedule, more particles, averaging, clipping. 6.14
- Grid approximation
- Evaluate prior $\times$ likelihood at $G$ grid points, add up the area, and divide. Fine for 1–2 parameters; it needs $G^d$ points for $d$ parameters (100 points and 10 parameters is $10^{20}$), which is why real models use MCMC or VI. 6.1 6.9
- Group (local) vs global parameters
- Group parameters $\theta_g$ belong to one segment; global parameters ($\mu$, $\tau$, a shared noise scale $\sigma$) are shared by all. Counts with $G$ groups: complete pooling 2, separate $G+1$, hierarchical $G+3$. 6.5
- Guide (autoguide)
- The variational family in NumPyro. An autoguide flattens all latent sites into one unconstrained vector $\mathbf z\in\mathbb R^d$ and puts a Gaussian on it; the guide type is the structure of the covariance. The guide is an approximation to the posterior, not the posterior. 6.13
- Guide dependence
- SVI's answer depends on the model, the family, the starting point, learning rate, steps, particles and seed. With several modes a Gaussian guide silently picks one, while several NUTS chains disagree loudly (high $\hat R$). Record the guide type with every reported result. 6.15
- Guide parameter counts
- Learned numbers for latent dimension $d$ and rank $r$: mean-field $2d$, low-rank $d(r+2)$, full-rank $d+d(d+1)/2$. Training memory is roughly bytes $\times$ numbers $\times 4$ (parameters, Adam's two moments, a gradient). 6.13 6.13
- Half-Normal / half-$t$ prior on $\tau$
- The usual hyperprior for a spread: a Normal or Student-t folded onto $\tau\ge0$ (HalfNormal($s$) puts about 95% of its mass below $2s$). With few groups $\tau$ is poorly learned, so this prior matters. 6.5
- Hamiltonian Monte Carlo (HMC)
- MCMC that adds a random momentum $p$ and simulates a frictionless puck on the landscape $U(\theta)=-\log\tilde p(\theta)$, using the gradient to glide far in one iteration. Accept with $\min(1,e^{-\Delta H})$. Needs continuous parameters and a differentiable log density. 6.10 6.10
- Hamiltonian, potential and kinetic energy
- $H(\theta,p)=U(\theta)+K(p)$ with $U=-\log\tilde p$ (height) and $K=\frac12p^\top M^{-1}p$ (motion). Exact motion conserves $H$, is reversible and volume-preserving; momentum is auxiliary, not a model parameter. 6.10
- Hierarchical (multilevel) model
- A model in two levels: group values $\theta_g\sim N(\mu,\tau^2)$ drawn from a population, then data $y_{gi}\sim p(y\mid\theta_g)$ from each group, with priors on $\mu$ and $\tau$. Groups are related, not identical. 6.5 6.2
- Hierarchical prior and hyperprior
- A prior with levels: group parameters $\theta_g\sim N(\mu,\tau^2)$ whose centre $\mu$ and spread $\tau$ (hyperparameters) get their own priors (hyperpriors). Because $\tau$ is learned, the data set the amount of shrinkage. 6.2 6.5
- Highest-density interval (HDI, HPDI)
- The region where the posterior density is above a water level, chosen to hold the stated mass. It is the shortest such interval for a one-humped posterior, can split into pieces for a multi-humped one, and depends on the scale. 6.4
- Hybrid workflow (SVI and NUTS together)
- Validate SVI against NUTS on a smaller problem, pick the cheapest guide that passes, run it at scale, and re-validate when the model changes. The guide can also help NUTS: start at its median (
init_to_value), precondition with its covariance, or useNeuTraReparamwith an AutoContinuous guide. 6.15 - Hyperparameter
- In a hierarchical model, an unknown quantity such as $\mu$ or $\tau$ that describes the population of groups and is learned from the data. Not a training setting like a learning rate; the fixed constants inside the hyperpriors are what you choose. 6.5 6.5
- Identifiability
- A model is identifiable if different parameter values always give different data distributions. Otherwise the likelihood is perfectly flat in some direction and the posterior there equals the prior; in $y\sim N(a+b,1)$ only $a+b$ is learned. 6.8
- Importance sampling
- Draw from an easy $q$ and re-weight by $w_s=\tilde p(\theta^{(s)})/q(\theta^{(s)})$; dividing by $\sum w_s$ cancels $p(D)$. The proposal must be wider than the target (heavier tails), never narrower; if $q$ is a poor match or the dimension is high, a few weights dominate and the weight ESS $1/\sum\bar w_s^2$ collapses (degeneracy). 6.9
- Informative prior
- A prior carrying substantial, specific outside information (past experiments, physics, business limits). It visibly affects the posterior unless the data are large. 6.2 6.2
- Intra-class correlation (ICC)
- $\tau^2/(\tau^2+\sigma^2)$: the share of an observation's total variance that is real between-group difference. The noise share of a group average is $(\sigma^2/n_g)/(\tau^2+\sigma^2/n_g)$. 6.5
- Intractable posterior
- The normalizing integral $p(D)=\int\tilde p(\theta)\,d\theta$ and every posterior average have no closed form, and brute-force integration is far too slow. Bayes' theorem still applies; the obstacle is computational. 6.9
- Jacobian (change of ruler)
- The stretching factor $|d\theta/d\eta|$ that densities pick up when you change parameter. It is why a prior flat on $p$ is a logistic hump on the log-odds, and why the MAP moves under reparameterization. 6.2
- JAX array (immutable)
- A fixed-shape, fixed-dtype block of numbers on a device that no operation changes in place. Update with
y = x.at[i].set(v)(also.add,.multiply), which returns a new array. Default float32; out-of-range indices are clamped (reads) or dropped (writes), never an error. 6.16 - JAX transformation
- A function that takes a function and returns a new function:
jit(compile),grad(derivative of a scalar output),vmap(batch),value_and_grad. They compose, e.g.jax.jit(jax.vmap(jax.grad(f))), and need pure functions. 6.16 - Jaxpr and XLA
- The jaxpr is the recorded program, a list of primitive operations on typed variables (
jax.make_jaxpr(f)(*args)shows it: operations and shapes, no values). XLA is the compiler that fuses and plans it into an executable for the device. 6.17 - Jeffreys prior
- A classical “non-informative” prior built to give the same inferences under any parameterization; for a Bernoulli rate it is Beta(½, ½). 6.2
- JIT compilation (jit)
- The pipeline Python $\to$ tracing with abstract values (shape, dtype) $\to$ jaxpr $\to$ StableHLO $\to$ XLA $\to$ executable, run on the first call for each new input signature and then cached. Python code runs only during tracing; the compiled program contains only array operations. 6.17
- Kernel (of a density)
- The part of a density that depends on $\theta$, after dropping constant factors. Conjugacy works because likelihood and prior kernels are built from the same blocks (powers of $\theta$ and $1-\theta$). 6.3 6.3
- KL divergence
- $KL(q\|p)=E_q[\log q(\theta)-\log p(\theta)]\ge0$, zero only when $q=p$. It is not symmetric (so a divergence, not a distance), is infinite if $q$ has mass where $p$ has none, and is unchanged by relabelling $\theta$. For two Normals: $\log\frac{\sigma_2}{\sigma_1}+\frac{\sigma_1^2+(\mu_1-\mu_2)^2}{2\sigma_2^2}-\frac12$. 6.11
- Laplace approximation
- Fit a Gaussian at the posterior mode: $N(\hat\theta, A^{-1})$ with $A=-\nabla^2\log\tilde p(\hat\theta)$ (in 1-D, sd $=1/\sqrt{-\ell''(\hat\theta)}$). Fastest and crude: it depends on the scale (use log or log-odds) and is blind to skew, bounds and second peaks. 6.9
- Laplace's rule of succession
- The posterior mean $(k+1)/(n+2)$ under a flat Beta(1, 1) prior. Its mode is $k/n$, not its mean. 6.3
- Latent dimension
- The number of unobserved scalar unknowns in a model, with plates multiplied out. For segments inside plates it grows like segments $\times$ variants, and it decides how expensive a full-rank guide becomes. 6.5 6.13 6.13
- Leapfrog integrator
- The numerical scheme HMC uses: a half kick of momentum, a drift of position, a half kick. $L$ steps cost $L$ gradients. It is reversible and volume-preserving, with energy error $O(\varepsilon^2)$ that explodes (a divergence) once $\varepsilon$ exceeds about 2 times the narrowest sd. 6.10
- Learning rate and optimizer (Adam, ClippedAdam)
- The learning rate is the step size: too small is slow, too large is jumpy or diverges. Adam adapts a step per parameter; ClippedAdam also clips each gradient value to ±
clip_norm(default 10) against rare huge gradients. 6.14 6.14 - Likelihood $p(D\mid\theta)$
- The data model read as a function of $\theta$ with the data held fixed: a score for how well each $\theta$ explains what we saw. It is not a probability distribution over $\theta$; only ratios matter. 6.1
- Location–scale family
- A distribution that is a standard shape shifted by a location and stretched by a scale (Normal, Student-t, Laplace, Cauchy), so $\theta=\mu+\tau z$ with $z$ from the standard version. This is what makes non-centering possible. 6.7 6.7
- LocScaleReparam and partial centering
reparam(model, config={"theta": LocScaleReparam(centered=0)})rewrites a location–scale site; $c=0$ is fully non-centered and $c=1$ unchanged.centered=Noneis a learnable 0.5 that SVI can learn but NUTS leaves at 0.5. 6.7- Logit-Normal and Beta hierarchies (rates)
- Two ways to put a population on conversion rates: $\text{logit}\,p_g\sim N(\mu,\tau^2)$ (where logistic($\mu$) is the median rate) or $p_g\sim\text{Beta}(\kappa\phi,\kappa(1-\phi))$ with mean $\phi$ and worth $\kappa$. The meaning of $\tau$ or $\kappa$ differs, so check which one your code builds. In the Beta version each segment's pooled rate is $\hat p_g=\frac{k_g+\kappa\phi}{n_g+\kappa}$: a starter pack of $\kappa$ visitors at rate $\phi$. 6.5 6.6
- Loss function and posterior expected loss
- A loss $L(\theta,d)$ says how costly it is to report $d$ when the truth is $\theta$; averaging it over the posterior gives the expected loss, and the best report minimizes it. 6.4 6.4
- Loss returned by svi.update ($-$ELBO)
svi.updatereturns(state, loss)with loss $=-\widehat{\text{ELBO}}$ for that step's draws. Improvement means the loss goes down; it is meaningless on its own scale (it contains $-\log p(D)$), so judge a smoothed trend. 6.12 6.14- Low-rank guide (AutoLowRankMultivariateNormal)
- $N(\boldsymbol\mu,WW^\top+D)$ with a $d\times r$ factor $W$ and diagonal $D$: $d(r+2)$ learned numbers (
auto_loc,auto_cov_factor,auto_scale). Sampling costs $O(dr)$. It keeps $r$ shared directions and treats the rest as independent, so it misses many separate pairwise ties and long chains. 6.13 - MAP (maximum a posteriori)
- The peak (mode) of the posterior: the single $\theta$ with the highest posterior density, equal to the MLE plus a penalty $-\log p(\theta)$. It is one point of the posterior, not the posterior, and it moves under a change of scale. 6.4 6.9
- Markov chain and stationary distribution
- A sequence of states in which the next depends only on the current one. A stationary distribution satisfies $\pi=\pi P$ (the chain stays there in distribution); an irreducible, aperiodic chain forgets its start and spends a share $\pi$ of its time in each state (ergodic theorem). 6.9
- Masked sum, mean and log-likelihood
- Fixed-shape replacements for a subset with mask $m_i\in\{0,1\}$: $\sum_i m_ix_i$, $\sum_i m_ix_i/\max(\sum_i m_i,1)$ and $\sum_i m_i\log p(y_i\mid\theta)$. Masked-out entries are still computed, so they must be finite ($0\times\infty=$ NaN). In NumPyro,
.mask(m)gives masked observations zero log-probability;obs_mask=instead treats them as missing values to impute. 6.17 - Mass matrix
- The matrix $M$ in the kinetic energy; $M^{-1}$ is adapted in warmup to the posterior variances (diagonal by default, full with
dense_mass=True). It acts as a rescaling of the coordinates. 6.10 6.10 - MCMC diagnostics checklist (print_summary)
- Stop at the first failure: divergences $=0$, then $\hat R\le1.01$, then ESS $\ge$ about 400, then traces and tree depth, then MCSE against the precision you need, then model checks.
print_summaryshows mean, std, median, 90% HPDI,n_eff,r_hat, divergences. 6.10 - MCMC (Markov chain Monte Carlo)
- Build a Markov chain whose stationary distribution is the posterior, using only ratios of $\tilde p$, so the evidence never appears. It returns correlated draws and is asymptotically exact, with Monte Carlo error. Say “asymptotically exact”, not “exact”. 6.9 6.9
- Mean-field family
- $q(\theta)=\prod_i q_i(\theta_i)$: independent factors (NumPyro's
AutoNormal). For a Normal posterior the reverse-KL optimum has exact means and the conditional variances $1/\Lambda_{ii}\le\Sigma_{ii}$; in 2-D the sd shrinks by $\sqrt{1-\rho^2}$ (0.44 at $\rho=0.9$). 6.11 - Mean-field guide (AutoNormal, AutoDiagonalNormal)
- A diagonal Gaussian, $\Sigma_q=\text{diag}(s_i^2)$, with $2d$ learned numbers. For a Gaussian posterior the means are exact and each variance is $1/\Lambda_{ii}$: too narrow along correlated directions, so sums along those directions look far too certain. 6.13
- Metropolis algorithm (random-walk)
- Propose $\theta'=\theta+\varepsilon z$, accept with probability $\min(1,\tilde p(\theta')/\tilde p(\theta))$ (computed in logs), otherwise record the current value again. Rejected steps are part of the output; $p(D)$ cancels. Metropolis–Hastings allows a non-symmetric proposal. 6.9
- Minibatch scaling ($N/B$)
- With a random minibatch of $B$ of $N$ points, multiply the per-point log-likelihood by $N/B$ (not the prior):
numpyro.plate("data", N, subsample_size=B)does it for you. Forgetting it acts as if you had only $B$ points. 6.12 - Minimum and maximum steps
- No stop is allowed before step $t_{min}$ (it protects the early transient); a hard maximum guarantees the run ends and should trigger a warning to investigate. 6.14
- Mixing
- How quickly an MCMC walker moves around the whole posterior. Good mixing means draws almost as useful as independent ones; a trace plot looks like a fuzzy caterpillar. 6.9 6.9
- Model, guide and latent variables
- The model is the joint $p(D,\theta)$; the guide is the variational family $q_\phi$ over the model's latent (unobserved) sample sites. Autoguides build it automatically. 6.12 6.13
- Monte Carlo ELBO and its noise
- Average $\log p(D,\theta_s)-\log q(\theta_s)$ over $S$ particles drawn from $q$: unbiased, noise falling like $1/\sqrt S$, cost growing like $S$. NumPyro's
Trace_ELBO(num_particles=S)defaults to $S=1$, so one step's loss is mostly noise. 6.12 - Monte Carlo error of a probability
- A probability estimated from $S$ independent draws has error $\sqrt{P(1-P)/S}\le 0.5/\sqrt S$ (0.008 for $S=4000$). For correlated MCMC draws use the effective sample size instead of $S$. 6.4 6.10
- Monte Carlo standard error (MCSE)
- The error of a posterior mean from the finite run: sd$/\sqrt{\text{ESS}}$; for a probability $\sqrt{q(1-q)/\text{ESS}}$. It is not the posterior sd, and sd$/\sqrt n$ must never be used for MCMC draws. 6.10
- Multiple chains
- Independent runs from dispersed starting points, e.g.
MCMC(..., num_chains=4)withget_samples(group_by_chain=True)giving shape (chains, draws, …). Pooling disagreeing chains hides the problem. 6.10 - No pooling
- Each group estimated alone, $\hat\theta_g=\bar y_g$ (the hierarchical model with $\tau\to\infty$). Unbiased but very noisy for small groups; it is not assumption-free. 6.6 6.6
- Non-centered parameterization
- Sample $z_g\sim N(0,1)$ and compute $\theta_g=\mu+\tau z_g$. Same model and same posterior; only the coordinates the algorithm explores change. Best when data per group are weak ($SE_g\gg\tau$); with strong data per group it becomes a thin curved ridge and centered is better (rule of thumb: $w_g=\tau^2/(\tau^2+SE_g^2)$ small, go non-centered). 6.7 6.7 6.7
- Normal-Normal (known $\sigma$)
- Prior $N(m_0, s_0^2)$ on a mean with known noise $\sigma$: precisions add, $1/s_n^2 = 1/s_0^2 + n/\sigma^2$, and the posterior mean is a precision-weighted average. The predictive is $N(m_n, \sigma^2 + s_n^2)$. 6.3
- NumPyro handlers
seed,trace,substitute,condition,mask: tools that turn a NumPyro model, a Python function, into a pure log-density or loss function that JAX can differentiate and compile. 6.16- NUTS (No-U-Turn Sampler)
- HMC that doubles the trajectory in a random direction until it starts to turn back (a U-turn), at most $2^{10}-1=1023$ steps by default (
max_tree_depth=10; hitting it often signals a hard geometry), then samples the next point from the trajectory in proportion to $e^{-H}$. It still has settings (warmup, chains,target_accept_prob) and still depends on the parameterization. 6.10 - NUTS reference
- A NUTS run on a problem small enough to afford (a subset, a few groups, a shorter history, fewer changepoints) with clean diagnostics, used to check SVI. Call it the reference, not the truth, and ignore SVI–NUTS gaps smaller than about 2 NUTS MCSE. 6.15 6.15
- Optimal scaling (0.234 and 0.44)
- For Gaussian-like targets the best random-walk step is about $2.38\,\sigma/\sqrt d$, with acceptance about 0.234 in high dimension and 0.44 in one. Starting points only: judge a sampler by ESS, not by acceptance. 6.9
- Optimize vs sample
- SVI optimizes a guide to maximize the ELBO and returns a formula you can draw from cheaply; NUTS builds a Markov chain and returns draws. Both need only the unnormalized log joint and its gradient, and neither needs the evidence. Draws from a guide describe the guide, not the posterior. 6.15
- Other autoguides (AutoDelta, AutoLaplaceApproximation, flows, AutoGuideList)
- AutoDelta is a point estimate (MAP); AutoLaplaceApproximation is a Gaussian at the mode from the Hessian; AutoIAFNormal and AutoBNAFNormal are normalizing flows that can be skewed or curved; AutoGuideList combines different guides for different blocks. 6.13
- Padding and masking (buckets)
- Pad arrays to a bucket size (64, 128, 256, …) and carry a boolean mask marking the real entries, writing every reduction with the mask, e.g. mean $=\sum_i m_ix_i/\sum_i m_i$. A few compiles then serve many sizes at under 50% wasted work. 6.17 6.17
- Paired draws (same draw index)
- Comparing $\theta_B^{(s)}$ with $\theta_A^{(s)}$ from the same joint draw $s$. Required whenever the posterior is correlated, as in hierarchical models; never shuffle or draw them separately. 6.4
- Parameter uncertainty vs observation noise
- The two parts of a forecast's variance: $Var(\tilde y\mid D) = E[Var(\tilde y\mid\theta)] + Var(E[\tilde y\mid\theta])$. Observation noise is the floor; parameter uncertainty shrinks with data and, for a trend, grows with the horizon. 6.1
- Pareto $\hat k$ (importance-sampling check)
- A diagnostic of how far $q$ is from the posterior, computed from importance weights $p/q$; a rule of thumb from the literature is that $\hat k\lt0.7$ is usable. 6.15
- Partial pooling
- Each group's estimate is a weighted average of its own mean and the population mean, $\hat\theta_g=w_g\bar y_g+(1-w_g)\mu$: the posterior mean of a hierarchical model. Small groups borrow more, large groups keep their own. 6.6
- Patience ($P$)
- The number of consecutive evaluations without a meaningful improvement that triggers a stop, counted in evaluations (waiting $P\cdot k$ steps). Needed because the ELBO is noisy: one non-improving evaluation is weak evidence, several in a row is strong. False stops fall roughly like $b^P$; cost grows linearly. 6.14
- Plate diagram (model graph)
- A drawing of the generative story: circles are random quantities, shaded circles observed data, squares constants, arrows “is drawn using”, and a plate means “repeat once per index”. The joint density has one factor per circle given its parents. 6.5
- Plug-in prediction
- Predicting with one point estimate $\hat\theta$ instead of the whole posterior. It ignores parameter uncertainty and is always too narrow. 6.1
- Posterior $p(\theta\mid D)$
- What you believe about $\theta$ after seeing the data: $p(\theta\mid D) = p(D\mid\theta)\,p(\theta)/p(D)$. It is a compromise between prior and likelihood, closer to whichever is sharper. 6.1
- Posterior contraction
- $1-Var_{\text{post}}/Var_{\text{prior}}$: near 1 means the parameter was learned from the data, near 0 means it is still the prior. A quick identifiability diagnostic. 6.8 6.8
- Posterior draws
- Random samples $\theta^{(1)},\dots,\theta^{(S)}$ from the posterior (from MCMC or from a fitted guide). Every summary, interval and probability is then a plain average or quantile of them. 6.4 6.4
- Posterior expectation (Monte Carlo estimate)
- Almost every question is an average over the posterior, $E[f(\theta)\mid D]\approx\frac1S\sum_s f(\theta^{(s)})$: a mean, a probability (share of draws), an interval (sorted draws). With independent draws the error is sd$/\sqrt S$, whatever the number of parameters. 6.9
- Posterior mean, median and mode
- Three one-number summaries, each best under a different loss: the mean for squared error, the median for absolute error, the mode (MAP) for all-or-nothing. For a right-skewed posterior mode $\lt$ median $\lt$ mean; the median alone is unchanged by increasing transformations. 6.4 6.4
- Posterior predictive check (PPC)
- Simulate replicated datasets $y^{\text{rep}}$ from the fitted model (a fresh posterior draw for each, same size and design as the real data) and compare them with the real data. A feature of the real data that stands out is something the model cannot produce; passing does not prove the model correct. 6.8
- Posterior predictive distribution
- The distribution of a new observation, $p(\tilde y\mid D) = \int p(\tilde y\mid\theta)\,p(\theta\mid D)\,d\theta$. Simulate it by drawing $\theta$ from the posterior, then $\tilde y$ from the likelihood. 6.1
- Posterior predictive p-value
- $P(T(y^{\text{rep}})\ge T(y)\mid y)$, estimated as the share of replicates at least as extreme as the real data. Near 0 or 1 flags misfit; it is conservative because the data are used twice, and it is not the probability the model is true. 6.8
- Practical threshold $\delta$
- The smallest gain worth shipping, fixed in advance and in the units of the parameter. $P(\theta_B-\theta_A\gt\delta\mid D)$ asks “better by enough to matter?” and can be near 0 while $P(\theta_B\gt\theta_A)$ is near 1. 6.4
- Precision
- The reciprocal of a variance, $1/\sigma^2$. In Normal updating precisions add, which is why more or sharper information always makes the posterior sharper. 6.3 6.6
- Precision-weighted average
- Combining two sources by weighting each with its precision (1/variance): $w_g=\frac{n_g/\sigma^2}{n_g/\sigma^2+1/\tau^2}$. Precisions add, so the posterior is sharper than either source; the population acts like $\sigma^2/\tau^2$ extra observations at $\mu$. 6.6
- Premature stopping
- Stopping while the model was still improving, because noise or a plateau longer than $P\cdot k$ steps hid the progress. It costs a worse guide (lower ELBO, and spreads or correlations that have not finished settling) and is invisible in the loss log unless you look. 6.14 6.14
- Prior $p(\theta)$
- What you believe about a parameter before seeing the data, written as a probability distribution (Beta(2, 18) means “about 10%, worth 20 visitors”). A value with prior density 0 keeps density 0 whatever the data say. 6.1
- Prior–data conflict
- The data land far in the prior's tail, so the posterior is a compromise matching neither. Check for it and report it; do not hide it. 6.2 6.8
- Prior predictive check
- Simulate datasets from the prior alone ($\theta\sim p(\theta)$, then $y\sim p(y\mid\theta)$) and ask whether they look like plausible worlds. It judges the whole model on the data scale, uses domain knowledge rather than the observed data, and is run before the fit. 6.2 6.2
- Prior sensitivity analysis
- Refit under a small set of reasonable alternative priors chosen beforehand (scales 2–3 times wider and narrower, heavier tails, a different centre) and report the range of the key outputs and whether the decision flips. Never pick the prior after seeing which answer it gives. 6.8
- Prior strength (prior worth)
- How many observations a prior is worth: $\alpha+\beta$ for a Beta($\alpha,\beta$), $\sigma^2/\tau^2$ for an $N(\mu_0,\tau^2)$ on a mean. Compare it with the data each parameter actually sees. 6.2 6.2 6.3
- PRNG key
- A small array that fully determines a JAX random draw. Same key and same call give the same numbers;
jax.random.split(key, n)makes independent keys; never reuse a key for two things that should be independent, and treat a split parent as used up. 6.16 - Probability that B beats A, $P(\theta_B \gt \theta_A \mid D)$
- The share of paired posterior draws with B ahead: the part of the $(\theta_A,\theta_B)$ cloud above the diagonal. It measures direction, not size, and is not $1-p$. 6.4
- Proportional ($\propto$) and the unnormalized posterior
- “Equal up to a constant that does not depend on $\theta$”. Posterior $\propto$ likelihood $\times$ prior; MCMC and VI only ever need this unnormalized log posterior $\log p(D\mid\theta) + \log p(\theta)$. 6.1
- Pseudo-counts
- Prior hyperparameters read as imaginary data: for the posterior mean, Beta($\alpha,\beta$) acts like $\alpha$ prior successes and $\beta$ prior failures, worth $\alpha+\beta$ pretend visitors (for the posterior mode it acts like $\alpha-1$ and $\beta-1$). 6.3
- Pure function
- Output depends only on explicit inputs and there are no side effects (no printing, global changes, hidden random state). Impure code often runs but gives quietly frozen results: globals become constants,
printfires once at trace time,np.randomis frozen. 6.16 - Pytree
- Any nesting of tuples, lists and dicts whose leaves are arrays: parameters, optimizer states and posterior-sample dictionaries.
jax.tree.map(f, tree)applies a function to every leaf. 6.16 - R-hat ($\hat R$, split R-hat)
- Split each chain in half and compare between-sequence variance $B$ with within-sequence variance $W$: $\hat R=\sqrt{\hat V/W}$. Rule of thumb: $\le 1.01$. Necessary, not sufficient, because chains can share the same blind spot. 6.10
- Relative ELBO improvement (rel)
- $\text{rel}_t=\dfrac{ELBO_t-ELBO_{best}}{|ELBO_{best}|+\epsilon}$, identical in the loss convention; divide by the absolute value because losses can be negative. A scale-free stopping statistic. 6.14 6.14
- Relative lift
- $L=\theta_B/\theta_A-1$, computed per posterior draw. It is right-skewed, so report its median and ETI next to the baseline; $E[\theta_B/\theta_A]\ne E[\theta_B]/E[\theta_A]$. 6.4
- Reparameterization trick
- Write the draw as a smooth function of the guide's parameters and parameter-free noise, $\theta=\mu+\sigma\varepsilon$ with $\varepsilon\sim N(0,1)$, so gradients flow through the draw ($\partial\theta/\partial\mu=1$, $\partial\theta/\partial\sigma=\varepsilon$). Low-noise; needs continuous (reparameterizable) latents. 6.12
- Reverse KL (mode-seeking, zero-forcing)
- $KL(q\|p)$ averages over $q$, so it punishes $q$ for putting mass where $p$ is thin. The fit stays inside a high-density region, can lock onto one mode, and tends to be too narrow. This is the direction VI minimizes. 6.11
- ROPE (region of practical equivalence)
- The interval $[-\delta,\delta]$ of differences too small to matter. $P(|\theta_B-\theta_A|\lt\delta\mid D)$ is the probability that the variants are practically the same. 6.4
- scan (loop with a carried state)
carry, ys = lax.scan(f, init, xs)withf(carry, x) -> (new_carry, y): a recursion with a fixed-shape carry, traced once and differentiable. Python loops unroll and compile slowly; scan runs in order, so vmap it to run many series. 6.16- Score-function estimator (REINFORCE)
- $E_q[f(\theta)\nabla_\phi\log q_\phi(\theta)]$: unbiased and valid even for discrete $\theta$, but it uses only the values of $f$, not its slope, and is much noisier. A baseline reduces the noise. 6.12
- Sensitivity curve
- A reported quantity plotted against a prior hyperparameter on a log axis. Steep means the prior decides; flat means the data decide. Typically steep for small data, scale parameters ($b$, $\tau$, holiday scales), dispersion/tail parameters and sparse effects. 6.8
- Sequential updating
- Today's posterior becomes tomorrow's prior: $p(\theta\mid D_1, D_2) \propto p(D_2\mid\theta)\,p(\theta\mid D_1)$. The result equals one big update, in any order, if $\theta$ stays fixed and the batches are independent given $\theta$. 6.1
- Shrinkage and the shrinkage factor
- The movement of an estimate from its raw average toward the centre: $\bar y_g-\hat\theta_g=B_g(\bar y_g-\mu)$ with $B_g=1-w_g=\frac{\sigma^2/n_g}{\sigma^2/n_g+\tau^2}$, exactly the noise share of the raw average. 6.6 6.6
- Shrinkage (regularizing) prior
- A prior centred on a reference value (usually 0 or a group mean) whose negative log acts as a penalty: Normal $\leftrightarrow$ ridge, Laplace $\leftrightarrow$ lasso, at the MAP. Its scale is the flexibility knob. 6.2
- Signed improvement
- With loss $\mathcal L=-\text{ELBO}$: $\Delta_t=ELBO_t-ELBO_{best}=\mathcal L_{best}-\mathcal L_t$, positive means better. Keep the sign for the decision; an absolute value would count a worsening as progress. 6.14
- Size-based guide selection rule
- Use full-rank while $d+d(d+1)/2$ fits a budget; otherwise use low-rank with $r$ from the spectrum (or $\text{round}(\sqrt d)$), capped so $d(r+2)\le B$. Any threshold is a design choice to be validated on a reduced model against full-rank or NUTS, not a law. 6.13
- Skeptical prior
- A prior on an effect centred at “no effect”, so a win has to be driven by this test's data. Informative priors on a baseline are common; priors on the effect being tested should be skeptical. 6.2
- Static arguments
- Arguments marked
static_argnumsorstatic_argnamesare treated as Python constants during tracing, so they may decide shapes and control flow, but each new value compiles a new program and they must be hashable. Fourier order and changepoint count are natural examples. 6.17 - Static shapes and recompilation
- The jit cache key contains the function, the arguments' pytree structure, shapes and dtypes, and static-argument values. A new key is retraced and recompiled; changed array values never recompile. Shapes must come from shapes, never from traced values. 6.17
- Stein's paradox (James–Stein)
- For 3 or more Normal means with known noise, a data-chosen shrinkage estimator has lower total expected squared error than the raw averages for every set of true means. The hierarchical model is its Bayesian version. 6.6
- Step size $\varepsilon$ (and why one cannot fit a funnel)
- The size of each simulated move in MCMC. The ideal step is proportional to the local width of the posterior; in a funnel that width changes with height, but NUTS adapts one global $\varepsilon$, so the neck is under-visited. 6.7 6.10
- Stochastic variational inference (SVI)
- Stochastic gradient ascent on the ELBO: draw particles from the guide, estimate the ELBO, differentiate by autodiff, take an Adam step, repeat. Model $=p(D,\theta)$, guide $=q_\phi$; only the guide's parameters are trained. 6.12
- Sufficient statistics
- The summary of the data that the posterior depends on: $(k, n)$ for the Beta-Binomial, counts for the Dirichlet-Multinomial. One user at a time, in batches or all at once gives the same posterior. 6.3
- Sum-to-zero constraint
- Forcing category effects to add to 0 (in NumPyro
s = s_raw - s_raw.mean()) so they cannot compete with the intercept. Predictions are unchanged but the level is identified. 6.8 - Support (of a parameter)
- The set of values a parameter is allowed to take (for a rate, $[0,1]$; for a scale, positive). A prior or a guide must put mass only inside it, which is why automatic guides work on an unconstrained scale (log, logit). 6.1 6.11 6.13
- SVI error vs NUTS error
- SVI: a fixed bias from the family gap plus optimization error, which more steps do not remove. Healthy NUTS: a vanishing bias plus Monte Carlo error sd$/\sqrt{\text{ESS}}$ that you can report. Compare guides for SVI; report MCSE for NUTS. 6.15
- SVI training loop
svi.init(key, *data)once, then repeatedstate, loss = svi.update(state, *data)(noisy ELBO, gradient, optimizer step), readingsvi.get_params(state).svi.runhas no early stopping, which is why a custom loop exists. 6.14- Test statistic (discrepancy) $T(y)$
- Any number computed from a dataset, such as the mean, variance, maximum, zero count, weekday gap or lag-1 autocorrelation of residuals. Choose a few tied to your decisions and not fitted directly by a parameter (a free mean always passes). 6.8 6.8
- Thinning
- Keeping every $m$-th draw. For a positively autocorrelated chain it never improves Monte Carlo precision (ESS 526 falls to 483 when thinning by 10 at $\rho=0.9$); it is justified only for storage or when each kept draw is expensive to process. 6.9
- Tolerance ($\tau$)
- The size of relative gain that counts as meaningful (rule of thumb $10^{-4}$ to $10^{-5}$ per evaluation). Translate it into nats, $\tau|ELBO|$, and compare it with the noise of one evaluation: far below the noise means noise drives the decision. 6.14
- Trace plot
- A parameter's value against iteration, one line per chain. Good: flat, fuzzy, overlapping chains. Bad: wander (step too small), stairs (too big), initial slope (burn-in), missing levels (modes), sticky stretches (funnel). It can show problems, never prove there are none. 6.9
- Tracer and tracing
- While recording a function, JAX passes placeholders (tracers) carrying only shape and dtype, no numbers. The recorded program is replayed on later calls. A NumPyro model's Python body therefore runs at trace time, not on every step. 6.16 6.16
- Two decisions: checkpoint vs patience
- “Is it the best so far?” (checkpoint) and “was it a meaningful gain?” (reset patience) are separate. A tiny new record is stored but does not reset the patience counter. 6.14
- Two readings of the ELBO
- Fit plus entropy, $E_q[\log p(D,\theta)]+H[q]$, and expected log-likelihood minus distance from the prior, $E_q[\log p(D\mid\theta)]-KL(q\|p(\theta))$. The ELBO contains the KL to the prior; its gap to $\log p(D)$ is the KL to the posterior. Two different KLs. 6.12
- Unconstrained latent vector and $d$
- $d$ counts unconstrained scalars, not sample statements: a site of shape (25,) adds 25, a $K$-category simplex adds $K-1$, a scale adds 1. In the original units a positive parameter's guide marginal is log-normal, skewed to the right. 6.13 6.13
- Unconstrained space (guide scale)
- Automatic guides fit Normals on an unconstrained scale (log for positive parameters, logit for probabilities) and map back, because $q$ must put no mass where $p=0$. A Normal on $\log\sigma$ is a log-normal on $\sigma$, so $e^{\text{loc}}$ is the median, not the mean. 6.11
- Under-dispersion
- The approximation is narrower than the posterior, so its “95%” intervals cover less than 95% (Student-t $\nu=3$: guide $\sigma$ 1.26 vs true sd 1.73, coverage 91%). A strong tendency of reverse KL with a too-simple family, not a theorem; check against NUTS or with coverage on held-out data. 6.11
- Unnormalized posterior $\tilde p(\theta)$
- Likelihood $\times$ prior, $\tilde p(\theta)=p(D\mid\theta)p(\theta)$: cheap to evaluate (and differentiate) at any $\theta$. MCMC and VI need only this, never $p(D)$. 6.9 6.1
- Variational family and variational parameters
- The family $\mathcal Q=\{q_\phi\}$ is the set of distributions we allow (all Normals, independent Normals, full-covariance Normals); the variational parameters $\phi$ pick one member (e.g. $\mu,\sigma$). They are the knobs of the approximation, not the model's parameters $\theta$; in NumPyro the guide is the family. 6.11
- Variational inference (VI)
- Turn inference into optimization: choose a family $q_\phi$ and tune $\phi$ to minimize $KL(q_\phi\,\|\,p(\theta\mid D))$, then use $q_{\phi^\star}$ in place of the posterior. In practice we maximize the equivalent ELBO. 6.11
- vmap (batching without loops)
- Write the function for one case and run it for many:
jax.vmap(f, in_axes=(0, None))maps over axis 0 and shares the other argument. It rewrites each operation into its batched version (not a loop); the whole batch is materialized in memory. 6.16 - Warmup (adaptation)
- The first, discarded part of a NUTS run during which NumPyro tunes the step size (target average acceptance 0.8 by default) and a scale per coordinate, then freezes them. 6.7 6.10
- Weak identification
- Technically identifiable, but the likelihood barely changes along some direction. Informative outside-data priors (last year's holiday lift, a pilot test) are legitimate identification aids if their source is documented. Symptoms: posterior correlations near $\pm1$, little contraction, slow high-autocorrelation chains, seed-dependent SVI, results that move with a prior scale. 6.8
- Weakly informative prior
- A prior deliberately wider than your real knowledge that still puts little mass on absurd values, so the data decide. Judge it in business units, on a known scale: $N(0, 0.5^2)$ on the log-odds lift puts B between 4% and 23% when A is 10%. 6.2
- When to use SVI or NUTS
- Choose by size ($N$, $d$), stakes (tails, correlations), shape (Gaussian, correlated, funnel, several modes) and schedule (once, daily, real-time). NUTS when affordable and details matter; SVI for scale, speed and many refits; slow NUTS usually signals bad geometry that also hurts SVI, so reparameterize first. 6.15
- Why shrinkage lowers error
- A raw average is unbiased but noisy; pulling it toward the centre adds a little bias and removes more variance. Average MSE falls from $s^2$ to $w\,s^2$; only groups far from the centre (beyond $\sqrt{2\tau^2+s^2}$) lose. 6.6
- Why VI uses the reverse KL, $KL(q\|p)$
- $KL(q\|p(\theta\mid D))=-\text{ELBO}+\log p(D)$ and $\log p(D)$ does not depend on $q$, so we need only the joint $p(\theta,D)$ and draws from our own $q$, never draws from the unknown posterior. 6.11
- Winner's curse
- The largest of many noisy estimates is on average too high, and the smallest, noisiest segments win most often. Rank pooled estimates, not raw rates. 6.6
Formula and code sheet
The key formulas of every chapter on one page, each with what it means in words and the trap that usually goes with it. Then the special tables: conjugate pairs, the guide family with exact NumPyro parameter counts, the SVI training-loop cheat sheet, the MCMC and NUTS diagnostics with their rules of thumb, a JAX and JIT cheat sheet (the boolean-masking fix, PRNG keys, static shapes) and the library conventions that silently change your answer. Notation: $\theta$ a parameter, $D$ the data, $\tilde y$ a new observation, $p(\cdot)$ a density or pmf, $N(\mu,\sigma^2)$ written with the variance (NumPyro and SciPy take the sd), $q_\phi$ a variational guide, $d$ the number of unconstrained latent numbers, $S$ the number of draws or particles. The checkout example is 50/500 (A) vs 60/500 (B). Every threshold is a rule of thumb unless a chapter proves it.
6.1–6.4 · Bayesian inference, priors, conjugate models and posterior decisions
6.1 · Bayesian inference: prior, likelihood, posterior, evidence, predictive
| Formula | In words | Main trap | See |
|---|---|---|---|
| $p(\theta\mid D) = \dfrac{p(D\mid\theta)\,p(\theta)}{p(D)}$, so posterior $\propto$ likelihood $\times$ prior | How well each $\theta$ explains the data, weighted by the prior and rescaled to area 1. A compromise, closer to whichever side is sharper. | "$\propto$" not "$=$"; a value with prior density 0 stays at 0 whatever the data say. | 6.1 |
| $p(D\mid\theta) = \prod_i p(y_i\mid\theta)$, $\;\log p(D\mid\theta) = \sum_i \log p(y_i\mid\theta)$ | The likelihood: data fixed, $\theta$ varies. | Not a distribution over $\theta$; only ratios matter. | 6.1 |
| $p(D) = \int p(D\mid\theta)\,p(\theta)\,d\theta = E_{\text{prior}}[p(D\mid\theta)]$ | The evidence: the likelihood averaged over the prior, one number per model. Bayes factor $= p(D\mid M_1)/p(D\mid M_2)$. | Not the likelihood at the best $\theta$; very sensitive to prior width; usually intractable. | 6.1 |
| Grid: $u_i = p(\theta_i)\,p(D\mid\theta_i)$, $\;p(D)\approx\sum_i u_i\Delta$, posterior $= u_i/(\sum_j u_j\Delta)$ | Exact enough in 1–2 dimensions. Beta(2, 18) prior + 3 of 20 gives Beta(5, 35). | A grid needs $G^d$ points: 100 per axis and 10 parameters is $10^{20}$. | 6.1 |
| $p(\theta\mid D_1, D_2) \propto p(D_2\mid\theta)\,p(\theta\mid D_1)$ | Today's posterior is tomorrow's prior; same result as one big update, in any order. | The prior counts once; it assumes a fixed $\theta$ (no drift); peeking affects decisions, not the posterior. | 6.1 |
| $p(\tilde y\mid D) = \int p(\tilde y\mid\theta)\,p(\theta\mid D)\,d\theta$; simulate: $\theta\sim$ posterior, then $\tilde y\sim$ likelihood | The posterior predictive: a forecast of new data. P(next user converts) $= E[\theta\mid D]$. | A different object from $P(\theta_B \gt \theta_A\mid D)$ (that one is about parameters). The plug-in prediction is always too narrow. | 6.1 |
| $Var(\tilde y\mid D) = E[Var(\tilde y\mid\theta)] + Var(E[\tilde y\mid\theta])$ | Observation noise + parameter uncertainty. Normal mean, flat prior: $\sigma^2 + \sigma^2/n$. | For a trend the parameter part grows with the horizon; the noise part is the floor. | 6.1 |
6.2 · Choosing priors and prior predictive checks
| Formula | In words | Main trap | See |
|---|---|---|---|
| Strength: Beta($\alpha,\beta$) is worth $\alpha+\beta$ observations; $N(\mu_0,\tau^2)$ on a mean is worth $\sigma^2/\tau^2$ | Informative, weakly informative, diffuse, flat: a ladder of how loudly a prior speaks. | "Vague" is not neutral; judge a prior in business units, on a known scale. | 6.2 6.2 |
| Beta by moments: $\alpha+\beta = m(1-m)/v - 1$; discount: Beta($a_0\alpha, a_0\beta$) | Mean 0.10, sd 0.015 gives Beta(39.9, 359.1), worth 399. | Needs $v \lt m(1-m)$; never build the prior from the same data; keep priors on effects skeptical. | 6.2 |
| $p_\eta(\eta) = p_p(p)\,|dp/d\eta|$; logit: $dp/d\eta = p(1-p)$ | A prior flat on $p$ is a logistic hump on the log-odds; $N(0,10^2)$ on the log-odds puts about 65% of the mass below 1% or above 99%. | No prior is flat on every scale: "uninformative" is not a property of a density. | 6.2 |
| $E[\theta\mid D] = \dfrac{n_0}{n_0+n}\,m_0 + \dfrac{n}{n_0+n}\,\dfrac kn$ | Prior mean and data rate, weighted by information. Beta(10, 90) + 30 of 200 gives 0.133. | "Lots of data" is per parameter: small segments and rare events stay prior-sensitive. | 6.2 |
| $-\log N(\theta\mid 0,\tau^2) = \theta^2/2\tau^2 + c$; $\;-\log\text{Laplace}(\theta\mid 0,b) = |\theta|/b + c$ | A shrinkage prior is a penalty: Normal $\leftrightarrow$ ridge, Laplace $\leftrightarrow$ lasso, at the MAP. The scale is the flexibility knob. | "Laplace prior $\Rightarrow$ sparse posterior" is false: only the MAP has exact zeros. | 6.2 |
| $\theta_g\sim N(\mu,\tau^2)$, $\mu\sim p(\mu)$, $\tau\sim p(\tau)$ | Hierarchical prior: $\tau$ says how similar the groups are; small $\tau$ means strong pooling. | With few groups $\tau$ is weakly identified, so its hyperprior matters. | 6.2 |
$p(y) = \int p(y\mid\theta)\,p(\theta)\,d\theta$; NumPyro Predictive(model, num_samples=…)(key, …) without posterior samples | Simulate $\theta\sim$ prior, then $y\sim$ likelihood, on the real time grid and regressors; check range, growth, wiggliness, noise. | Never tune the prior to the observed data; prior scales are relative to the data scaling; "too tight" is also a failure. | 6.2 6.2 |
6.3 · Conjugate models: Beta-Binomial and Dirichlet-Multinomial
| Formula | In words | Main trap | See |
|---|---|---|---|
| $\theta\sim\text{Beta}(\alpha,\beta)$, $k\mid\theta\sim\text{Bin}(n,\theta)$ $\Rightarrow$ $\theta\mid k\sim\text{Beta}(\alpha+k,\ \beta+n-k)$ | Add successes to $\alpha$, failures to $\beta$. Sufficient statistics $(k, n)$; batching and order never matter. | NumPyro: Beta(concentration1=α, concentration0=β), with concentration1 counting successes. | 6.3 |
| Evidence $p(k) = \binom nk \dfrac{B(\alpha+k,\ \beta+n-k)}{B(\alpha,\beta)}$ | The beta-binomial probability of the observed count (SciPy betabinom(n, a, b).pmf(k)); flat prior gives $1/(n+1)$. | Constants cancel from the posterior but not from the evidence. | 6.3 |
| $E[\theta\mid k] = w\,m_0 + (1-w)\dfrac kn$, $\;w = \dfrac{\alpha+\beta}{\alpha+\beta+n}$; $\;Var = \dfrac{m(1-m)}{\alpha+\beta+n+1}$ | The prior is worth $\alpha+\beta$ pretend visitors. Prior mean 10%, worth 50, 8 000 users per variant: $w\approx0.006$. | Flat prior: the mode is $k/n$ but the mean is $(k+1)/(n+2)$. | 6.3 |
| $\theta_A\mid D$ and $\theta_B\mid D$ independent Betas; gap sd $=\sqrt{sd_A^2 + sd_B^2}$ | The A/B model: same prior, one Beta per variant, decisions use the gap. | The posterior of the gap or the lift is not a Beta; get it from draws. Overlapping intervals do not mean no difference. | 6.3 |
| Predictive of $m$ new users: BetaBinomial($m,\alpha',\beta'$); variance $m\bar p(1-\bar p)\dfrac{\alpha'+\beta'+m}{\alpha'+\beta'+1}$ | Same mean as the plug-in Binomial but wider: parameter plus observation uncertainty. | "Beta-binomial" names a model and a distribution; say which. | 6.3 |
| Dirichlet($\boldsymbol\alpha$) + counts $\mathbf c$ $\Rightarrow$ Dirichlet($\boldsymbol\alpha+\mathbf c$); one share: Beta($\alpha_k+c_k,\ \alpha_0+n-\alpha_k-c_k$) | One pseudo-count per category; posterior mean $\frac{\alpha_k+c_k}{\alpha_0+n}$. | The shares sum to 1: whole-mix questions need joint draws, not separate Betas. | 6.3 |
| Gamma($a$, rate $b$) + Poisson counts, exposures $t_i$ $\Rightarrow$ Gamma($a+\sum y_i,\ b+\sum t_i$) | Events go into the shape, exposure into the rate; predictive Negative Binomial. | SciPy gamma(a, scale=1/b); the Poisson assumes no overdispersion. | 6.3 |
| $\dfrac1{s_n^2} = \dfrac1{s_0^2} + \dfrac n{\sigma^2}$, $\;m_n = w\,m_0 + (1-w)\bar y$; predictive $N(m_n,\ \sigma^2 + s_n^2)$ | Precisions add; the mean is precision-weighted. The prior is worth $\sigma^2/s_0^2$ observations. | Weight by precision, not 50/50; NumPyro Normal(loc, scale) takes the sd. | 6.3 |
6.4 · Credible intervals and posterior decisions
| Formula | In words | Main trap | See |
|---|---|---|---|
| Squared loss $\to$ mean; absolute loss $\to$ median; all-or-nothing $\to$ mode (MAP) | Each one-number summary is best for a different cost. Right skew: mode $\lt$ median $\lt$ mean. | The mode and the mean change under a change of scale; the median does not. | 6.4 |
$P(L\le\theta\le U\mid D) = 0.95$; from draws np.quantile(x, [0.025, 0.975]) | A direct probability statement about $\theta$, given the model, prior and data. | It depends on the prior and is not a coverage guarantee for every fixed $\theta$. | 6.4 |
| ETI $=[F^{-1}(\gamma/2),\,F^{-1}(1-\gamma/2)]$; HDI $=\{\theta: p(\theta\mid D)\ge c\}$ | Equal tails (scale-free) vs shortest interval (equal density at the ends). Beta(3, 40): ETI [0.015, 0.162], HDI [0.008, 0.145]. | HDIs can split into pieces and depend on the scale; say which interval and which mass. | 6.4 |
| $P(\theta_B \gt \theta_A\mid D)\approx\frac1S\sum_s\mathbf 1[\theta_B^{(s)} \gt \theta_A^{(s)}]$; MC error $\sqrt{P(1-P)/S}$ | The share of paired draws above the diagonal. $S = 4\,000$: error at most 0.008. | Direction, not size; not $1-p$; use the same draw index for every variant. | 6.4 |
| $P(\theta_B-\theta_A \gt \delta\mid D)$; ROPE: $P(|\theta_B-\theta_A| \lt \delta\mid D)$ | Better by enough to matter, with $\delta$ fixed in advance. Huge test, 10.0% vs 10.2%: $P(B \gt A) = 0.98$ but $P(\text{gap} \gt 0.5\text{ pt}) = 0.001$. | "98% likely better" is about direction, not importance. | 6.4 |
| Lift $L^{(s)} = \theta_B^{(s)}/\theta_A^{(s)} - 1$ per draw | Report the median and ETI, next to the baseline. $P(L \gt 0) = P(\theta_B \gt \theta_A)$. | Never divide interval ends or average segment lifts; $E[\theta_B/\theta_A]\ne E[\theta_B]/E[\theta_A]$; convert from standardized units first. | 6.4 |
| Expected loss of choosing B: $E[\max(\theta_A-\theta_B,\,0)\mid D]$ | Weighs mistakes by their size, in the metric's units. A rule fixes $\delta$, the threshold, $\varepsilon$ and a maximum $n$ in advance. | Peeking and many metrics raise the rate of bad ships even though the posterior stays valid; simulate the rule as you will run it. | 6.4 |
| Confidence: $P_D(\theta\in CI\mid\theta)\approx0.95$. Credible: $P(\theta\in CrI\mid D)=0.95$ | Recipe over repeated samples vs belief given the data. Wilson [0.0944, 0.1514] $\approx$ flat CrI [0.0944, 0.1515] for 60 of 500. | Similar numbers with weak priors and big $n$, different meanings always. | 6.4 |
6.5–6.8 · Hierarchical models, pooling, parameterization and model checking
6.5 · Hierarchical models
| Formula | In words | Main trap | See |
|---|---|---|---|
| $\theta_g\mid\mu,\tau\sim N(\mu,\tau^2)$, $\;y_{gi}\mid\theta_g\sim p(y\mid\theta_g)$, plus $p(\mu)$, $p(\tau)$ | Two levels: group values drawn from a population, data drawn from each group. $\tau$ = real spread between groups, $\sigma$ = noise inside a group. | An observed group average is not the group's true value; it carries noise $\sigma^2/n_g$. | 6.5 |
| Joint $= p(\mu)p(\tau)p(\sigma)\prod_g\big[N(\theta_g\mid\mu,\tau^2)\prod_i p(y_{gi}\mid\theta_g,\sigma)\big]$ | One factor per circle of the plate diagram given its parents. Unknowns with $G$ groups: complete pooling 2, separate $G+1$, hierarchical $G+3$. | A "hyperparameter" here is a learned quantity with its own prior, not a training setting. | 6.5 |
| $\bar y_g\mid\mu,\tau\sim N(\mu,\ \tau^2+\sigma^2/n_g)$; $\;\tau\sim$ HalfNormal / half-$t$ | The data see $\tau$ only through how far the group averages spread beyond their noise. | Few groups: $\tau$ is poorly learned and the hyperprior matters; check it by prior predictive simulation. | 6.5 |
| Exchangeable: $p(\theta_1..\theta_G)$ unchanged by reordering; shared-centre correlation $\dfrac{Var(\mu)}{Var(\mu)+\tau^2}$ | Interchangeable before the data: roughly iid given $(\mu,\tau)$ (de Finetti). Known differences go in as covariates: $\theta_g\sim N(\mu+\beta x_g,\tau^2)$. | Not "identical" and not "independent". Do not pool control and treatment into one population. | 6.5 |
| $Var(y)=\sigma^2+\tau^2$, $\;\text{ICC}=\dfrac{\tau^2}{\tau^2+\sigma^2}$, $\;Var(\bar y_g)=\tau^2+\dfrac{\sigma^2}{n_g}$ | Within plus between. Rough $\hat\tau^2 = s^2_{\bar y} - \sigma^2/n$. | That estimate can be negative; the Bayesian posterior of $\tau$ stays $\ge 0$. Per-group scaling would make $\tau^2$ vanish: use one global scaler. | 6.5 |
| Rates: $\text{logit}\,p_g\sim N(\mu,\tau^2)$, or $p_g\sim\text{Beta}(\kappa\phi,\ \kappa(1-\phi))$ with $Var=\dfrac{\phi(1-\phi)}{\kappa+1}$ | Put the Normal on the natural scale: logit for rates, log for counts. $\kappa$ = "worth $\kappa$ visitors". | $\tau$ is not in percentage points; logistic($\mu$) is the median rate, not the mean. | 6.5 |
| Fake-data check: fix $\to$ simulate $\to$ fit $\to$ compare $\to$ repeat | Cheapest insurance for the whole pipeline. Truth drawn from the prior each time: simulation-based calibration. | It tests the inference, not whether real data follow the model. | 6.5 |
6.6 · Pooling and shrinkage
| Formula | In words | Main trap | See |
|---|---|---|---|
| Complete: $\hat\theta_g=\bar y$ ($\tau=0$). None: $\hat\theta_g=\bar y_g$ ($\tau=\infty$). Partial: $w_g\bar y_g+(1-w_g)\mu$ | Three settings of one model; partial pooling is the posterior mean of the hierarchical model. | The weight is not a taste choice and differs by segment; "no pooling" is not assumption-free. | 6.6 |
| $w_g=\dfrac{n_g/\sigma^2}{n_g/\sigma^2+1/\tau^2}=\dfrac{\tau^2}{\tau^2+\sigma^2/n_g}=\dfrac{n_g}{n_g+\sigma^2/\tau^2}$; posterior variance $\dfrac1{n_g/\sigma^2+1/\tau^2}$ | Precisions add. The population is worth $\sigma^2/\tau^2$ extra observations at $\mu$. With $\sigma=12$, $\tau=4$: $n=1, 9, 36, 81$ give $w=0.1, 0.5, 0.8, 0.9$. | NumPyro's Normal takes the sd; the precision uses the variance. | 6.6 |
| Shrinkage $B_g=1-w_g=\dfrac{\sigma^2}{\sigma^2+n_g\tau^2}$; movement $=B_g(\bar y_g-\mu)$ | A segment is shrunk by the fraction of its average that is noise. Small groups move more. | Shrunk estimates understate the true spread of the groups; look at $\tau$ or posterior draws. | 6.6 |
| $\tau\to0$: $w_g\to0$; $\tau\to\infty$: $w_g\to1$; $p_{\text{eff}}=\sum_g\partial\hat\theta_g/\partial\bar y_g\in[1,G]$ | "Do lifts differ by segment?" is a question about the posterior of $\tau_\delta$. | Look at that posterior before showing any segment breakdown. | 6.6 |
| $\hat\mu(\tau)=\dfrac{\sum_j v_jy_j}{\sum_j v_j}$, $v_j=\dfrac1{\tau^2+s_j^2}$; full Bayes averages over $p(\tau\mid y)$ | Empirical Bayes plugs in $\hat\tau$; full Bayes keeps $Var(E[\theta_j\mid\tau,y])$. Eight schools, school A: 11.4 $\pm$ 8.3. | Empirical Bayes is too narrow with few groups, badly so when $\hat\tau$ is at 0. | 6.6 |
| $\text{MSE}_{\text{partial}}(d)=w^2s^2+(1-w)^2d^2$; averaged: $w\,s^2\lt s^2$ | Shrinkage adds a little bias and removes more variance. Losers: $|d| \gt \sqrt{2\tau^2+s^2}$, a minority. | Not every segment improves; the claim is about total error. | 6.6 |
| Rates: $\hat p_g=\dfrac{k_g+\kappa\phi}{n_g+\kappa}=w_g\dfrac{k_g}{n_g}+(1-w_g)\phi$, $\;w_g=\dfrac{n_g}{n_g+\kappa}$ | With $\phi=10\%$, $\kappa=100$: 2/5 $\to$ 11.4%, 30/200 $\to$ 13.3%, 1300/10 000 $\to$ 12.97%. | Winner's curse: the top raw segment is usually a small lucky one; rank pooled estimates. | 6.6 |
6.7 · Centered vs non-centered parameterization
| Formula | In words | Main trap | See |
|---|---|---|---|
| Centered: $\theta_g\sim N(\mu,\tau^2)$. Non-centered: $z_g\sim N(0,1)$, $\theta_g=\mu+\tau z_g$ | Same model and posterior; only the coordinates the algorithm explores differ. | If the two fits disagree, a fit is broken, not the model. | 6.7 |
| Funnel: $p(\theta,s)=N(s;0,1.5^2)\,N(\theta;0,e^{2s})$ with $s=\log\tau$ | The room for $\theta_g$ is proportional to $\tau$: narrow neck, wide mouth. Toy widths $\pm0.27$ vs $\pm14.5$ at $\log\tau=\mp2$. | It appears when data per group are weak. Diagnose with a pairs plot of $\theta_g$ against $\log\tau$. | 6.7 |
| Random-walk acceptance $\dfrac2\pi\arctan(2\tau/\varepsilon)$; leapfrog stable only if $\varepsilon \lt 2\tau$ | One global step size cannot fit a width that changes with height. | NUTS adapts one global $\varepsilon$ (and a scale per coordinate), not a local one. | 6.7 |
Divergence: energy error $\Delta H \gt 1000$ (NumPyro max_delta_energy); extra_fields=("diverging",) | The simulation broke in a region of high curvature; draws are biased toward large $\tau$. | $\hat R=1.00$ and a good ESS do not cancel divergences: all chains can miss the same neck. | 6.7 |
| Rule of thumb: $w_g=\tau^2/(\tau^2+SE_g^2)$ small $\to$ non-centered; near 1 $\to$ centered | Eight groups: $\times1$ data, non-centered wins (0 vs 54 divergences); $\times100$ data, centered wins (ESS about 1 040 vs 126). | "Strong data" is about $SE_g$ against $\tau$, not total dataset size. | 6.7 |
| Gaussian guide in the toy funnel: best $\log\tau\sim N(0,0.64^2)$ vs truth $N(0,1.5^2)$, KL $\approx0.85$ | A Gaussian guide has the same width at every height, so it understates the uncertainty in $\tau$. Non-centered toy: KL $=0$. | SVI has no divergence alarm. Same model means same $\log p(D)$: the form with the higher final ELBO has the smaller KL. | 6.7 |
reparam(model, config={"theta": LocScaleReparam(centered=0)}) | Creates the latent theta_decentered; theta becomes deterministic. $c=0$ non-centered, $c=1$ centered. | Default centered=None is a learnable 0.5 for SVI only; under NUTS it stays at 0.5. | 6.7 |
6.8 · Checking a Bayesian model
| Formula | In words | Main trap | See |
|---|---|---|---|
| PPC: $\theta^{(s)}\sim p(\theta\mid y)$, $\;y^{\text{rep}(s)}\sim p(y\mid\theta^{(s)})$, same size and design as $y$ | Can the fitted model reproduce the data? Real data that stand out show a missing feature (Poisson: 0–2 zero days per fake month vs 7 real). | Passing a PPC does not prove the model; use a fresh posterior draw for every replicate. | 6.8 |
| $p_B=P\big(T(y^{\text{rep}})\ge T(y)\mid y\big)\approx$ share of replicates at least as extreme | Near 0 or 1 flags misfit on that statistic. Pick a few statistics tied to decisions: variance, zero days, max, group spread, lag-1 autocorrelation. | Conservative (data used twice), not P(model true); a statistic a parameter fits directly always passes. | 6.8 6.8 |
| Prior sensitivity: refit under pre-chosen reasonable priors; report the range of the key outputs | 3/20 vs 8/20: $P(B \gt A)$ = 0.957 / 0.963 / 0.931 / 0.774 (flat, Jeffreys, weak, strong): fragile; with 10$\times$ data all about 1.000. | Never pick the prior after seeing which answer it gives. Steep for small data, scales ($b$, $\tau$), dispersion, sparse effects. | 6.8 6.8 |
| Identifiable: $p(y\mid\theta_1)=p(y\mid\theta_2)\ \forall y\Rightarrow\theta_1=\theta_2$. For $y\sim N(a+b,1)$: sd($a+b$) $=0.1$, sd($a-b$) $=14.1$ (the prior), corr $=-0.9999$ | Along a flat direction the posterior is the prior; only the identified combination ($a+b$) is learned. | Symptoms: correlations near $\pm1$, little contraction, slow chains, seed-dependent SVI. | 6.8 |
| Normal priors, co-occurring effects with sum $u$: $E[\beta_h\mid D]=\dfrac{s_h^2}{s_h^2+s_x^2+se^2}\,u$ | Holiday plus promotion always together: the priors write the split (priors 10/10 give 9.8/9.8; priors 20/5 give 18.6/1.2). | Report the sum; forecasts break when the components decouple. | 6.8 |
Contraction $1-Var_{\text{post}}/Var_{\text{prior}}$; sum-to-zero s = s_raw - s_raw.mean() | Near 1 = learned from data, near 0 = still the prior. Fixes: separating data, constraints, defensible priors, reparameterize, merge. | Document the source of any informative prior used as an identification aid. | 6.8 |
6.9–6.10 · Approximate inference, MCMC, Hamiltonian Monte Carlo and NUTS
6.9 · Why approximate inference? MCMC from scratch
| Formula | In words | Main trap | See |
|---|---|---|---|
| $p(\theta\mid D)=\tilde p(\theta)/p(D)$, $\;\tilde p(\theta)=p(D\mid\theta)\,p(\theta)$, $\;p(D)=\int\tilde p\,d\theta$ | The shape is cheap (and differentiable), the total is not. Grid cost $k^d$: 100 points, 10 parameters is $10^{20}$ evaluations. | "Intractable" does not mean "no formula known"; Bayes' theorem still applies. | 6.9 |
| $E[f(\theta)\mid D]\approx\frac1S\sum_s f(\theta^{(s)})$; error $\text{sd}/\sqrt S$ (probability: $\sqrt{q(1-q)/S}$) | Probabilities are shares, intervals are sorted draws; the error does not depend on the number of parameters. | More draws do not narrow the posterior; correlated MCMC draws need the ESS instead of $S$. | 6.9 |
| Laplace: $N(\hat\theta, A^{-1})$, $A=-\nabla^2\log\tilde p(\hat\theta)$; 1-D sd $=1/\sqrt{-\ell''(\hat\theta)}$ | A bell at the peak. Beta(5, 35): $\hat\theta=0.105$, sd 0.050 (exact mean 0.125, sd 0.052). | Centred at the mode, not the mean; depends on the scale (use log or log-odds); blind to skew and second peaks. | 6.9 |
| Importance sampling: $w_s=\tilde p(\theta^{(s)})/q(\theta^{(s)})$, $\;E_p[f]\approx\frac{\sum w_sf}{\sum w_s}$, $\;\text{ESS}_w=1/\sum\bar w_s^2$ | Re-weight draws from an easy $q$; $p(D)$ cancels. | The proposal must be wider than the target; the weights collapse onto a few draws as the dimension grows. | 6.9 |
| $\pi_{t+1}=\pi_tP$; stationary: $\pi=\pi P$; detailed balance: $\pi_iP_{ij}=\pi_jP_{ji}$ | An irreducible, aperiodic chain forgets its start and spends a share $\pi$ of its time in each state. MCMC designs a chain whose $\pi$ is the posterior. | Detailed balance is sufficient, not necessary; draws are correlated, not independent. | 6.9 |
| Metropolis: $\theta'=\theta+\varepsilon z$; accept if $\log u \lt \log\tilde p(\theta')-\log\tilde p(\theta)$, else record $\theta$ again | Uphill always, downhill sometimes. $p(D)$ cancels in the ratio. | Rejected steps are part of the output; it is not an optimizer. | 6.9 |
| Step size: $\varepsilon\approx2.38\,\sigma/\sqrt d$, acceptance $\approx0.234$ (0.44 in 1-D) | Rules of thumb for Gaussian-like targets; tune during warmup, then freeze. | High acceptance is not good mixing; judge by ESS. | 6.9 |
| $\text{ESS}=n/(1+2\sum_k\rho_k)$; AR(1): $n(1-\rho)/(1+\rho)$ | $\rho=0.9$, $n=10\,000$: ESS 526; thinning by 10 gives 483. | Thinning never raises precision; autocorrelation is inefficiency, not bias. | 6.9 |
Burn-in / warmup: MCMC(NUTS(model), num_warmup=500, num_samples=1000) | 1 500 iterations per chain run, 1 000 returned; warmup also tunes the sampler. | Convergence is checked (traces, several chains, $\hat R$), never proven; burn-in cannot fix a missed mode. | 6.9 6.9 |
6.10 · Hamiltonian Monte Carlo, NUTS and MCMC diagnostics
| Formula | In words | Main trap | See |
|---|---|---|---|
| Random walk: distance $\approx\varepsilon\sqrt n$, so crossing $D$ costs $(D/\varepsilon)^2$ steps | Momentum makes it about $D/\varepsilon$. HMC uses $\nabla\log\tilde p$ (autodiff; no $p(D)$ needed). | The gradient bends the path; HMC is not an optimizer. It needs continuous, differentiable parameters. | 6.10 |
| $H=U+K$, $\;U(\theta)=-\log\tilde p(\theta)$, $\;K(p)=\tfrac12p^\top M^{-1}p$, $\;p\sim N(0,M)$ afresh each iteration | A frictionless puck on the landscape. Exact motion conserves $H$, is reversible and volume-preserving. | Momentum is not a model parameter; $U$ changes along the path, $H$ does not. | 6.10 |
| Leapfrog: $p\mathrel{+}=\tfrac\varepsilon2\nabla\log\tilde p$; $\theta\mathrel{+}=\varepsilon p$; $p\mathrel{+}=\tfrac\varepsilon2\nabla\log\tilde p$ | $L$ steps cost $L$ gradients; energy error $O(\varepsilon^2)$, bounded while $\varepsilon \lesssim 2\times$ the narrowest sd. | Euler drifts; a too-big $\varepsilon$ explodes, which is a divergence. | 6.10 |
| Accept with $\min(1,e^{-\Delta H})$, $\;\Delta H=H_1-H_0$ | $\Delta H=0.1\to0.905$, $0.5\to0.607$, $2\to0.135$. Far moves with high acceptance. | Acceptance measures simulation accuracy, not exploration; judge by ESS per gradient. | 6.10 |
| NUTS: double the trajectory until $(q^+-q^-)\cdot p^-\lt 0$ or $(q^+-q^-)\cdot p^+\lt 0$; steps $2^j-1$; default max depth 10 (1 023 steps) | Then sample a point from the trajectory in proportion to $e^{-H}$. | NUTS still has settings (warmup, chains, target_accept_prob) and still depends on the parameterization. | 6.10 |
Warmup: step size by dual averaging toward target_accept_prob (default 0.8); $M^{-1}\approx$ posterior variances (diagonal by default; dense_mass=True for full) | Windows: 75 fast, slow windows 25, 50, 100, …, 50 fast. $\varepsilon$ and $M$ are frozen afterwards. | A diagonal $M$ does not fix correlations; a tiny adapted $\varepsilon$ signals a hard geometry. | 6.10 |
| $\hat R=\sqrt{\hat V/W}$, $\;\hat V=\frac{n-1}nW+\frac Bn$ (split chains) | Do the chains agree? Rule of thumb $\le1.01$. | Necessary, not sufficient: chains can share the same blind spot. | 6.10 |
| $\text{ESS}=mn/\hat\tau$; $\;\text{MCSE}=\text{sd}/\sqrt{\text{ESS}}$; probability: $\sqrt{q(1-q)/\text{ESS}}$ | ESS 400 and $q=0.95$: MCSE 0.011, so 0.95 vs 0.94 is not a real difference. | MCSE is not the posterior sd; never use sd$/\sqrt n$ on MCMC draws. | 6.10 |
6.11–6.15 · Variational inference, the ELBO, SVI, guides, the training loop and SVI vs NUTS
6.11 · Variational inference and KL divergence
| Formula | In words | Main trap | See |
|---|---|---|---|
| $\phi^\star=\arg\min_\phi KL\big(q_\phi\,\|\,p(\theta\mid D)\big)$ | Fit a distribution instead of sampling. $\phi$ are the knobs of $q$ (e.g. $\mu,\sigma$), not the model parameters $\theta$. | The best member of the family is not the posterior; approximation error plus optimization error. | 6.11 6.11 |
| $KL(q\|p)=E_q[\log q-\log p]\ge0$; two Normals: $\log\frac{\sigma_2}{\sigma_1}+\frac{\sigma_1^2+(\mu_1-\mu_2)^2}{2\sigma_2^2}-\frac12$ | Zero only if $q=p$. $N(0,1)$ vs $N(0,2^2)$: 0.318 one way, 0.807 the other. | Not symmetric; infinite if $q$ has mass where $p$ has none (so fit on the unconstrained scale). | 6.11 |
| Reverse $E_q[\log q/p]$ (VI) vs forward $E_p[\log p/q]$ | Reverse: mode-seeking, zero-forcing, one mode, too narrow. Forward: mass-covering; for a Normal $q$ it matches mean and variance. | Missing half the mass costs only $\log2\approx0.69$ in reverse KL; two modes $\pm3$: reverse gives $N(2.98, 1.02^2)$. | 6.11 |
| $KL(q\|p(\theta\mid D))=-\text{ELBO}(q)+\log p(D)$ | $\log p(D)$ does not depend on $q$, so VI needs only the joint and draws from its own $q$. | The forward KL would need posterior draws, which VI does not have. | 6.11 |
| Mean-field: $q=\prod_iq_i$; Normal posterior gives $q_i=N(\mu_i,\ 1/\Lambda_{ii})$; 2-D: sd $\times\sqrt{1-\rho^2}$, $KL=-\tfrac12\log(1-\rho^2)$ | Exact means, conditional variances. $\rho=0.9$: sd $\times0.44$, coverage 61%. Student-t $\nu=3$: guide $\sigma$ 1.26 vs sd 1.73. | A strong tendency to under-disperse, not a law: check against NUTS or with held-out coverage. The family is to blame, not the optimizer. | 6.11 6.11 |
6.12 · The ELBO and stochastic variational inference
| Formula | In words | Main trap | See |
|---|---|---|---|
| $\log p(D)=\underbrace{E_q[\log p(D,\theta)-\log q(\theta)]}_{\text{ELBO}}+KL(q\|p(\theta\mid D))$ | A fixed total split into a computable part and a gap. ELBO $\le\log p(D)$, equal only at the posterior. | ELBO is not the evidence; comparing ELBOs across different models is unreliable. | 6.12 |
| $\text{ELBO}=E_q[\log p(D,\theta)]+H[q]=E_q[\log p(D\mid\theta)]-KL(q\|p(\theta))$ | Fit plus entropy; or data fit minus distance from the prior (the VAE loss, with a minus sign). | Two different KLs: to the prior inside the ELBO, to the posterior in the gap. | 6.12 |
| $\widehat{\text{ELBO}}=\frac1S\sum_s[\log p(D,\theta_s)-\log q(\theta_s)]$, $\;\theta_s\sim q$ | Unbiased; sd $\propto1/\sqrt S$; cost $\propto S$. Trace_ELBO(num_particles=S), default $S=1$. | One step's loss is mostly noise; smooth before judging. | 6.12 |
| $\theta=\mu+\sigma\varepsilon$, $\varepsilon\sim N(0,1)$: $\nabla_\phi E_q[f]=E_\varepsilon[\nabla_\phi f(\mu+\sigma\varepsilon)]$ | Low-noise gradients: $\partial\theta/\partial\mu=1$, $\partial\theta/\partial\sigma=\varepsilon$. Full-rank: $\theta=m+L\varepsilon$. | Needs reparameterizable (continuous) latents; the score-function estimator is unbiased but far noisier. | 6.12 |
state, loss = svi.update(state, *data), with loss $=-\widehat{\text{ELBO}}$ | Sample, score, differentiate, Adam step. Improvement means the loss goes down. | It returns the loss, not the ELBO; its scale is meaningless (contains $-\log p(D)$), only its smoothed trend is. | 6.12 |
| Minibatch: $\frac NB\sum_{i\in\mathcal B}\log p(y_i\mid\theta)+\log p(\theta)-\log q(\theta)$ | Unbiased; NumPyro plate("data", N, subsample_size=B) (full $N$) applies the factor. | Forgetting $N/B$ acts as if you had $B$ points (posterior $\sqrt{N/B}$ too wide); scale data terms only. | 6.12 |
6.13 · Variational guides (counts and the full table are in the guide table below)
| Formula | In words | Main trap | See |
|---|---|---|---|
| $q(\mathbf z)=N(\boldsymbol\mu,\Sigma_q)$ over the flattened unconstrained latent vector $\mathbf z\in\mathbb R^d$ | The guide type is the structure of $\Sigma_q$. $d$ counts unconstrained scalars (shape (25,) adds 25; a $K$-simplex adds $K-1$). | The guide is not the posterior; a positive parameter's guide marginal is log-normal, skewed to the right. | 6.13 |
| Mean-field optimum on a Gaussian posterior: $s_i^2=1/\Lambda_{ii}\le\Sigma_{ii}$ | Correlation 0.9 gives 0.19 instead of 1; the variance of a sum along the correlated direction comes out about ten times too small. | Too narrow along the long axis, too wide along the short axis: a contrast can be off either way. | 6.13 |
| Full-rank: $\mathbf z=\boldsymbol\mu+L\boldsymbol\varepsilon$, $\Sigma_q=LL^\top$, $\log|\Sigma_q|=2\sum_i\log L_{ii}$ | Learning $L$ guarantees a valid covariance; $O(d^2)$ memory and work per sample. | auto_scale_tril is $L$, not $\Sigma$. $O(d^3)$ appears only to factorize a dense matrix. Full-rank is still Gaussian, not exact. | 6.13 |
| Low-rank: $\mathbf z=\boldsymbol\mu+W\boldsymbol\varepsilon_r+D^{1/2}\boldsymbol\varepsilon_d$, $\Sigma_q=WW^\top+D$ | $r$ shared directions plus private variances; sample $O(dr)$, density $O(dr^2+r^3)$ (Woodbury). | Misses many separate pairwise ties and long chains. Pick $r$ from the eigenvalues above the floor, not "90% of variance". | 6.13 |
| Selection rule: full-rank if $d+d(d+1)/2\le B$; else low-rank, $r$ from the spectrum (or $\text{round}(\sqrt d)$), $d(r+2)\le B$ | Memory, not accuracy, forces the switch. Validate on a reduced model against full-rank or NUTS; raise $r$ until the conclusions stop changing. | Any threshold is a design choice, not a law. Converging is not the same as an adequate family. | 6.13 |
6.14 · The custom SVI training loop (the loop itself is in the loop cheat sheet below)
| Formula | In words | Main trap | See |
|---|---|---|---|
| $\text{rel}_t=\dfrac{ELBO_t-ELBO_{best}}{|ELBO_{best}|+\epsilon}$; meaningful if $\text{rel}_t \gt \tau$ | Scale-free under multiplying the ELBO by $c \gt 0$. Rule of thumb $\tau=10^{-4}$ to $10^{-5}$ per evaluation. | Divide by $|\cdot|$ because losses can be negative; $\epsilon$ makes it absolute near zero (threshold $\tau\epsilon$). | 6.14 6.14 |
| Patience: stop if $\text{rel}\le\tau$ for $P$ evaluations in a row and step $\ge t_{min}$ | Waits $P\cdot k$ steps; false stops fall roughly like $b^P$, cost grows linearly. | Counted in evaluations, not steps; plateaus longer than $P\cdot k$ cause premature stops. | 6.14 |
| Window mean: sd $\sigma/\sqrt k$, lag $(k-1)/2$. EMA($\beta$): sd $\sigma\sqrt{\tfrac{1-\beta}{1+\beta}}$, lag $\tfrac\beta{1-\beta}$ | Judge averages, not single steps; or a fixed-key multi-particle Trace_ELBO(num_particles=M).loss(...). | A smoothed loss lags the current parameters; convert to floats only at evaluations. | 6.14 |
6.15 · SVI vs NUTS
| Formula | In words | Main trap | See |
|---|---|---|---|
| SVI cost $\approx T\times K\times B$; NUTS cost $\approx C\times(W+S)\times\bar L\times N$ | Steps $\times$ particles $\times$ rows vs chains $\times$ iterations $\times$ leapfrog steps $\times$ rows. Leapfrog steps $\le2^{\text{depth}}-1$. | "SVI is faster" is a statement about large models; measure ESS per second for NUTS. | 6.15 |
| SVI error = family gap + optimization error. NUTS error $\approx\text{sd}/\sqrt{\text{ESS}}$ + non-convergence bias | More SVI steps never close the family gap; a healthy NUTS bias vanishes. | A small MCSE never proves convergence; a flat ELBO is not accuracy. | 6.15 |
| For the same model and data: $\text{ELBO}(q_1)-\text{ELBO}(q_2)=KL(q_2\|p)-KL(q_1\|p)$ | A higher average ELBO (beyond seed-to-seed noise) means a guide closer in KL. | Only valid for the same model and data. | 6.15 |
| $\frac1S\sum_sf(\theta^{(s)})\to E[f\mid D]$ as $S\to\infty$ (valid, ergodic chain) | "Asymptotically exact": finite chains have Monte Carlo error, convergence risk, and it is the posterior of your model. | Never say "NUTS gives the exact posterior". | 6.15 |
6.16–6.17 · JAX fundamentals and JIT compilation
6.16 · JAX: arrays, pure functions, grad, vmap, scan, PRNG keys
| Idea | In words | Main trap | See |
|---|---|---|---|
y = x.at[i].set(v) (also .add, .multiply, .min, .max) | Arrays are immutable; an update returns a new array. Default float32 (about 7 digits); 64-bit needs jax_enable_x64. | Out-of-range indices are clamped (reads) or dropped (writes), never an error. | 6.16 |
| Pure: output depends only on inputs, no side effects | JAX traces once with tracers (shape and dtype, no numbers) and replays the recording. | Globals freeze at trace time, print fires once, np.random is frozen. Fix: arguments, jax.debug.print, PRNG keys. | 6.16 |
$\text{state}_{t+1}=\text{update}(\text{state}_t,\text{data})$; jax.tree.map(f, tree) | Pytrees (nested dicts, tuples, lists of arrays) carry parameters and states. Keeping the best state is keeping a reference. | In-place libraries (e.g. PyTorch) need a deep copy; JAX with donate_argnums needs jax.tree.map(jnp.copy, ...). | 6.16 |
jax.jit(jax.vmap(jax.grad(f))) | Transformations take a function and return a function; they compose and need pure functions. | grad(vmap(f)) fails (vector output); grad needs a scalar output and float inputs. | 6.16 6.16 |
| Central difference $\dfrac{f(x+h)-f(x-h)}{2h}=f'(x)+\dfrac{h^2}6f'''(x)+O(\varepsilon|f|/h)$ | Check custom gradients in float64 with check_grads (from jax.test_util). Five days, $\mu=95$, $\sigma=10$: $\partial/\partial\mu=\sum(y_i-\mu)/\sigma^2=0.25$. | The tiniest $h$ drowns in rounding error (float32 $\varepsilon\approx1.2\times10^{-7}$). | 6.16 |
jax.vmap(f, in_axes=(0, None))(draws, t) | Write for one case, run for many: output gets a leading batch axis; $S$ draws give $(S,T)$ forecasts. Not a loop. | Shared arguments need None; the whole batch is materialized, so chunk huge batches. | 6.16 |
carry, ys = lax.scan(f, init, xs), with f(carry, x) -> (new_carry, y) | A recursion with a fixed-shape carry, traced once, differentiable. AR(1), $\varphi=0.5$, $\varepsilon=[1,0,2,-1]$ gives $[1, 0.5, 2.25, 0.125]$. | Python loops unroll and compile slowly; scan is sequential, so vmap it across series. | 6.16 |
k1, k2 = jax.random.split(key); jax.random.fold_in(key, i) | Same key and same call give the same numbers. Separate keys for svi.init, training and Predictive; record the root seed. | Never reuse a key for two things that should be independent; a split parent is used up. | 6.16 |
NumPyro model $=$ traced function; handlers.seed, mcmc.run(key), svi.init(key), Predictive(...)(key) | Handlers turn the model into a pure loss; JAX then applies grad, vmap and jit. Data, scaler statistics, prior scales and switches are arguments. | Python logic on data inside the model runs at trace time only; use numpyro.deterministic to record values. | 6.16 |
6.17 · JIT compilation and model size
| Idea | In words | Main trap | See |
|---|---|---|---|
| Python $\to$ tracing $\to$ jaxpr $\to$ StableHLO $\to$ XLA $\to$ executable, cached by function, shapes, dtypes, static args | jax.make_jaxpr(f)(*args) shows the recording: operations and shapes, no values. | Python code runs only while tracing; the compiled program contains only array operations. | 6.17 |
| $T\approx T_{\text{compile}}+N\,T_{\text{step}}$ per signature; break-even $N^*=C/(e-r)$ | Compile cost $C$, eager step $e$, compiled step $r$. Measured on one CPU: jitted svi.update 0.43–0.75 s first call, 0.1–0.2 ms later; eager 40–70 ms. | Asynchronous dispatch: time after a warm-up call with .block_until_ready(), and report compile and step time separately. | 6.17 |
| Cache key $=$ function $+$ pytree structure $+$ shapes $+$ dtypes $+$ static args | A new key retraces and recompiles; new array values never do. Find recompiles with jax_log_compiles. | Static arguments must be hashable, and each new value compiles a program. Dtype changes recompile. | 6.17 |
| Pad to buckets (64, 128, 256, …) and carry a mask: mean $=\sum_im_ix_i/\sum_im_i$ | A few compiles serve many sizes at under 50% wasted work. | Shapes can never depend on traced values (jnp.zeros(n) needs a concrete n). | 6.17 |
Python if on a setting: fine. On traced data: jnp.where or lax.cond. Loops: fori_loop, scan, while_loop | 5 000 Python-loop steps unroll into 10 000 equations (1.8 s compile); fori_loop keeps one body (about 12–16 ms). | where(x > 0, log(x), 0) has a NaN gradient at 0: use a double where. | 6.17 |
x[mask] $\to$ NonConcreteBooleanIndexError under jit | Don't cut, mask: $\sum m_ix_i$, $\sum m_ix_i/\max(\sum m_i,1)$, jnp.where(mask, x, 0), NumPyro .mask(m); or select in NumPy before tracing. | Masked-out entries must be finite ($0\times\infty=$ NaN); obs_mask= means "missing, impute", not "ignore". | 6.17 |
| Loop cost $\approx N\,t_{\text{step}}+(N/k)\,t_{\text{sync}}$ | Read losses every $k$ steps; put data on the device once. Measured (2 000 steps): reading every step 235–240 ms, every 100 steps 141 ms. | float(x), print(x), np.asarray(x) all wait for the device. | 6.17 |
| $d$ $=$ number of scalar latents. Guides: $2d$, $d+d(d+1)/2$, $d(r+2)$; Adam keeps about 3 $\times$ the parameters | Example $d=69$, $r=10$: 138, 2 484, 828. Doubling $d$: mean-field $\times2$, full-rank step $\times4$, a dense Cholesky $\times8$, low-rank $\times2$. | Full-rank memory is $O(d^2)$, not $O(d^3)$; the cubic cost is factorizing a given covariance. | 6.17 6.17 |
Conjugate pairs: prior, likelihood, posterior, posterior predictive
Chapter 6.3's table, one row per pair. "Hyperparameters in, hyperparameters out": the update is arithmetic, with no integral, no sampling and no optimization. The posterior-predictive column answers "what does the next observation look like?" with the parameter integrated out. Most real models (hierarchical pooling, Student-t or Negative Binomial likelihoods, Laplace priors, logistic and NB regression) are not conjugate, which is why both projects run SVI; the conjugate sub-models remain exact tests for the inference code.
| Likelihood (data) | Conjugate prior | Posterior | Posterior predictive | Notes |
|---|---|---|---|---|
| Bernoulli / Binomial($n,\theta$): $k$ successes | Beta($\alpha,\beta$) | Beta($\alpha+k,\ \beta+n-k$) | beta-binomial (stats.betabinom(m, a, b)) | Conversion metrics. NumPyro Beta(concentration1=α, concentration0=β). 6.3 6.3 |
| Categorical / Multinomial($n,\mathbf p$): counts $\mathbf c$ | Dirichlet($\boldsymbol\alpha$) | Dirichlet($\boldsymbol\alpha+\mathbf c$) | Dirichlet-multinomial (stats.dirichlet_multinomial) | Categorical metrics. One share: Beta($\alpha_k+c_k,\ \alpha_0+n-\alpha_k-c_k$); $K=2$ is the Beta-Binomial. 6.3 |
| Poisson($\lambda t_i$): counts $y_i$, exposures $t_i$ | Gamma($a$, rate $b$) | Gamma($a+\sum y_i,\ b+\sum t_i$) | Negative Binomial: NB2(mean $a'/b'$, concentration $a'$) for exposure 1; SciPy nbinom(a', b'/(b'+1)) | Count metrics. SciPy gamma(a, scale=1/b); NumPy rng.gamma(a, 1/b); NumPyro Gamma(a, rate=b). 6.3 |
| Exponential($\lambda$): waiting times $y_i$ | Gamma($a$, rate $b$) | Gamma($a+n,\ b+\sum y_i$) | Lomax (Pareto II): SciPy lomax(a', scale=b') | Same Gamma family as above. 6.3 |
| Normal($\mu,\sigma^2$), $\sigma$ known: $\bar y$, $n$ | Normal($m_0, s_0^2$) on $\mu$ | Normal: precisions add, $\frac1{s_n^2}=\frac1{s_0^2}+\frac n{\sigma^2}$; precision-weighted mean | Normal($m_n,\ \sigma^2+s_n^2$) | The partial-pooling formula of 6.6 with the population as the prior. NumPyro Normal takes the sd. 6.3 |
| Normal($\mu,\sigma^2$), $\mu$ known | Inverse-Gamma($a,b$) on $\sigma^2$ | Inv-Gamma($a+\tfrac n2,\ b+\tfrac12\sum(y_i-\mu)^2$) | Student-t ($2a'$ degrees of freedom, centre $\mu$, scale $\sqrt{b'/a'}$) | Unknown noise level. 6.3 |
| Normal($\mu,\sigma^2$), both unknown | Normal-Inverse-Gamma | Normal-Inverse-Gamma | Student-t | This is where the Student-t predictive comes from. 6.3 |
| Not conjugate (compute the posterior): logistic and Poisson regression with Normal priors, Student-t likelihoods, Negative Binomial with unknown dispersion, Laplace priors, hierarchical models with unknown population parameters, mixtures. Tools: grid (1–2 parameters), Laplace approximation, NUTS, SVI (6.9). Choose a likelihood for the data, not for conjugacy. | ||||
The guide family: exact NumPyro parameter counts
An autoguide flattens every latent site into one unconstrained vector of length $d$ and puts a Gaussian on it; the guide type is the structure of its covariance. Counts below are the learned numbers, verified with svi.init in NumPyro 0.22 (the optimizer sees the full-rank guide's Cholesky factor as $d(d+1)/2$ free numbers, although auto_scale_tril is stored as a $d\times d$ matrix). With NumPyro's default rank=None, the low-rank guide uses $r=\text{round}(\sqrt d)$.
| Mean-field | Low-rank + diagonal | Full-rank | |
|---|---|---|---|
| NumPyro class | AutoNormal (per-site parameters), AutoDiagonalNormal (two vectors) | AutoLowRankMultivariateNormal(model, rank=r) | AutoMultivariateNormal |
| Covariance $\Sigma_q$ | $\text{diag}(s_1^2,\dots,s_d^2)$ | $WW^\top+D$, $W$ is $d\times r$ | $LL^\top$, $L$ lower-triangular |
| Learned numbers | $2d$ | $d(r+2)$ | $d+d(d+1)/2$ |
| Stored as | loc $(d)$, scale $(d)$ | auto_loc $(d)$, auto_cov_factor $(d,r)$, auto_scale $(d)$ | auto_loc $(d)$, auto_scale_tril $(d,d)$ |
| Growth | $O(d)$ | $O(dr)$ | $O(d^2)$ |
| Work per sample | $O(d)$ | $O(dr)$ to sample; density $O(dr^2+r^3)$ | $O(d^2)$ (matrix-vector product, triangular solve) |
| Captures | each parameter's own spread, no correlations | $r$ shared directions plus own spreads | every pairwise correlation |
| Misses | all correlations (variance $1/\Lambda_{ii}$, too narrow along ridges) | patterns needing more than $r$ directions: many separate pairs, long chains of neighbour ties | non-Gaussian shapes only (skew, heavy tails, funnels, several modes) |
| Training memory (rough) | bytes $\times$ learned numbers $\times4$ (parameters, Adam's two moments, one gradient); 4 bytes in float32, 8 in float64; full-rank adds $b\,d^2$ for the materialized $L$ | ||
| Where taught | 6.13 | 6.13 | 6.13 |
Worked numbers (computed from the three formulas; $r$ is the NumPyro default $\text{round}(\sqrt d)$ unless stated):
| Latent dimension $d$ | Mean-field $2d$ | Low-rank, default $r$ | Low-rank, $r=10$ | Full-rank $d+d(d+1)/2$ |
|---|---|---|---|---|
| $d=10$ | 20 | 50 ($r=3$) | 120 | 65 |
| $d=25$ | 50 | 175 ($r=5$) | 300 | 350 |
| $d=50$ | 100 | 450 ($r=7$) | 600 | 1 325 |
| $d=69$ (the worked example of Chapter 6.17) | 138 | 690 ($r=8$) | 828 | 2 484 |
| $d=100$ | 200 | 1 200 ($r=10$) | 1 200 | 5 150 |
| $d=250$ | 500 | 4 500 ($r=16$) | 3 000 | 31 625 |
| $d=1000$ | 2 000 | 34 000 ($r=32$) | 12 000 | 501 500 |
Say it right. Full-rank is $O(d^2)$ in storage and per-step work, not $O(d^3)$: the cubic cost belongs to factorizing a dense covariance you are given. Low-rank is cheaper because it stores less covariance, not because it is more accurate. All three guides are Gaussian in unconstrained space and are fitted by minimizing $KL(q\|p)$, so none of them is "the posterior". See 6.13 for the size-based selection rule (a design choice to validate on a reduced model) and 6.17 for measured costs.
SVI training-loop cheat sheet
From Chapter 6.14. The step is jitted; every decision about stopping lives in plain Python, outside the compiled function, because it needs concrete numbers.
| Piece | Rule | Why | See |
|---|---|---|---|
| Signs | svi.update returns (state, loss) with loss $=-\widehat{\text{ELBO}}$ of that step's draws. Lower loss is better. Improvement $\Delta_t=ELBO_t-ELBO_{best}=\mathcal L_{best}-\mathcal L_t$ (positive $=$ better). | Keep the sign for the decision. An absolute value would count a worsening as progress and reset patience. | 6.14 |
| Relative improvement | $\text{rel}_t=\dfrac{ELBO_t-ELBO_{best}}{|ELBO_{best}|+\epsilon}=\dfrac{\mathcal L_{best}-\mathcal L_t}{|\mathcal L_{best}|+\epsilon}$; meaningful if $\text{rel}_t \gt \tau$ (rule of thumb $\tau=10^{-4}$ to $10^{-5}$ per evaluation). | Scale-free: the ELBO's size grows with the data. Divide by $|\cdot|$ because the loss can be negative. $\epsilon$ guards zero and makes the rule absolute (threshold $\tau\epsilon$) near zero. | 6.14 |
| Patience | Stop when $P$ evaluations in a row had $\text{rel}\le\tau$ and step $\ge t_{min}$. Always stop at the maximum number of steps. | One evaluation is noisy; several in a row are strong evidence. False stops fall roughly like $b^P$; cost is about $P\cdot k$ steps. | 6.14 |
| Evaluate every $k$ steps | Mean of the last $k$ losses (sd $\sigma/\sqrt k$, lag $(k-1)/2$), or a fixed-key multi-particle Trace_ELBO(num_particles=M).loss(...). | Smoothing, and one host sync every $k$ steps instead of every step. | 6.14 |
| Best-state checkpoint | If cur < best_loss: store svi.get_params(state), step and loss. Return the checkpoint at the end (also after a NaN). Store the whole svi_state to resume. | The last iterate is not the best one; JAX arrays are immutable, so a reference is a snapshot. | 6.14 |
| Two separate decisions | "Best so far?" stores a checkpoint. "Meaningful gain?" resets patience. A tiny new record is stored but does not reset patience. | Keeps the checkpoint greedy and the stopping rule conservative. | 6.14 |
| Reading a finished run | Compare the stop step with the best step (a gap of about $P\cdot k$ is normal); look at the smoothed curve for plateaus; re-run with a stricter rule. | "Stopped by patience" is not "converged"; "hit max steps" is not success; "converged" is not "accurate". | 6.14 6.15 |
The decision core of the loop (the surrounding jax.jit(svi.update), window and logging lines are in Chapter 6.14):
for step in range(1, max_steps + 1):
state, loss = update(state, *args) # loss = -ELBO estimate (lower is better)
window.append(loss)
if step % eval_every: # evaluate only every k steps
continue
cur = float(jnp.mean(jnp.stack(window))) # smoothed: mean loss of the last k steps
window = []
if not np.isfinite(cur): # NaN or inf: stop and keep the best state
break
rel = (best_loss - cur) / (abs(best_loss) + 1e-8) if np.isfinite(best_loss) else np.inf
if cur < best_loss: # best so far -> checkpoint it
best_loss, best_params, best_step = cur, svi.get_params(state), step
bad = 0 if rel > rel_tol else bad + 1 # only a MEANINGFUL gain resets patience
if step >= min_steps and bad >= patience:
break
MCMC and NUTS diagnostics: rules of thumb
Chapter 6.10's checklist: stop at the first failure and fix it. Every threshold is a convention or rule of thumb, not a test. Passing every check means "no evidence of sampling problems", not "exact posterior"; none of this checks whether the model fits the data (6.8).
| Order | Diagnostic | Rule of thumb (as labelled in the chapters) | NumPyro / ArviZ | If it fails | See |
|---|---|---|---|---|---|
| 1 | Divergences | 0. Any divergence needs investigating: it means bias, not noise. | extra_fields=("diverging",), mcmc.get_extra_fields()["diverging"].sum(); print_summary prints "Number of divergences" | Reparameterize (non-centered), then raise target_accept_prob (0.9, 0.95, 0.99), then reconsider priors. Never delete divergent draws. | 6.10 6.7 |
| 2 | Split $\hat R$ | $\le1.01$ for every parameter and key derived quantity (Vehtari et al. 2021; the older 1.1 is too lenient). Necessary, not sufficient. | NumPyro r_hat is split $\hat R$ (numpyro.diagnostics.split_gelman_rubin); ArviZ az.rhat / az.summary use rank-normalized split $\hat R$ | Run longer, check for modes, reparameterize. | 6.10 |
| 3 | Effective sample size | Bulk-ESS and tail-ESS at least about 100 per chain, about 400 in total for 4 chains. | NumPyro n_eff (numpyro.diagnostics.effective_sample_size); ArviZ ess_bulk (centre), ess_tail (interval ends) in az.summary | Run longer or improve the geometry. | 6.10 |
| 4 | Traces, rank plots, tree depth | Overlapping fuzzy caterpillars; num_steps not stuck at the maximum ($2^{10}-1=1\,023$). | az.plot_trace, az.plot_rank; extra_fields=("num_steps",) | A hard geometry (very different scales or strong correlations): reparameterize or standardize; dense_mass=True also captures correlations (at $O(d^2)$ memory). | 6.9 6.10 |
| 5 | E-BFMI (energy) | Not below about 0.3. It compares the squared successive changes of the energy with the spread of the energy itself. Listed in the 6.15 diagnostics table, not taught in 6.10. | extra_fields=("energy",), then np.sum(np.diff(e)**2) / np.sum((e - e.mean())**2); ArviZ az.bfmi | Low values mean the sampler explores energy levels slowly (heavy tails or a hard geometry): reparameterize and re-check. | 6.15 |
| 6 | MCSE | $\text{sd}/\sqrt{\text{ESS}}$; for a probability $\sqrt{q(1-q)/\text{ESS}}$. Make it small compared with the precision your decision needs. ESS 400, $q=0.95$: 0.011. | ArviZ az.mcse, mcse_mean in az.summary | More draws, or report fewer digits. | 6.10 |
| 7 | Then the model | Posterior predictive checks and prior sensitivity (6.8). | Predictive(model, samples) | Change the model, not the sampler. | 6.8 |
Setup facts worth knowing: num_chains=4 with numpyro.set_host_device_count(4) (before JAX starts) to run chains in parallel; mcmc.get_samples(group_by_chain=True) gives shape (chains, draws, …); NUTS defaults are target_accept_prob=0.8, max_tree_depth=10, dense_mass=False, init_to_uniform; warmup draws are dropped for you. print_summary(prob=0.9) shows mean, std, median, the 90% HPDI ends labelled 5.0% and 95.0%, n_eff, r_hat.
Diagnostics for SVI (6.15)
| Diagnostic | What it checks | Rule of thumb |
|---|---|---|
| ELBO (loss) trace, smoothed | The optimizer has stopped improving | Relative change below your tolerance for your patience window. A flat ELBO means "optimizer done", not "posterior correct". |
| Several seeds and starting points | Optimization stability, local optima | Decision quantities agree across runs. |
| ELBO of different guides, same model and data | Which guide is closer in KL | A higher average ELBO (beyond seed-to-seed noise) wins. |
| Importance-sampling check (Pareto $\hat k$) | How far $q$ is from the posterior, from weights $p/q$ | $\hat k \lt 0.7$ usable (a rule of thumb from the literature). |
| Comparison with a clean NUTS run on a smaller problem | The family gap, directly | Differences small compared with the posterior sd; ignore gaps under about 2 NUTS MCSE. |
| Posterior predictive checks | Does the fitted model reproduce the data? | No systematic misfit. |
JAX and JIT cheat sheet: masking, keys, static shapes
The three things that bit people most in Chapters 6.16 and 6.17. The snippets were run as written (JAX 0.11.2, NumPyro 0.22.0).
| Situation | Wrong | Right | See |
|---|---|---|---|
| Pick a subset (a variant's or segment's rows) inside traced code | x[mask] or x[mask].mean() under jit: the output length depends on the data, so JAX raises NonConcreteBooleanIndexError. | Keep the shape fixed and mask arithmetically: jnp.where(mask, x, 0.0).sum() / jnp.maximum(mask.sum(), 1); in NumPyro dist.X(...).mask(m) or handlers.mask(mask=m). Or select in NumPy/pandas before tracing and accept one compile per shape. | 6.17 |
| Randomness | One key used twice, or np.random inside a jitted function (frozen at trace time). | k1, k2 = jax.random.split(key); fold_in(key, i) per step; a separate branch for svi.init and for Predictive; record the root seed. | 6.16 |
| Varying sizes | Jitting on the exact per-segment arrays: one compile per distinct size. | Pad to a few bucket sizes (64, 128, 256, …) and carry a mask; make shape-deciding settings (Fourier order, changepoint count) static. | 6.17 |
| Branching on data | Python if x > 0 on a traced value (TracerBoolConversionError). | jnp.where (both branches computed; double-where for log and division), lax.cond, or a plain if on a setting known before tracing. | 6.17 |
| Timing | Timing the first call, or timing without waiting (asynchronous dispatch). | Warm up once, then time with f(x).block_until_ready(); report compile time and step time separately. | 6.17 |
import jax, jax.numpy as jnp
from functools import partial
@jax.jit
def masked_mean(x, mask): # fixed shapes: x[mask] would fail here
return jnp.where(mask, x, 0.0).sum() / jnp.maximum(mask.sum(), 1)
key = jax.random.PRNGKey(0)
k_init, k_pred = jax.random.split(key) # same key + same call -> same numbers
@partial(jax.jit, static_argnames="order") # static: decides shapes, must be hashable
def fourier(t, order): # a new value of `order` compiles a new program
return jnp.stack([jnp.sin(k * t) for k in range(1, order + 1)])
safe_log = lambda x: jnp.where(x > 0, jnp.log(jnp.where(x > 0, x, 1.0)), 0.0) # "double where": gradient 0, not NaN at 0
Library conventions that silently change your answer
None of these raise an error. Each one gives a number that looks fine and answers a different question from the one you meant. Where a convention depends on how your own code is written, check which one it calls.
| What | Convention | See |
|---|---|---|
| Beta | NumPyro Beta(concentration1=α, concentration0=β): concentration1 counts successes. SciPy stats.beta(a, b) has the same $a=\alpha$, $b=\beta$. | 6.3 |
| Gamma | The chapters write Gamma($a$, rate $b$). NumPyro Gamma(concentration, rate); SciPy gamma(a, scale=1/b); NumPy rng.gamma(a, 1/b). | 6.3 |
| Normal and scale parameters | NumPyro Normal(loc, scale) takes the standard deviation, not the variance; the maths in these guides writes $N(\mu,\sigma^2)$ with the variance. Precisions use the variance. | 6.3 6.6 |
| Negative Binomial | Gamma-Poisson predictive is NB2 with mean $a'/b'$ and concentration $a'$: NumPyro NegativeBinomial2(mean, concentration), SciPy nbinom(n=a', p=b'/(b'+1)). Know which class and which prior (on $\alpha$ or $1/\alpha$) your code uses. | 6.3 |
svi.update | Returns (state, loss) with loss $=-$ELBO of that step (one particle by default: Trace_ELBO(num_particles=1)). svi.run has no early stopping. | 6.12 6.14 |
| Autoguide initialization | init_to_uniform: unconstrained locations drawn uniformly in $(-2,2)$; every scale starts at init_scale=0.1. A full-rank factor starts at $0.1\,I$; a low-rank factor starts at zero (no correlations). | 6.14 6.13 |
| Low-rank default | AutoLowRankMultivariateNormal(model, rank=None) uses $r=\text{round}(\sqrt d)$. | 6.13 |
LocScaleReparam | centered=None (the default) is a learnable 0.5 for SVI; NUTS does not learn it. For MCMC pass centered=0 explicitly. | 6.7 |
| Guide draws on constrained parameters | A Normal guide on $\log\sigma$ gives a log-normal $\sigma$: exp(loc) is the median, not the mean. | 6.11 |
| Minibatches | numpyro.plate("data", N, subsample_size=B) needs the full $N$ and applies $N/B$ to the data terms only. | 6.12 |
| Masks | .mask(m) and handlers.mask give masked observations zero log-probability; obs_mask= treats them as missing values to impute. | 6.17 |
| MCMC output | get_samples() returns only the post-warmup draws, pooled over chains; group_by_chain=True for diagnostics. NumPyro r_hat is split $\hat R$; ArviZ uses rank-normalized split $\hat R$. | 6.9 6.10 |
| float32 | JAX defaults to float32 (about 7 digits); tiny differences in the ELBO are rounding noise. 64-bit needs jax.config.update("jax_enable_x64", True) at the very start. | 6.16 |
Interview question bank
Fifty questions an interviewer could ask about this guide's material while you walk them through your Bayesian A/B framework or your Prophet-style forecasting model with its custom SVI loop. Most come straight from the chapters' "Say it right" boxes: posterior vs posterior predictive, flat is not uninformative, credible vs confidence intervals, partial pooling, centered vs non-centered, "NUTS is exact", KL direction and under-dispersion, full-rank vs low-rank guides, relative-ELBO stopping with patience, best-state checkpointing, JIT recompilation and boolean masking under jit.
How to use this bank. Read the question, answer it out loud in about a minute, and only then open the model answer. A strong answer usually has four parts: (1) a one-sentence definition in plain words, (2) a small number or picture, (3) where it lives in your project, (4) the trap you avoid. Where an answer mentions something about your own code (the guide-selection threshold, the stopping tolerance, which Negative Binomial class you call), check what your code actually does before you quote it. If an answer feels shaky, follow its link back to the chapter.
Bayesian inference, priors and conjugate models (6.1–6.3)
1. In your A/B framework, name the five Bayesian objects for one variant's conversion rate and say each in plain words.
The parameter $\theta$ is the variant's true conversion rate, an unknown number. The prior $p(\theta)$ is what I believe before the test, for example a Beta distribution. The likelihood $p(D\mid\theta)$ says how probable the observed $k$ conversions among $n$ users would be for each possible $\theta$ (data fixed, $\theta$ varying). The posterior $p(\theta\mid D)$ is my belief after the data: likelihood times prior, rescaled to area 1. The rescaling number is the evidence $p(D)$, the likelihood averaged over the prior. Finally the posterior predictive $p(\tilde y\mid D)$ describes new data, such as conversions among the next users, by averaging over every $\theta$ the posterior still allows. For a Beta prior and Binomial data the posterior is exactly Beta$(\alpha+k, \beta+n-k)$. A tiny check: candidate rates 5%, 10%, 15%, equal priors, 3 buyers in 20 visitors give posterior weights 0.12, 0.39, 0.49, and the chance the next visitor buys is 0.119, not 0.15. Chapter 6.1
2. Is the posterior just "likelihood times prior"? Do you ever compute the evidence in your NumPyro models?
It is proportional to likelihood times prior; the constant that makes it integrate to 1 is the evidence $p(D)=\int p(D\mid\theta)p(\theta)\,d\theta$. I never compute it. NumPyro builds $\log p(D\mid\theta)+\log p(\theta)$ from the sample statements, and SVI and NUTS use only that unnormalized log posterior and its gradient, because the constant disappears when you take a ratio (Metropolis, HMC) or a gradient of a log, and it is a constant in the ELBO. The evidence does matter for two things: Bayes factors between models, which swing strongly with prior width, and the ELBO, which is a lower bound on $\log p(D)$. Chapter 6.1 · evidence
3. What is the difference between the posterior and the posterior predictive? Which one is P(B beats A)?
The posterior describes the parameters, so $P(\theta_B \gt \theta_A\mid D)$ is a posterior probability: the share of paired draws where B is ahead. The posterior predictive describes future data. It adds the observation noise of the likelihood on top of the parameter uncertainty: draw $\theta$ from the posterior, then draw $\tilde y$ from the likelihood, repeat, and read means, quantiles and exceedance probabilities. In my forecasting model every forecast is a posterior predictive (Normal, Student-t or Negative Binomial noise on top of the trend, season, holiday and regressor terms). A plug-in prediction that uses one point estimate ignores parameter uncertainty and is always too narrow: with a Beta(5, 35) posterior and 50 new users the predictive sd is 3.46, the plug-in one 2.34. Chapter 6.1
4. A colleague says, "I used an uninformative flat prior." What do you say?
Flat on which scale? A constant density can be flat on at most one scale, because densities pick up a Jacobian when you change variables. Uniform(0, 1) on a conversion rate is a logistic hump on the log-odds, and a wide $N(0, 10^2)$ on the log-odds puts about 65% of its mass below 1% or above 99% on the rate. So I do not say "uninformative". I use weakly informative priors on a sensible scale (which is one reason for a global scaler), translate them into business units, and confirm them with a prior predictive simulation. Flat is also a prior, and it is not neutral. Chapter 6.2
5. How do you choose priors, and how do you know they are sensible? What would a prior predictive check show in your forecasting model?
For the A/B framework, the control rate can carry an informative Beta prior built from past tests (it is worth $\alpha+\beta$ pretend users, which I compare with the traffic), while the treatment effect gets a skeptical, weakly informative prior so a win comes from this test's data. In the forecasting model, $\delta_j\sim\text{Laplace}(0,b)$ says "most candidate slope changes do nothing, a few matter", and Normal priors regularize the Fourier, holiday and regressor coefficients. The check is to simulate whole series from the priors, with the real time index, holidays and regressors but no $y$, and ask whether they are plausible worlds: negative demand, absurd growth, wildly wiggly trends. In Chapter 6.2's example, vague trend priors made about 94% of simulated years absurd and weakly informative ones 1–2%. Prior scales only mean something relative to the data scaling, so I check what scaling the code applies. It uses domain knowledge, not the observed data. Chapter 6.2 · forecasting check
6. Derive the Beta-Binomial update, and say how strong a Beta(2, 18) prior is against 8 000 users per variant, and against a 40-user segment.
The likelihood is proportional to $\theta^k(1-\theta)^{n-k}$ and the prior to $\theta^{\alpha-1}(1-\theta)^{\beta-1}$. The binomial coefficient and the Beta function do not contain $\theta$, so they cancel from the posterior; the exponents add, and the kernel is that of Beta$(\alpha+k,\ \beta+n-k)$. The prior is worth $\alpha+\beta=20$ pretend users and the posterior mean is a weighted average, $w\,m_0+(1-w)\,k/n$ with $w=20/(20+n)$. At $n=8000$, $w\approx0.0025$, so the prior is irrelevant; for a 40-user segment $w=1/3$, so the prior (or, in a hierarchical model, the population) does real work. In NumPyro remember Beta(concentration1=α, concentration0=β), where concentration1 counts successes. Chapter 6.3 · weighted average
7. "A Bayesian estimate is biased toward the prior, so it is worse." Respond.
It is a compromise between the prior mean and the data rate, weighted by information. The pull adds a little bias but removes variance, which usually lowers the average error in small samples (the bias–variance trade-off of Guide 2), and it vanishes as $n$ grows because the weight on the prior is $(\alpha+\beta)/(\alpha+\beta+n)$. The same logic is why partial pooling beats raw segment averages. What I do owe the reader is the prior's worth next to the traffic: "prior mean 10%, worth 50 users, 8 000 per variant" gives a prior weight of about 0.006, so nobody needs to worry. Chapter 6.3
8. What is a conjugate prior? Which parts of your projects are conjugate, and why do you still need SVI?
A prior family is conjugate to a likelihood if the posterior stays in the family for every dataset, so updating just changes hyperparameters, with no integral or sampling. In the A/B framework a single variant's Beta-Binomial (conversions) and Dirichlet-Multinomial (categorical metrics) are exact, and Gamma-Poisson would be for a Poisson rate. As soon as I add hierarchical partial pooling across segments, a Student-t likelihood, or a non-conjugate prior, the evidence has no formula; the forecasting model, with Laplace priors on the $\delta_j$, many coefficients and Normal, Student-t or NB likelihoods with learned scale or dispersion, is not conjugate either. That is why both projects use SVI. The conjugate pieces are still valuable as exact unit tests: SVI on one variant should reproduce Beta$(\alpha+k,\beta+n-k)$ closely. Conjugacy is a convenience, not evidence that the prior is right. Chapter 6.3 · pairs
9. How does the Dirichlet-Multinomial update work, and why do whole-mix questions need joint draws?
Dirichlet$(\boldsymbol\alpha)$ plus counts $\mathbf c$ gives Dirichlet$(\boldsymbol\alpha+\mathbf c)$, one pseudo-count per category; the posterior mean of share $k$ is $(\alpha_k+c_k)/(\alpha_0+n)$, a weighted average of the prior share and the data share. A single share on its own is a Beta, Beta$(\alpha_k+c_k,\ \alpha_0+n-\alpha_k-c_k)$, so "did the Enterprise share rise?" can use that marginal. But the shares sum to 1 and are negatively correlated, so a question about the whole mix needs joint draws, with the same draw index for every category and variant. With two categories it reduces to the Beta-Binomial. Chapter 6.3
Posterior decisions and credible intervals (6.4)
10. How do you report an A/B result from posterior draws? What is the difference between P(B beats A) and "better by at least δ"? Someone says "B is 98% likely better, so it matters."
From paired posterior draws I report $P(\theta_B \gt \theta_A\mid D)$ (the share of the cloud above the diagonal, with Monte Carlo error $\sqrt{P(1-P)/S}$, at most 0.008 for 4 000 draws), $P(\theta_B-\theta_A \gt \delta\mid D)$ with $\delta$ the smallest gain worth shipping, fixed in advance and converted with the global scaler where needed, the probability of practical equivalence $P(|\theta_B-\theta_A| \lt \delta\mid D)$, and the posterior of the gap with an interval. The first is about direction, the second about size. A huge test with 10.0% vs 10.2% can give $P(B \gt A)=0.98$ while $P(\text{gap} \gt 0.5\text{ pt})=0.001$: almost surely positive, almost surely too small to matter. Chapter 6.4 · $P(B \gt A)$
11. Equal-tailed or highest-density interval: which do you report and why?
They agree for a symmetric one-humped posterior and differ otherwise. The equal-tailed interval cuts equal probability from both tails (just the 2.5% and 97.5% quantiles of the draws) and transforms with any increasing change of scale, so I can convert units afterwards. The highest-density interval is the shortest interval with that mass for a one-humped posterior, keeps equal density at its ends, can split into pieces for a multi-humped one, is noisier when computed from draws, and is not invariant to reparameterization. For a skewed small-sample rate (Beta(3, 40): ETI [0.015, 0.162], HDI [0.008, 0.145]) the difference is visible. I say which interval and which mass. Chapter 6.4
12. Your dashboard shows one number per variant. Is it the MAP, the posterior mean or the median? And is a Laplace prior on the changepoint slopes the same as asking for sparse changepoints?
From posterior or guide draws, the natural number is the mean or the median of the draws, not the MAP. They are different summaries of the same posterior, each best under a different cost: the mean for squared error, the median for absolute error, the mode (MAP) for "exactly right or nothing". For a symmetric posterior they agree; for a skewed one (a rare conversion rate, a relative lift) they spread apart, with mode below median below mean for a right skew. The median is the only one that survives an increasing change of scale, while the mode and the mean change, so a logit-scale guide and a rate-scale report have different modes. I say which summary I show. The same distinction answers the Laplace question: with $\delta_j\sim\text{Laplace}(0,b)$ the MAP is a soft-thresholding (lasso-like) solution with exact zeros, but the full posterior is a continuous density whose mean and median are shrunk toward zero and not exactly zero. A fit that approximates the posterior therefore gives "sparse-ish" slope changes: most tiny, none exactly zero. Chapter 6.4 · Chapter 6.2 · Guide 2, 5.3
13. What is the difference between a 95% confidence interval and a 95% credible interval?
A confidence interval comes from a procedure that captures the fixed true value in 95% of repeated experiments; it is the recipe that has the 95%, not this interval. A credible interval is a direct statement: given my model, prior and data, $\theta$ lies in it with probability 0.95. It depends on the prior, and it is calibrated on average over the prior, not at every fixed $\theta$. With a weak prior and lots of data they are numerically close (Bernstein–von Mises: for 60 of 500 a Wilson interval and a flat-prior credible interval agree to within 0.0001), but with small samples they differ, and a Wald interval can go negative (2 of 41) where a credible interval cannot. If a stakeholder asks "how often would this be wrong?" that is a frequentist question, so I answer by simulating the whole framework on synthetic experiments with known effects. Chapter 6.4 · Guide 2, 5.8
14. How do you compute and report a relative lift from posterior draws? Why not divide the interval ends?
Per draw: $L^{(s)}=\theta_B^{(s)}/\theta_A^{(s)}-1$, with the same draw index for both variants, then report the median and an equal-tailed interval next to the baseline. A ratio of two uncertain numbers is right-skewed, so its mean is above its median, and $E[\theta_B/\theta_A]\ne E[\theta_B]/E[\theta_A]$; dividing interval ends or averaging segment lifts is wrong. For the 50/500 vs 60/500 checkout example the median lift is 19.8% with a 95% interval of about $-15.7\%$ to $+70.5\%$: positive on average, far from certain. If the metric was standardized, I transform the draws back to raw units first, because a relative lift of standardized values (which can be near 0 or negative) is meaningless. Chapter 6.4
15. Describe a Bayesian decision rule for shipping a variant. Does peeking matter if you are Bayesian?
I write the rule into the analysis before launch: the primary metric, $\delta$ in that metric's units, a probability threshold, a maximum sample size, and guardrail checks, and I report the same few numbers every time from the same paired draws: $P(B \gt A)$, $P(\text{gap} \gt \delta)$, the practical-equivalence probability and the expected loss $E[\max(\theta_A-\theta_B,0)\mid D]$. The posterior itself is a valid summary at any moment, but how often a rule makes bad calls depends on how it is used: checking every day and stopping at the first crossing, or scanning many metrics and segments, raises the rate of shipping variants that are not better. So peeking is still a design question; I simulate the rule on A/A and A/B data at my real traffic and look schedule. Chapter 6.4 · Guide 2, 5.11
Hierarchical models, pooling and parameterization (6.5–6.7)
16. Explain hierarchical partial pooling across segments to a product manager. What does it assume?
Each segment has its own true value $\theta_g$, but the segments are drawn from a common population, $\theta_g\sim N(\mu,\tau^2)$, and the model learns the population centre $\mu$ and spread $\tau$ from all segments together. A small segment's noisy average is pulled toward the population; a large segment mostly keeps its own. The pull is not a taste choice: it is the posterior mean of a model, and its strength is set by how noisy the segment is compared with how different segments really are. The key assumption is exchangeability: before seeing data I would say the same thing about every segment, so they are interchangeable draws from one population. If I know something that separates segments (device, market size), I add it as a predictor so that only the left-over differences are treated as interchangeable. And control and treatment are not exchangeable groups: their difference is exactly what the experiment measures. Chapter 6.5 · pooling
17. "Partial pooling just averages each segment with the global mean." Correct this, with the formula.
It is a precision-weighted average whose weight depends on the segment: $\hat\theta_g=w_g\bar y_g+(1-w_g)\mu$ with $w_g=\dfrac{\tau^2}{\tau^2+\sigma^2/n_g}=\dfrac{n_g}{n_g+\sigma^2/\tau^2}$. Precisions add, and the population acts like $\sigma^2/\tau^2$ extra observations sitting at $\mu$. With $\sigma=12$, $\tau=4$ and segments of $n=1, 9, 36, 81$ the weights on the segment's own mean are $0.1, 0.5, 0.8, 0.9$. A quick check in a code review: take a segment's raw average and its standard error, compute $w=\tau^2/(\tau^2+SE^2)$ with the fitted $\tau$, and the reported estimate should sit about the fraction $w$ of the way from the population mean to the raw average. Complete pooling is the limit $\tau\to0$, no pooling is $\tau\to\infty$. Chapter 6.6
18. What is a "hyperparameter" in your hierarchical model? Is it like a learning rate?
No. In machine learning a hyperparameter is a training setting you tune. In a hierarchical model $\mu$ and $\tau$, the centre and spread of the segment population, are unknown quantities with their own priors (hyperpriors), and the posterior says what the data tell us about them. Group (local) parameters $\theta_g$ belong to one segment; global parameters ($\mu$, $\tau$, a shared noise scale) are shared. The only things I fix by hand are the constants inside the hyperpriors, such as a HalfNormal scale on $\tau$, and I check those with a prior predictive simulation, because with few segments the data say little about $\tau$ and its hyperprior matters. Chapter 6.5 · hyperpriors
19. Why does shrinkage reduce error? Who can lose, and what is the winner's curse for segments?
A raw segment average is unbiased but noisy. Pulling it toward the population adds a little bias and removes more variance: the error of a raw average is $s^2=\sigma^2/n$, and averaged over the population the partially pooled error is $w\,s^2$, smaller. Only segments truly far from the centre (beyond $\sqrt{2\tau^2+s^2}$) lose, a minority, and total error falls. The winner's curse is that the largest of many noisy estimates is on average too high, and the smallest, noisiest segments win most often, so "the treatment helped most in this 40-user segment" is usually luck. I answer "which segment benefited most?" with pooled lifts and $P(\text{lift}_g \gt 0\mid D)$, not the largest raw lift. Shrunk estimates understate the true spread of the segments, so for that I look at $\tau$ itself. Chapter 6.6 · rates and winner's curse
20. Why one global scaler for all groups instead of standardizing each group separately?
Because per-group standardization erases exactly what the model is trying to estimate. If every group is scaled by its own mean and sd, every group's average becomes 0, the between-group variance $\tau^2$ in the data is 0, and the model can only conclude "no segment differences". The scaler changes units, not which unknowns are explored. One global mean and sd keeps the group differences visible, and it makes one prior scale (say $N(0,1)$ on effects, a HalfNormal on $\tau$) mean the same thing in every segment. The ratio of $\tau^2$ to $\sigma^2$ is unchanged by a common rescaling. I convert back to raw units by multiplying by the global sd before reporting, and I never compute relative lifts of standardized values. Chapter 6.5 · identifiability
21. What is the difference between centered and non-centered parameterization, and when is each better?
Same model, different coordinates. Centered: $\theta_g\sim N(\mu,\tau^2)$, and the sampler or guide moves the $\theta_g$. Non-centered: $z_g\sim N(0,1)$ and $\theta_g=\mu+\tau z_g$, and it moves the $z_g$. With weak data per group the centered posterior of $\theta_g$ against $\log\tau$ is a funnel: the room for $\theta_g$ is proportional to $\tau$, so one step size cannot fit both the neck and the mouth. Non-centering removes $\tau$ from the prior part and makes the weak-data posterior round. With strong data per group the likelihood forces $\mu+\tau z_g$ to match the data and the non-centered posterior becomes a thin curved ridge, so centered is better. A rule of thumb: $w_g=\tau^2/(\tau^2+SE_g^2)$ small, go non-centered; near 1, centered. In the chapter's eight-group example, non-centered won with ordinary data (0 vs 54 divergences) and centered won with 100 times more data. In NumPyro: a $z$ site by hand or LocScaleReparam(centered=0). Chapter 6.7 · when centered is better
22. You get a few divergences, but R-hat is 1.00 and the ESS looks fine. Is the fit okay? And does the funnel matter if you fit with SVI?
No. $\hat R$ compares chains with each other; if every chain misses the same neck of the funnel they agree, and $\hat R$ looks perfect. A divergence is a direct signal: the leapfrog simulation's energy error blew up (over 1000 in NumPyro), usually because the step is too big for the local curvature, so that region is under-explored and the estimates are biased, typically toward too-large $\tau$. I look at where the divergences cluster (plot a group effect against $\log\tau$), switch the group effects to non-centered, and only then consider a higher target_accept_prob. With SVI there is no alarm at all: a Gaussian guide has the same width at every height, so in centered coordinates it silently understates the uncertainty in $\tau$, which mis-sets how much each segment is shrunk. Writing the effects non-centered and comparing the final ELBO of both forms is cheap insurance. Chapter 6.7 · SVI and the funnel
Posterior predictive checks, prior sensitivity and identifiability (6.8)
23. What is a posterior predictive check? Which statistics would you check for your forecasting model and for a Poisson count metric in the A/B framework?
I draw a posterior (or guide) sample $\theta^{(s)}$, simulate a replicated dataset with exactly the same size and design as the real one, and compare: real data that stand out among the fakes reveal a feature the model cannot produce. A fresh $\theta^{(s)}$ for each replicate, so parameter uncertainty is included. I choose a few statistics tied to decisions and not fitted directly by a parameter (a free mean always passes). For the forecasting model: zero-demand days (Negative Binomial vs Normal), the largest days and the 99th percentile (Student-t or NB tails), the weekday pattern, behaviour around holidays, and the lag-1 autocorrelation of residuals. For a Poisson metric in the A/B framework, the variance and the zero count in the control group, since Poisson assumes variance equals the mean; for the hierarchical model, the spread of the observed segment rates. Passing does not prove the model is right. Chapter 6.8 · choosing statistics
24. Is a posterior predictive p-value the Bayesian version of a p-value?
It looks similar but means something different. A classical p-value is computed under a fixed null and, when the null is true, is uniformly distributed. The posterior predictive p-value $P(T(y^{\text{rep}})\ge T(y)\mid y)$ averages over the posterior, which was fitted to the same data, so it is not uniform and tends to sit closer to 0.5: it is conservative. I use it as a misfit diagnostic for a chosen feature: values near 0 or 1 say the real data are extreme for the model on that feature (in the chapter's orders example, a Poisson model gave 0.000 for the variance and zero-day count while a Negative Binomial gave 0.50 and 0.38). It is not an error rate and not the probability that the model is true, and a moderately small value such as 0.05 is a real warning. Chapter 6.8
25. How do you do a prior sensitivity analysis, and where does it matter most in your projects?
Before looking at results I list the quantities I will act on (a decision probability, an interval, a forecast quantile), choose a small set of reasonable alternative priors (scales a factor of 2–3 wider and narrower, a heavier-tailed family, a different defensible centre, flat as one option), refit under each, and report the range and whether the decision flips. It matters most with little data per parameter, with scale parameters, and with dispersion and sparse effects. In the A/B framework: small segments and early looks (3/20 vs 8/20 gave $P(B \gt A)$ from 0.77 to 0.96 across priors, while with ten times the data all were about 1.00), and the prior on the between-segment spread $\tau$ when there are few segments. In the forecasting model: the Laplace scale $b$ on the $\delta_j$, holiday effects seen once or twice, and the Negative Binomial concentration or Student-t $\nu$, which (if learned) are mostly informed by a few extreme days. I never choose the prior after seeing which answer it gives. Chapter 6.8 · where priors bite
26. In your forecasting model a holiday always coincides with a promotion regressor. What happens to their coefficients, and what do you do?
In an additive model, components that always co-occur are not separately identified; only their sum is. The likelihood is flat along the "holiday up, promotion down" direction, so the posterior shows a strong negative correlation between the two effects and posterior sds close to the prior sds, and the split between them is written by the priors (with a combined effect of 20 and priors 10/10 versus 20/5, the split came out 9.8/9.8 versus 18.6/1.2). Forecasts are fine while they co-occur and wrong when they decouple. I report the combined effect, run a sensitivity analysis of the split, and if the separate effect matters I get data where they vary (a promotion on a non-holiday week) or use a documented informative prior from outside data. The same logic applies to trend versus yearly seasonality over a short history, and to a changepoint candidate placed right on a holiday. Chapter 6.8 · identifiability
Approximate inference, MCMC, HMC and NUTS (6.9–6.10)
27. Why can't you just compute the posterior? What are your options?
Bayes' theorem applies to every model. The obstacle is computational: $\log p(D\mid\theta)+\log p(\theta)$ is cheap to evaluate at any $\theta$, but the evidence and every posterior average are high-dimensional integrals with no closed form, and a grid costs $k^d$ (100 points and 10 parameters is $10^{20}$ evaluations). The options are MCMC (correlated draws, asymptotically exact, needs many evaluations and diagnostics), variational inference (a fitted distribution, fast, limited by the guide family), the Laplace approximation (a Gaussian at the mode, fastest and crudest) and importance sampling (weights that collapse in high dimension). Only conjugate models, such as one variant's Beta-Binomial, avoid the problem. My projects use SVI, which is fast and scales to many parameters (and, in the forecasting model, fits inside a JIT-compiled loop), and I defend it by checking against NUTS on a smaller version of the model. Chapter 6.9 · the menu
28. Explain Metropolis in a few sentences. Why does the evidence cancel? What are burn-in and thinning for?
Propose a move $\theta'=\theta+\varepsilon z$; accept it with probability $\min(1,\tilde p(\theta')/\tilde p(\theta))$, where $\tilde p$ is likelihood times prior; if rejected, record the current value again. The evidence $p(D)$ is the same in numerator and denominator, so it cancels. The rule satisfies detailed balance, so the posterior is the chain's stationary distribution, and after the chain forgets its start the visited states are correlated draws from it. Rejected steps are part of the output. Burn-in discards the early draws that still remember the starting point (modern samplers combine it with tuning and call it warmup). Thinning only saves storage: at $\rho=0.9$ and 10 000 draws the ESS is 526, and thinning by 10 lowers it to 483. Autocorrelation makes MCMC inefficient, not wrong, and the step size is judged by ESS, not acceptance (rules of thumb: 0.234 in high dimension, 0.44 in one). Chapter 6.9 · autocorrelation
29. Why is HMC more efficient than random-walk Metropolis? Is it because it accepts more proposals?
Not simply. A random walk with step $\varepsilon$ travels about $\varepsilon\sqrt n$ after $n$ steps, so crossing a distance $D$ costs about $(D/\varepsilon)^2$ steps; with momentum it costs about $D/\varepsilon$. HMC adds a random momentum, simulates a frictionless puck on the landscape $-\log\tilde p$ with the leapfrog integrator (using the gradient, which JAX provides by autodiff, and not needing the evidence), and accepts the end point with probability $\min(1,e^{-\Delta H})$. Because leapfrog nearly conserves energy, the end point can be far from the start and still be accepted (an energy error of 0.1 is accepted with probability 0.905), so successive draws are far less correlated. The honest measure is ESS per gradient evaluation, not acceptance rate. HMC needs continuous, differentiable parameters. Chapter 6.10 · why random walks are slow
30. What does NUTS add to HMC, and what does warmup tune?
Plain HMC needs a step size $\varepsilon$ and a trajectory length. NUTS doubles the trajectory in a random direction until it starts to turn back (a U-turn), up to $2^{10}-1=1023$ steps by default, then samples the next point from the trajectory in proportion to $e^{-H}$. During warmup it tunes the step size by dual averaging toward a target acceptance (target_accept_prob, default 0.8) and a mass matrix from the posterior variances (diagonal by default, full with dense_mass=True), then freezes both and records draws. It is not magic: it still has settings, it still depends on the parameterization (a centered funnel gives divergences), and a tiny adapted step or a tree depth stuck at the maximum signals a hard geometry. Chapter 6.10 · warmup
31. Which MCMC diagnostics do you check, in what order, and with what thresholds?
I stop at the first failure and fix it. First divergences: zero, otherwise reparameterize (non-centered), then raise target_accept_prob. Then split $\hat R$ at most 1.01 for every parameter and key derived quantity (the older 1.1 is too lenient), remembering that it is necessary but not sufficient. Then effective sample size, bulk and tail, at least about 100 per chain and about 400 in total. Then traces or rank plots that look like overlapping fuzzy caterpillars and a tree depth not stuck at the maximum; E-BFMI not below about 0.3 is listed with the SVI-vs-NUTS diagnostics in Chapter 6.15. Then the Monte Carlo standard error, sd$/\sqrt{\text{ESS}}$, small compared with the precision my decision needs: with ESS 400 and $q=0.95$ it is 0.011, so "0.95 vs 0.94" is not a real difference. Only then posterior predictive checks. All thresholds are rules of thumb. Chapter 6.10 · $\hat R$ · ESS
32. Someone says, "NUTS gives the exact posterior." Do you agree?
No. NUTS is asymptotically exact: as the number of draws grows, averages converge to the true posterior expectations. Any finite run has Monte Carlo error, about sd$/\sqrt{\text{ESS}}$, and may have convergence problems (stuck chains, divergences in a funnel, low ESS), which is why I check $\hat R$, ESS and divergences. It is also the posterior of my model, not the truth about the world, so I still need predictive checks, and floating-point arithmetic adds a little more. The diagnostics can reveal problems, never prove their absence. Even so it is a stronger guarantee than SVI's, whose error is a fixed family gap that more steps cannot remove. When I validate SVI against NUTS I call the NUTS run the reference, not the truth, and quote its MCSE. Chapter 6.15 · Chapter 6.10
Variational inference, the ELBO and guides (6.11–6.13)
33. What does variational inference minimise, and why that direction of the KL? What does it imply for your uncertainty?
VI picks a family $q_\phi$ (the guide) and minimises $KL(q_\phi\|p(\theta\mid D))$, the reverse KL. It uses that direction because $KL(q\|p)=-\text{ELBO}+\log p(D)$ and $\log p(D)$ is a constant in $q$: I need only the joint $p(\theta,D)$ and draws from my own $q$, never draws from the unknown posterior. The consequence is that the KL averages over $q$, so it punishes putting mass where the posterior has little and forgives missing mass: the fit is mode-seeking and zero-forcing, can lock onto one mode, and its intervals tend to be too narrow. The forward KL, which averages over $p$, is mass-covering (for a Normal $q$ it matches the mean and variance). So I treat SVI's uncertainty as an approximation to check against NUTS on a smaller version, not as the truth. Chapter 6.11 · why reverse
34. Does variational inference always underestimate the variance?
It is a strong tendency, not a guarantee, and it comes from the family, not from the optimizer. With reverse KL and a family that cannot match the posterior, VI usually under-estimates the spread most along correlations (mean-field), in heavy tails and when it drops a mode, and it can bias the means too. For a Gaussian posterior with correlation $\rho$, the best mean-field fit has exact means and each variance equal to the conditional variance $1/\Lambda_{ii}$, so each sd shrinks by $\sqrt{1-\rho^2}$ (0.44 at $\rho=0.9$, coverage 61%), while the variance of a sum or difference can come out too small or too large depending on the direction. In an A/B framework that pushes $P(\theta_B \gt \theta_A\mid D)$ toward 0 or 1 too early; in the forecasting model, forecast bands, which add many correlated pieces, are off. A full-rank or low-rank guide fixes correlations but not skewness or funnels. Chapter 6.11 · mean-field guides
35. Derive log p(D) = ELBO + KL in a few lines. Why do we maximise the ELBO?
For any $q$ that is positive wherever the posterior is: $\log p(D)=E_q[\log p(D)]=E_q\big[\log\frac{p(D,\theta)}{p(\theta\mid D)}\big]$, using Bayes' rule inside the log. Multiply and divide by $q(\theta)$: $=E_q[\log p(D,\theta)-\log q(\theta)]+E_q[\log q(\theta)-\log p(\theta\mid D)]$. The first term is the ELBO and the second is $KL(q\|p(\theta\mid D))\ge0$. So the ELBO is a lower bound on $\log p(D)$, with equality only at the posterior, and since $\log p(D)$ is fixed, maximising the ELBO is the same as minimising the KL. It needs only the joint and draws from $q$. Two readings are useful: fit plus entropy, and expected log-likelihood minus $KL(q\|\text{prior})$, which is how the Laplace prior's shrinkage enters SVI. Note these are two different KLs. Chapter 6.12 · two readings
36. What does svi.update return? Why is the ELBO estimate noisy, and what is the reparameterization trick?
svi.update returns (state, loss) with loss equal to the negative ELBO estimate for that step's draws, so improvement means the loss goes down. The estimate is a Monte Carlo average of $\log p(D,\theta_s)-\log q(\theta_s)$ over $S$ particles drawn from the guide: unbiased, with noise falling like $1/\sqrt S$, and NumPyro's Trace_ELBO uses one particle by default, so a single step's loss is mostly noise and I judge a smoothed trend. The reparameterization trick writes the draw as a smooth function of the guide's parameters and parameter-free noise, $\theta=\mu+\sigma\varepsilon$ with $\varepsilon\sim N(0,1)$ (for a full-rank guide $m+L\varepsilon$), so autodiff can send gradients through the sample with low noise. The alternative, the score-function estimator, works for discrete latents but is much noisier. Chapter 6.12 · reparameterization
37. Compare mean-field, full-rank and low-rank guides: parameter counts and what each can capture.
All three put a Gaussian on the flattened unconstrained latent vector of length $d$; they differ in the covariance. Mean-field (AutoNormal) is diagonal: $2d$ learned numbers, no correlations. Full-rank (AutoMultivariateNormal) is $LL^\top$ with a learned Cholesky factor: $d+d(d+1)/2$ numbers, every pairwise correlation, $O(d^2)$ memory. Low-rank (AutoLowRankMultivariateNormal) is $WW^\top+D$ with $W$ of size $d\times r$: $d(r+2)$ numbers (location, factor, diagonal scale), with NumPyro's default $r=\text{round}(\sqrt d)$; it keeps the $r$ strongest shared directions and treats the rest as independent, so it misses many separate pairwise ties and long chains. The three counts were verified with svi.init in NumPyro 0.22; for $d=50$ they are 100, 1 325 and 600 (with $r=10$). None is exact: all are Gaussian in unconstrained space and fitted by reverse KL. Chapter 6.13 · counts and memory
38. How does your forecasting model choose between a full-rank and a low-rank guide, and how would you validate that choice?
The choice is based on model size, the latent dimension $d$: every grid or PELT changepoint adds a $\delta_j$, every Fourier order adds two coefficients per seasonality, each holiday and regressor adds one, and the likelihood adds its scale (plus $\nu$ or a concentration). While the full-rank cost $d+d(d+1)/2$ fits a budget I take the richer guide, because it captures every trade-off (slope vs changepoint adjustments, holiday vs weekday); beyond it I switch to low-rank with $d(r+2)$ numbers and $r$ chosen from the spectrum of a smaller full-rank fit or a default like $\text{round}(\sqrt d)$. Memory, not accuracy, forces the switch. The threshold is a design choice I state as such, and I validate it on a smaller version of the model (shorter history, fewer changepoints): fit full-rank, low-rank and, if affordable, NUTS, compare the ELBO and the forecast intervals and changepoint summaries, and raise $r$ until they stop changing. I check the exact threshold and rank my code uses before quoting them. Chapter 6.13
39. "A full-rank guide is O(d³), so it is unusable." Correct or not?
Not quite. A full-rank guide stores a triangle of $d(d+1)/2$ numbers, so memory and the work per sample are $O(d^2)$: because NumPyro learns the Cholesky factor $L$ directly, sampling is a triangular matrix-vector product and the log-determinant is $2\sum\log L_{ii}$. The cubic cost, about $d^3/3$, appears only when a dense covariance has to be factorized or inverted (for example a Laplace approximation). Still, quadratic growth is what ends full-rank: at $d=1000$ it learns about half a million numbers against 12 000 for rank 10; at $d=20\,000$ about 200 million. Measured in the chapter on one CPU, going from $d=1000$ to 2 000 made a full-rank step about 4.8 times slower and a low-rank step 1.9 times. That is the engineering reason for a size-based switch. Chapter 6.17 · full-rank
The custom SVI loop, and SVI vs NUTS (6.14–6.15)
40. Walk me through your custom SVI training loop.
The loop has this shape (Chapter 6.14 builds a minimal version; the exact settings are the ones to check in my own code). First svi.init with a PRNG key, which runs model and guide once and creates the parameters. Then a jit-compiled svi.update called in a Python loop: the first call traces and compiles, later calls reuse the compiled program, and each returns (state, loss) with loss equal to the negative ELBO estimate. Every $k$ steps I evaluate on a smoothed loss (the mean of the last $k$ losses, or a fixed-key multi-particle ELBO), converting to a Python float only then. I compute the signed relative improvement over the best loss so far, $(\mathcal L_{best}-\mathcal L_{cur})/(|\mathcal L_{best}|+\epsilon)$. Two separate decisions follow: if it is the best so far I checkpoint the parameters; if the gain exceeds the tolerance I reset the patience counter, otherwise I add one. After a minimum number of steps, $P$ evaluations without a meaningful gain stop the run, and a maximum step count always does. A NaN stops it too. I return the checkpoint, not the last state. I would add logging of the best step and a low-noise evaluation if checkpoint quality mattered. Chapter 6.14
41. Why a relative ELBO criterion rather than an absolute threshold? What is the ε for, and what about the sign?
The ELBO is a sum over observations, so its size depends on the data and the model: the ELBO of a two-year daily series and of a ten-year one differ roughly by the ratio of their lengths. A fixed number of nats would be too strict for one and too loose for the other. The relative improvement $\text{rel}_t=\frac{ELBO_t-ELBO_{best}}{|ELBO_{best}|+\epsilon}$ is unchanged if the whole objective is multiplied by a positive constant, so one tolerance means the same "converged" everywhere (a rule of thumb is $10^{-4}$ to $10^{-5}$ per evaluation, compared with the evaluation noise). $\epsilon$ prevents division by zero, and near zero it turns the rule into an absolute one with threshold $\tau\epsilon$. On the sign: NumPyro reports the loss, the negative ELBO, so improvement is best loss minus current loss. I keep the sign and divide by the absolute value, because an absolute value in the numerator would count a worsening as progress. The rule is not shift-free: rescaling $y$ shifts the ELBO and moves the effective threshold. Chapter 6.14 · signs
42. What is patience for? What is premature stopping, and how would you detect it?
Patience exists because the ELBO estimate is noisy: one evaluation without improvement is weak evidence, several in a row is strong evidence. Requiring $P$ consecutive evaluations without a relative gain above $\tau$ makes noise-caused stops roughly geometrically unlikely in $P$, at a linear cost of about $P\cdot k$ extra steps, and a minimum number of steps protects the early transient. Premature stopping is stopping while the model was still improving meaningfully, because noise or a plateau longer than $P\cdot k$ steps hid the progress (a late group of changepoint slopes starting to move can produce exactly that plateau-then-drop shape). Its cost is a worse guide (lower ELBO, and spreads or correlations that have not finished settling), and it is invisible in the loss log unless I look. To detect it I compare the stop step with the best step (a gap of about $P\cdot k$ is normal), inspect the smoothed curve, and re-run once with a stricter rule (smaller $\tau$ or larger $P$) to see whether the ELBO or the reported intervals change. "Stopped by patience" is not "converged". Chapter 6.14 · the whole rule
43. Why checkpoint the best state? Does that return the optimal parameters?
With a constant learning rate SVI's parameters jitter around the optimum and the objective is a noisy estimate, so the last state is not the best one and training can degrade late (a blow-up or a NaN after many good steps). At each evaluation I store the parameters if the smoothed loss is the lowest so far and return them at the end, also after a NaN. It is insurance more than a big gain on smooth runs. It does not return the optimal parameters: the minimum of noisy values is optimistic (a winner's curse), so the checkpoint's measured loss is slightly better than its true loss, less so with smoothing or a fixed-key multi-particle evaluation. In JAX arrays are immutable, so keeping a reference to svi.get_params(state) is already a snapshot (copy only if the update is jitted with donate_argnums). I store the step, loss, seed and settings with it, and the whole svi_state if I want to resume. Chapter 6.14
44. When would you use SVI and when NUTS?
I decide on four questions: size (rows $N$ and latent dimension $d$), stakes (do tails, correlations or intervals drive the decision?), shape (Gaussian, correlated, funnel, several modes) and schedule (once, daily, real time). NUTS is the better tool when it is affordable and the details matter: no family gap, and its failures show up in $\hat R$, ESS and divergences; it is also the right reference while a model is being developed. SVI wins for huge $N$ (minibatches, which plain NUTS lacks), thousands of latents, many refits and latency budgets, because a fitted guide is instant to sample. A forecasting model that is refit as data arrive, with many correlated latents and a JIT-compiled loop, suits SVI, and an A/B platform with many metrics and experiments benefits from one fast engine. The honest addition is how the SVI answers are checked, using NUTS on individual experiments or on a smaller version of the forecasting model. A slow NUTS usually signals bad geometry that also hurts SVI, so I reparameterize before switching. Chapter 6.15
45. How do you know your SVI posterior is good enough? "The ELBO converged" is not an answer.
A flat smoothed ELBO tells me the optimizer is done. It says nothing about the family gap, because $\log p(D)=\text{ELBO}+KL$ and $\log p(D)$ is unknown. So I compare: guides on the same data (a higher average ELBO beyond seed-to-seed noise means closer in KL), several seeds and starts, posterior predictive checks, and a NUTS run on a problem small enough to afford (a subset, a few segments, a shorter history, fewer changepoints) whose diagnostics are clean. I compare the quantities my decisions depend on, such as means, 90% intervals, $P(\theta_B \gt \theta_A\mid D)$ and forecast quantiles, relative to the posterior sd and ignoring gaps smaller than about two NUTS MCSEs, pick the cheapest guide that passes, run it at scale and re-validate when the model, data or guide change. A subset posterior is usually a harder test for a Gaussian guide, but some problems appear only at full scale, hence a smaller version of the full model as a second check. Chapter 6.15 · diagnostics for each
JAX and JIT (6.16–6.17)
46. What is a pure function, and why does JAX care? What happens to a print or a global variable inside a jitted function?
A pure function's output depends only on its explicit inputs, and it has no side effects (no printing, no global changes, no hidden random state). jit, grad and vmap trace the function once with tracers that carry only shape and dtype, record the operations, and replay the recording, which is valid forever only if the function is pure. Otherwise it often does not crash but gives quietly wrong results: a global is baked in as a constant, a print fires once at trace time (use jax.debug.print), and np.random is frozen. In my NumPyro models that means data, prior scales and switches such as the likelihood family are arguments, not globals edited later, and Python logging inside the model fires only at trace time; I record values with numpyro.deterministic. The chapter measured it: in 100 jitted updates of a small Beta-Binomial model the Python body ran once, with a tracer in place of the number. Chapter 6.16 · NumPyro as traced functions
47. How do PRNG keys work in JAX and NumPyro, and why are they explicit?
JAX functions are pure, so there is no hidden global random state: randomness is an explicit input. A key is a small array that fully determines a random draw; the same key and the same call give the same numbers, and jax.random.split(key, n) makes new independent keys. The rules are to never reuse a key for two things that should be independent, to treat a split parent as used up, and to pass keys as arguments. That keeps results reproducible under jit, vmap and parallel chains. In practice: one key for svi.init, let the SVI state carry and split its own key during training, and a separate key for Predictive, so changing the number of forecast draws cannot change the fitted parameters; in an A/A or power simulation, one key per simulated experiment, or the "independent" experiments are copies of each other. I record the root seed with the guide type, optimizer and iteration count. Chapter 6.16
48. What does jit do? Why is the first call slow, and what triggers recompilation?
On the first call for each new input signature JAX traces the Python function with abstract values (shape and dtype), records a jaxpr, lowers it to XLA and compiles an executable, which is cached under a key made of the function, the argument pytree structure, shapes, dtypes and static arguments. The first call pays trace plus compile (roughly 0.4–0.75 s for an SVI step on one CPU in the chapter's measurement); later calls take a fraction of a millisecond against tens of milliseconds eager. A new shape, dtype or static-argument value recompiles; new array values never do. If I refit the forecasting model with one more day of history, the data shape changes and each refit compiles once, which is fine; what hurts is recompiling inside a run or once per segment size. I time with a warm-up call and block_until_ready() because dispatch is asynchronous, and I report compile time and step time separately. Chapter 6.17 · compile once, run fast
49. You hit a problem with boolean masking in the A/B framework. What went wrong, and how do you fix it?
Selecting a variant's or segment's rows with x[mask] produces an array whose length depends on the data. A compiled program must know every shape in advance, so inside traced code (a jitted step, an SVI update, a NumPyro model) JAX raises NonConcreteBooleanIndexError; and if the mask was computed earlier from concrete data, every new subset size forces a recompile. It is not a JAX bug, it is a consequence of compiling for shapes. There are two clean designs, and I explain whichever my code uses. (a) Do all boolean selection in NumPy or pandas outside the traced code and accept one compile per distinct shape, fine for a handful of segments. (b) Keep fixed shapes: pass a group index and a mask into the model and use masked sums, jnp.where(mask, x, 0).sum() / jnp.maximum(mask.sum(), 1), or NumPyro's .mask(m) and handlers.mask, so one compiled program serves every experiment (padding to a few bucket sizes helps). In hand-written arithmetic masked-out entries must still be finite, since $0\times\infty$ is NaN (NumPyro's mask swaps padded observations for a safe value, but padded covariates that feed the mean are still your job); the chapter checked that the masked log density on padded data equals the unpadded one. Chapter 6.17
50. Where do grad, vmap and scan show up in your projects?
Every SVI step is a value_and_grad of the Monte Carlo ELBO with respect to the guide's parameters, handed to Adam; for a custom likelihood, such as a hand-written Negative Binomial parameterization, a one-time float64 gradient check against central differences catches sign and parameterization bugs. vmap is how a one-draw function becomes batched array code: the posterior predictive for $T$ future days is an $(S,T)$ array, NumPyro's Predictive(parallel=True) is a vmap over draws, and any per-draw decision quantity (lift, expected loss, $P(\theta_B-\theta_A \gt \delta)$) or an analysis across many segments is naturally vectorized. scan is for truly recursive pieces, such as an AR term on residuals or a local-level trend; my piecewise-linear trend can be written without a loop (slope from a changepoint matrix times $\delta$), so it is not needed there. NumPyro's own svi.run runs its loop as one lax.scan; a custom loop with early stopping cannot, because the stopping decision is made in Python between steps. Chapter 6.16 · vmap · scan
Where to go next
This guide taught you to write uncertainty as a distribution and to compute it: priors, likelihoods and posteriors, conjugate models, credible intervals and decisions, hierarchical pooling, model checking, and then MCMC, NUTS, variational inference, the ELBO, SVI, guides, your training loop, JAX and JIT. Here is what it used from Guides 1 and 2, where Guide 4 picks it up, and a short plan for the rest of the series.
How this guide connects to the other three
Each row starts from an idea and points to where it is used next. Links to other guides open them at the right chapter.
Coming from Guide 1 · Probability & Data
| From Guide 1 | Used in this guide for |
|---|---|
| Bayes' theorem, base rates, $P(A\mid B)\ne P(B\mid A)$ (4.3) | Bayesian inference as the same rule applied to a parameter (6.1) |
| Bernoulli, Binomial, Poisson, Negative Binomial (4.7, 4.8) | Likelihoods and conjugate pairs (6.3); replicated count data in posterior predictive checks (6.8) |
| Normal, Student-t, Laplace (4.9) | Normal-Normal updating (6.3); Laplace as a shrinkage prior (6.2) |
| Beta and Dirichlet, pseudo-counts (4.11) | Priors worth $\alpha+\beta$ observations and the Beta-Binomial and Dirichlet-Multinomial updates (6.3) |
| Law of total variance, within vs between (4.6) | Observation noise vs parameter uncertainty (6.1); $\tau^2$ vs $\sigma^2$ in hierarchical models (6.5) |
| LLN and CLT, the $1/\sqrt n$ error (4.13) | Monte Carlo error of posterior averages, MCSE and ESS (6.9, 6.10) |
| Covariance and correlation (4.15) | Posterior correlations, mean-field vs full-rank and low-rank guides (6.11, 6.13) |
| Standardization and the global scaler (4.18) | One prior scale meaning the same everywhere, and not erasing $\tau$ (6.2, 6.5, 6.8) |
Coming from Guide 2 · Estimation, Inference & Experiments
| From Guide 2 | Used in this guide for |
|---|---|
| Likelihood, MLE vs MAP, MAP is not a posterior (5.2) | The likelihood inside Bayes' theorem (6.1); mean, median and mode as different summaries (6.4) |
| Bias–variance, MSE, shrinkage (5.1) | The pull toward the prior (6.3) and why partial pooling lowers error (6.6) |
| Regularization as a prior, ridge and lasso (5.3) | Shrinkage and regularizing priors (6.2); Laplace posteriors that are shrunk but not sparse |
| Confidence intervals (5.8) | Credible vs confidence intervals (6.4) |
| Hypothesis tests, experiment design, peeking and multiple metrics (5.6, 5.10, 5.11) | Decision rules, practical thresholds $\delta$ and simulating a rule's error rates (6.4) |
| Regression, GLMs, the design matrix, collinearity (5.13, 5.14) | Identifiability and which component gets the credit (6.8); Negative Binomial and log links in the forecasting likelihood |
| Covariance matrices, the multivariate Normal, Cholesky factors (5.15) | Full-rank and low-rank Gaussian guides and their parameter counts (6.13) |
| KL divergence in t-SNE (5.17) | KL and the direction that variational inference minimizes (6.11) |
Going to Guide 4 · Time Series & Bayesian Forecasting
| From this guide | Used next for |
|---|---|
| Priors, shrinkage and the Laplace prior (6.2, 6.6) | Laplace priors on trend changes (7.10); the changepoint grid and PELT (7.8, 7.9) |
| Prior predictive checks and prior sensitivity (6.2, 6.8) | Fourier seasonality, holidays and regressors with sensible priors (7.11, 7.12) |
| Identifiability: trend, season, holiday and regressor competing for the same bump (6.8) | The additive forecasting model (7.7) and its components |
| Posterior predictive distributions and checks (6.1, 6.8) | Predictive distributions and forecast checks (7.14); coverage and calibration (7.16); residual diagnostics (7.17) |
| Likelihood choice and the Negative Binomial parameterizations (6.3, 6.8) | Normal, Student-t and Negative Binomial forecast likelihoods (7.13) |
| SVI, the guide family and the training loop (6.12–6.15) | Model complexity across both projects (7.18); reproducibility and monitoring (7.19) |
| JAX, JIT and numerical care (6.16, 6.17) | Numerical stability and production habits (7.19) |
| Everything above | The capstone: one set of ideas, two projects (7.20) |
A short study plan
- Close this guide properly. Reread the notebook boxes of the P0 chapters: Bayesian inference and the predictive (6.1), conjugate models (6.3), credible intervals and decisions (6.4), hierarchical models and shrinkage (6.5, 6.6), MCMC and NUTS (6.9, 6.10), VI, KL and the ELBO (6.11, 6.12), guides (6.13), the training loop (6.14) and JAX/JIT (6.16, 6.17). Then answer the question bank out loud, one minute per question, before you open the model answers.
- Do the checkout example both ways from memory. Beta posteriors for 50/500 and 60/500, $P(\theta_B \gt \theta_A\mid D)$, the gap with an interval, the median lift, and $P(\text{gap} \gt \delta)$ for a $\delta$ you choose; then the frequentist numbers from Guide 2. Explain why they roughly agree and why the sentences differ.
- Simulate your own decision rule. Take the rule your A/B framework uses, simulate A/A and A/B data at your real traffic and look schedule, and write down its false-ship rate and its power (6.4, Guide 2, 5.11).
- Hierarchical models by hand and by machine. Compute the weight $w_g=\tau^2/(\tau^2+\sigma^2/n_g)$ for a few segment sizes; simulate segments with known lifts (most zero, a few positive); fit centered and non-centered versions and compare divergences and the final ELBO (6.6, 6.7).
- Check an SVI fit. On a small model run NUTS with four chains and pass the checklist (diagnostics table), then fit mean-field, full-rank and low-rank guides and compare decision quantities with the NUTS reference (6.13, 6.15).
- Write the loop from memory. Thirty lines:
svi.init, a jitted update, an evaluation every $k$ steps, the signed relative improvement, the checkpoint, the patience counter. Then change $\tau$, $P$ and $k$ and predict what happens before you run it (loop cheat sheet). - Break and fix things in JAX. Reproduce the boolean-masking error and both fixes, time the first call against later calls, and split keys for an A/A simulation (JAX cheat sheet).
- Guide 4, your forecasting model component by component. Chapters 7.1–7.6 give the time-series basics, 7.7–7.13 go through the model piece by piece, 7.14–7.17 cover predictive distributions, evaluation and diagnostics, and 7.18–7.20 tie both projects together. If you want to see where the ideas of this guide are used first, start with 7.10, 7.13 and 7.14.
- Every week. Run one "Code it" block, change one number and predict the result before you run it. Write one paragraph explaining an idea from your projects to an interviewer, in plain words.
How to make it stick
- Simulate before you trust a formula. Twenty lines of NumPy check most claims in this guide: the Beta-Binomial update, the shrinkage weight, the funnel, the mean-field variance factor $\sqrt{1-\rho^2}$, the false-ship rate of a decision rule.
- Always say "of what". Which posterior summary (mean, median, mode), which interval and which mass, which KL direction, which guide family, which NB or Beta parameterization, which scale (raw or standardized).
- Separate "the optimizer is done" from "the posterior is right". A flat ELBO tells you about the first; comparisons with other guides and with NUTS tell you about the second.
- Say "asymptotically exact", never "exact". And say "approximation within the guide family" for SVI.
- Check the geometry before you blame the sampler. Divergences, a funnel, collinear components and a bad parameterization hurt NUTS and SVI alike.
- Keep your notebook. The "Write this in your notebook" boxes, copied by hand, are the fastest revision sheet you can have.
Well-known resources
- Gelman, Carlin, Stern, Dunson, Vehtari & Rubin, Bayesian Data Analysis (3rd ed.): the reference for hierarchical models, posterior predictive checks and computation; and Gelman et al. (2020), "Bayesian Workflow" on how the pieces fit together.
- McElreath, Statistical Rethinking: a gentle, simulation-first introduction to priors, hierarchical models, divergences and non-centered parameterization.
- Betancourt (2017), "A Conceptual Introduction to Hamiltonian Monte Carlo" and Hoffman & Gelman (2014), "The No-U-Turn Sampler": HMC, the leapfrog integrator and NUTS.
- Vehtari, Gelman, Simpson, Carpenter & Bürkner (2021), "Rank-normalization, folding, and localization: an improved $\hat R$": the source of the $\hat R \le 1.01$ and ESS rules of thumb.
- Blei, Kucukelbir & McAuliffe (2017), "Variational Inference: A Review for Statisticians", Hoffman, Blei, Wang & Paisley (2013), "Stochastic Variational Inference" and Kingma & Welling (2014), "Auto-Encoding Variational Bayes" (the reparameterization trick): VI, the ELBO and SVI.
- The NumPyro and JAX documentation (in particular JAX's "Thinking in JAX" and "The Sharp Bits"): the final word on autoguides,
Predictive, handlers, jit, PRNG keys and static shapes. - Taylor & Letham (2018), "Forecasting at Scale": the Prophet paper, if you want the model your forecasting project is modelled on.
Companion guides
This guide stands on Guide 1 · Probability & Data and Guide 2 · Estimation, Inference & Experiments, and on three earlier guides: the Linear Algebra guide (Cholesky factors, eigenvalues), the Calculus guide (gradients and automatic differentiation) and the Optimization guide (Adam, step sizes, noisy gradients). Next stop: Guide 4 · Time Series & Bayesian Forecasting.