如何简化NumPy中带循环的np.tensordot表达式 实现无循环并行运算
解决方案
可以用单行代码实现,以下提供几种不同的实现方式,运算结果均和你原有的for循环逻辑完全一致,全量向量化支持并行执行:
- 方案1:使用
@批量矩阵乘法(最简洁,性能最优)
你原有逻辑等价于对每个样本的a做转置后和对应b做矩阵乘法,直接沿样本维度批量运算即可:c = a.transpose(0,2,1) @ b - 方案2:使用
np.einsum(可读性最高,维度对应逻辑清晰)
直接通过维度标记指定要收缩的轴即可:c = np.einsum('ijk,ijl->ikl', a, b) - 方案3:使用
np.tensordot实现
需要额外处理维度对齐,相对前两种写法更繁琐:c = np.tensordot(a, b, axes=([1], [1])).diagonal(axis1=0, axis2=2).T
你可以通过np.allclose(新运算结果, 原for循环输出)验证结果完全匹配。
内容的提问来源于stack exchange,提问作者Thomas Wagenaar
相关产品推荐
相关产品推荐

