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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 10:45:03