Python绘制高斯混合模型(GMM)聚类椭圆图代码参数错误咨询
代码错误原因及修正方案
1 现有代码的核心错误点
- 构造精度矩阵数组的语法错误:
Covs = np.array(inv_cov1,inv_cov2)不符合numpy数组构造规则,np.array()第一个参数为待转换的数组序列,你需要将两个精度矩阵放在列表中作为入参,否则第二个参数会被识别为dtype参数引发类型报错。 - 补充注意:
precisions_init接收的是精度矩阵(协方差矩阵的逆),如果你代码中的inv_cov1、inv_cov2本身就是提前计算好的精度矩阵可直接使用,若你实际持有的是协方差矩阵,需要先通过np.linalg.inv()求逆后再传入。
2 修正后的可运行代码
from sklearn.mixture import GaussianMixture import pandas as pd import numpy as np import matplotlib.pyplot as plt from matplotlib.patches import Ellipse # 样本数据 X = np.array([[2, 4], [-1, -4], [-1, 2], [4, 0]]) # 精度矩阵(协方差的逆) inv_cov1= [[1,0],[0,1]] inv_cov2= [[0.5,0],[0,0.5]] weights = np.array([0.7,0.3]) means = np.array([[2, 4], [-1, -4]]) # 修正:将两个精度矩阵放入列表再构造numpy数组 Covs = np.array([inv_cov1, inv_cov2]) # 初始化高斯混合模型 gm = GaussianMixture(n_components=2, random_state=0, reg_covar=0, weights_init=weights, means_init=means, precisions_init=Covs).fit(X)
3 绘制聚类簇及高斯椭圆的代码
def draw_ellipse(position, covariance, ax=None, **kwargs): ax = ax or plt.gca() # 协方差矩阵分解得到轴方向和长度 if covariance.shape == (2, 2): U, s, Vt = np.linalg.svd(covariance) angle = np.degrees(np.arctan2(U[1, 0], U[0, 0])) width, height = 2 * np.sqrt(s) else: angle = 0 width, height = 2 * np.sqrt(covariance) # 绘制3σ范围的椭圆 for nsig in range(1, 4): ax.add_patch(Ellipse(position, nsig * width, nsig * height, angle, **kwargs)) # 绘制散点和椭圆 plt.figure(figsize=(8, 6)) plt.scatter(X[:, 0], X[:, 1], c=gm.predict(X), s=50, cmap='viridis') w_factor = 0.2 / gm.weights_.max() for pos, covar, w in zip(gm.means_, gm.covariances_, gm.weights_): draw_ellipse(pos, covar, alpha=w * w_factor) plt.xlabel('x') plt.ylabel('y') plt.title('GMM Clustering Result with Gaussian Ellipses') plt.show()
内容的提问来源于stack exchange,提问作者Martim Correia
相关产品推荐
相关产品推荐

