兼容Numba njit的Numpy无for循环点积累积和计算方案
Numba兼容的Numpy向量化优化方案
首先我们可以先对原式做代数化简,降低计算量:
原式每一项为 xi[i] @ R[i] @ R[i].T @ xi[i],根据矩阵乘法结合律,等价于 (xi[i] @ R[i]) @ (xi[i] @ R[i]).T,也就是长度为3的向量自身的内积,等于该向量所有元素的平方和,无需计算开销更大的N×N中间矩阵R[i] @ R[i].T。
以下两种实现均在Numba支持的Numpy函数范围内,可直接在@njit装饰的函数中使用:
方案1:基于einsum实现
import numpy as np h = np.sum(np.einsum('ij,ijk->ik', xi, R) ** 2)
np.einsum('ij,ijk->ik', xi, R)会批量计算每个i对应的xi[i] @ R[i],输出形状为(N, 3)- 对输出逐元素平方后求和,即可得到最终结果h
方案2:基于批量矩阵乘法实现
如果更习惯matmul写法,可使用等价实现:
h = np.sum((xi[:, np.newaxis, :] @ R).squeeze() ** 2)
- 给xi增加维度得到形状为(N, 1, N)的数组,和形状为(N, N, 3)的R做批量矩阵乘法,得到形状为(N, 1, 3)的结果
- squeeze去掉多余维度后平方求和,和方案1结果完全一致
两种实现均为纯向量化操作,无Python层面循环,计算复杂度从原实现的O(N³)降低到O(N²),在N=100的场景下性能明显优于原for循环实现,且无需开启并行不会产生额外开销。
内容的提问来源于stack exchange,提问作者Francesco Musso
相关产品推荐
相关产品推荐

