如何仅用Numpy无循环实现3D高斯分布的矩阵运算
用Numpy矢量化操作替代嵌套循环计算3D高斯分布中的c值
你可以用Numpy的矢量化点积操作直接替代嵌套循环,不需要逐元素遍历计算。核心思路是利用Numpy对多维数组的广播和维度求和能力,一次性完成所有位置的向量点积计算。
替代方案说明
原来的循环是对每个(i,j)位置,计算b[i,j]和a[i,j]的点积。由于b和a的形状都是(81,81,2),我们可以通过以下几种方式实现矢量化计算:
元素相乘后沿最后一维求和(最直观推荐):
点积的本质是对应元素相乘再求和,直接对b和a逐元素相乘,然后沿最后一个维度(向量维度)求和即可得到c。使用
np.einsum:
通过爱因斯坦求和符号明确指定维度的对应关系,精准实现点积计算。调整维度后用矩阵乘法:
将a扩展为(81,81,2,1),与b做矩阵乘法后提取结果。
修改后的完整代码
import numpy as np import matplotlib.pyplot as plt R = np.arange(-4,4+1e-9,0.1) X,Y = np.meshgrid(R,R) x = np.stack((X,Y),axis=2) mu = np.array([[-0.5],[-0.5]]) cov = np.array([[1.,0],[0,0.5]]) a = x - mu.T b = np.matmul(a, np.linalg.inv(cov)) # 替代嵌套循环的矢量化计算(三种方式任选其一) c = np.sum(b * a, axis=2) # 方式1:简洁高效,优先选择 # c = np.einsum('ijk,ijk->ij', b, a) # 方式2:明确维度对应逻辑 # c = np.matmul(b, a[..., np.newaxis])[..., 0] # 方式3:矩阵乘法实现 P_xw1 = 1/(np.sum(np.exp(-0.5*c)))*np.exp(-0.5*c) fig = plt.figure(figsize=(9,9)) ax = plt.axes(projection='3d') ax.scatter(X,Y,P_xw1,s=2) plt.show()
效果验证
上述三种矢量化方式的计算结果和原嵌套循环完全一致,同时运行效率远高于循环(尤其是当网格规模更大时,矢量化的优势会更明显)。
内容的提问来源于stack exchange,提问作者Jules
相关产品推荐
相关产品推荐

