Python如何高效计算含缺失值的多元正态分布PDF
带缺失值的多元正态分布PDF计算方案
scipy.stats.multivariate_normal 原生不支持自动识别、忽略NaN缺失值,传入含缺失值的观测向量直接返回nan是预期行为,不存在可直接开启的内置掩码处理参数。
最高效的实现方式完全基于多元正态分布的固有性质:任意维度子集服从对应维度的多元正态边际分布,只需要提取非缺失维度对应的均值子向量、协方差子矩阵,即可直接复用scipy原生接口计算,无额外计算开销,结果和手动降维计算完全等价。
实现逻辑
对任意含NaN的观测向量x,按如下规则提取参数即可:
- 标记
x中所有非NaN值的位置索引为有效索引 - 边际分布均值:原均值向量中,有效索引对应位置的元素构成的子向量
- 边际分布协方差:原协方差矩阵中,行、列索引均属于有效索引的子矩阵
- 边际观测值:
x中有效索引对应位置的非NaN值
可复用实现与样例
from scipy.stats import multivariate_normal as mvnorm import numpy as np def mvnorm_pdf_missing(x, mean, cov, allow_singular=False): x = np.asarray(x) mean = np.asarray(mean) cov = np.asarray(cov) valid_mask = ~np.isnan(x) # 提取边际分布对应参数 marg_mean = mean[valid_mask] marg_cov = cov[np.ix_(valid_mask, valid_mask)] marg_x = x[valid_mask] # 调用scipy原生接口计算 return mvnorm.pdf(marg_x, marg_mean, marg_cov, allow_singular=allow_singular) # 对应测试样例 means = [0.0, 0.0, 0.0] cov = np.array([[1.0, 0.2, 0.2], [0.2, 1.0, 0.2], [0.2, 0.2, 1.0]]) x = [0.5, -0.2, np.nan] print(mvnorm_pdf_missing(x, means, cov))
运行代码输出的结果,和手动将三维分布降维到前两个有效维度计算的PDF值完全一致。
性能优化提示
- 如果批量计算的所有观测向量缺失位置完全固定,可以提前一次性提取对应边际分布的均值和协方差,初始化单个
mvnorm对象重复调用pdf方法,避免重复提取子矩阵的开销 - 如果不同观测向量的缺失位置不一致,必须逐向量按上述逻辑提取对应参数计算,不要跨缺失模式复用分布对象
内容的提问来源于stack exchange,提问作者kallie
相关产品推荐
相关产品推荐

