如何加速/并行计算多组多元正态分布的PDF值?
问题描述
我有三个数组:
- N×3的点数组(
points) - N×3的均值数组(
means) - N×3×3的协方差数组(
covs)
每个点对应一个独立的3维多元高斯分布(对应一组均值和协方差)。目前只能通过Python for循环逐个计算每个点的PDF值,代码如下:
import numpy as np from scipy.stats import multivariate_normal # 生成示例数据 N = 5 # 示例用小数值,实际场景可增大 points = np.random.rand(N, 3) means = np.random.rand(N, 3) covs = np.array([np.eye(3) for _ in range(N)]) # 用单位矩阵作为示例协方差 # 存储PDF值的数组 pdf_values = np.zeros(N) # 逐个计算PDF for i in range(N): pdf_values[i] = multivariate_normal.pdf(points[i], mean=means[i], cov=covs[i]) print("Points:\n", points) print("Means:\n", means) print("Covariances:\n", covs) print("PDF Values:\n", pdf_values)
尝试直接把整个数组传入multivariate_normal.pdf不被支持,希望找到避免Python循环的加速方法,非SciPy实现也可以。
解决方案
方法1:NumPy完全向量化实现
直接基于多元高斯PDF公式手动实现批量计算,所有操作依赖NumPy的底层C实现,完全消除Python循环:
import numpy as np def vectorized_multivariate_gaussian_pdf(points, means, covs): D = points.shape[1] # 维度固定为3 # 批量计算协方差行列式的平方根 sqrt_dets = np.sqrt(np.linalg.det(covs)) # 批量计算协方差的逆矩阵 inv_covs = np.linalg.inv(covs) # 计算每个点与对应均值的差 diffs = points - means # 用 einsum 高效计算批量马氏距离 mahalanobis = np.einsum('ni,nij,nj->n', diffs, inv_covs, diffs) # 计算最终PDF值 normalization = 1.0 / ((2 * np.pi) ** (D/2) * sqrt_dets) pdf = normalization * np.exp(-0.5 * mahalanobis) return pdf # 测试代码 N = 5 points = np.random.rand(N, 3) means = np.random.rand(N, 3) covs = np.array([np.eye(3) for _ in range(N)]) pdf_values_vec = vectorized_multivariate_gaussian_pdf(points, means, covs) print("向量化计算结果:\n", pdf_values_vec)
方法2:Numba JIT加速循环
如果需要贴近原逻辑,用Numba对循环进行JIT编译,把Python循环转换成机器码执行,大幅降低循环开销:
import numpy as np from numba import jit @jit(nopython=True) def numba_accelerated_pdf(points, means, covs): N = points.shape[0] D = points.shape[1] pdf_values = np.zeros(N) for i in range(N): diff = points[i] - means[i] inv_cov = np.linalg.inv(covs[i]) det = np.linalg.det(covs[i]) mahalanobis = diff @ inv_cov @ diff.T pdf_values[i] = (1.0 / ((2 * np.pi) ** (D/2) * np.sqrt(det))) * np.exp(-0.5 * mahalanobis) return pdf_values # 测试代码 pdf_values_numba = numba_accelerated_pdf(points, means, covs) print("Numba加速结果:\n", pdf_values_numba)
注:第一次运行会有编译开销,后续重复调用速度极快。
方法3:JAX原生批量支持
如果可以引入JAX库,它的multivariate_normal.pdf原生支持批量输入(每个样本对应独立的均值和协方差),还能自动利用GPU加速:
import jax import jax.numpy as jnp from jax.scipy.stats import multivariate_normal # 直接批量计算 pdf_values_jax = multivariate_normal.pdf(points, mean=means, cov=covs) print("JAX计算结果:\n", pdf_values_jax)
内容的提问来源于stack exchange,提问作者Valeria
相关产品推荐
相关产品推荐

