资讯动态

期望最大化算法原理与实现:从高斯混合模型到机器学习实践

发布时间:2026/8/24 9:31:06 来源:尧图企业网站定制
1. 项目概述从“Task3 EM||datawhale”看算法学习与实践的融合最近在算法学习社区里经常能看到类似“Task3 EM||datawhale”这样的标题。这通常不是一个具体的软件项目而更像是一个学习任务或挑战的标识符。拆解来看“Task3”指向一个系列学习任务的第三部分“EM”极大概率指的是经典的期望最大化算法而“datawhale”则是一个知名的开源数据科学学习组织。所以这个标题背后很可能是一个由Datawhale发起的、围绕EM算法展开的集中式学习与实践项目。对于任何希望深入理解概率模型、无监督学习乃至更广泛机器学习理论的朋友来说EM算法都是一个绕不开的“硬骨头”兼“里程碑”。它不像一些现成的深度学习框架调用起来那么直观但其思想之精妙应用之广泛从高斯混合模型聚类到隐马尔可夫模型训练乃至很多含有隐变量的概率图模型参数估计都离不开它。今天我就结合自己多次讲授和实战应用EM算法的经验来一次深度的拆解。我们不止步于公式推导更要弄明白每一个步骤的物理意义并亲手用代码实现一个完整的案例。你会发现理解了EM再看很多复杂的模型会有一种“拨云见日”的感觉。无论你是正在参加Datawhale学习活动的学员还是对机器学习核心算法有追求的独立学习者这篇内容都将带你从原理到实践彻底吃透EM算法。2. EM算法核心思想与原理深度剖析2.1 问题场景为什么我们需要EM算法想象一个经典的场景你想对一组数据进行聚类但你不知道这些数据具体来自几个类别也不知道每个类别的分布参数。比如我们测量了一批产品的尺寸这批产品实际上由A、B两条生产线混合生产每条线的产品尺寸服从不同的正态分布。现在我们只有混合后的尺寸数据并不知道每个数据点具体来自哪条线我们的目标是估计出两条生产线各自产品尺寸的均值μ和方差σ²。这里的“来自哪条生产线”就是一个隐变量我们无法直接观测。这就是EM算法大显身手的典型场景含有隐变量的概率模型参数估计。如果我们能观测到隐变量即知道每个数据点的归属那么直接用极大似然估计就能轻松求出参数。但现实是我们观测不到。EM算法提供了一种巧妙的迭代思路既然直接求解困难我们就“猜”完再“证”循环往复逐步逼近最优解。2.2 E步与M步迭代优化的艺术EM算法的名字就揭示了它的两个核心步骤期望步和最大化步。E步计算期望在当前参数估计下计算隐变量分布的条件概率期望。回到生产线的例子E步就是基于当前猜测的A、B两条线的参数μ_A, σ_A², μ_B, σ_B²对于每一个数据点计算它“有多大可能性来自A线有多大可能性来自B线”。这个可能性是一个概率值在学术上我们计算的是完全数据的对数似然函数关于隐变量后验分布的期望。简单说就是给每个数据点都分配一个“软标签”标明它属于各个类别的“责任”有多大而不是非此即彼的“硬分配”。注意这里的“软标签”或“责任”是理解EM与K-Means等硬聚类区别的关键。K-Means每次迭代中一个点只属于一个簇而在EM框架下的高斯混合模型一个点以一定概率属于所有簇这使得模型更灵活能处理重叠的簇。M步最大化期望有了E步计算出的“责任”分布我们就可以“假装”自己知道了隐变量只不过是以概率形式知道的。然后基于这个“完整”的数据原始数据隐变量分布重新用极大似然估计来更新模型参数。对于生产线例子更新均值μ时不再是所有数据点平均而是用每个数据点乘以它属于该类的“责任”权重再进行加权平均。方差σ²的更新也类似。M步的目标是找到一组新的参数使得E步中计算出的那个期望值变得更大。因为那个期望值是真实对数似然的一个下界提升这个下界也就间接地提升了我们真正关心的、但难以直接优化的真实似然函数。2.3 算法收敛性为什么这样做是有效的一个自然的疑问是这样交替执行E和M一定能找到最优解吗EM算法被证明能够保证单调递增的性质。即每一次迭代模型的对数似然函数值都不会下降通常会增加直到收敛到一个局部极值点。这得益于数学上优美的Jensen不等式的应用。你可以把目标函数含隐变量的似然函数想象成一个形状复杂的山脉我们直接攀登优化很困难。EM算法则是在每次迭代中构造一个当前位置的“下界曲面”这个曲面在当前位置与目标函数相切且处处低于或等于目标函数。然后我们去优化这个更简单、更规则的“下界曲面”M步找到它的最高点。由于下界提升了目标函数本身的值也必然被抬升。然后我们在新的位置再构造一个新的、更紧致的下界继续优化。如此反复就像踩着不断抬升的垫脚石一步步爬上山峰。当然EM找到的通常是局部最优解而非全局最优。因此初始值的选取非常重要实践中我们通常会随机初始化多次选择似然函数值最高的那次作为最终结果。3. 高斯混合模型EM算法的经典练兵场理论讲得再多不如亲手实现一遍。高斯混合模型是阐释EM算法最直观、应用也最广泛的模型。下面我们就以GMM为例展示完整的EM算法流程和代码实现。3.1 模型定义与符号说明假设我们有N个观测数据点X {x₁, x₂, ..., x_N}我们认为这些数据由K个高斯分布混合生成。模型参数为π_k第k个高斯分布的混合系数满足 Σπ_k 1 π_k ≥ 0。可以理解为选择第k个高斯分布的先验概率。μ_k第k个高斯分布的均值向量。Σ_k第k个高斯分布的协方差矩阵。为简化我们常假设各分量是独立的即Σ_k为对角阵甚至退化为σ_k² * I。隐变量Z {z₁, z₂, ..., z_N}其中z_i是一个K维的one-hot向量表示数据点x_i来自哪一个高斯分量。3.2 完整EM算法推导与实现1. 初始化参数随机或使用K-Means的结果初始化{π_k, μ_k, Σ_k}。2. E步计算响应度计算数据点x_i由第k个高斯分布生成的后验概率即“责任”γ(z_ik)γ(z_ik) p(z_ik | x_i, θ) [π_k * N(x_i | μ_k, Σ_k)] / [Σ_{j1}^K π_j * N(x_i | μ_j, Σ_j)]其中N(x | μ, Σ) 是多变量高斯分布的概率密度函数。用Python实现E步的核心代码如下import numpy as np from scipy.stats import multivariate_normal def e_step(X, pi, mu, sigma): E步计算每个数据点对每个高斯分量的响应度责任 参数 X: 数据矩阵形状 (N, D) pi: 混合系数形状 (K,) mu: 均值矩阵形状 (K, D) sigma: 协方差矩阵列表长度为K每个元素为 (D, D) 矩阵 返回 gamma: 响应度矩阵形状 (N, K) N, D X.shape K len(pi) gamma np.zeros((N, K)) # 计算每个高斯分量下所有数据点的概率密度 for k in range(K): # 使用多元高斯分布这里假设协方差矩阵是满秩的 # 实际中需处理数值稳定性例如给协方差矩阵对角线加一个小的正则项 try: gamma[:, k] pi[k] * multivariate_normal(meanmu[k], covsigma[k]).pdf(X) except np.linalg.LinAlgError: # 如果协方差矩阵奇异添加一个微小单位阵 sigma_reg sigma[k] 1e-6 * np.eye(D) gamma[:, k] pi[k] * multivariate_normal(meanmu[k], covsigma_reg).pdf(X) # 归一化每个数据点的责任之和为1 row_sums gamma.sum(axis1, keepdimsTrue) row_sums[row_sums 0] 1e-12 # 防止除零 gamma gamma / row_sums return gamma3. M步更新参数利用E步得到的“责任”γ(z_ik)我们更新模型参数。这相当于用加权极大似然估计来更新每个高斯分量。更新每个分量的有效样本数N_k Σ_{i1}^N γ(z_ik)更新混合系数π_k_new N_k / N更新均值μ_k_new (1/N_k) * Σ_{i1}^N γ(z_ik) * x_i更新协方差Σ_k_new (1/N_k) * Σ_{i1}^N γ(z_ik) * (x_i - μ_k_new) * (x_i - μ_k_new)^TM步的Python实现def m_step(X, gamma): M步根据响应度更新模型参数 参数 X: 数据矩阵形状 (N, D) gamma: 响应度矩阵形状 (N, K) 返回 pi_new: 新的混合系数形状 (K,) mu_new: 新的均值矩阵形状 (K, D) sigma_new: 新的协方差矩阵列表 N, D X.shape K gamma.shape[1] # 有效样本数 N_k gamma.sum(axis0) # 形状 (K,) # 更新混合系数 pi_new N_k / N # 更新均值 mu_new np.zeros((K, D)) for k in range(K): mu_new[k] np.dot(gamma[:, k], X) / N_k[k] # 更新协方差 sigma_new [] for k in range(K): # 计算中心化的数据 diff X - mu_new[k] # 形状 (N, D) # 加权协方差计算gamma[:, k, np.newaxis] * diff 用于对每个样本加权 weighted_diff gamma[:, k, np.newaxis] * diff # 形状 (N, D) sigma_k np.dot(weighted_diff.T, diff) / N_k[k] # 形状 (D, D) # 确保协方差矩阵是正定的添加一个小的正则项 sigma_k 1e-6 * np.eye(D) sigma_new.append(sigma_k) return pi_new, mu_new, sigma_new4. 迭代与收敛重复执行E步和M步直到模型参数的变化小于某个阈值或者对数似然函数的变化不再显著。计算对数似然函数用于监控收敛def compute_log_likelihood(X, pi, mu, sigma): 计算当前参数下的对数似然函数值 N, D X.shape K len(pi) likelihood np.zeros((N, K)) for k in range(K): try: likelihood[:, k] pi[k] * multivariate_normal(meanmu[k], covsigma[k]).pdf(X) except np.linalg.LinAlgError: sigma_reg sigma[k] 1e-6 * np.eye(D) likelihood[:, k] pi[k] * multivariate_normal(meanmu[k], covsigma_reg).pdf(X) # 对每个样本求其属于各个分量的概率和然后取log再对所有样本求和 log_likelihood np.sum(np.log(likelihood.sum(axis1) 1e-12)) return log_likelihood主循环def gmm_em(X, K, max_iter100, tol1e-4): 高斯混合模型的EM算法主函数 参数 X: 数据 K: 高斯分量个数 max_iter: 最大迭代次数 tol: 收敛阈值对数似然变化 返回 估计的参数和似然历史 N, D X.shape # 1. 初始化参数 # 使用K-Means中心作为均值初始值是一个好策略 from sklearn.cluster import KMeans kmeans KMeans(n_clustersK, n_init10).fit(X) mu_init kmeans.cluster_centers_ pi_init np.ones(K) / K # 均匀初始化混合系数 sigma_init [np.cov(X.T) for _ in range(K)] # 用全局协方差初始化每个分量 pi, mu, sigma pi_init, mu_init, sigma_init log_likelihood_history [] for i in range(max_iter): # E步 gamma e_step(X, pi, mu, sigma) # M步 pi, mu, sigma m_step(X, gamma) # 计算对数似然 llh compute_log_likelihood(X, pi, mu, sigma) log_likelihood_history.append(llh) # 检查收敛对数似然变化小于阈值 if i 0 and abs(log_likelihood_history[-1] - log_likelihood_history[-2]) tol: print(f迭代 {i1} 次后收敛。) break return pi, mu, sigma, gamma, log_likelihood_history3.3 实战演示与可视化让我们用一个二维的合成数据来演示整个过程并可视化EM算法的迭代过程。import matplotlib.pyplot as plt from sklearn.datasets import make_blobs # 1. 生成模拟数据 X, y_true make_blobs(n_samples500, centers3, cluster_std[0.8, 0.5, 1.0], random_state42) # 2. 运行EM算法 K 3 pi, mu, sigma, gamma, llh_history gmm_em(X, K, max_iter50) # 3. 可视化迭代过程以第一次和最后一次迭代为例 def plot_gmm(X, mu, sigma, gamma, title): plt.figure(figsize(10, 4)) # 子图1数据点与高斯分量中心 plt.subplot(1, 2, 1) plt.scatter(X[:, 0], X[:, 1], cgamma.argmax(axis1), cmapviridis, alpha0.6, s30) plt.scatter(mu[:, 0], mu[:, 1], cred, markerX, s200, labelGMM Centers) for k in range(K): # 绘制高斯分布的置信椭圆2σ eigvals, eigvecs np.linalg.eigh(sigma[k]) angle np.degrees(np.arctan2(*eigvecs[:, 0][::-1])) width, height 2 * np.sqrt(eigvals) # 2倍标准差 ell plt.matplotlib.patches.Ellipse(xymu[k], widthwidth, heightheight, angleangle, alpha0.3, colorred) plt.gca().add_patch(ell) plt.title(f{title} - 聚类结果) plt.xlabel(Feature 1) plt.ylabel(Feature 2) plt.legend() # 子图2每个数据点的责任分布软分配 plt.subplot(1, 2, 2) # 选取一个分量例如第一个的责任值来着色 plt.scatter(X[:, 0], X[:, 1], cgamma[:, 0], cmapReds, alpha0.6, s30) plt.colorbar(labelResponsibility to Cluster 0) plt.scatter(mu[:, 0], mu[:, 1], cblack, markerX, s200) plt.title(f{title} - 对簇0的响应度) plt.xlabel(Feature 1) plt.tight_layout() plt.show() # 为了展示迭代过程我们手动运行两次并绘图实际中应记录每次迭代结果 # 初始化后迭代0次 pi_init, mu_init, sigma_init initialize_parameters(X, K) gamma_init e_step(X, pi_init, mu_init, sigma_init) plot_gmm(X, mu_init, sigma_init, gamma_init, 初始化 (迭代 0)) # 收敛后最终结果 plot_gmm(X, mu, sigma, gamma, f收敛后 (迭代 {len(llh_history)})) # 4. 绘制对数似然收敛曲线 plt.figure(figsize(8, 4)) plt.plot(range(1, len(llh_history)1), llh_history, markero, linestyle-) plt.xlabel(迭代次数) plt.ylabel(对数似然值) plt.title(EM算法收敛过程对数似然单调上升) plt.grid(True, alpha0.3) plt.show()通过可视化你可以清晰地看到初始化时高斯分量的中心可能位置不佳椭圆很大数据点的“责任”分配混乱。迭代过程中中心点逐渐移动到数据密集区域协方差椭圆调整形状以适应簇的分布数据点的“责任”变得越来越清晰一个点对某个分量的责任接近1对其他接近0。收敛后每个高斯分量很好地捕捉到了一个真实的簇对数似然曲线趋于平缓表明算法已收敛。4. EM算法实现中的关键细节与调优经验纸上得来终觉浅绝知此事要躬行。实现EM算法时有几个细节处理不好轻则模型效果差重则程序直接崩溃。4.1 数值稳定性协方差矩阵的奇异性处理这是实现EM算法尤其是GMM时最常见的“坑”。在M步计算协方差矩阵时如果某个高斯分量分配到的有效样本数N_k很少或者分配到的点几乎共线计算出的协方差矩阵可能奇异或病态导致在下一步E步计算多元高斯概率密度时出现数值错误例如计算行列式或求逆矩阵。解决方案正则化在每次更新协方差矩阵后为其对角线添加一个很小的正数。这是最常用且简单有效的方法。sigma_k epsilon * np.eye(D) # epsilon通常取1e-6约束协方差形式根据数据特征假设协方差矩阵为对角阵各特征独立甚至为球形σ²I。这大大减少了参数数量降低了奇异性风险。scikit-learn的GaussianMixture类就提供了covariance_type参数full,tied,diag,spherical。下限约束确保N_k不至于太小。可以在M步更新后检查所有N_k如果某个N_k小于某个阈值如1e-3则重置该分量的参数例如重新随机初始化或者将其混合系数π_k设为零并重新归一化其他系数。4.2 初始化的艺术好的开始是成功的一半EM算法对初始值敏感容易陷入局部最优。糟糕的初始化可能导致某个分量“饿死”N_k - 0。收敛到一个无意义的解如两个分量几乎重合。有效的初始化策略K-Means初始化使用K-Means算法的结果来初始化均值μ_k。K-Means的初始化方式本身就能较好地使初始中心点分散在数据空间中。混合系数π_k可以初始化为均匀分布协方差Σ_k可以初始化为全局数据的协方差矩阵或者由K-Means得到的簇内协方差。多次随机重启这是最鲁棒的方法。随机初始化多组参数分别运行EM算法至收敛最后选择对数似然函数值最大的那组参数作为最终模型。scikit-learn的GaussianMixture中的n_init参数就是干这个的。基于网格或先验知识的初始化如果对数据分布有一定先验认知可以手动指定初始均值的位置。4.3 分量数K的选择奥卡姆剃刀原则高斯混合模型中分量数K是一个超参数需要预先指定。如何选择领域知识如果你知道数据大概由几个过程生成就直接指定。信息准则这是更通用的方法。在模型训练后计算赤池信息准则或贝叶斯信息准则。这两个准则都在模型似然度上增加了对模型复杂度的惩罚参数越多惩罚越大。选择使AIC或BIC最小的K。from sklearn.mixture import GaussianMixture aic_scores, bic_scores [], [] K_range range(1, 8) for k in K_range: gmm GaussianMixture(n_componentsk, covariance_typediag, n_init5).fit(X) aic_scores.append(gmm.aic(X)) bic_scores.append(gmm.bic(X)) # 绘制AIC/BIC曲线选择“肘部”点BIC的惩罚项通常比AIC更重因此在样本量较大时更倾向于选择更简单的模型。可视化与交叉验证对于低维数据可以尝试不同的K并可视化聚类结果。也可以使用轮廓系数等内部评估指标但需注意这些指标不一定与概率模型的拟合优度完全一致。4.4 收敛判断与停止准则迭代何时停止除了设置最大迭代次数外常用的收敛判断标准有对数似然变化当相邻两次迭代的对数似然值之差小于一个预设的阈值tol如1e-4或1e-3时认为已收敛。这是最常用的标准。参数变化监控模型参数如均值μ的变化量。当变化量的范数小于阈值时停止。但参数变化有时不如似然变化稳定。提前停止如果对数似然在连续多次迭代中不再显著提升例如变化小于一个更宽松的阈值也可以提前停止避免不必要的计算。实操心得在调试阶段建议将每次迭代的对数似然和参数变化都打印或记录下来并绘制收敛曲线。这不仅能帮你判断算法是否收敛还能发现潜在问题如震荡、不升反降等。5. 超越GMMEM算法的广泛应用与变体理解了GMM中的EM你就掌握了EM算法的核心范式。这种“在缺失数据下进行最大似然估计”的思想可以推广到无数场景。5.1 隐马尔可夫模型HMM是EM算法在该领域通常称为Baum-Welch算法的另一个经典应用。在HMM中隐变量是隐藏状态序列。E步利用前向-后向算法计算在给定观测序列和当前模型参数下隐藏状态处于某个状态的概率以及状态转移的概率。M步则利用这些概率相当于“软计数”来更新HMM的初始状态分布、状态转移矩阵和观测概率矩阵的参数。用于语音识别、自然语言处理中的词性标注等序列标注任务。5.2 主题模型如概率潜在语义分析pLSA和潜在狄利克雷分配LDA模型。隐变量是文档-词对背后的主题分配。E步计算给定文档和词的情况下主题的后验分布。M步更新主题-词分布和文档-主题分布的参数。这是文本挖掘和无监督特征提取的利器。5.3 含缺失数据的一般模型任何概率模型只要其数据有缺失部分无论是真的缺失还是人为引入的隐变量理论上都可以尝试用EM算法来估计参数。E步求的是缺失数据关于已知数据和当前参数的条件期望M步则是基于“补全”的数据进行极大似然估计。5.4 EM算法的变体与改进在线EM算法适用于数据流式到达的场景每次只用一个小批量数据或单个数据来更新参数避免重新处理全部数据。变分EM算法当E步中隐变量的后验分布p(Z|X,θ)难以精确计算时如在复杂的贝叶斯网络中可以用一个更简单的分布q(Z)来近似它并优化两者之间的KL散度。这就是变分推断的核心思想它使得EM算法能应用到更复杂的模型上。蒙特卡洛EM算法当E步的期望无法解析计算时可以用蒙特卡洛方法如MCMC采样来近似这个期望。M步则基于采样的样本来更新参数。6. 常见问题排查与性能调优实战记录在实际编码和调试EM算法时你几乎一定会遇到下面这些问题。这里是我的排查笔记。6.1 问题对数似然值出现NaN或Inf可能原因与排查协方差矩阵奇异这是最常见原因。在计算多元高斯密度时需要对协方差矩阵求逆并计算行列式。奇异矩阵会导致计算失败。检查在计算概率密度前打印或检查协方差矩阵的条件数np.linalg.cond(sigma[k])。如果条件数极大如1e15则矩阵病态。解决实施强制正则化见4.1节。确保每次更新协方差后都添加epsilon * np.eye(D)。责任度γ出现零在E步计算后某个数据点对所有分量的责任γ都为零可能由于数值下溢导致。随后在M步计算N_k时会出现除以零的错误。检查在E步归一化后检查gamma矩阵是否有全零的行。解决在计算每个分量的概率密度时确保即使很小也有一个非零的下界如1e-300或者在归一化前给所有概率加一个极小的平滑项。数值下溢概率值连乘可能导致数值下溢至零。我们总是在对数空间进行计算。解决实现一个log_sum_exp函数来稳定地计算对数空间中的求和。这是数值计算中的经典技巧。def log_sum_exp(log_probs): 稳定地计算 log(sum(exp(log_probs))) max_val np.max(log_probs, axis1, keepdimsTrue) return max_val.squeeze() np.log(np.sum(np.exp(log_probs - max_val), axis1))在E步中先计算每个数据点属于每个分量的对数概率包括混合系数的log然后使用log_sum_exp来计算归一化的分母再通过指数和对数运算得到最终的gamma整个过程都在对数空间附近进行非常稳定。6.2 问题算法收敛速度慢需要很多次迭代可能原因与排查初始化太差初始中心点挤在一起需要很多轮迭代才能分开。解决采用K-Means初始化。或者在多次随机重启时不仅选择似然最高的结果也检查其收敛速度。学习率或加速技巧标准EM算法可以看作是一种坐标上升法有时收敛较慢。解决可以考虑使用EM算法的加速变体如基于外推的加速方法。一个简单的启发式是如果连续几次迭代中参数更新方向大致相同可以尝试一个稍大的步长。但实现需谨慎。数据尺度差异大如果特征之间的量纲或数值范围差异巨大协方差矩阵的条件数会很大影响数值稳定性和收敛速度。解决标准化你的数据。对于GMM通常建议对每个特征进行零均值、单位方差的标准化。这不会改变数据的聚类结构但能极大改善算法的数值行为。6.3 问题模型过拟合或分量“退化”现象某个高斯分量的协方差矩阵对角线元素变得非常小趋于零该分量几乎只对一个或极少数数据点有高责任对数似然会变得异常高但这通常不是有意义的解。原因与解决 这是GMM使用完全协方差矩阵时的一个已知问题。当某个分量“抓住”了几个非常接近的点时其协方差可以缩到极小使得在这几个点上的概率密度极大从而拉高整体似然。贝叶斯方法/正则化采用贝叶斯视角为参数引入先验分布如对协方差矩阵使用逆Wishart先验。这等价于在M步的协方差更新公式中添加一个正则化项。# 在M步协方差更新中加入一个先验项 sigma_k (np.dot(weighted_diff.T, diff) psi) / (N_k[k] nu) # 其中 psi 是先验的尺度矩阵nu 是自由度通常 psi epsilon * I, nu epsilon约束协方差如前所述使用diag或spherical等更简单的协方差形式从根本上减少过拟合的自由度。下限约束强制协方差矩阵的对角线元素不小于一个最小值如min_covar1e-3。6.4 性能优化技巧向量化计算避免在E步和M步中对每个数据点和每个分量使用for循环。利用numpy的广播机制进行向量化计算可以带来数十倍的速度提升。例如计算所有数据点对所有分量的概率密度可以构造一个(N, K, D)的数组进行操作。提前计算常数对于高维数据计算多元高斯密度时协方差矩阵的逆和行列式是固定的。可以在M步更新协方差后就为每个分量预计算其协方差矩阵的逆矩阵和行列式的对数存储在模型中供E步直接使用。使用专业库对于生产环境或大规模数据直接使用高度优化的库如scikit-learn的GaussianMixture是明智之举。它们用Cython实现并集成了上述所有数值稳定性和性能优化技巧。7. 从理论到拓展EM算法的本质与相关算法对比最后我们跳出代码再回头品味一下EM算法的思想并看看它与其它算法的联系这能帮助你更深刻地理解机器学习。7.1 EM与梯度下降的异同两者都是迭代优化算法但路径不同梯度下降直接对目标函数如负对数似然求梯度沿着梯度反方向更新参数。它需要目标函数可导且学习率的选择很关键。EM算法通过引入隐变量构造了一个易于优化的“代理函数”Q函数并在每次迭代中优化这个代理函数。它不需要直接计算原目标函数的梯度但要求E步的期望可以计算或近似。对于某些问题EM的收敛速度可能比梯度下降更快因为它每次迭代都利用了模型的结构信息。简单来说梯度下降是“硬着头皮往下走”而EM是“先铺好一条好走的路再沿着路走”。7.2 EM与K-Means软聚类与硬聚类K-Means可以看作是GMM在特定极端假设下的一个特例。假设GMM中每个高斯分量的协方差矩阵为εIε趋向于0且先验分布π_k均匀。那么在E步中每个数据点会以概率1分配给离它最近的中心点所对应的分量因为其他分量的概率密度在ε-0时趋于0这就是K-Means的“硬分配”。在M步中更新均值就变成了计算分配给该分量的所有点的平均值与K-Means完全相同。因此K-Means是一种“硬EM”算法。GMM提供了更灵活的“软分配”能刻画簇的重叠、不同形状和大小但计算也更复杂。7.3 如何决定使用EM还是其他算法如果你的数据有明显的簇结构且簇形状近似球形、大小相近K-Means更简单、更快。如果你的数据簇形状复杂、有重叠或者你需要一个概率式的输出例如一个点属于各个簇的概率那么GMMEM算法是更好的选择。如果你的模型天然含有隐变量且完整数据的似然函数易于优化那么EM算法通常是首选的推导和优化框架。当模型非常复杂E步无法精确计算时需要考虑变分EM或蒙特卡洛EM。理解EM算法不仅仅是学会了一个聚类工具更是掌握了一种处理“缺失数据”和“潜在变量”的强大建模思想。这种思想贯穿于机器学习的许多高级主题中。当你下次遇到一个复杂的概率图模型时不妨想想这里有没有隐变量能不能用EM的思想来推导一下它的学习算法

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价