如何使用np.einsum实现维度为(time, variable, space)的3D矩阵乘法?
关于用np.einsum实现3D批量矩阵乘法的问题解答
可以使用np.einsum实现该运算。
报错原因
你原先的写法np.einsum("tvs,tvs->tvv", M_np, M_np)触发报错有两个核心问题:
- 不符合einsum语法规则:输出维度的下标必须全局唯一,不能重复出现相同的下标
v - 下标逻辑和运算逻辑不匹配:你要实现的是矩阵乘法,需要对空间维度
s做收缩求和,同时保留两个变量维度作为输出的后两维,你写的两个输入都用tvs下标,没有区分两个输入的变量维度,也没有正确声明收缩逻辑。
正确实现方式
提供两种等价的写法:
写法1:配合transpose使用,维度对应更直观
# 写法和matmul的逻辑完全对应,第一个输入维度(t,v,s),转置后的输入维度(t,s,v) np.einsum("tvs,tsv->tvv", M_np, M_np.transpose(0, 2, 1)).shape # 输出为 (100, 64, 64)
写法2:无需提前转置,直接在einsum内完成维度匹配
# 给第二个输入的变量维度单独声明下标w,收缩公共的空间维度s,输出维度为(t,v,w) np.einsum("tvs,tws->tvw", M_np, M_np).shape # 输出为 (100, 64, 64)
结果一致性验证
两种einsum写法的结果和原matmul运算结果完全一致:
import numpy as np M_np = np.random.random((100, 64, 50)) matmul_res = np.matmul(M_np, M_np.transpose(0, 2, 1)) einsum_res1 = np.einsum("tvs,tsv->tvv", M_np, M_np.transpose(0, 2, 1)) einsum_res2 = np.einsum("tvs,tws->tvw", M_np, M_np) print(np.allclose(matmul_res, einsum_res1)) # 输出 True print(np.allclose(matmul_res, einsum_res2)) # 输出 True
内容的提问来源于stack exchange,提问作者Tommy Lees
相关产品推荐
相关产品推荐

