Python实战:用scikit-learn轻松搞定高斯混合模型(GMM)聚类与异常检测
Python实战:用scikit-learn轻松搞定高斯混合模型(GMM)聚类与异常检测
你是否曾面对一堆看似毫无规律的数据点,试图从中找出隐藏的“群落”?或者,你是否需要从海量交易记录中,精准地揪出那些行为异常的“少数派”?对于数据分析师和机器学习工程师来说,这类任务几乎每天都在发生。传统的硬聚类方法,比如K-means,虽然快,但总觉得有点“武断”——它强行给每个点贴上一个唯一的标签,却忽略了数据点之间模糊的归属关系。而今天我们要深入探讨的高斯混合模型,则提供了一种更优雅、更符合直觉的解决方案。它不再简单地说“你属于A类”,而是告诉你“你有60%的概率属于A类,30%的概率属于B类,还有10%的概率是个特例”。这种概率化的视角,不仅让聚类结果更细腻,更自然地延伸到了密度估计和异常检测的领域。在Python生态中,scikit-learn库为我们封装了强大且易用的GaussianMixture类,让我们能够绕过复杂的数学推导,直接聚焦于解决实际问题。这篇文章,我将带你从零开始,手把手实现GMM,并重点分享它在异常检测场景下的实战技巧与避坑指南。
1. 从“硬分”到“软分”:理解GMM的核心思想
在开始敲代码之前,我们有必要花点时间,在脑海里构建起GMM的直观图像。这能帮助你在调参和解读结果时,不至于迷失在数学公式中。
想象一下,你面前有一片星空图。用K-means的眼光看,你会试图找到几个最亮的点作为“星团中心”,然后把每颗星星划归到离它最近的那个中心所在的星团。星星的归属是非此即彼的。但GMM的视角截然不同:它认为这片星空是由几团发光的气体云(高斯分布)混合而成的。每一团气体云都有自己的亮度中心(均值)和扩散形态(协方差)。任何一颗星星,都同时受到这几团气体云的影响,只是影响的程度不同。靠近某团云中心的星星,受其影响自然就大;位于几团云交界处的星星,则可能同时受到多个云的显著影响。
这种“混合”的思想,就是GMM的基石。它用一个加权和来模拟整个数据集的概率分布:
P(x) = Σ (权重_k * 高斯分布_k(x))
这里的权重_k代表了第k团气体云(即第k个高斯分量)在整个星空中的“重要性”或“比例”。所有权重之和为1。高斯分布_k(x)则描述了在第k团云内部,数据点x出现的可能性。
注意:GMM是一个生成式模型。这意味着它试图学习数据是如何“生成”出来的。一旦模型训练完成,我们理论上可以按照学到的分布(几个高斯分布的混合)来生成新的、与原始数据类似的数据点。这个特性在数据增强、仿真等场景中非常有用。
那么,计算机如何从一堆数据点中,反推出这些看不见的“气体云”的参数(权重、均值、协方差)呢?这就要靠期望最大化算法。EM算法就像一个不断自我修正的探索过程:
- E步(Expectation):在现有对“气体云”的猜测下,计算每个数据点“属于”每一团云的概率(责任值)。
- M步(Maximization):根据上一步算出的归属概率,重新估算每一团“气体云”的参数(让它的中心移到属于它的那些点的概率加权中心,调整它的形状以匹配这些点的分布)。
- 重复E步和M步,直到参数基本稳定不变。
这个过程保证了模型的对数似然(可以理解为模型对现有数据的“解释力”)在每一步迭代中都不会下降,最终收敛到一个局部最优解。
2. 环境搭建与数据准备:为GMM实战铺路
理论聊得差不多了,是时候打开你的Python环境了。我强烈建议使用conda或venv创建一个独立的虚拟环境,避免包版本冲突。核心依赖库非常简单:
pip install numpy pandas scikit-learn matplotlib seaborn
numpy&pandas: 数据操作的基石。scikit-learn: 今天的主角,提供了GaussianMixture实现。matplotlib&seaborn: 可视化神器,帮助我们直观理解数据和模型结果。
接下来,我们准备数据。为了清晰地演示GMM的能力,我设计一个包含两个特性簇和一个异常点集的二维数据集。这样我们既能观察聚类,也能直观地看到异常检测的效果。
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
# 设置随机种子,确保结果可复现
np.random.seed(2024)
# 生成两个主要的簇
n_samples = 1000
# 簇1:以(2, 2)为中心,协方差矩阵为[[1, 0.5], [0.5, 1]]
cluster_1 = np.random.multivariate_normal(mean=[2, 2], cov=[[1, 0.5], [0.5, 1]], size=n_samples//2)
# 簇2:以(-1, -1)为中心,协方差矩阵为[[0.8, -0.3], [-0.3, 0.8]]
cluster_2 = np.random.multivariate_normal(mean=[-1, -1], cov=[[0.8, -0.3], [-0.3, 0.8]], size=n_samples//2)
# 生成一些潜在的异常点(均匀分布在边缘区域)
n_anomalies = 30
anomalies = np.random.uniform(low=[-5, -5], high=[5, 5], size=(n_anomalies, 2))
# 合并数据
X = np.vstack([cluster_1, cluster_2, anomalies])
# 为后续可视化标记来源(0: 簇1, 1: 簇2, 2: 异常点)
true_labels = np.array([0]*(n_samples//2) + [1]*(n_samples//2) + [2]*n_anomalies)
# 可视化原始数据
plt.figure(figsize=(10, 8))
sns.scatterplot(x=X[:, 0], y=X[:, 1], hue=true_labels, palette='viridis', alpha=0.7, edgecolor='k')
plt.title('原始数据分布(含两个主簇及散落异常点)')
plt.xlabel('特征 1')
plt.ylabel('特征 2')
plt.legend(title='来源', labels=['主簇 1', '主簇 2', '潜在异常'])
plt.grid(True, alpha=0.3)
plt.show()
运行这段代码,你会得到一张散点图,清晰地展示出两个椭圆状的密集区域,以及一些散布在四周的离群点。我们的目标就是让GMM识别出这两个主簇,并量化每个点属于这两个主簇的“典型程度”。
3. 模型训练与聚类:让scikit-learn替你完成数学
有了数据,训练GMM模型在scikit-learn中堪称一行代码的艺术。但在这行代码之前,有几个关键参数需要我们决策:
n_components: 即K值,你猜测数据由几个高斯分布混合而成。这是我们面临的首要挑战。对于这个示例,我们“作弊”地知道是2。但在真实场景,你需要借助一些方法。covariance_type: 协方差矩阵的类型,决定了每个高斯分量的“形状”约束。'full': 每个分量有自己独立的、任意形状的椭圆(全协方差矩阵)。最灵活,参数最多,可能过拟合。'tied': 所有分量共享同一个全协方差矩阵。所有簇形状、方向相同,只是位置不同。'diag': 每个分量的协方差矩阵是对角矩阵。簇的形状是轴对齐的椭圆。'spherical': 每个分量的协方差矩阵是标量乘以单位矩阵。簇是圆形。
init_params: 初始化参数的方法,'kmeans'(默认)通常是不错的选择。max_iter&tol: 控制EM算法的迭代次数和收敛阈值。
让我们先尝试用n_components=2和covariance_type='full'来拟合模型。
from sklearn.mixture import GaussianMixture
# 初始化并训练GMM模型
gmm = GaussianMixture(n_components=2, covariance_type='full', random_state=2024)
gmm.fit(X) # X是我们的训练数据
# 查看训练好的模型参数
print("混合权重 (各分量的重要性):", gmm.weights_.round(3))
print("\n各分量的均值中心:")
for i, mean in enumerate(gmm.means_):
print(f" 分量 {i}: {mean.round(3)}")
print("\n各分量的协方差矩阵 (形状):")
for i, cov in enumerate(gmm.covariances_):
print(f" 分量 {i}:\n{cov.round(3)}\n")
输出会显示两个高斯分量的权重、中心坐标和协方差矩阵。你会发现,权重之和为1,两个中心点大致位于我们生成的两个簇的中心附近,协方差矩阵也近似于我们生成数据时设定的值。
现在,我们可以用这个训练好的模型来做软聚类预测。predict方法给出的是每个点最可能属于的分量索引(硬标签),而predict_proba方法则给出了我们更关心的归属概率。
# 预测硬聚类标签(最大概率对应的分量)
hard_labels = gmm.predict(X)
# 预测软聚类概率(每个点属于各个分量的概率)
soft_proba = gmm.predict_proba(X)
# 我们取前5个点看看软聚类结果
print("前5个数据点的软聚类概率:")
for i in range(5):
print(f" 点 {i}: {soft_proba[i].round(3)} -> 最可能属于分量 {hard_labels[i]}")
你会看到,位于簇中心的数据点,其属于对应分量的概率接近1;而位于两个簇边界附近的数据点,其概率可能更平均,比如[0.55, 0.45],这真实反映了其“模糊”的归属状态。
如何选择最佳的K值?
在未知真实K的情况下,我们可以借助信息准则。scikit-learn的GaussianMixture模型在训练后会自动计算两个值:
- AIC: 赤池信息准则,在模型复杂度和拟合优度之间寻求平衡。
- BIC: 贝叶斯信息准则,对模型复杂度的惩罚比AIC更重。
通常,我们选择使AIC或BIC最小的K值。下面是一个简单的循环搜索:
n_components_range = range(1, 11)
aic_scores = []
bic_scores = []
for n in n_components_range:
gmm_temp = GaussianMixture(n_components=n, covariance_type='full', random_state=2024)
gmm_temp.fit(X)
aic_scores.append(gmm_temp.aic(X))
bic_scores.append(gmm_temp.bic(X))
# 可视化
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.plot(n_components_range, aic_scores, 'bo-')
plt.xlabel('混合分量数量 K')
plt.ylabel('AIC')
plt.title('AIC准则选择K')
plt.grid(True)
plt.subplot(1, 2, 2)
plt.plot(n_components_range, bic_scores, 'ro-')
plt.xlabel('混合分量数量 K')
plt.ylabel('BIC')
plt.title('BIC准则选择K')
plt.grid(True)
plt.tight_layout()
plt.show()
# 找出最优K
optimal_k_aic = n_components_range[np.argmin(aic_scores)]
optimal_k_bic = n_components_range[np.argmin(bic_scores)]
print(f"根据AIC,最优K值为: {optimal_k_aic}")
print(f"根据BIC,最优K值为: {optimal_k_bic}")
对于我们的示例数据,AIC和BIC曲线通常会在K=2或K=3处出现拐点。BIC由于惩罚更重,更倾向于选择更简单的模型(K=2)。这是一个非常重要的模型选择工具。
4. 从密度估计到异常检测:GMM的杀手级应用
GMM训练完成后,我们不仅得到了一个聚类器,更得到了一个对整个数据空间概率分布的估计模型。score_samples(X)方法返回的是每个样本在模型下的对数似然。这个值越大,说明该样本点出现在这个概率分布下的可能性越高,即越“正常”;反之,值越小(越负),则说明该点越“不可能”由这个混合分布生成,即越可能是异常点。
这就是GMM用于异常检测的基本原理:设定一个对数似然阈值,低于该阈值的样本被判为异常。
# 计算所有样本的对数似然
log_densities = gmm.score_samples(X)
# 可视化对数似然的分布
plt.figure(figsize=(10, 6))
plt.hist(log_densities, bins=50, edgecolor='black', alpha=0.7)
plt.xlabel('样本的对数似然值')
plt.ylabel('频数')
plt.title('所有样本对数似然值分布')
plt.axvline(x=np.percentile(log_densities, 5), color='red', linestyle='--', label='5%分位数(异常阈值候选)')
plt.legend()
plt.grid(True, alpha=0.3)
plt.show()
从直方图可以看到,大部分样本的对数似然值集中在一个较高的区间,而左侧有一条长长的“尾巴”,这些就是似然值极低的点——我们的候选异常点。
如何设定阈值?这没有绝对标准,通常取决于业务对异常的定义(比如希望捕捉尾部5%的数据)。我们可以使用百分位数来设定:
# 将对数似然最低的5%样本标记为异常
threshold = np.percentile(log_densities, 5)
is_anomaly = log_densities < threshold
print(f"对数似然阈值: {threshold:.3f}")
print(f"被标记为异常点的数量: {is_anomaly.sum()}")
print(f"异常点占比: {is_anomaly.sum() / len(X):.2%}")
# 可视化异常检测结果
plt.figure(figsize=(10, 8))
# 先画正常点
normal_points = X[~is_anomaly]
plt.scatter(normal_points[:, 0], normal_points[:, 1], c='lightblue', alpha=0.6, edgecolor='k', label='正常点', s=30)
# 再高亮异常点
anomaly_points = X[is_anomaly]
plt.scatter(anomaly_points[:, 0], anomaly_points[:, 1], c='red', alpha=1.0, edgecolor='darkred', label='检测出的异常点', s=50, marker='x')
# 可以画出高斯分量的等高线,更直观
x = np.linspace(X[:, 0].min()-1, X[:, 0].max()+1, 200)
y = np.linspace(X[:, 1].min()-1, X[:, 1].max()+1, 200)
X_grid, Y_grid = np.meshgrid(x, y)
XX = np.array([X_grid.ravel(), Y_grid.ravel()]).T
Z = gmm.score_samples(XX)
Z = Z.reshape(X_grid.shape)
plt.contour(X_grid, Y_grid, Z, levels=np.linspace(Z.min(), Z.max(), 10), linewidths=1, colors='gray', alpha=0.5)
plt.title('GMM异常检测结果可视化')
plt.xlabel('特征 1')
plt.ylabel('特征 2')
plt.legend()
plt.grid(True, alpha=0.3)
plt.show()
在这张结果图上,红色的“x”点就是模型识别出的异常点。它们大多远离两个主簇的密集区域。等高线图描绘了模型估计的概率密度,颜色越深的区域密度越高。异常点恰恰落在了密度非常低的区域。
实战技巧:处理高维与复杂数据
在实际项目中,数据往往不止二维,且分布可能更加复杂。这里有几个提升GMM异常检测效果的经验:
-
特征缩放至关重要:GMM对特征的尺度非常敏感。一个取值范围在0-1的特征,与一个取值范围在0-10000的特征,其方差会主导整个模型。务必使用
StandardScaler或RobustScaler对数据进行标准化。from sklearn.preprocessing import StandardScaler scaler = StandardScaler() X_scaled = scaler.fit_transform(X) # 在X_scaled上训练GMM -
协方差类型的选择:
'full'最灵活但也最容易过拟合,特别是在样本量少、维度高时。如果数据维度很高(比如>50),'diag'或'spherical'是更稳定、计算更快的选择。你可以通过交叉验证比较不同covariance_type下的AIC/BIC。 -
应对非高斯分布:如果单个簇的分布明显不是高斯状(例如,极度偏斜或有多峰),单个高斯分量就无法很好拟合。这时可以考虑增加
n_components,用多个高斯分量的组合来近似一个复杂形状的簇。但这会提高模型复杂度和过拟合风险。 -
阈值设定的业务化:不要僵化地使用5%分位数。这个阈值应该与你的业务成本相关联。例如,在金融欺诈检测中,误报(将正常交易判为欺诈)和漏报(放过了欺诈交易)的成本不同。你需要根据精确率-召回率曲线来选择一个在业务上最优的阈值。
5. 超越基础:GMM的高级用法与调优策略
掌握了基础应用后,我们可以探索一些更高级的用法,让GMM在复杂场景下发挥更大威力。
5.1 使用贝叶斯高斯混合模型自动确定K值
scikit-learn还提供了BayesianGaussianMixture类,它通过引入狄利克雷过程先验,可以让模型在训练过程中自动决定需要多少个分量(K值)。这对于完全未知聚类数量的场景非常有用。
from sklearn.mixture import BayesianGaussianMixture
# 设置一个较大的n_components上限,模型会自动将不必要的分量权重趋于0
bgmm = BayesianGaussianMixture(n_components=10, max_iter=1000, random_state=2024)
bgmm.fit(X_scaled) # 假设X已经标准化
print("学习到的混合权重:", bgmm.weights_.round(3))
# 权重接近0的分量可以被认为是“未使用”的
effective_components = np.sum(bgmm.weights_ > 0.01)
print(f"有效的分量数量(权重>0.01): {effective_components}")
贝叶斯方法通常更稳健,但计算开销也更大。它输出的权重可以直观地告诉你哪些分量是重要的。
5.2 协方差矩阵正则化:提升数值稳定性
当某个簇的样本数很少,或者特征间高度相关时,计算出的协方差矩阵可能接近奇异(不可逆),导致数值计算问题。GaussianMixture提供了reg_covar参数,用于在协方差矩阵的对角线上添加一个很小的值,确保其可逆。
# 添加一个较小的正则化项
gmm_reg = GaussianMixture(n_components=3, covariance_type='full', reg_covar=1e-6, random_state=2024)
gmm_reg.fit(X_scaled)
在特征很多或数据很“薄”的时候,适当增加reg_covar(如1e-3)是很好的实践。
5.3 与监督学习结合:使用GMM特征增强
GMM本身是无监督的,但其输出可以作为特征,输入到有监督的模型中,这被称为特征工程。例如,在分类或回归问题中,除了原始特征,你还可以加入:
- 样本属于每个高斯分量的概率(
predict_proba输出)。 - 样本的对数似然值(
score_samples输出)。 - 样本到每个高斯分量中心的马氏距离。
这些新特征可能携带了原始特征中未显式表达的、关于数据分布结构的信息,有时能显著提升监督模型的性能。
# 假设我们有一个分类任务,X_train是特征,y_train是标签
gmm_feat = GaussianMixture(n_components=5, random_state=2024).fit(X_train)
# 为训练集生成GMM特征
train_proba = gmm_feat.predict_proba(X_train)
train_log_dens = gmm_feat.score_samples(X_train).reshape(-1, 1)
# 组合原始特征和GMM特征
X_train_augmented = np.hstack([X_train, train_proba, train_log_dens])
# 现在可以用X_train_augmented去训练你的分类器(如随机森林、XGBoost)
5.4 可视化与诊断:理解你的模型
一个模型光能用还不够,我们还需要理解它。除了前面画等高线,还可以可视化每个高斯分量的置信椭圆。
def plot_gmm_ellipses(gmm, X, ax=None):
"""在给定坐标轴上绘制GMM各分量的置信椭圆"""
if ax is None:
ax = plt.gca()
colors = ['navy', 'darkorange', 'green', 'purple', 'brown'] # 定义颜色
for n, color in enumerate(colors[:gmm.n_components]):
# 获取分量的均值和协方差
mean = gmm.means_[n]
cov = gmm.covariances_[n]
# 计算椭圆的角度和轴长
v, w = np.linalg.eigh(cov)
u = w[0] / np.linalg.norm(w[0])
angle = np.arctan2(u[1], u[0])
angle = 180 * angle / np.pi # 转换为度
v = 2. * np.sqrt(2.) * np.sqrt(v) # 2*sqrt(2)*标准差对应约95%置信区间
# 绘制椭圆
ell = mpl.patches.Ellipse(mean, v[0], v[1], 180 + angle, color=color, alpha=0.3)
ell.set_clip_box(ax.bbox)
ax.add_artist(ell)
ax.scatter(mean[0], mean[1], marker='*', s=100, color=color, edgecolor='k')
# 使用函数
fig, ax = plt.subplots(figsize=(10, 8))
ax.scatter(X[:, 0], X[:, 1], alpha=0.5, s=20)
plot_gmm_ellipses(gmm, X, ax=ax) # gmm是之前训练的模型
ax.set_title('GMM各分量置信椭圆(95%)')
plt.show()
这张图能让你清晰地看到模型学到的每个高斯分量的位置、方向和覆盖范围,对于判断模型是否捕捉到了数据真实结构非常有帮助。
6. 避坑指南:GMM实战中的常见问题与解决方案
在实际应用GMM时,我踩过不少坑。这里总结几个最常见的问题和应对策略,希望能帮你节省时间。
问题一:EM算法收敛慢或不收敛。
- 可能原因:初始值太差;数据尺度差异巨大;
n_components设置不合理。 - 解决方案:
- 尝试不同的
init_params方法(如'random')并设置不同的random_state多次运行。 - 务必进行特征标准化。
- 检查AIC/BIC曲线,确认K值选择是否合理。过大的K值会导致模型复杂,收敛困难。
- 尝试不同的
问题二:模型将大量数据点识别为异常。
- 可能原因:阈值设置过于严格;数据本身包含多个分布模式,但K值设置太小,导致模型无法拟合所有正常模式,从而将某些正常模式也判为异常。
- 解决方案:
- 调整阈值百分位数(如从5%调到2%或1%)。
- 增加
n_components,让模型能更好地描述正常数据的多样性。 - 可视化异常点,结合业务知识判断它们是否真的是异常。有时“异常”可能代表了一个未被发现的新簇。
问题三:高维数据下性能不佳。
- 可能原因:“维数灾难”。在高维空间,数据变得极其稀疏,高斯分布的估计变得不稳定,且计算协方差矩阵开销巨大。
- 解决方案:
- 使用
covariance_type='diag'或'spherical'简化模型。 - 先使用PCA、t-SNE或UMAP等降维技术,在低维空间应用GMM。但要注意,降维可能会扭曲密度估计。
- 考虑使用专门为高维异常检测设计的算法,如Isolation Forest或Local Outlier Factor作为补充或替代。
- 使用
问题四:处理类别型特征或混合类型数据。
- 核心限制:GMM本质上是为连续型数值特征设计的。直接对类别特征进行one-hot编码后使用GMM,效果通常很差,因为one-hot向量的分布不符合高斯假设。
- 解决方案:
- 对于混合类型数据,可以考虑使用能处理此类数据的模型,如聚类采用K-Prototypes,异常检测采用专门算法。
- 如果必须用GMM,一个变通方法是:仅对连续特征使用GMM进行密度估计。对于类别特征,可以单独计算其频率或使用其他模型,最后将两者的异常得分以某种方式(如加权平均)结合起来。但这需要仔细的设计和验证。
最后,记住没有“银弹”。GMM在数据近似高斯混合分布时是利器,但在其他情况下可能力不从心。我的习惯是,在启动一个异常检测项目时,会同时用GMM、Isolation Forest、LOF等几种算法跑一遍基线,对比它们的ROC曲线和业务解释性,再决定深入优化哪一个。工具终究是为目标服务的,理解其原理和局限,才能让它真正为你所用。
更多推荐



所有评论(0)