哈喽,大家好~
咱们今天来聊聊高斯混合模型~
比如说,在看一堆点,这些点其实是由几种“类型”的数据混合产生的。每种类型的数据在二维平面上大致呈现一个“高斯(钟形)”云团,但每个云团的位置、形状(椭圆大小、方向)可以不一样。
我们要做两件事:
-
估计出这些云团“有多少个”(通常假定给定 个),以及每个云团的参数(中心、形状、这个云团在整体数据中的占比)。 -
给每个点算一个“属于每个云团的概率”,从而得到 soft 分类(或者选概率最大的做硬分类)。
GMM 的本质就是把数据看做是多个高斯分布的加权和。数学上,混合模型的概率密度是:
其中 是第 个高斯分布的权重(占比,满足 ), 是均值, 是协方差矩阵。
训练 GMM 常用的方法是 EM(Expectation-Maximization)算法,分两步迭代:
-
E 步(期望步):根据当前参数,算出每个样本属于每个高斯成分的“责任度”或“后验概率”: -
M 步(最大化步):用这些责任度去更新参数,使得在这些软“标签”下数据的似然最大化:
重复 E 和 M,直到对数似然不再显著增加或达到最大迭代数。对数似然为:
要点与注意事项:
-
GMM 是生成式模型:它建模的是数据的生成分布,而不是直接学一个判别边界。 -
GMM 给出的不是硬标签,而是每个点属于每个成分的概率(很有用)。 -
如果协方差是全矩阵(full covariance),模型更灵活但参数多,容错能力差;对角矩阵则更简洁但可能拟合得不好。 -
数值稳定性:计算概率时通常用对数和 log-sum-exp 避免下溢。协方差矩阵可能奇异,需要加一个小的对角项(regularize)。 -
EM 可能收敛到局部极值,初始化很重要(kmeans 初始化较常用)。
案例实现
下面给出完整代码,生成随机的数据、实现 EM、可视化等等~
import torch
import numpy as np
import matplotlib.pyplot as plt
from matplotlib.patches import Ellipse
torch.manual_seed(0)
np.random.seed(0)
device = torch.device("cpu")
# 1. 数据集:二维,4 个成分,椭圆形,混合
def sample_gaussian(mean, cov, n):
return np.random.multivariate_normal(mean, cov, size=n)
# 定义 4 个成分的真实参数(用于模拟)
true_means = [
[0.0, 0.0],
[4.0, 4.0],
[-3.5, 3.0],
[3.5, -3.0]
]
true_covs = [
[[0.8, 0.6], [0.6, 1.5]],
[[1.2, -0.4], [-0.4, 0.5]],
[[0.3, 0.0], [0.0, 0.3]],
[[0.6, 0.2], [0.2, 1.0]],
]
true_weights = [0.25, 0.35, 0.15, 0.25] # 和为 1
N = 5200
data = []
labels = []
for k, w in enumerate(true_weights):
n_k = int(N * w)
s = sample_gaussian(true_means[k], true_covs[k], n_k)
data.append(s)
labels.append(np.ones(n_k, dtype=int) * k)
# 合并并打乱
X = np.vstack(data)
y_true = np.concatenate(labels)
perm = np.random.permutation(len(X))
X = X[perm]
y_true = y_true[perm]
X_torch = torch.from_numpy(X).float().to(device)
# 2. 可视化函数:绘制数据、成分椭圆、密度等
colors = ['#e41a1c', '#377eb8', '#4daf4a', '#984ea3', '#ff7f00', '#ffff33']
def plot_scatter(ax, X, labels=None, title=""):
if labels is None:
ax.scatter(X[:,0], X[:,1], s=10, color="#2b8cbe", alpha=0.8)
else:
for k in np.unique(labels):
mask = labels == k
ax.scatter(X[mask,0], X[mask,1], s=12, color=colors[k%len(colors)], alpha=0.8, label=f"c{k}")
ax.legend()
ax.set_title(title)
ax.set_aspect('equal', 'box')
def draw_ellipse(ax, mean, cov, color='k', alpha=0.3, linewidth=2):
# 画协方差椭圆(1 std)
vals, vecs = np.linalg.eigh(cov)
order = vals.argsort()[::-1]
vals = vals[order]
vecs = vecs[:, order]
angle = np.degrees(np.arctan2(vecs[1,0], vecs[0,0]))
width, height = 2 * np.sqrt(vals) # 2 * std
ell = Ellipse(xy=mean, width=width, height=height, angle=angle,
edgecolor=color, facecolor=color, alpha=alpha, linewidth=linewidth)
ax.add_patch(ell)
# 3. GMM EM 算法(PyTorch 实现)
def log_multivariate_normal(x, mu, cov):
# x: (N,D), mu: (D,), cov: (D,D)
D = x.shape[1]
# 为数值稳定,加 epsilon
eps = 1e-6
cov_eps = cov + eps * torch.eye(D, device=x.device)
# 使用 Cholesky 分解求解马氏距离
L = torch.linalg.cholesky(cov_eps) # (D,D)
diff = x - mu.unsqueeze(0) # (N,D)
# solve L u = diff^T -> u = L^{-1} diff^T
sol = torch.cholesky_solve(diff.unsqueeze(-1), L) # (N,D,1)
quad = torch.sum(diff.unsqueeze(-1) * sol, dim=(1,2)) # (N,)
logdet = 2.0 * torch.sum(torch.log(torch.diagonal(L)))
log_norm = -0.5 * (D * np.log(2*np.pi) + logdet)
return log_norm - 0.5 * quad # (N,)
def gmm_em(X, K=4, max_iters=100, tol=1e-4, verbose=True):
N, D = X.shape
# 初始化:权重均匀,均值随机从数据中选取,协方差设为数据协方差
pi = torch.ones(K, device=X.device) / K
idx = torch.randperm(N)[:K]
mu = X[idx].clone() # (K,D)
# 初始协方差用全体数据协方差
emp_cov = torch.from_numpy(np.cov(X.cpu().numpy(), rowvar=False)).float().to(X.device)
covs = torch.stack([emp_cov.clone() for _ in range(K)]) # (K,D,D)
ll_hist = []
for it in range(max_iters):
# E-step: 计算对数概率,避免下溢(log-space)
log_pi = torch.log(pi + 1e-12) # (K,)
log_prob = torch.zeros((N, K), device=X.device)
for k in range(K):
log_prob[:, k] = log_multivariate_normal(X, mu[k], covs[k])
# log gamma numerator = log_pi + log_prob
log_gamma = log_pi.unsqueeze(0) + log_prob # (N,K)
# 归一化:log-sum-exp
logsum = torch.logsumexp(log_gamma, dim=1, keepdim=True) # (N,1)
log_gamma_norm = log_gamma - logsum # (N,K)
gamma = torch.exp(log_gamma_norm) # (N,K)
# M-step
N_k = gamma.sum(dim=0) + 1e-8 # (K,)
pi = (N_k / N)
mu = (gamma.t() @ X) / N_k.unsqueeze(1) # (K,D)
for k in range(K):
diff = X - mu[k].unsqueeze(0) # (N,D)
# weighted covariance
gamma_k = gamma[:, k].unsqueeze(1) # (N,1)
cov_k = (gamma_k * diff).t() @ diff # (D,D)
covs[k] = cov_k / N_k[k] + 1e-6 * torch.eye(D, device=X.device) # (正定化)
# 计算对数似然
ll = torch.sum(logsum)
ll_hist.append(ll.item())
if verbose and (it % 10 == 0 or it == max_iters-1):
print(f"Iter {it:03d}, loglik = {ll.item():.4f}")
# 收敛性判断
if it > 0 and abs(ll_hist[-1] - ll_hist[-2]) < tol:
if verbose:
print("Converged.")
break
return {
'pi': pi.detach().cpu().numpy(),
'mu': mu.detach().cpu().numpy(),
'covs': covs.detach().cpu().numpy(),
'gamma': gamma.detach().cpu().numpy(),
'll_hist': ll_hist
}
# 4. 运行 EM
res = gmm_em(X_torch, K=4, max_iters=200, tol=1e-5, verbose=True)
pi_est = res['pi']
mu_est = res['mu']
covs_est = res['covs']
gamma = res['gamma']
ll_hist = res['ll_hist']
# 预测硬标签(最大后验)
y_pred = np.argmax(gamma, axis=1)
# 5. 绘图:至少 4 张图(原始数据、真实标签、学习结果含椭圆、混合密度/决策边界、对数似然曲线)
fig, axes = plt.subplots(2, 3, figsize=(18, 12))
axes = axes.ravel()
# 图 1:原始点云(无标签)
plot_scatter(axes[0], X, labels=None, title="图1:原始数据(无标签)")
# 图 2:真实标签(生成时用的成分)
plot_scatter(axes[1], X, labels=y_true, title="图2:真实成分标签(模拟用)")
# 图 3:GMM 学习后的硬分类(每个点用后验最大成分上色)
plot_scatter(axes[2], X, labels=y_pred, title="图3:GMM 硬分类(后验最大)")
# 画出估计的成分椭圆
for k in range(len(pi_est)):
draw_ellipse(axes[2], mu_est[k], covs_est[k], color=colors[k%len(colors)], alpha=0.25, linewidth=2)
axes[2].scatter(mu_est[k][0], mu_est[k][1], marker='x', color=colors[k%len(colors)], s=80)
# 图 4:混合密度等高线(contourf),用网格计算 p(x)
ax = axes[3]
xx, yy = np.meshgrid(np.linspace(X[:,0].min()-1, X[:,0].max()+1, 200),
np.linspace(X[:,1].min()-1, X[:,1].max()+1, 200))
grid = np.stack([xx.ravel(), yy.ravel()], axis=1)
grid_t = torch.from_numpy(grid).float().to(device)
# 计算所有成分的 pdf 并加权
pdf_vals = np.zeros(grid.shape[0])
for k in range(len(pi_est)):
# 使用 numpy 来计算 pdf(方便)
mu_k = mu_est[k]
cov_k = covs_est[k]
invcov = np.linalg.inv(cov_k)
det = np.linalg.det(cov_k)
diff = grid - mu_k
exponent = -0.5 * np.sum(diff @ invcov * diff, axis=1)
norm = 1.0 / (2*np.pi * np.sqrt(det) + 1e-12)
pdf_vals += pi_est[k] * norm * np.exp(exponent)
pdf_vals = pdf_vals.reshape(xx.shape)
cf = ax.contourf(xx, yy, pdf_vals, levels=50, cmap='plasma')
axes[3].set_title("图4:估计混合密度等高线(越亮概率越大)")
plt.colorbar(cf, ax=ax)
# 图 5:决策边界(后验最大成分)
ax = axes[4]
# 对网格点计算后验
log_prob_grid = np.zeros((grid.shape[0], len(pi_est)))
for k in range(len(pi_est)):
mu_k = mu_est[k]
cov_k = covs_est[k]
invcov = np.linalg.inv(cov_k)
det = np.linalg.det(cov_k)
diff = grid - mu_k
exponent = -0.5 * np.sum(diff @ invcov * diff, axis=1)
norm = 1.0 / (2*np.pi * np.sqrt(det) + 1e-12)
log_prob_grid[:, k] = np.log(pi_est[k] + 1e-12) + np.log(norm + 1e-12) + exponent
z_grid = np.argmax(log_prob_grid, axis=1)
Z = z_grid.reshape(xx.shape)
ax.contourf(xx, yy, Z, alpha=0.25, levels=len(pi_est), colors=colors[:len(pi_est)])
plot_scatter(ax, X, labels=y_pred, title="图5:后验最大成分的决策区域与样本点")
# 画出椭圆
for k in range(len(pi_est)):
draw_ellipse(ax, mu_est[k], covs_est[k], color=colors[k%len(colors)], alpha=0.25, linewidth=2)
ax.scatter(mu_est[k][0], mu_est[k][1], marker='x', color=colors[k%len(colors)], s=80)
# 图 6:对数似然曲线
ax = axes[5]
ax.plot(ll_hist, color='#2c7fb8', linewidth=2)
ax.set_title("图6:训练过程中的对数似然(log-likelihood)")
ax.set_xlabel("迭代次数")
ax.set_ylabel("对数似然")
plt.tight_layout()
plt.show()
# 6. 简单评估:用多数投票将 GMM 成分对齐到真实标签
from collections import defaultdict
mapping = {}
for k in range(len(pi_est)):
# 该成分内样本真实标签的众数
mask = (y_pred == k)
if mask.sum() == 0:
mapping[k] = -1
continue
labels_in_k = y_true[mask]
counts = np.bincount(labels_in_k)
mapping[k] = np.argmax(counts)
# 计算精度(按对齐后)
y_map = np.array([mapping.get(k, -1) for k in y_pred])
acc = (y_map == y_true).mean()
print("简单多数投票对齐后的聚类准确率(非严格意义上的 supervised accuracy): {:.3f}".format(acc))
原始数据:
仅展示点云的分布,帮助读者直观理解:数据并非严格分离,有明显重叠区域与不同方向的椭圆形分布。
是否存在明显簇、是否有离群点、簇的形状是否为椭圆而非圆形(决定是否用全协方差)。
真实成分标签:
展示我们用于模拟生成数据时的“真实”成分分配。用来作为参考,便于评估 GMM 的效果。
比较不同簇的大小、形状、方向。便于后续与估计结果对比。
GMM 学习后的硬分类:
把每个点按 GMM 计算的后验概率最大值进行着色(即 argmax_k p(k|x))。同时画出估计出的每个高斯的均值(叉号)和 1 个标准差的椭圆。
椭圆中心是否与真实均值接近;椭圆方向和形状(协方差)是否能捕捉到簇的真实散布;重叠区域点是否被合理分配。
估计混合密度等高线:
在平面上绘制混合模型的概率密度大小(颜色越亮代表 p(x) 越高)。这是 GMM 作为生成模型的直接输出——数据在哪些区域更可能出现。
高密度区域是否落在点云簇的中心;密度的铺展是否合逻辑(例如一个成分被学得过宽或过窄都会在等高线中显现)。
后验最大成分的决策区:
在平面上为每个格点计算 p(k|x),然后按哪一个成分最大来划分“决策区域”,并叠加样本点与椭圆。
决策边界如何穿过点云,是否与直觉一致;如果边界复杂且非线性,说明 GMM 用概率分布建模可以给出更柔性的划分。
这张图对理解 GMM 在“做分类”时的行为非常直观:GMM 给分类边界,但它是基于每个成分概率的比较,而不是直接学习一个判别模型。
对数似然:
每次 EM 更新后记录的训练对数似然(log-likelihood)。EM 原则上会使对数似然单调不降。
曲线是否收敛、是否很快达到瓶颈(过早收敛可能是局部最优或初始化不好)、是否有震荡(表明数值不稳定)。
用这张图可以判断是否需要更多迭代或更改初始化策略。
总结
总的来说,GMM 是一个非常直观且强大的生成式模型,适合描述由多个高斯分布混合产生的数据,能输出后验概率并给出数据的概率密度。
EM 是标准的求解方法,局部极值和初始化敏感是其主要缺点。

