scipy.stats.multivariate_normal.pdf与norm.pdf差异及标量密度获取
区分scipy.stats.multivariate_normal.pdf() 和 scipy.stats.norm.pdf()
先把核心差异说透:
scipy.stats.norm.pdf()是一元正态分布的密度函数,不管输入是单个值还是一维数组,它都会把每个元素当成独立的一元正态变量计算密度,输出形状和输入完全一致。scipy.stats.multivariate_normal.pdf()是多元正态分布的密度函数,专门处理多变量的联合密度——它会通过协方差矩阵考虑变量间的相关性,输入单个向量时输出标量密度值,输入多个向量时输出对应每个向量的密度标量数组。
为什么你的jax代码输出一致?
你用jax得到两者结果相同,大概率是构造了独立多元正态分布(协方差矩阵是对角阵,且每个维度的均值、方差和一元正态完全匹配)。这种情况下,多元正态的联合密度刚好等于各维度一元密度的乘积——比如你在jax里可能做了jax.numpy.prod(norm.pdf(x), axis=-1),自然和multivariate_normal的结果一模一样,但这只是特殊情况。
用scipy获取多元正态归一化密度的正确方式
直接用multivariate_normal.pdf(),指定好均值向量、协方差矩阵,输入目标向量就能得到标量密度值:
from scipy.stats import multivariate_normal, norm import numpy as np # 定义带相关性的二维多元正态分布 mu = np.array([0, 0]) cov = np.array([[1, 0.6], [0.6, 1]]) # 单个向量输入,得到标量密度 x = np.array([1.0, 1.0]) mv_density = multivariate_normal.pdf(x, mean=mu, cov=cov) print(mv_density) # 输出约0.1093 # 对比独立情况的等价性:协方差为对角阵时 cov_diag = np.array([[1, 0], [0, 1]]) mv_density_diag = multivariate_normal.pdf(x, mean=mu, cov=cov_diag) norm_prod_density = np.prod(norm.pdf(x)) print(np.isclose(mv_density_diag, norm_prod_density)) # 输出True,独立场景下两者等价
关键提醒
当变量之间存在相关性(协方差矩阵非对角元素不为0)时,multivariate_normal.pdf()的结果和norm.pdf()乘积的结果会完全不同——这才是两者最核心的区别:多元正态分布考虑变量间的关联,而一元正态只处理单一维度的独立概率。
内容的提问来源于stack exchange,提问作者paul
相关产品推荐
相关产品推荐

