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

如何加速/并行计算多组多元正态分布的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 11:43:19