PyTorch矩阵相等但乘法结果不同的问题排查及求助
问题:相同矩阵与同一输入相乘结果不一致?
问题现象
明明w1和w2矩阵数值完全相等,但与同一输入i1相乘后结果却不相等,为何会出现这种情况?
>>> torch.all(w1 == w2) tensor(True) >>> torch.all(i1 @ w1.T == i1 @ w2.T) tensor(False)
环境验证
已确认所有矩阵的数据类型、计算设备完全一致:
>>> i.dtype, w1.dtype, w2.dtype (torch.float32, torch.float32, torch.float32) >>> i.device, w1.device, w2.device (device(type='cpu'), device(type='cpu'), device(type='cpu')) >>> i.shape, w1.shape, w2.shape (torch.Size([3, 4096]), torch.Size([14336, 4096]), torch.Size([14336, 4096]))
差值分析
计算乘积差值后发现存在微小的数值差异:
>>> diff = (i @ w1.T) - (i @ w2.T) >>> diff tensor([[ 1.6391e-07, -5.3644e-07, 3.8743e-07, ..., 4.7684e-07, -5.9605e-08, -3.8743e-07], [ 3.5763e-07, 2.9802e-08, -4.7684e-07, ..., 8.6427e-07, 8.9407e-08, -1.1921e-07], [ 3.5763e-07, 3.4273e-07, -4.4703e-08, ..., 2.1607e-07, 2.3656e-07, 7.4506e-08]], grad_fn=<SubBackward0>) >>> torch.max(torch.abs(diff)) tensor(7.1526e-06, grad_fn=<MaxBackward1>) >>> torch.all(w1 - w2 == 0) tensor(True)
尝试解决
参考相关问题设置了确定性参数,但没有效果:
>>> torch.manual_seed(0) >>> torch.backends.cudnn.deterministic = True >>> torch.backends.cudnn.benchmark = False
问题背景
该问题出现在调试Mixtral-8x7B-v0.1两种实现差异的过程中:i对应current_state,w1为expert_mlp.W_gate.T,w2为HuggingFace实现中的对应矩阵。
原因分析
- 浮点数精度固有局限:float32类型仅能保留6-7位有效数字,大规模矩阵乘法涉及大量累加操作,不同的计算顺序(如并行计算时的累加顺序差异)会导致微小舍入误差累积,最终产生可观测的差值。
- PyTorch底层优化路径差异:即使输入矩阵数值完全一致,不同实现的矩阵乘法可能触发PyTorch不同的底层优化分支(比如CPU上MKL库的不同并行策略),导致计算过程中的舍入方式略有不同,进而出现结果差异。
- 计算图隐式差异:从差值的
grad_fn可以看出,两个矩阵可能属于不同的计算图分支,虽然当前数值相等,但反向传播相关的标记可能影响了前向计算的优化选择(该因素影响相对较小)。
解决方案
- 使用近似相等判断:避免用
torch.all(a == b)判断浮点数计算结果一致性,改用torch.allclose(a, b, rtol=1e-5, atol=1e-8),通过设置合理的相对误差和绝对误差阈值,适配浮点数计算的实际特性。 - 统一计算路径:在CPU环境下,可尝试设置
torch.set_num_threads(1)禁用并行计算,强制单线程累加,减少因并行导致的累加顺序差异;也可以切换到float64 dtype进行计算,提升精度以观察差异是否消失(会增加计算开销)。 - 严格核对矩阵一致性:使用
torch.equal(w1, w2)替代torch.all(w1 == w2),该函数会更严谨地检查矩阵的形状和所有元素的数值一致性,排除潜在的隐式差异。
内容的提问来源于stack exchange,提问作者Joel Burget
相关产品推荐
相关产品推荐

