如何对mx2x2与mx2x1嵌套数组执行批量点积运算?
解决批量嵌套数组的点积运算问题
我来帮你搞定这个批量矩阵点积的需求!针对你这种m×2×2的张量A和m×2×1的张量B,要对每一组对应的2×2矩阵和2×1向量做点积,其实不用纠结tensordot的轴参数,有几种更直观且高效的方法:
方法1:用np.matmul(或@运算符)
np.matmul天生支持批量矩阵乘法,会自动识别前面的m作为批量维度,对每一组A[i]和B[i]执行标准的矩阵-向量乘法:
import numpy as np A = np.array([[[10, -1], [-1, 10]], [[30, 4], [5, 10]]]) B = np.array([[[5],[2]], [[3],[4]]]) # 执行批量点积,得到m×2×1的结果 res_matmul = np.matmul(A, B) print(res_matmul) # 输出: # [[[48] # [15]] # # [[106] # [55]]] # 如果想要更简洁的m×2形式,直接压缩最后一个维度 res_squeezed = res_matmul.squeeze(axis=-1) print(res_squeezed) # 输出: # [[ 48 15] # [106 55]]
这种方法效率极高,哪怕m是250000这样的大数值,也能借助numpy的向量化操作快速完成计算。
方法2:用np.einsum(灵活的张量运算工具)
如果你想更清晰地控制维度的对应关系,np.einsum是绝佳选择,它的语法可以直观描述张量间的运算逻辑:
# 得到m×2×1的结果 res_einsum = np.einsum('mij,mjk->mik', A, B) print(res_einsum) # 和matmul结果一致 # 直接得到m×2的简洁结果,省略最后一个维度 res_einsum_squeezed = np.einsum('mij,mjk->mi', A, B) print(res_einsum_squeezed) # 输出: # [[ 48 15] # [106 55]]
语法解释:'mij,mjk->mik'表示对每个批量m,将A的第2个维度(j)和B的第1个维度(j)相乘求和,最终保留m、i、k三个维度;如果写成'mij,mjk->mi',则直接去掉最后一个维度k,一步得到你想要的简洁形式。
方法3:手动广播求和(原理演示)
如果你想理解底层逻辑,可以用广播和求和手动实现,不过这种方法不如前两种简洁高效,仅作原理参考:
res_manual = np.sum(A * B[:, :, 0, np.newaxis], axis=2)[:, :, np.newaxis] print(res_manual) # 同样得到预期的m×2×1结果
总结
- 优先推荐
np.matmul或@运算符,最直观且性能最优,适合批量矩阵乘法场景; np.einsum适合更复杂的张量运算需求,语法灵活易懂;- 两种方法都能完美支持任意大的
m值,完全满足你的业务场景。
内容的提问来源于stack exchange,提问作者Athena
相关产品推荐
相关产品推荐

