残差网络深入解析:从原理到PyTorch实现与训练调优
先开门见山说个现象很多做深度学习的同学模型一开始训得好好的loss降到某个程度之后怎么调都下不去甚至把网络层数加得更深效果反而变差了。这时候十有八九不是代码写错了而是你踩到了网络结构的“退化问题”。而解决这个问题最经典的手段就是今天要聊的残差。残差或者说得更直白一点残差网络ResNet是深度学习里值得反复咀嚼的基础模块。它不光是2015年ImageNet竞赛的冠军结构更是如今绝大多数CNN骨干网络的“地基”包括后面的DenseNet、ResNeXt、EfficientNet甚至Transformer里的Add Norm多少都带点残差思想的影子。这篇内容不打算堆公式我会从“为什么需要残差”讲起再拆到残差块的工程实现细节最后直接给出一份可以照着敲的PyTorch实现和训练调试经验。无论你是刚入门深度学习、准备期末复习还是已经在跑自己的项目遇到loss死活降不动的情况这篇文章都值得读完。1. 残差到底解决了什么问题1.1 加深网络不是万能药先讲退化问题在残差网络出现之前主流观点是“网络越深表达能力越强”所以大家拼命把VGG往深了堆。VGG16、VGG19再往上叠几层问题就来了训练集上的准确率不仅没提高反而明显下降。这就很奇怪了按道理说深网络应该包含浅网络的全部能力至少不应该更差吧这个现象不是过拟合。过拟合是训练集好、测试集差而这里的退化问题是训练集本身的表现就变差了。用大白话说一个20层的网络理论上它的前10层完全可以复制10层网络的参数后10层做成恒等映射那效果至少不应该比10层差。但实际训练中让神经网络通过梯度下降去学习“恒等映射”这件事本身就非常困难。为什么会困难核心原因在于非线性激活函数。每一层卷积、ReLU、BN叠下来输入分布的几何结构一再被扭曲。你想让网络在某几层输出等于输入哪怕是零矩阵都行但SGD随机梯度下降在这么高维的非凸优化空间里很难精确定位到这个“什么都不做”的解。结果就是网络在试图学习一个更复杂的映射时反而把原本已经学好的特征给破坏了。我当时第一次读到这个观点时脑子里的画面是你让一个人去做一张精确的复印件他越努力加细节反而越容易把原件涂花。而残差网络做的事就是告诉网络“你不需要学习完整的复印件你只需要告诉我复印件和原件差在哪儿”。这个思路转换一下子把学习难度降了下来。1.2 残差块的核心思想让网络学习增量残差网络的基本单元很简单。它不直接让某个卷积层去拟合期望的映射 H(x)而是让卷积层去拟合一个残差映射 F(x) H(x) - x然后把原始输入 x 再通过一条“捷径”shortcut传到后面两者相加得到最终的输出。用公式写就是y F(x, {W_i}) x这里的 F(x) 可以是两层卷积、BN、ReLU的组合x 是这一块的输入。关键就在于那个“ x”。如果 H(x) 等于恒等映射最理想那网络只需要让 F(x) 学习到0即可。哪怕网络初始化成比较小的权重一开始输出也接近0那叠加 x 之后至少不会比浅层特征差太多。这种“默认不做事也能保住下限”的设计简直太聪明了。我自己理解残差还有一个角度它其实是在做“增量学习”。网络在最开始几轮训练时主干特征被快速提取出来后续层只需要在这个基础上做修正。修正的幅度往往不会太大所以梯度的量级也比较可控不容易出现梯度爆炸或梯度消失。1.3 为什么恒等映射能改善反向传播从反向传播的角度看残差连接等于给梯度开了一条“高速公路”。假设第 l 层的输出是 x_{l1} x_l F(x_l, W_l)那么反向传播时对 x_l 的梯度可以写成∂Loss / ∂x_l ∂Loss / ∂x_{l1} · (1 ∂F / ∂x_l)括号里的这个 1 非常关键。它意味着哪怕后面若干层的 ∂F/∂x_l 很小梯度也不会完全消失因为总有一条路径让梯度原封不动地传回输入。这个机制比单纯靠Batch Normalization控制数值分布要直接得多也解释了我们后面要聊到的“为什么残差网络能堆到上百层”。另外有个细节容易忽略上式里“1”的存在是建立在“捷径上没有额外操作”的前提下的。如果你在捷径上插了一个卷积层或者一个dropout、一个BN这个“1”就变成了别的操作对输入的导数梯度传导的优势会打折扣。所以工程上只要维度能对上捷径上尽量不做多余操作只有维度对不上时才动用 1x1 卷积这类降维工具。2. 残差块的工程实现与细节解析2.1 残差块的两种经典设计残差网络最有意思的设计细节之一是它针对不同深度采用了不同的残差块。我们先说BasicBlock这是ResNet-18和ResNet-34用的。它由两个3x3卷积组成每个卷积后面接BatchNorm和ReLU输入 x 经过这两层卷积后得到 F(x)再与捷径上的 x 相加最后过一个ReLU。整个结构小巧适合层数不太深的网络。再说Bottleneck这是ResNet-50及更深网络的标配。它做三件事先用1x1卷积把通道数压下来比如原来是256个通道先压到64再经过一个3x3卷积提取空间特征最后用1x1卷积把通道数恢复到256。这样一来3x3卷积的计算量被大幅削减。用公式看就是通道数从256到64再到256的过程中间那个瓶颈状的窄通道就是“bottleneck”名称的来源。Bottleneck最大的价值是省计算量。如果ResNet-50换成BasicBlock参数量和FLOPs会暴涨训练速度也慢得多。所以当你决定自己设计网络时记住这个原则浅层网络用BasicBlock深层网络用Bottleneck。2.2 shortcut与维度匹配的两种办法残差块里最容易出bug的环节就是输入输出维度对不上。F(x) 的卷积层可能改变空间尺寸或者通道数这时候 x 不能直接加进去需要做维度匹配。主流做法有两种。第一种zero padding。在捷径上给 x 补0让通道数填充到目标维度。这种方法不需要额外参数但补0会引入一部分“无效信息”实验表现通常略差一些。所以它更多地出现在早期论文的消融实验里实际项目中用得少。第二种1x1卷积投影。在捷径上放一个1x1卷积步长设为2如果需要减半空间尺寸把通道数也调整到目标值。这个1x1卷积虽然增加了少量参数但能学到的特征组合更灵活。论文里的ResNet实现在每一个“维度变化”发生的地方都会用这种 projection shortcut 来匹配尺寸。还有一个容易被坑的点是下采样。当残差块里的主卷积设置 stride2 时3x3卷积会把特征图的宽高减半这时捷径上的1x1卷积也必须设置 stride2否则两张特征图的空间尺寸对不上相加时直接形状不匹配。这个错误在初学者代码里出现率极高建议写代码时先打印一下每一层的输出尺寸再写下一个模块。2.3 BN、激活函数与残差相加的先后顺序残差块的内部结构不同效果可以差出很多。最容易混淆的就是ReLU加在哪里的问题。Originally也就是2015年那版论文的BasicBlock结构是“Conv - BN - ReLU - Conv - BN”然后与 shortcut 相加最后再过一个 ReLU。这个结构一直用得很广实现起来也简单。但后来PreAct ResNet2016年论文改成了“BN - ReLU - Conv - BN - ReLU - Conv”把激活放到卷积之前残差相加之后不再接ReLU。实验结果证明这种预激活结构在某些任务上收敛更快正则化效果也稍微好一点。为什么相加之后再过一个ReLU不好因为相加后输出经过ReLU会强制把负值截断为0但残差学习本身允许输出有负值。如果你把残差相加后的结果直接传给下一层信息保留更完整。所以后来很多新结构的写法都倾向于“相加之后直接输出不再额外加ReLU”。我做实验时比较喜欢的组合是每个卷积后面接BN卷积输出过ReLU激活最后一个卷积的输出先做BN然后和捷径相加不加ReLU。如果一定要过ReLU我会在相加之后加一个但要注意它的输入分布可能和纯粹的卷积输出很不一样调参时得多留个心眼。另外说一句BN的实现细节残差块里的BN一般放在卷积之后、激活之前和常规卷积块的顺序一致。它配合残差让数值稳定性的提升非常明显。有个经验是训练很深层网络时BN的 momentum 可以稍微调低一点比如 0.05 或 0.1让batch统计量变化更平滑但这属于调参偏好后面再展开讲。3. 从零搭建一个残差网络PyTorch实战3.1 基础卷积块与残差块的代码实现理论说再多不如上手写一遍。这里我用PyTorch写一个可以放心用的残差网络实现。先说清楚环境PyTorch版本2.0以上Python3.8CUDA不管用不用都行CPU也能跑通小网络验证逻辑。我们先把BasicBlock实现出来import torch import torch.nn as nn class BasicBlock(nn.Module): expansion 1 def __init__(self, in_channels, out_channels, stride1, downsampleNone): super().__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, stride1, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_channels) self.downsample downsample def forward(self, x): identity x out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) if self.downsample is not None: identity self.downsample(x) out identity out self.relu(out) return out这里有个细节两个卷积层都设了 biasFalse。因为后面紧跟着 BatchNormBN会做平移操作卷积的偏置参数是多余的留着反而多占显存、增加过拟合风险。这个习惯建议从第一天写网络就养成。然后是Bottleneck结构是1x1压通道、3x3提空间、1x1恢复通道class Bottleneck(nn.Module): expansion 4 def __init__(self, in_channels, out_channels, stride1, downsampleNone): super().__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size1, stride1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_channels) self.conv3 nn.Conv2d(out_channels, out_channels * self.expansion, kernel_size1, stride1, biasFalse) self.bn3 nn.BatchNorm2d(out_channels * self.expansion) self.relu nn.ReLU(inplaceTrue) self.downsample downsample def forward(self, x): identity x out self.relu(self.bn1(self.conv1(x))) out self.relu(self.bn2(self.conv2(out))) out self.bn3(self.conv3(out)) if self.downsample is not None: identity self.downsample(x) out identity out self.relu(out) return out注意Bottleneck里第三层输出通道是 out_channels * expansion对Bottleneck来说 expansion4。比如第一个残差块输出通道是64那么最终输出就是256。这个 “expansion” 是ResNet系列里让人容易懵的点看代码时一定要留意。3.2 搭建ResNet-18/34主干网络有了残差块搭ResNet主体是顺手的事。核心流程先过一层7x7卷积把输入降采样再接一个3x3的max pooling然后按层堆叠不同数量的残差块最后做一个全局平均池化和全连接分类。ResNet-18 的层配置是 [2, 2, 2, 2]ResNet-34 是 [3, 4, 6, 3]ResNet-50 则是 [3, 4, 6, 3] 但每个块换成Bottleneck。我把通用框架写出来想换哪个网络改一下配置即可class ResNet(nn.Module): def __init__(self, block, layers, num_classes1000): super().__init__() self.in_channels 64 self.conv1 nn.Conv2d(3, 64, kernel_size7, stride2, padding3, biasFalse) self.bn1 nn.BatchNorm2d(64) self.relu nn.ReLU(inplaceTrue) self.maxpool nn.MaxPool2d(kernel_size3, stride2, padding1) self.layer1 self._make_layer(block, 64, layers[0]) self.layer2 self._make_layer(block, 128, layers[1], stride2) self.layer3 self._make_layer(block, 256, layers[2], stride2) self.layer4 self._make_layer(block, 512, layers[3], stride2) self.avgpool nn.AdaptiveAvgPool2d((1, 1)) self.fc nn.Linear(512 * block.expansion, num_classes) def _make_layer(self, block, out_channels, blocks, stride1): downsample None if stride ! 1 or self.in_channels ! out_channels * block.expansion: downsample nn.Sequential( nn.Conv2d(self.in_channels, out_channels * block.expansion, kernel_size1, stridestride, biasFalse), nn.BatchNorm2d(out_channels * block.expansion), ) layers [] layers.append(block(self.in_channels, out_channels, stride, downsample)) self.in_channels out_channels * block.expansion for _ in range(1, blocks): layers.append(block(self.in_channels, out_channels)) return nn.Sequential(*layers) def forward(self, x): x self.maxpool(self.relu(self.bn1(self.conv1(x)))) x self.layer1(x) x self.layer2(x) x self.layer3(x) x self.layer4(x) x self.avgpool(x) x torch.flatten(x, 1) x self.fc(x) return x def resnet18(num_classes10): return ResNet(BasicBlock, [2, 2, 2, 2], num_classesnum_classes) def resnet34(num_classes10): return ResNet(BasicBlock, [3, 4, 6, 3], num_classesnum_classes) def resnet50(num_classes10): return ResNet(Bottleneck, [3, 4, 6, 3], num_classesnum_classes)搭好之后强烈建议做一件事用一张随机输入过一遍前向把每一层的输出shape打出来确认没有维度问题。x torch.randn(2, 3, 224, 224) model resnet18(num_classes10) out model(x) print(out.shape) # 期望 (2, 10)出现 shape 不匹配时不要直接跳过把 _make_layer 里的 downSample 条件重看一遍大多数问题都出在 stride 和通道数判断上。3.3 关键参数与训练配置建议网络搭好后训练策略直接决定你能不能复现“残差网络又快又稳”的体验。这里给出一组我用下来比较顺手的配置适用于CIFAR-10、ImageNet这类图像分类任务。优化器SGD with momentum 就够了。Adam在超深网络上偶尔会出现收敛不稳定尤其是batch size偏小的时候。momen tum设0.9weight decay设1e-4到5e-4之间。Weight decay太大会把BN的gamma也压得厉害反而降低模型表达能力。学习率策略初始0.1起步batch size为256时比较合适。如果显存小batch size减半学习率也建议减到0.05左右。训练时用余弦退火或者ReduceLROnPlateau我个人习惯用多步长衰减在epoch 30/60/80分别把学习率除以10对CIFAR-10这类小数据集非常有效。Batch Size建议不要低于64。残差网络对BatchNorm里的batch统计量比较敏感batch太小时BN估计的均值和方差噪音大训练前期容易震荡。如果你的GPU只够跑16或32的batch要么用更小的模型要么把BN的momentum调低比如0.05。数据增强随机裁剪加随机水平翻转对分类任务是标配。CIFAR-10上我会pad到36再随机crop到32基本能把项目准确率往上抬一个点。这里不涉及太玄学的东西就是“多样本喂进去模型见过更多形态泛化必然好”。4. 训练残差网络的踩坑记录与排查技巧4.1 loss不下降、震荡、梯度爆炸的排查思路残差网络虽然比纯深网络稳得多但把它用到新任务目标检测、分割、自监督时还是会遇到各种训练问题。我把这几年高频踩的坑整理成一个速查表。现象常见原因解决方向loss一直不降多半卡在1.x学习率太大或BN没收敛先跑一个很小的模型验证数据流降低学习率到1e-3loss震荡剧烈曲线像锯齿学习率过大、batch size太小学习率降为原来的1/10或增大batch size验证集准确率上不去训练集也不低网络层数太深、没有有效正则加入dropout或更强的weight decay训练集过拟合早测试集崩dropout缺失、数据增强不足加入RandomAugment、CutMix等增强loss变成NaN初始化过大、BN结构写错检查shortcut路径缩小初始化范围训练到后期准确率回落学习率衰减策略过于激进改用余弦退火避免学习率断崖下降一个我屡试不爽的排查方法先把训练集裡的一个batch拿出来反复拟合它观察到这个batch的loss能稳定降到很小比如1e-2以下那说明代码逻辑没有大问题如果连单batch都过拟合不了那问题大概率在模型结构和数据加载流程。4.2 学习率、初始化、BN的调参经验很多同学一上来就用ImageNet预训练权重然后微调自己的小数据集这个方法在迁移学习场景很靠谱。但如果你需要从头训练一个残差网络初始化就是第一个坑。PyTorch的nn.Conv2d默认初始化还算合理但我想强调一个经验当网络特别深时把残差块最后一个BN层初始化成零gamma效果会让训练稳定不少。这个技巧最初来自“Bag of Tricks for Image Classification”这类论文后来在目标检测里也被大量使用。原因很简单初始化时让残差块一开始输出为0整个模块等价于恒等映射网络初期的表达压力大幅降低梯度也不会因为深层叠加而爆炸。另一个容易忽略的点是学习率适配数据集大小。ImageNet上0.1的初始学习率很好但换到CIFAR-10这种只有5万张小图的数据集0.1经常在一开始就把loss冲飞。我一般会先跑10个step打印初始loss如果loss在5以上多半就是学习率大了直接降到0.01再试。BN的调参也要看训练习惯。固定batch size训练时默认momentum0.1问题不大。但如果你用Gradient Accumulation模拟较大batch或者跨卡训练时batch被切得很小BN统计量容易乱这时候有两种解法一是把momentum压低到0.05二是干脆用Ghost BN、Sync BN这类专门处理小batch的技术。比较前沿的做法还有把BN换成分组归一化GroupNorm但那是另一套体系了不是在讲残差先不展开。4.3 用残差网络的思路诊断你的模型有个很有意思的用法残差网络的“恒等路径优先”思想不仅能用来搭网络还能用来诊断模型是不是真的学到了东西。方法是自己造一个小实验把输入 x 直接接到输出如果模型的预测效果和纯一些的线性模型差不多说明网络可能根本没有好好利用深层的残差特征只是在做恒等映射而已。反过来如果你在某个残差块之后把 shortcut 断开或者把某一层的残差输出强制置零观察指标变化就能定位出哪一层对当前任务贡献最大。我有一次做一个小物体检测项目发现模型训练loss一直徘徊不掉用这个方法逐块排查最后定位到第3个stage的残差块输出异常原来是下采样时 stride 设置成了2而 shortcut 没有跟着改导致后续层拿到的特征图尺寸不对训练时模型干脆学废了。从这个角度说理解残差的本质不只是为了应付面试或期末试题也是实际排障的利器。5. 残差思想在更多结构里的延伸5.1 从ResNet到TransformerAdd Norm残差思想最成功的跨领域传播就是Transformer里的 Add Norm。Transformer的每一层子层比如多头注意力、前馈网络都是先算出一个输出再和输入相加并做层归一化本质上就是残差连接。所以你会发现很多深度学习框架的API设计中Transformer Block 会自带残差路径不需要每一步都手动加。但如果你自己实现Transformer千万别漏掉这条捷径否则深层模型训几天loss都降不下来。我见过太多新手复现Attention时把残差加在Softmax之前结构看起来像梯度行为却一塌糊涂。5.2 残差学习在扩散模型、Neural ODE等方向的变化扩散模型里模型学习的目标是预测加入的噪声或者直接预测原始图像与噪声图像之间的差值这本身就是一种“预测残差”的思路。Neural ODE则更进一步把残差连接看成欧拉法的逐步积分每一层学的是当前状态的变化率而不是完整状态。这些思想虽然比ResNet复杂得多但底层逻辑一脉相承。如果你对“残差修正”这个概念有过接触也会发现很多领域会用残差作为修正项加到粗预测结果上。这不只是深度学习专用技巧本质上它是在承认“模型不可能一次预测完美但预测误差往往比预测本身更容易学”这跟信号处理里“先拟合趋势再拟合残差”的做法是通的。5.3 什么时候别用残差残差不是万能药。如果任务本身非常浅层比如一个只有两三层的MLP残差纯属多余如果输入输出维度不匹配而你又不想引入额外参数那残差反而会增加代码复杂度。另一个情况是推理阶段的显存优化需求强于训练精度残差路径会多一份显存占用一些部署场景可能会选择把残差折叠掉或干脆去掉。不过绝大多数情况下如果你的网络达到8层以上我建议默认加残差。它带来的收益远大于那一点点参数开销这句话基本可以成立到我目前做过的所有视觉和NLP任务上。最后分享一个个人经验初学残差时别急着去啃各种变形网络先把BasicBlock的代码敲三遍把shortcut为什么这样设计想明白再去看Bottleneck、PreAct、DenseNet这些变体你会觉得它们全都顺理成章。踩过几次坑之后我越来越相信深度学习里真正值钱的东西不在模型有多花哨而在于你对你用的那个基本模块理解得够不够透。残差这个基石值得多花几天去磨。