You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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:广播 + 元素乘法 + 求和

逻辑直观,容易理解:

  1. 给decay_kernel插入h维度,让它的形状变成(b,1,f,t,s),实现与decay_q的h维度广播匹配;
  2. 给decay_q插入t维度,形状变成(b,h,f,1,s),实现与decay_kernel的t维度广播匹配;
  3. 对两个张量做元素乘法;
  4. 沿着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做了高度优化,适合这类内积求和场景:

  1. 对decay_kernel做维度重排,将s维度提前,调整为(b,s,f,t)(原形状(b,f,t,s) → permute(0,3,1,2));
  2. 对decay_q做维度重排,调整为(b,s,h,f)(原形状(b,h,f,s) → permute(0,3,1,2));
  3. 用torch.matmul做批量矩阵乘法:每个b和s对应的h×f矩阵与f×t矩阵相乘,得到h×t的结果,此时输出形状为(b,s,h,t);
  4. 最后再做一次维度重排,将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.04 20:16:05