如何将torch.einsum("bfts,bhfs->bhts")转换为matmul等原生方法?
转换
torch.einsum("bfts,bhfs->bhts")为原生PyTorch操作 完全可行,先拆解这个einsum的计算逻辑,再给出两种高效的原生实现方式:
第一步:明确einsum的计算逻辑
这个表达式是对两个张量的f维度做内积求和,最终输出维度为(b,h,t,s)。具体来说,对于每个batch b、每个h/t/s位置,计算所有f上的元素乘积之和:output[b,h,t,s] = sum_{f} decay_kernel[b,f,t,s] * decay_q[b,h,f,s]
方法1:广播 + 元素乘法 + 求和
逻辑直观,容易理解:
- 给
decay_kernel插入h维度,让它的形状变成(b,1,f,t,s),实现与decay_q的h维度广播匹配; - 给
decay_q插入t维度,形状变成(b,h,f,1,s),实现与decay_kernel的t维度广播匹配; - 对两个张量做元素乘法;
- 沿着
f维度求和,得到目标输出。
代码示例:
# 输入张量:decay_kernel (b,f,t,s), decay_q (b,h,f,s) expanded_kernel = decay_kernel.unsqueeze(1) # shape: (b,1,f,t,s) expanded_q = decay_q.unsqueeze(3) # shape: (b,h,f,1,s) output = (expanded_kernel * expanded_q).sum(dim=2) # sum over f维度
方法2:维度重排 + 矩阵乘法(性能更优)
PyTorch的torch.matmul做了高度优化,适合这类内积求和场景:
- 对
decay_kernel做维度重排,将s维度提前,调整为(b,s,f,t)(原形状(b,f,t,s)→permute(0,3,1,2)); - 对
decay_q做维度重排,调整为(b,s,h,f)(原形状(b,h,f,s)→permute(0,3,1,2)); - 用
torch.matmul做批量矩阵乘法:每个b和s对应的h×f矩阵与f×t矩阵相乘,得到h×t的结果,此时输出形状为(b,s,h,t); - 最后再做一次维度重排,将
s维度移到最后,得到(b,h,t,s)的目标形状。
代码示例:
# 输入张量:decay_kernel (b,f,t,s), decay_q (b,h,f,s) kernel_permuted = decay_kernel.permute(0, 3, 1, 2) # shape: (b,s,f,t) q_permuted = decay_q.permute(0, 3, 1, 2) # shape: (b,s,h,f) matmul_result = torch.matmul(q_permuted, kernel_permuted) # shape: (b,s,h,t) output = matmul_result.permute(0, 2, 3, 1) # shape: (b,h,t,s)
你可以用torch.allclose对比原生einsum输出和上述两种实现的结果,验证正确性。
内容的提问来源于stack exchange,提问作者majstor
相关产品推荐
相关产品推荐

