如何可视化高斯混合模型(GMM)中的EM算法迭代步骤?
GMM-EM算法迭代过程可视化实现思路
核心逻辑
EM算法的迭代本质是**E步(计算后验概率)和M步(更新模型参数:权重、均值、协方差)**的循环,可视化的关键就是记录每一轮迭代后的模型参数与聚类分配结果,逐帧绘制后合成动画。
具体实现步骤
1. 自定义训练循环,记录每轮参数
sklearn自带的GaussianMixture.fit()会直接跑完所有迭代,无法获取中间状态,所以需要手动控制迭代过程:
- 初始化GMM时设置
max_iter=1和warm_start=True,这样每次调用fit()只会执行一轮EM迭代,且会基于上一轮的参数继续优化 - 每轮迭代后,记录当前的聚类标签、均值、协方差、权重,同时可以把3D均值通过降维转换到2D空间(方便后续可视化)
2. 降维处理(针对3D数据)
你的数据是3维的,直接做3D动画可读性差,建议用PCA降维到2D:
- 用
sklearn.decomposition.PCA对原始数据做降维,同时把每轮的均值也做同样的降维转换,保证坐标空间一致
3. 逐帧绘制迭代状态
每一轮的可视化需要包含以下元素:
- 降维后的样本点,按当前聚类分配上色
- 每个高斯分量的均值点(用特殊标记如叉号突出)
- 对应高斯分量的置信椭圆(用协方差矩阵计算椭圆的大小、角度,一般选95%置信区间)
4. 合成动画
用matplotlib.animation.FuncAnimation把所有帧串联成动画,可保存为GIF或视频格式,直观展示参数更新和聚类变化的过程。
完整代码示例
import numpy as np import pandas as pd import matplotlib.pyplot as plt from matplotlib.patches import Ellipse from sklearn.mixture import GaussianMixture from sklearn.decomposition import PCA from matplotlib.animation import FuncAnimation # 生成你提供的合成数据 a = np.random.normal(loc=[2,2,2], scale=1.0, size=(100,3)) b = np.random.normal(loc=[5,5,5], scale=1.0, size=(100,3)) c = np.random.normal(loc=[7,7,7], scale=1.0, size=(100,3)) data = np.concatenate((a,b,c), axis=0) df = pd.DataFrame(data, columns=['x', 'y', 'z']) # 3D数据降维到2D pca = PCA(n_components=2) data_2d = pca.fit_transform(df) # 初始化GMM并记录每轮参数 gm = GaussianMixture(n_components=3, random_state=213, max_iter=1, warm_start=True) history = [] # 手动迭代10轮(可根据收敛情况调整) for _ in range(10): gm.fit(df) # 转换均值和协方差到PCA空间 means_2d = pca.transform(gm.means_) covs_2d = [pca.components_ @ cov @ pca.components_.T for cov in gm.covariances_] history.append({ 'labels': gm.predict(df), 'means': means_2d, 'covs': covs_2d }) # 定义单帧绘制函数 def plot_frame(frame_idx): plt.cla() current = history[frame_idx] # 绘制样本点 plt.scatter(data_2d[:,0], data_2d[:,1], c=current['labels'], cmap='viridis', alpha=0.6) # 绘制每个高斯分量的均值和置信椭圆 colors = ['#ff4444', '#4444ff', '#44aa44'] for i in range(3): mean = current['means'][i] cov = current['covs'][i] # 计算椭圆参数(基于协方差的特征值/特征向量) vals, vecs = np.linalg.eigh(cov) order = vals.argsort()[::-1] vals, vecs = vals[order], vecs[:, order] # 95%置信区间对应的缩放系数 scale = np.sqrt(5.991) width, height = 2 * scale * np.sqrt(vals) angle = np.degrees(np.arctan2(*vecs[:,0][::-1])) # 添加椭圆到画布 ellipse = Ellipse(xy=mean, width=width, height=height, angle=angle, color=colors[i], alpha=0.3) plt.gca().add_patch(ellipse) # 标记均值点 plt.scatter(mean[0], mean[1], marker='x', color=colors[i], s=120, linewidth=2) plt.title(f"EM Iteration {frame_idx+1}") plt.xlabel("PCA Component 1") plt.ylabel("PCA Component 2") plt.xlim(data_2d[:,0].min()-1, data_2d[:,0].max()+1) plt.ylim(data_2d[:,1].min()-1, data_2d[:,1].max()+1) # 生成并保存动画 fig = plt.figure(figsize=(8,6)) ani = FuncAnimation(fig, plot_frame, frames=len(history), interval=600) ani.save('gmm_em_iteration.gif', writer='pillow') plt.show()
额外提示
- 可以加入对数似然的判断,当相邻两轮的似然变化小于设定阈值时停止迭代,避免不必要的帧
- 如果坚持做3D可视化,可以用
matplotlib的3D轴,结合mpl_toolkits.mplot3d.art3d.Poly3DCollection绘制椭球,但复杂度会更高 - 调整
interval参数可以控制动画的播放速度
内容的提问来源于stack exchange,提问作者Jean-Paul Azzopardi
相关产品推荐
相关产品推荐

