PyTorch中如何将指定torch.einsum操作等价转换为torch.matmul实现
运算逻辑说明
你使用的einsum表达式bi,bij->bij没有定义任何需要归约求和的维度(即不存在仅出现在输入侧、未出现在输出侧的下标),因此它执行的不是矩阵乘法,而是带广播的逐元素乘法:
- 对每个批次下标
b、每个行下标i,取degree_matrix_hat[b,i]的标量值 - 将该标量与edge_features_reduced
[b,i,:](即第b个样本下第i行的所有j位置元素)相乘 - 最终输出形状和edge_features_reduced完全一致,为
(64, 115, 115)
你之前直接调用torch.matmul得到错误结果,是因为PyTorch的matmul对高维张量执行批量矩阵乘法规则:会将输入的最后两个维度视为矩阵维度,前面的维度视为可广播的批次维度。你传入的degree_matrix_hat形状为(64,115),与3维张量edge_features_reduced运算时会自动补前导维度做广播,最终触发了不符合预期的批量矩阵乘,输出形状变成(64,64,115),和目标逻辑完全无关。
等价实现方案
方案1:广播逐元素乘(最推荐,效率最高)
只需要给degree_matrix_hat在末尾补充一个长度为1的维度,使其形状变为(64,115,1),即可直接通过逐元素乘触发广播,得到和原einsum完全一致的结果:
# 给最后一维补维度后做逐元素乘 result = degree_matrix_hat.unsqueeze(-1) * edge_features_reduced # 验证结果一致性 print(torch.allclose(result, ein1)) # 输出True
方案2:基于matmul的对角矩阵乘实现
如果你必须使用torch.matmul实现,可以利用对角矩阵左乘的性质:对角矩阵左乘另一个矩阵时,等价于将对角线上的元素分别乘到被乘矩阵的对应行上,和目标逻辑完全匹配。你需要先将每个批次的度向量转为对角矩阵,再执行批量矩阵乘:
# 将(64,115)的度向量转为(64,115,115)的批量对角矩阵 diag_mat = torch.diag_embed(degree_matrix_hat) result_matmul = torch.matmul(diag_mat, edge_features_reduced) # 验证结果一致性 print(torch.allclose(result_matmul, ein1)) # 输出True
注意:该方案需要构造额外的对角矩阵,存在大量无意义的乘0计算,运行效率远低于方案1,无特殊需求不建议使用。
内容的提问来源于stack exchange,提问作者Marcos Santana
相关产品推荐
相关产品推荐

