如何在NumPy中实现高维度自定义规则的矩阵乘法?
简洁实现方法
这里提供两种NumPy原生方法来实现你需要的操作,避免手动索引计算:
方法1:使用np.einsum(最直观)
einsum可以通过索引直接描述张量运算,完美匹配你的需求:
import numpy as np # 假设A和B已定义 C = np.einsum('ijkl,jml->iml', A, B) # C的shape为(2, 20, 5),其中C[0]对应手动计算的C0,C[1]对应C1
索引说明:
ijkl对应A的维度:i=2、j=2、k=1、l=5jml对应B的维度:j=2、m=20、l=5->iml表示对j维度求和,最终保留i(2)、m(20)、l(5)维度,完全符合预期。
方法2:广播+求和(基础NumPy操作)
利用广播机制让A和B维度对齐,再对目标维度聚合:
# 扩展维度实现广播对齐 broadcasted = A[:, :, None] * B[None, :, :] # 对第二个维度(j维度)求和得到结果 C = broadcasted.sum(axis=1)
维度变化说明:
A[:, :, None]的shape为(2, 2, 1, 5),B[None, :, :]的shape为(1, 2, 20, 5)- 广播后两者统一为
(2, 2, 20, 5),每个A[i,j]与B[j]的乘积结果对应手动计算的单项 - 对
axis=1求和后,得到(2, 20, 5)的最终结果,与手动计算逻辑完全一致
验证结果
可以通过以下代码确认两种方法的结果和手动计算一致:
assert np.allclose(C[0], A[0,0]*B[0] + A[0,1]*B[1]) assert np.allclose(C[1], A[1,0]*B[0] + A[1,1]*B[1])
内容的提问来源于stack exchange,提问作者MichaelW
相关产品推荐
相关产品推荐

