使用sklearn的GaussianMixture拟合双高斯分布结果异常排查
问题核心错误点
- 你混淆了样本特征和样本频数的含义:代码里的y是每个x坐标对应的出现频次,不是和x并列的特征维度,GMM拟合需要的是「按频次重复后的x值样本集合」,而非把x和y堆叠后拉平的混合序列。
- 传入GMM的训练数据完全错误:你把
x和y堆叠后调用reshape(-1,1),相当于把所有x、y值混在一起变成了无意义的一维序列,训练的是这个混合序列的分布,自然得不到预期的-15、+15附近的均值。
修正后的代码
import numpy as np import matplotlib.pyplot as plt from sklearn.mixture import GaussianMixture from scipy.stats import norm def main(): x = np.arange(-28,28,2) y = np.array([0,1,2,3,4,5,5,5,4,3,1,1,0,0,0,0,0,0,1,2,3,3,3,2,1,0,0,1]) # 按频数y生成待拟合的x样本集合 data = np.repeat(x, y) # 训练GMM gmm = GaussianMixture( n_components=2, covariance_type='spherical', random_state=42 # 固定随机种子保证结果可复现 ).fit(data.reshape(-1,1)) weights = gmm.weights_ means = gmm.means_ cov = gmm.covariances_ print("权重:", weights) print("均值:", means.flatten()) print("方差:", cov) # 绘制拟合结果 gd0 = weights[0] * norm.pdf(x, means[0][0], np.sqrt(cov[0])) gd1 = weights[1] * norm.pdf(x, means[1][0], np.sqrt(cov[1])) plt.plot(x, y, 'ro', label='原始数据') plt.plot(x, gd0, label='高斯分量1') plt.plot(x, gd1, label='高斯分量2') plt.plot(x, gd0+gd1, label='拟合总和', linestyle='--') plt.legend() plt.savefig("tmp2.png") plt.show() if __name__ == "__main__": main()
修正后效果
运行代码后会得到均值在[-16, 12]附近的两个高斯分量,完全符合你预期的左右两个峰值的拟合要求。
内容的提问来源于stack exchange,提问作者Sakurai.JJ
相关产品推荐
相关产品推荐

