PyTorch float32张量小于比较运算结果异常求助
float32张量比较异常的原因及解决方法
问题复现
执行以下PyTorch代码时出现异常:
import torch t = torch.load(r"value.pt") print(t.shape, t.dtype) #t = t.double() for i in range(t.shape[0]): print(i, "%.20f" % (t[i].sum(-1)-1)) print((t.sum(-1)-1).abs()<1e-6) print("%.8e"%(t[35].sum()-1), (t[35].sum(-1)-1).abs()<1e-6, (t[34:50].sum(-1)-1).abs()<1e-6, (t[34:40].sum(-1)-1).abs()<1e-6)
输出结果显示:
torch.Size([100, 1600]) torch.float32 ... 33 -0.00000008132246875903 34 0.00000014945180737413 35 0.00000053211988415569 36 -0.00000006957179721212 37 -0.00000010645544534782 38 -0.00000000481304596178 ... tensor([ True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, False, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True, True], device='cuda:0') 7.15255737e-07 tensor(True, device='cuda:0') tensor([ True, False, True, True, True, True, True, True, True, True, True, True, True, True, True, True], device='cuda:0') tensor([True, True, True, True, True, True], device='cuda:0')
核心问题:
- 第35行的
sum-1绝对值约5.3e-7(小于1e-6),但整张量比较时显示False - 不同切片/索引方式下,同一行的比较结果不一致
原因分析
1. float32的精度限制
float32单精度浮点数仅能提供6-7位有效十进制数字,而1e-6刚好处于其精度边缘。当对1600个float32元素累加时,累加过程中的舍入误差会逐步累积,最终的sum结果可能和理论值存在微小偏差。这种偏差在某些计算路径下可能刚好超过1e-6的阈值,导致比较结果为False。
2. CUDA并行计算的非确定性
在CUDA设备上,张量的sum操作是通过并行归约实现的。不同的张量范围(整张量、切片、单元素)会触发不同的并行计算策略,导致元素的累加顺序不同。由于浮点数加法不满足结合律((a+b)+c ≠ a+(b+c)在精度有限时),不同的累加顺序会产生细微不同的sum结果,这就是为什么同一行在整张量比较和单独索引比较时结果不一致。
解决方法
1. 切换到更高精度的数据类型
将张量转换为float64(double)类型,其拥有15-17位有效数字,累加误差会被大幅降低,足以避免这种阈值附近的判断异常:
t = t.double()
2. 调整比较阈值
避免使用刚好卡在float32精度边缘的阈值,根据实际情况适当放宽,比如将1e-6调整为1e-5,或者基于float32的机器epsilon(约1.19e-7)设置合理容差:
# 用1e-5替代1e-6 print((t.sum(-1)-1).abs() < 1e-5)
3. 高精度累加
即使保持张量为float32,也可以在累加时指定更高精度的计算类型,减少累加过程中的误差:
# 累加时用float64计算,结果再转回float32 sum_result = t.sum(-1, dtype=torch.float64) - 1 print(sum_result.abs() < 1e-6)
4. 切换到CPU计算
CPU上的归约计算顺序更稳定,不会因为并行策略不同产生结果差异,适合需要精确一致结果的场景:
t_cpu = t.cpu() print((t_cpu.sum(-1)-1).abs() < 1e-6)
内容的提问来源于stack exchange,提问作者user2224350
相关产品推荐
相关产品推荐

