如何加速Python中Numpy高维数组的批量dot点积运算
矩阵批量运算加速方案
你当前代码的性能瓶颈完全来自Python层遍历10000个样本的循环,numpy本身的向量化运算可以直接替代这层循环,所有运算下沉到C层执行,速度会有数量级的提升,以下是可行方案:
- 方案1:使用
np.matmul/@运算符直接批量计算
numpy的矩阵乘法运算符默认支持前导维度的批量广播,只要维度匹配不需要额外处理,代码可以直接简化为:
def get_P(X, Z): # X shape: (n_sample, n_m, n_t, n_f) # Z shape: (n_sample, n_m, n_f, n_t) return X @ Z
这种写法最简洁,底层自动做批量并行运算,完全去掉了Python循环。
- 方案2:使用
np.einsum自定义运算路径
如果你需要更灵活的维度匹配控制,可以用爱因斯坦求和函数,开启优化参数后性能非常可观:
def get_P(X, Z): return np.einsum('ijkl,ijlm->ijkm', X, Z, optimize=True)
optimize=True会自动选择最优的运算调度逻辑,对于大张量的计算效率提升尤其明显。
- 进阶GPU加速方案
如果你的设备有NVIDIA显卡,可以用CuPy库完全替换numpy依赖,代码不需要做其他修改,就能调用GPU做大规模并行计算,相比CPU运算速度还能再提升几十到上百倍。
以上优化方案都可以把原来预计几小时的运算耗时压缩到分钟甚至秒级。
内容的提问来源于stack exchange,提问作者Linyun He
相关产品推荐
相关产品推荐

