Bayesian Gaussian Mixture Model¶
Overview¶
A finite Gaussian mixture model treats each observed value as a draw from one of \(K\) Gaussian components. The components share three per-component vector parameters: the mixing weights on the simplex, the component locations, and the component scales. A Dirichlet prior governs the mixing weights, and independent priors govern the locations and scales.
Rather than sample a discrete per-row component label and integrate it out by hand, this example scores each row directly against the mixture through the MixtureNormal likelihood. The per-row marginal is the closed-form convex combination
so the model carries no discrete latent: MixtureNormal integrates the component assignment out analytically, evaluating the marginal by log-sum-exp over the \(K\) components.
QVR source¶
# Bayesian Gaussian Mixture Model
#
# A finite Gaussian mixture with K components. Per-component means,
# scales, and mixing weights are latents on the Component plate; each
# observed row is drawn from the resulting mixture. The per-row
# component assignment is integrated out in closed form by the
# MixtureNormal likelihood, so the model carries no discrete latent.
#
# Generative structure:
#
# probs ~ Dirichlet(alpha) Component-simplex mixing weights
# mu_k ~ Normal(0, 5) per-component mean
# sigma_k ~ HalfNormal(1) per-component scale
# r_n ~ MixtureNormal(probs, mu, sigma) observed row
#
# Per-row marginal likelihood (closed form):
#
# p(r_n) = sum_k probs[k] * Normal(r_n; mu[k], sigma[k]).
#
# MixtureNormal takes three per-component vector parameters (the
# mixing weights, the locations, and the scales), each shared across
# every row of the Resp plate, and scores each row against the
# K-component mixture they define.
#
# Resp is the plate the rows are observed over; Weights is the
# value space of the returned mixing weights, a point of the
# Component simplex embedded in R^K.
composition log_prob [level=algebra]
object Component : FinSet 3
object Resp : FinSet 100
object Weights : Real 3
program gmm(alpha : Real) : Resp -> Weights
sample probs <- Dirichlet(alpha) [over=Component]
sample mu : Component <- Normal(0.0, 5.0)
sample sigma : Component <- HalfNormal(1.0)
observe r : Resp <- MixtureNormal(probs, mu, sigma)
return probs
export gmm
Walkthrough¶
composition log_prob [level=algebra] selects the log-probability semiring so the program's Score effect accumulates log-densities additively. object Component : FinSet 3 and object Resp : FinSet 100 declare the two discrete plates: \(K = 3\) mixture components and \(N = 100\) observed rows, and object Weights : Real 3 is the value space of the returned mixing weights, a point of the component simplex embedded in \(\mathbb{R}^K\). program gmm(alpha : Real) : Resp -> Weights parameterises the program by the Dirichlet concentration and declares that what it returns is a mixing-weight vector rather than a row index.
The three sample steps draw the shared per-component parameters:
sample probs <- Dirichlet(alpha) [over=Component]draws the mixing weights as a single point on theComponentsimplex;over=Componentnames the family's event axis (Dirichlet event-rank 1).sample mu : Component <- Normal(0.0, 5.0)draws the \(K\) component locations as plate-bound continuous latents, one per component.sample sigma : Component <- HalfNormal(1.0)draws the \(K\) positive component scales the same way.
observe r : Resp <- MixtureNormal(probs, mu, sigma) scores each observed row against the \(K\)-component mixture the three shared vectors define. MixtureNormal takes the mixing weights, the locations, and the scales as three per-component vectors, each broadcast across every row of the Resp plate, and returns the per-row marginal \(\sum_k \mathrm{probs}[k]\,\mathcal{N}(r_n;\,\mu[k],\,\sigma[k])\) in closed form. No component-assignment latent is sampled: the mixture likelihood integrates it out.
return probs projects the program's joint kernel onto the mixing-weight site.
Try it¶
The short fits below demonstrate the API. Assess convergence with multiple chains and diagnostics before interpreting a posterior.
Generating synthetic data¶
import torch
from quivers.dsl import load
torch.manual_seed(0)
prog = load("docs/examples/source/mixture_model.qvr")
fit = prog.gmm(alpha=1.0)
model = fit.morphism
K, N = 3, 100
true_probs = torch.tensor([0.4, 0.35, 0.25])
true_mu = torch.tensor([-3.0, 0.0, 3.0])
true_sigma = torch.tensor([0.5, 0.7, 0.4])
comps = torch.distributions.Categorical(true_probs).sample((N,))
r = torch.distributions.Normal(true_mu[comps], true_sigma[comps]).sample()
observations = {"r": r, "probs": true_probs}
x_in = torch.zeros(N, 1)
The synthetic rows are drawn by sampling a component per row and then a value from that component's Gaussian, but the model never sees the component labels: only the rows r and the fixed mixing weights probs enter observations. The per-component locations mu and scales sigma remain unobserved latents and are recovered by SVI.
SVI fit¶
from quivers.inference import AutoNormalGuide, ELBO, SVI
torch.manual_seed(1)
guide = AutoNormalGuide(model, observed_names={"r", "probs"})
optim = torch.optim.Adam(
list(model.parameters()) + list(guide.parameters()), lr=1e-1,
)
svi = SVI(model, guide, optim, ELBO(num_particles=1))
losses = [svi.step(x_in, observations) for _ in range(1200)]
recovered = sorted(guide.loc_mu.detach().flatten().tolist())
print(f"initial loss: {losses[0]:.2f}")
print(f"final loss: {losses[-1]:.2f}")
print("recovered means: [" + ", ".join(f"{m:.2f}" for m in recovered) + "]")
print("true means: [-3.00, 0.00, 3.00]")
The variational locations guide.loc_mu hold the recovered component means. Because the mixture ELBO is multimodal, the fit is sensitive to initialisation and to the mixing weights held fixed at probs; a sharper separation between components makes the locations easier to recover.
NUTS posterior¶
from quivers.inference import MCMC, NUTSKernel
torch.manual_seed(2)
kernel = NUTSKernel(step_size=0.05, max_tree_depth=3, target_accept=0.8)
mc = MCMC(kernel, num_warmup=15, num_samples=15, num_chains=1)
result = mc.run(model, x_in, observations)
print(f"acceptance: {float(result.acceptance_rates.mean()):.2f}")
print(f"divergences: {int(result.divergence_counts.sum())}")
Categorical perspective¶
The per-row likelihood MixtureNormal(probs, mu, sigma) is a Kleisli arrow \(\mathsf{Resp} \to \mathsf{Resp}\) in the Giry monad. It is a finite convex combination of the \(K\) Gaussian component measures, weighted by the categorical measure \(\mathrm{probs}\) on Component: the Giry-monad mixture operation that draws a component from \(\mathrm{Categorical}(\mathrm{probs})\), then a value from the chosen Gaussian, and keeps the marginal on the value. Equivalently, MixtureNormal is the pushforward of the joint component-and-value measure along the projection \(\mathsf{Component} \times \mathbb{R} \to \mathbb{R}\), which is exactly the closed-form marginal \(\sum_k \mathrm{probs}[k]\,\mathcal{N}(\cdot;\,\mu[k],\,\sigma[k])\). Because the component index is integrated out inside the likelihood rather than sampled, the program carries no discrete latent to marginalise and every site it holds is continuous.
See also¶
- Latent Dirichlet Allocation, the grouped discrete-mixture generalisation whose per-word topic assignment is integrated out by a scoped
marginalizeblock.