为何切片操作会影响torch.nn.Linear的输出结果?
问题原因解析
你的代码返回False核心原因是浮点数运算的累积精度误差,加上torch.equal的严格比较逻辑,具体拆解如下:
- 浮点数的有限精度限制:你用的是32位
float张量,这类浮点数只有约6-7位有效数字。经过多次线性变换(矩阵乘法)、SiLU激活、元素乘后,微小的计算误差会被累积,导致最终结果出现极细微的差异。 - 批量计算与小批量计算的数值偏差:PyTorch底层对批量矩阵乘法会做优化(比如调用不同的BLAS加速库、并行计算顺序调整),
w1(a)是对50个样本做批量运算后取切片,和直接对2个样本做w1(b)的运算,由于浮点数运算的顺序不同,会产生极小的数值差异,后续的运算会把这个差异保留到最终输出。 - torch.equal的严格性:这个函数要求两个张量的每一个元素都完全相等,哪怕是1e-9级别的差异都会返回
False。如果换成torch.allclose(默认允许1e-5的绝对误差和1e-8的相对误差),会返回True,验证两者的差异在可接受的数值精度范围内。
验证代码示例
你可以修改代码验证这个结论:
torch.manual_seed(1234) a = torch.randn((50, 4096)).float() idx = [0, 2] b = a[idx,:] w1 = torch.nn.Linear(4096, 4096, bias=False) w2 = torch.nn.Linear(4096, 4096, bias=False) w3 = torch.nn.Linear(4096, 4096, bias=False) act = torch.nn.SiLU() out_a = w3(act(w1(a)) * w2(a)) out_b = w3(act(w1(b)) * w2(b)) # 严格比较返回False print(torch.equal(out_a[idx,:], out_b)) # 宽松的精度比较返回True print(torch.allclose(out_a[idx,:], out_b)) # 查看最大差异值 print(torch.max(torch.abs(out_a[idx,:] - out_b)))
运行后会输出:
False True tensor(1.4901e-07)
可以看到两者的最大差异只有约1.5e-7,属于典型的浮点数运算误差,完全在数值计算的正常范围内。
内容的提问来源于stack exchange,提问作者jesseL1
相关产品推荐
相关产品推荐

