1. 从一次“翻车”说起:深层网络为什么不香了
先说个我自己的事儿。前阵子帮朋友调一个图像分类模型,他基线用的是18层的网络,效果马马虎虎,准确率在验证集上死活上不去。我寻思着按经验把网络加深到56层总该有个提升吧,结果训练完一看,不仅没涨,反而掉了两个多点。当时我盯着Loss曲线看了半天,训练集上的Loss居然也比浅层网络高。这就说明不是过拟合的问题,是网络本身“学不动”了。
这就是残差网络出现之前,深度学习最尴尬的那个阶段。你堆再多的层,理论上表达能力更强,但实际训练中梯度传播越来越难,网络越深反而越差。后来我认真去翻ResNet那篇论文,才知道这里面有一个被很多人忽略的细节,那就是残差连接并不是为了“让网络更深”而强行加上的技巧,它本质上是在改变化优问题的地形。这几年我在CV、PINNs、各种回归任务里反复用到残差思想,越用越觉得,它其实是深度学习里最朴素也最容易被低估的一个突破。
这篇是这个系列的第三篇,前面聊过一些别的机缘巧合,这次就好好说说残差这件事:从残差块的结构、恒等映射的数学直觉,到工程实现里的维度匹配问题,再到残差思想怎么从图像分类一路“出圈”到物理信息神经网络(PINNs)的残差修正,最后把我踩过的坑一并列出来。
2. 残差块和恒等映射:那个“抄近路”的设计到底妙在哪
2.1 从“直连”到“绕道”:残差块的结构拆解
先看最基本的残差块长什么样。假设你的网络某一层想学的是一个映射 H(x),传统网络就是让这一层的权重直接去拟合 H(x)。残差块的做法是,把这层想学的东西拆成两部分:恒等映射 x 和残差 F(x),让这一层的实际输出变成:
y = F(x) + x
这里的 F(x) 通常是两个卷积层(或者两个全连接层)加激活函数的组合,x 就是输入本身。这个结构里,从输入 x 到输出 y 有一条“公路”直接穿过,这条捷径在论文里叫 shortcut connection,中文常叫捷径连接或短连接。
我当时看到这个结构的第一反应是:这不就是把输入加了个旁路吗,能有多大区别?真正让我改变想法的是一次实验。我在同样的数据集上分别训练了一个34层的普通网络和一个34层的残差网络,普通网络的训练Loss卡在1.2左右下不去,残差网络却能一路降到0.8以下。训练曲线摆在那里,你不得不承认,加一条“绕道”的旁路,整个学习过程就顺畅多了。
这里面最关键的逻辑在于,拟合目标变了。普通网络要直接学 H(x),残差网络只需要学 H(x) - x,也就是输入和输出之间的差值。大多数情况下,输出和输入差别没那么大,逼近一个接近零的小函数,比从零开始逼近一个复杂的映射容易得多。这就像是让你直接画出整幅画,和让你在原图上轻轻描几笔修改,难度完全不是一个档次的。
2.2 恒等映射的数学直觉:为什么多一条路就变好了
从数学上看,残差块等价于把原来的一层变换改成了这样一个形式。如果不加残差连接,某一层网络的梯度在反向传播时要经过权重矩阵的连乘,层数一深,梯度要么指数级衰减到零,要么指数级爆炸。加了残差连接之后,反向传播的路径里多了一条“高速公路”,梯度可以通过恒等映射这条路径几乎无损地传回浅层。
这就像一条大河,原来的主干道水流湍急,容易断流;现在你在旁边开了一条缓坡支流,即使主干道堵了,水也能沿着支流源源不断地流回去。这个支流对梯度的意义就是,它让浅层的参数始终能收到足够强度的更新信号,不会因为层数增加而陷入“学不到东西”的死循环。
不过有一点必须说清楚,恒等映射的数学优势不是“梯度不衰减了”,而是“梯度的衰减有了兜底”。F(x)+x 结构里,梯度在恒等路径上可以保持为1的系数传播,但同时它也会经过 F(x) 这条路径,经过两层卷积。所以严格来说,残差网络让梯度有了两条可选择的路径,一条无损,一条有损,网络在训练过程中会自己去选择更优的那条路线。这也是为什么残差网络对初始化不那么敏感的原因之一。
2.3 结构图里那些被忽略的细节:维度、步长和1x1卷积
很多人看残差网络结构图时只看到了那个“弯过去的弧线”,但真正动手实现时会发现,问题全出在shortcut的维度匹配上。
最基本的残差块要求输入和输出的维度完全一致,这样才能直接做加法。但在真实网络里,卷积层会改变特征图的通道数,或者通过步长为2的卷积来降低特征图的空间尺寸。这时候x和F(x)的shape就不一样了,直接相加会报错。论文里给了两个方案:
- 如果只是通道数不同,空间尺寸一样,可以用零填充,把x的通道数补到和F(x)一致,然后直接相加。
- 如果空间尺寸也变了,就要在shortcut上加一个1x1卷积,步长设为2,把x的尺寸和通道都调整到和F(x)一致。
我当时写代码时偷懒,全部用1x1卷积去匹配维度,结果模型参数暴增,训练速度明显变慢。后来才发现,在空间尺寸不变的层之间,直接做恒等映射才是最优的,只有当下采样时才需要引入1x1卷积。这个细节,看结构图是看不出来的,只有自己动手搭一遍才能体会。
3. 残差计算的工程实现:从理论到能跑的代码
3.1 一个标准的残差块长什么样
说了这么多理论,直接上代码。我用PyTorch写一个最常用的BasicBlock,对应ResNet18和ResNet34的基础模块。这个块的设计思路很清晰:两个3x3卷积,中间夹一个ReLU,最后把输入加上去,再做一次ReLU。
import torch import torch.nn as nn import torch.nn.functional as F class BasicBlock(nn.Module): def __init__(self, in_channels, out_channels, stride=1): super(BasicBlock, self).__init__() self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(out_channels) # 当输入输出维度不一致时,用1x1卷积调整 self.shortcut = nn.Sequential() if stride != 1 or in_channels != out_channels: self.shortcut = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False), nn.BatchNorm2d(out_channels) ) def forward(self, x): identity = x out = self.conv1(x) out = self.bn1(out) out = F.relu(out) out = self.conv2(out) out = self.bn2(out) # 残差连接:F(x) + x out += self.shortcut(identity) out = F.relu(out) return out这里有几个关键点。第一个是bias=False,因为卷积后面跟了BatchNorm,BN层自带偏置项,卷积层再加bias就重复了,而且会影响BN的统计量计算。第二个是BatchNorm的位置,放在卷积之后、ReLU之前,这是论文里的原始设计顺序,实际效果也最稳定。第三个是shortcut的构建方式,直接用Sequential包起来,维度一致的情况下就是空模块,forward里相当于直接加原输入。
3.2 维度匹配的三种处理方式对比
实际搭网络时,你最常遇到的问题就是维度对不上。我总结了一下,一共有三种处理方式,各有各的适用场景,我把它们整理成了表格:
| 方式 | 实现方法 | 适用场景 | 优缺点 |
|---|---|---|---|
| 零填充 | 在x上补零到目标通道数 | 通道数增加且空间尺寸不变 | 不增加参数,但新增通道全是0,信息量有限 |
| 1x1卷积 | 用1x1卷积调整x的通道和尺寸 | 空间尺寸和通道都变化 | 增加少量参数,能学到shortcut上的变换,最常用 |
| 恒等+下采样 | 空间尺寸用平均池化或最大池化调整,通道用1x1卷积 | 通道数翻倍且尺寸减半 | 下采样时没有额外计算量,适合大网络 |
我的使用经验是:当空间尺寸减半、通道数翻倍的过渡层,直接用1x1卷积加stride=2的方式一步到位;当仅通道数变化时,优先用零填充,让恒等映射保持真正的“恒等”。不管你选哪种方式,核心思路都是让x能够以尽可能少的参数变换去匹配F(x)的shape,而不是让shortcut变成一个复杂的学习模块。
3.3 训练残差网络的几个心得
残差网络的训练和普通网络不太一样,如果你是从零开始自己实现,我建议你注意下面几点:
学习率的选择。ResNet论文里用的是初始学习率0.1,配合批大小256。如果你自己训练,批大小减半时学习率最好也减半,不然容易在训练初期就出现Loss震荡。我实际试过,在CIFAR-10上用批大小128训练ResNet18,学习率0.1太大,直接不收敛,降到0.05才稳定下来。
BatchNorm的动量参数。残差网络里BN层的momentum默认是0.1,这个值在大多数情况下没问题,但是当你的batch size偏小(比如小于32)时,batch内统计量波动大,建议把momentum调到0.05甚至0.01。否则你会在验证集上看到一种诡异的现象:训练Loss正常下降,验证Loss却来回跳。
初始化也要注意。PyTorch默认的初始化方式对残差网络基本够用,但如果你追求更好的效果,可以在每个卷积层上用Kaiming初始化,并把最后一个BN层的gamma初始化为0。这样在训练初期,残差块的输出就等于输入本身,相当于网络从恒等映射开始学起,稳定性会更好。这个技巧在训练超深网络时尤其明显。
4. 残差思想如何“出圈”:从残差网络到PINNs残差修正
4.1 物理信息神经网络里的“残差”是什么
如果说残差网络是残差思想在图像领域的代表作,那PINNs(Physics-Informed Neural Networks)就是残差思想在科学计算领域的一次漂亮应用。我第一次接触PINNs时觉得很新奇,因为它不是用神经网络去拟合数据,而是用神经网络去近似求解偏微分方程。
PINNs的核心思路是这样的:假设你要求解一个方程,方程在区域内有一个形式F(u)=0,你用一个神经网络来表示未知解u(x),网络的输出记为u_hat(x)。如果u_hat真的是精确解,那么把u_hat代入原方程,F(u_hat)应该处处为零。但神经网络一开始是随机初始化的,它输出的东西代进方程当然不为零,这个不为零的误差,就是方程残差。
所以PINNs的Loss函数里最重要的一项,就是要把这个方程残差压到最小。你看,这背后的逻辑和残差网络是惊人的一致:不直接去逼近目标函数本身,而是去逼近目标函数和当前估计之间的差值。残差网络逼近的是H(x)-x,PINNs逼近的是F(u)-0,本质上都是“残差修正”的思路。
4.2 方程残差怎么变成损失函数
这一步是PINNs能不能work的关键。我拿一个最简单的泊松方程举个例子,在一维区间内:
u_xx = f(x)
假设边界条件是u(0)=0,u(1)=0。用PINNs求解时,你需要构建这样一个损失函数:
Loss = Loss_residual + Loss_bc
Loss_residual是在区域内随机采样一批点,把网络输出代入方程,计算u_hat_xx - f(x)的均方误差。Loss_bc则是在边界点上计算u_hat(x)和给定边界值之间的均方误差。优化的时候,Adam优化器会同时减小这两项,最终网络输出的函数在区域内满足方程、在边界上满足边界条件。
我在最初实现的时候犯过一个低级错误:直接在Loss里同时加了两项,但没注意两个Loss的量级差异。区域内的残差Loss一般比边界Loss大好几个数量级,结果训练出来的函数像个跷跷板,边界对齐了但区域内严重不满足方程。后来我把两项Loss都做了归一化处理,或者给边界Loss乘以一个较大的权重系数,才稳定住训练过程。
PINNs里的残差还有一个变体用法,就是渐进式残差修正。你可以先在一个较粗的网格上训练,让网络学到解的“大致形状”,然后不断加密采样点,用新的稠密点计算残差,继续训练。这相当于用残差在网络当前的解上打补丁,一次比一次精细。这种做法的收敛速度比一上来就用密集采样点快很多,而且不容易陷入局部最优。
4.3 残差块在PINNs中的实际用法
PINNs的骨干网络通常是全连接网络,但你同样可以把残差连接塞进去。常见的做法是在隐藏层之间加入残差连接,比如输入x经过第一层线性变换得到h1,然后第二层和第三层组成一个残差块,输出h3 = h2 + h1。
为什么PINNs也需要残差连接?因为PINNs的损失函数包含高阶导数项,微分算子会放大高频噪声。网络层数加深后,梯度在二阶甚至三阶导数传播中衰减得非常快,这时候残差连接就起到了“梯度保底”的作用,让深层的网络参数也能得到更新。我在求解一个带有边界层的对流扩散方程时,对比过带残差连接和不带残差连接的PINNs,收敛速度差了大概一倍,精度也明显提升。
另外,PINNs里还有一种“残差修正”不是指网络结构,而是指迭代求解的思路。你可以先用低阶数值方法得到一个近似解,然后用神经网络去拟合这个近似解与真解之间的残差,最终精度可以显著提升。这个思路和机器学习里的Boosting很像,本质上也是把大问题拆成一层一层的残差来逼近。遇到难解的方程时,这个技巧往往比单纯提高神经网络的容量管用得多。
5. 实操中的常见问题与排查技巧实录
5.1 残差网络训练不收敛的排查顺序
我自己的排查经验是有一套固定顺序的,按这个顺序查,大多数问题都能定位到。
先看Loss是不是nan或者特别大。如果一开始就是nan,大概率是学习率太高,或者数据没有归一化。残差网络对输入数据的尺度很敏感,输入像素值如果都在0到255之间不做归一化,第一批数据更新就会让权重爆炸。
再看训练集上的Loss是否下降。如果训练集Loss完全不动,说明梯度根本没有有效传播。这时候可以先去掉全部残差连接,只用普通网络跑一遍,如果普通网络能收敛,说明问题出在残差连接的实现上。我遇到过最乌龙的一次是shortcut里的1x1卷积写错了stride,导致特征图空间尺寸对不上,forward里add时广播出了奇怪的shape,程序没报错但结果全乱了。
最后检查验证集的Loss。训练集Loss收敛但验证集Loss不降,这通常是过拟合或者数据增强不够。残差网络因为参数多、拟合能力强,在中小数据集上过拟合是常态,不能指望它天生就有很强的泛化性。
5.2 残差振荡和梯度问题怎么处理
训练残差网络时,你可能遇到训练Loss在某个值附近来回震荡,既不下降也不发散。这种情况在加了残差连接后又用了较大学习率的模型里很常见。
根本原因是残差连接把梯度路径缩短了,浅层参数能收到很大的梯度信号,如果学习率偏大,权重更新幅度就过大,导致loss波动。处理办法很简单,调低学习率,或者采用warmup策略,前几个epoch用较小的学习率把网络参数稳定下来,再逐步加大。
有一种更隐蔽的振荡来自恒等映射路径上的“信息稀释”。如果F(x)的输出总是很小,接近零,那残差块的输出基本就等于x,网络退化成了一条直线,深度等于白加。这时候你可以检查一下最后一个ReLU之前的输出分布,如果大部分值都集中在零附近,说明需要调整初始化或者加入更大的权重衰减。
5.3 一些不容易注意到的细节坑
这里列几个我踩过的坑,不一定是大问题,但会浪费你不少排查时间。
第一个坑是卷积层的padding方式。残差网络推荐用ZeroPad2d或者padding=1,而不是reflect padding。reflect padding在边缘处的梯度不太稳定,容易让模型在图像边缘处出现条带伪影。尤其是做超分辨率和分割任务时,这个差异会在输出图上看到明显的边缘效应。
第二个坑是BatchNorm的batch size。残差网络对BN的batch size很敏感,如果batch size小于16,BN的统计量噪声会很大,导致shortcut路径上传递的梯度也变得不稳定。我通常会把BN层换成GroupNorm来绕开这个问题,虽然效果略差一点,但在小batch下稳定得多。
第三个坑是残差连接上的激活函数位置。有些改进版本会把ReLU放在残差块的末尾激活之后,也就是先相加再激活,这没问题。但也有版本把ReLU放在相加之前,也就是激活后再相加,这个顺序改了之后,恒等路径上的信息就不再是真正的恒等了,输出的范围会被ReLU截断到非负区间。如果你的任务需要网络输出负数(比如回归任务),这个顺序会导致系统性的偏置,要格外小心。
6. 残差思想扩展:还有哪些领域可以用这一招
6.1 从网络结构到优化器视角的残差
很多人在用残差网络时只把它当成一个“层数更深不会退化”的工具,但如果你换一个角度看,残差连接其实是一种优化器层面的技巧。它等价于在梯度下降的过程中加入了一个“惯性项”,让更新方向不完全取决于当前梯度,而是保留了一部分之前的状态。
这件事在强化学习里尤其明显。用残差结构做策略网络和价值网络时,我发现它比普通多层感知机更容易学到稳定的策略。原因也很简单,强化学习的回报信号稀疏且噪声大,如果网络直接去拟合最终回报,梯度方差很大;而残差结构让网络只去拟合“当前估计和真实回报的差值”,这个差值相对稳定,信号也干净得多。
类似的用法还出现在时序预测里。预测股票或者流量数据时,如果让网络直接预测未来值,难度很高;但如果让网络预测“未来值和当前值的差”,然后用当前值加上这个差值作为最终预测,效果会好很多。这也是残差思想的典型应用。
6.2 动态残差修正:在线学习的新思路
这几年还有一个让我觉得眼前一亮的方向,是把残差修正做成动态的。传统残差连接是在网络结构里写死的,但动态残差修正的思路是,在模型的预测结果上叠加一个专门的修正模块,这个模块的输入是当前预测和反馈之间的误差,输出是修正量。
这个做法在PINNs里体现得比较充分,也是热词里提到的“pinns残差修正”。具体来说,你先训练一个粗模型,得到预测值,然后拿真实值或高精度数值结果减去粗模型预测值,得到一组残差数据,再训练一个小网络去拟合这组残差。预测时,把粗模型的输出和残差网络的输出加起来,就是最终结果。
我做一个流体仿真加速项目时就用到了这个思路。传统CFD算一个工况要几分钟,我只算了少量工况作为训练数据,先训练一个粗模型,然后用残差网络拟合数值误差,最终把预测精度提升了一个数量级以上,而推理时间几乎没增加。这种方式很适合“拿不到足够标签但又有先验粗模型”的场景,算是残差思想最实用的一种变体。
7. 写在最后的一点个人体会
从ResNet提出到现在,残差已经不是一个新概念了。但每次我在新任务里重新用到它,还是会感慨这个设计的巧妙:它没有引入复杂的数学技巧,没有改变损失函数的形式,更没有增加多少参数,只是给数据流和梯度流开了一条“旁路”,就解决了深层网络训练的根本难题。
我个人的体会是,残差思想真正的价值不在于“网络需要加深”这个结论,而在于一种看待问题的角度。当你发现某个模型怎么也学不动的时候,与其强行增加模型的表达力,不如先看看它当前的输出和理想值之间差了什么,把这个差作为学习目标,往往会更轻松。这个思路适用于网络结构设计,适用于PINNs里的方程约束,也适用于很多工程系统里的误差修正。
如果你正在做模型训练调优,遇到“加层没提升、加深就退化”的情况,我建议你先别急着换模型架构,把残差连接加上之后观察一下梯度回传的变化,大概率会看到不一样的结果。这就是这个小小的“捷径”带来的机缘,希望你也能撞见属于自己的那一次。