You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何可视化高斯混合模型(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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 10:17:23