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

为何切片操作会影响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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 11:29:55