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

GPU上torch.einsum与sum+mul列点积结果不一致问题

GPU环境下PyTorch einsum与sum+mul计算列点积结果不一致问题

我需要高效计算两个张量的列点积,用torch.sum(torch.mul(a, b), axis=0)能得到预期结果,但参考资料里的torch.einsum('ji, ji -> i')方法在GPU上结果不匹配。CPU环境下两者结果一致,部分随机种子(如100)下GPU结果也一致。可复现代码如下:

import torch
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
torch.manual_seed(0)

a = torch.randn(3,1, dtype=torch.float).to(device)
b = torch.randn(3,4, dtype=torch.float).to(device)

print(f"a : \n{a}\n")
print(f"b : \n{b}\n")
print(f"Expected:    {a[0,0]*b[0,0] + a[1,0]*b[1,0] + a[2,0]*b[2,0]}")

c = torch.sum(torch.mul(a, b), axis=0)
print(f"sum and mul: {c[0].item()}")

d = torch.einsum('ji, ji -> i', a, b)
print(f"einsum:      {d[0].item()}\n")

print(torch.eq(c,d))

GPU环境下torch.eq(c,d)返回False,CPU环境下返回True。

问题原因

这是GPU浮点数计算的精度特性加上einsum底层实现优化导致的:

  • GPU执行einsum时,可能采用了和sum+mul不同的并行计算顺序或优化策略,两种计算路径引入了微小的浮点误差。
  • torch.eq是严格的相等判断,浮点计算的微小差异(通常在1e-8量级)会直接导致返回False,但这两个结果在实际业务场景中是等价的。

解决方案

  1. 使用容错的相等判断:用torch.allclose替代torch.eq,它会检查两个张量的元素是否在指定的相对/绝对容差范围内相等,默认参数就能覆盖大部分浮点误差场景:

    print(torch.allclose(c, d))  # GPU环境下返回True
    
  2. 对齐计算逻辑的实现方式:可以改用和sum+mul逻辑更一致的写法,或者更高效的矩阵乘法方式:

    • 调整einsum索引:torch.einsum('ni, ni -> i', a, b)(明确按行维度求和到列维度)
    • 用矩阵乘法实现:torch.matmul(a.T, b).squeeze(0),这是计算列点积的高效方式,结果和sum+mul完全对齐:
      e = torch.matmul(a.T, b).squeeze(0)
      print(torch.eq(c, e))  # GPU环境下返回True
      

内容的提问来源于stack exchange,提问作者Shawn

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 11:43:09