用Python手把手实现高斯混合模型(GMM):从理论到代码实战

在机器学习的概率模型领域,高斯混合模型(Gaussian Mixture Model, GMM)因其强大的聚类能力和概率解释性而备受青睐。不同于传统的K-means等硬聚类方法,GMM通过概率分布来描述数据,能够捕捉更复杂的数据结构。本文将带您从零开始,用Python完整实现一个GMM模型,并通过可视化手段深入理解EM算法的迭代过程。

1. 高斯混合模型的核心概念

高斯混合模型本质上是由多个高斯分布线性叠加而成的概率模型。假设我们有一个包含K个高斯分量的GMM,其概率密度函数可以表示为:

import numpy as np
from scipy.stats import multivariate_normal

def gmm_pdf(X, weights, means, covariances):
    """
    计算GMM在X处的概率密度
    参数:
        X: 数据点 (n_samples, n_features)
        weights: 各高斯分量的权重 (n_components,)
        means: 各高斯分量的均值 (n_components, n_features)
        covariances: 各高斯分量的协方差矩阵 (n_components, n_features, n_features)
    返回:
        概率密度值 (n_samples,)
    """
    n_samples = X.shape[0]
    pdf = np.zeros(n_samples)
    for weight, mean, cov in zip(weights, means, covariances):
        pdf += weight * multivariate_normal(mean, cov).pdf(X)
    return pdf

GMM的关键优势在于:

  • 能够拟合任意复杂的数据分布
  • 提供软聚类能力(每个数据点属于各个簇的概率)
  • 具有坚实的概率理论基础

表1:GMM与K-means的主要区别

特性 GMM K-means
聚类类型 软聚类 硬聚类
形状假设 椭圆簇 球形簇
收敛速度 较慢 较快
异常值敏感度 较低 较高
概率解释

2. EM算法原理与实现

期望最大化(EM)算法是训练GMM的核心方法,它通过交替执行以下两步来优化模型参数:

  1. E步(Expectation): 计算每个数据点属于各个高斯分量的后验概率
  2. M步(Maximization): 基于E步的结果更新模型参数

让我们实现完整的EM算法:

class GMM:
    def __init__(self, n_components, max_iter=100, tol=1e-6):
        self.n_components = n_components
        self.max_iter = max_iter
        self.tol = tol
        
    def fit(self, X):
        n_samples, n_features = X.shape
        
        # 1. 初始化参数
        self.weights_ = np.ones(self.n_components) / self.n_components
        self.means_ = X[np.random.choice(n_samples, self.n_components, replace=False)]
        self.covariances_ = [np.eye(n_features) for _ in range(self.n_components)]
        
        log_likelihood = 0
        for iteration in range(self.max_iter):
            # 2. E步:计算后验概率
            responsibilities = self._e_step(X)
            
            # 3. M步:更新参数
            self._m_step(X, responsibilities)
            
            # 计算对数似然判断收敛
            new_log_likelihood = self._compute_log_likelihood(X)
            if abs(new_log_likelihood - log_likelihood) < self.tol:
                break
            log_likelihood = new_log_likelihood
    
    def _e_step(self, X):
        n_samples = X.shape[0]
        responsibilities = np.zeros((n_samples, self.n_components))
        
        for k in range(self.n_components):
            responsibilities[:, k] = self.weights_[k] * multivariate_normal(
                self.means_[k], self.covariances_[k]).pdf(X)
        
        responsibilities /= responsibilities.sum(axis=1, keepdims=True)
        return responsibilities
    
    def _m_step(self, X, responsibilities):
        n_samples, n_features = X.shape
        
        # 更新权重
        self.weights_ = responsibilities.sum(axis=0) / n_samples
        
        # 更新均值
        self.means_ = np.zeros((self.n_components, n_features))
        for k in range(self.n_components):
            self.means_[k] = np.sum(responsibilities[:, k][:, np.newaxis] * X, axis=0) / responsibilities[:, k].sum()
        
        # 更新协方差
        self.covariances_ = [np.zeros((n_features, n_features)) for _ in range(self.n_components)]
        for k in range(self.n_components):
            diff = X - self.means_[k]
            self.covariances_[k] = np.dot(responsibilities[:, k] * diff.T, diff) / responsibilities[:, k].sum()
    
    def _compute_log_likelihood(self, X):
        log_likelihood = 0
        for k in range(self.n_components):
            log_likelihood += self.weights_[k] * multivariate_normal(
                self.means_[k], self.covariances_[k]).pdf(X)
        return np.log(log_likelihood).sum()

注意:在实际应用中,为了防止数值下溢,对数似然的计算通常采用log-sum-exp技巧进行优化。

3. 实战:从数据生成到模型评估

为了更好地理解GMM的工作原理,我们将创建一个合成数据集并可视化整个训练过程。

3.1 数据生成

import matplotlib.pyplot as plt
from sklearn.datasets import make_blobs

# 生成包含3个簇的合成数据
X, y_true = make_blobs(n_samples=500, centers=3, cluster_std=[1.0, 2.5, 0.5], random_state=42)

plt.figure(figsize=(8, 6))
plt.scatter(X[:, 0], X[:, 1], c=y_true, s=40, cmap='viridis')
plt.title("原始数据分布")
plt.show()

3.2 训练过程可视化

让我们修改GMM类,使其能够记录每次迭代的结果:

class VisualGMM(GMM):
    def fit(self, X):
        # 初始化记录变量
        self.history = {
            'means': [],
            'covariances': [],
            'weights': [],
            'log_likelihood': []
        }
        
        # 调用父类的fit方法
        super().fit(X)
    
    def _m_step(self, X, responsibilities):
        super()._m_step(X, responsibilities)
        
        # 记录当前参数
        self.history['means'].append(self.means_.copy())
        self.history['covariances'].append([cov.copy() for cov in self.covariances_])
        self.history['weights'].append(self.weights_.copy())
    
    def _compute_log_likelihood(self, X):
        ll = super()._compute_log_likelihood(X)
        self.history['log_likelihood'].append(ll)
        return ll

3.3 可视化训练过程

def plot_gmm(gmm, X, iteration):
    plt.figure(figsize=(8, 6))
    plt.scatter(X[:, 0], X[:, 1], s=40, color='gray', alpha=0.3)
    
    for k in range(gmm.n_components):
        mean = gmm.history['means'][iteration][k]
        cov = gmm.history['covariances'][iteration][k]
        
        # 绘制等高线
        x, y = np.mgrid[-10:10:.1, -10:10:.1]
        pos = np.dstack((x, y))
        rv = multivariate_normal(mean, cov)
        plt.contour(x, y, rv.pdf(pos), levels=3, colors='blue', alpha=0.5)
        
        # 绘制均值点
        plt.scatter(mean[0], mean[1], marker='x', s=100, linewidths=3, color='red')
    
    plt.title(f"Iteration {iteration+1}, Log-likelihood: {gmm.history['log_likelihood'][iteration]:.2f}")
    plt.xlim(-10, 10)
    plt.ylim(-10, 10)
    plt.show()

# 训练并可视化
gmm = VisualGMM(n_components=3)
gmm.fit(X)

for i in range(len(gmm.history['means'])):
    plot_gmm(gmm, X, i)

4. 模型评估与选择

4.1 确定最佳组件数量

选择GMM中合适的高斯分量数量(K)是一个重要问题。我们可以使用以下方法:

  1. 贝叶斯信息准则(BIC):

    from sklearn.mixture import GaussianMixture
    
    n_components_range = range(1, 8)
    bic = []
    
    for n in n_components_range:
        model = GaussianMixture(n_components=n, random_state=42)
        model.fit(X)
        bic.append(model.bic(X))
    
    plt.plot(n_components_range, bic, 'o-')
    plt.xlabel("Number of components")
    plt.ylabel("BIC score")
    plt.title("BIC for selecting number of components")
    plt.show()
    
  2. 轮廓系数:

    from sklearn.metrics import silhouette_score
    
    silhouette_scores = []
    
    for n in n_components_range:
        model = GaussianMixture(n_components=n, random_state=42)
        labels = model.fit_predict(X)
        silhouette_scores.append(silhouette_score(X, labels))
    
    plt.plot(n_components_range, silhouette_scores, 'o-')
    plt.xlabel("Number of components")
    plt.ylabel("Silhouette score")
    plt.title("Silhouette score for selecting number of components")
    plt.show()
    

4.2 与K-means的对比

from sklearn.cluster import KMeans

kmeans = KMeans(n_clusters=3, random_state=42)
kmeans_labels = kmeans.fit_predict(X)

gmm = GaussianMixture(n_components=3, random_state=42)
gmm_labels = gmm.fit_predict(X)

fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 5))
ax1.scatter(X[:, 0], X[:, 1], c=kmeans_labels, s=40, cmap='viridis')
ax1.set_title("K-means Clustering")
ax2.scatter(X[:, 0], X[:, 1], c=gmm_labels, s=40, cmap='viridis')
ax2.set_title("GMM Clustering")
plt.show()

表2:GMM与K-means在不同场景下的表现对比

数据特征 GMM表现 K-means表现
簇大小差异大 优秀 一般
非球形簇 优秀
重叠簇 优秀 一般
噪声数据 稳健 敏感
计算效率 较慢 快速

5. 高级主题与优化技巧

5.1 协方差矩阵约束

GMM的性能很大程度上取决于协方差矩阵的类型选择。常见的约束类型包括:

  • 'full': 每个分量有自己的任意协方差矩阵
  • 'tied': 所有分量共享同一个协方差矩阵
  • 'diag': 每个分量的协方差矩阵是对角矩阵
  • 'spherical': 每个分量的协方差矩阵是标量乘以单位矩阵
covariance_types = ['full', 'tied', 'diag', 'spherical']
bic_scores = []

for cov_type in covariance_types:
    model = GaussianMixture(n_components=3, covariance_type=cov_type, random_state=42)
    model.fit(X)
    bic_scores.append(model.bic(X))

plt.bar(covariance_types, bic_scores)
plt.xlabel("Covariance type")
plt.ylabel("BIC score")
plt.title("BIC for different covariance types")
plt.show()

5.2 初始化策略优化

GMM对初始参数敏感,sklearn提供了两种初始化方法:

  1. k-means初始化(默认)
  2. 随机初始化
init_methods = ['kmeans', 'random']
log_likelihoods = []

for init in init_methods:
    model = GaussianMixture(n_components=3, init_params=init, random_state=42)
    model.fit(X)
    log_likelihoods.append(model.score(X))

plt.bar(init_methods, log_likelihoods)
plt.xlabel("Initialization method")
plt.ylabel("Log-likelihood")
plt.title("Performance of different initialization methods")
plt.show()

5.3 处理高维数据

对于高维数据,GMM可能会遇到"维度灾难"。可以采用以下策略:

  1. 主成分分析(PCA)降维:

    from sklearn.decomposition import PCA
    
    # 生成高维数据
    X_high_dim, _ = make_blobs(n_samples=500, n_features=10, centers=3, random_state=42)
    
    # 降维
    pca = PCA(n_components=2)
    X_reduced = pca.fit_transform(X_high_dim)
    
    # 在降维后的空间拟合GMM
    gmm = GaussianMixture(n_components=3, random_state=42)
    gmm.fit(X_reduced)
    
  2. 使用对角协方差矩阵:

    gmm_diag = GaussianMixture(n_components=3, covariance_type='diag', random_state=42)
    gmm_diag.fit(X_high_dim)
    

在实际项目中,我发现对于文本数据等超高维场景,先使用主题模型(LDA)或自动编码器降维,再应用GMM通常会获得更好的效果。而对于图像数据,CNN特征提取后的低维表示配合GMM往往能发现更有意义的视觉模式。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐