使用SciPy计算多元正态分布时的维度匹配错误排查与修正
问题排查与修正方案
错误原因
- 参数顺序完全错误:
scipy.stats.multivariate_normal的构造函数第一个参数是均值(mean),第二个是协方差(cov),原代码把数据点x放在了第一个参数位置,导致程序误将x(形状(2,2))当作均值向量,认为其长度应为4,触发ValueError。 - 输入维度不匹配:测试传入的
test_mu是二维数组(1,2)、test_sigma是三维数组(1,2,2),但multivariate_normal要求均值是一维向量、协方差是二维矩阵。 - 误用构造函数:直接实例化
multivariate_normal并不会计算概率密度,需要调用.pdf()方法来获取数据点对应的多元正态分布概率值。
修正后的函数
from scipy.stats import multivariate_normal import numpy as np def calc_gamma(x, mu, sigma): # 压缩mu和sigma的多余维度,转为一维向量和二维矩阵 mu_1d = mu.squeeze() sigma_2d = sigma.squeeze() # 计算每个数据点的多元正态分布概率密度 Gamma = multivariate_normal.pdf(x, mean=mu_1d, cov=sigma_2d) return Gamma
测试验证
传入给定测试值后,会返回两个数据点对应的概率密度值:
test_data = np.array([[0, 0],[1,2]]) test_mu = np.array([[1, 1]]) test_sigma = np.array([[[1, 0], [0, 1]]]) test_gamma = calc_gamma(test_data, test_mu, test_sigma) # 输出结果示例:array([0.05855018, 0.05855018])
内容的提问来源于stack exchange,提问作者aflatoon
相关产品推荐
相关产品推荐

