PyTorch中矩阵滚动后矩阵乘法结果不一致的原因及解决方法
问题:为什么在PyTorch中对矩阵执行roll操作后,矩阵乘法的结果差异如此之大?
说明:虽然浮点数乘法存在精度问题,但此处的差距略大,且该差异与滚动步长相关。
代码示例:
x = torch.rand((1, 5)) y = torch.rand((5, 1)) print("%.10f"%torch.matmul(x,y)) >>> 1.2710412741 print("%.10f"%torch.matmul(torch.roll(x, 1, 1), torch.roll(y, 1, 0))) >>> 1.2710412741 print("%.10f"%torch.matmul(torch.roll(x, 2, 1), torch.roll(y, 2, 0))) >>> 1.2710413933
请问上述问题的成因是什么?如何获得更一致的计算结果?
成因分析
这是浮点数加法的实际非交换性导致的:
- 从数学理论上,矩阵乘法的结果应完全一致,因为加法满足交换律,仅改变求和项顺序不影响最终值。但在有限精度的浮点数运算中,加法顺序会直接影响舍入误差的累积。
- 执行
roll操作后,矩阵乘法的求和项顺序发生了变化。例如原求和顺序为x0y0 + x1y1 + x2y2 + x3y3 + x4y4,滚动步长为2时,求和顺序变为x3y3 + x4y4 + x0y0 + x1y1 + x2y2。 - 浮点数加法中,不同的累加顺序会产生不同的中间结果:若先累加大数再累加小数,小数的贡献可能被大数的精度截断;反之则能保留更多有效信息。当求和项的数值差异较大时,这种误差会被放大,最终导致结果出现可见差异。
解决方法
要获得更一致的计算结果,可尝试以下方案:
- 切换更高精度的浮点数类型:将张量的
dtype从默认的float32改为float64(双精度)。双精度浮点数拥有更多有效位数,能大幅降低舍入误差的影响。示例:x = torch.rand((1,5), dtype=torch.float64) y = torch.rand((5,1), dtype=torch.float64) - 统一求和顺序:若业务逻辑允许,手动固定求和项的顺序(如按数值大小排序后累加),避免因
roll操作改变求和顺序引入误差。 - 使用数值稳定的求和方法:例如在求和时指定更高精度的临时 dtype,或实现Kahan求和算法来抵消累加过程中的舍入误差。
内容的提问来源于stack exchange,提问作者opflow
相关产品推荐
相关产品推荐

