PyTorch中Conv2d输出出现无穷大的原因咨询
Conv2d输出出现无穷大的原因分析
问题背景
输入张量x形状为[1,256,60,120],Conv2d层定义如下:
import torch.nn as nn conv2d = nn.Conv2d( 256, 256, kernel_size=2, stride=2, bias=False, )
注:原代码末尾多了一个逗号,会导致conv2d成为元组而非Conv2d层对象,调用conv2d(x)会报错,需先修正该语法问题。
已知参数:
x.max() = tensor(5140., device='cuda:0', dtype=torch.float16)x.min() = tensor(0., device='cuda:0', dtype=torch.float16)conv2d.weight.max() = tensor(1.5796, device='cuda:0')conv2d.weight.min() = tensor(-0.8045, device='cuda:0')
部分场景下conv2d(x).isinf().any()返回True,原因如下:
核心原因:半精度浮点数上溢
torch.float16(半精度浮点数)的最大可表示正值为65504.0,而卷积运算的本质是输入元素与对应权重相乘后累加:
- 单个输出像素的计算涉及
256(输入通道数) × 2×2(卷积核尺寸) = 1024个乘积项的求和。 - 极端情况:当输入元素取最大值
5140.0,权重取最大值1.5796时,单个乘积项的值约为5140 × 1.5796 ≈ 8119.14,1024个这样的项累加结果约为8119.14 × 1024 ≈ 831万,远超float16的最大值65504.0,直接触发数值上溢,结果变为inf。
其他辅助因素
- 无偏置设计:由于未使用偏置项,卷积输出完全由输入与权重的乘积累加决定,缺少偏置的“缓冲”,进一步加剧了大输入值带来的上溢风险。
- 输入数值过大:输入张量
x的最大值达到5140.0,已经接近float16范围的十分之一,再经过大量乘积累加后极易突破上限。
内容的提问来源于stack exchange,提问作者Liubove
相关产品推荐
相关产品推荐

