Numpy多维数组乘法报形状不匹配ValueError错误的原因及解决求助
报错原因解释
核心是两个问题共同导致的计算失败:
- 参与计算的变量存在pandas DataFrame/Series类型,而非纯numpy数组。从报错栈的
pandas/core/generic.py in __array_ufunc__调用路径可以看到,numpy的matmul触发了pandas自定义的右矩阵乘法逻辑__rmatmul__,该方法会对输入做多余的转置操作,把原本符合维度匹配要求的数组转成了(200,1,2808702)的错误维度,最终出现维度对齐失败。 - 原本纯numpy数组的维度逻辑是成立的:
(2808702, 1, 200) @ (200, 200)广播后得到(2808702, 1, 200),再和(2808702, 200, 1)相乘得到(2808702, 1, 1),reshape后可以得到目标一维结果,本身维度设计没有问题,问题完全来自pandas类型的混入。
修复方案
只需要先把所有参与矩阵运算的变量转为纯numpy数组,再执行计算即可:
# 将所有输入转为numpy数组,排除pandas类型干扰 data1_np = data1.to_numpy() S_np = S.to_numpy() data2_np = data2.to_numpy() # 执行矩阵乘法,输出结果维度为(2808702,) results = np.matmul(np.matmul(data1_np, S_np), data2_np).reshape(dd_kpy.shape[0])
也可以用更简洁的爱因斯坦求和写法,避免多层matmul嵌套,逻辑更清晰:
results = np.einsum('ni,ij,nj->n', data1_np.squeeze(), S_np, data2_np.squeeze())
内容的提问来源于stack exchange,提问作者DT Ngo
相关产品推荐
相关产品推荐

