大数跨境

导师:“你懂 GMM 吗?”我:“不懂,但 Codex 跑出来了!”导师:“代码能跑不算懂,EM 两步你得讲明白!”

导师:“你懂 GMM 吗?”我:“不懂,但 Codex 跑出来了!”导师:“代码能跑不算懂,EM 两步你得讲明白!” 机器学习和人工智能AI
2026-08-14
1

哈喽,大家好~

咱们今天来聊聊高斯混合模型~

比如说,在看一堆点,这些点其实是由几种“类型”的数据混合产生的。每种类型的数据在二维平面上大致呈现一个“高斯(钟形)”云团,但每个云团的位置、形状(椭圆大小、方向)可以不一样。

我们要做两件事:

  • 估计出这些云团“有多少个”(通常假定给定   个),以及每个云团的参数(中心、形状、这个云团在整体数据中的占比)。
  • 给每个点算一个“属于每个云团的概率”,从而得到 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.00.0],
    [4.04.0],
    [-3.53.0],
    [3.5-3.0]
]
true_covs = [
    [[0.80.6], [0.61.5]],
    [[1.2-0.4], [-0.40.5]],
    [[0.30.0], [0.00.3]],
    [[0.60.2], [0.21.0]],
]
true_weights = [0.250.350.150.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(23, figsize=(1812))
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()+1200),
                     np.linspace(X[:,1].min()-1, X[:,1].max()+1200))
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, -1for 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 是标准的求解方法,局部极值和初始化敏感是其主要缺点。

【声明】内容源于网络
0
0
机器学习和人工智能AI
让我们一起期待 AI 带给我们的每一场变革!推送最新行业内最新最前沿人工智能技术!
内容 381
粉丝 0
机器学习和人工智能AI 让我们一起期待 AI 带给我们的每一场变革!推送最新行业内最新最前沿人工智能技术!
总阅读5.0k
粉丝0
内容381