1. 深层网络为什么越堆越差:一个被误读的退化现象
如果你手上有一份CIFAR-10或者ImageNet的数据,随手搭一个卷积网络,你可能会有一种很自然的直觉:层数越多,模型容量越大,拟合能力越强,效果应该越好。我在最早接触ResNet之前也是这么想的,结果做实验的时候被现实按在地上摩擦——把一个20层的普通卷积网络加到56层,训练误差反而更高,测试误差也跟着涨。这个现象有个专门的名字,叫网络退化(degradation),它是ResNet和残差块这套设计要解决的核心问题,也是我今天想跟你从头捋清楚的东西。残差块看起来只是加了一条"旁路",但这条线背后牵扯到优化难度、梯度传播、恒等映射的可学习性,值得拆开一层层看。这篇文章适合刚入门深度学习、写过几行PyTorch、但对着ResNet源码里那条identity感到困惑的人,也适合想重新理解"为什么陈年老结构至今还活跃在预训练模型里"的从业者。
1.1 退化问题不是过拟合,别把它们混为一谈
先把这个最容易搞混的点说清楚。很多人的第一反应是:56层训练误差比20层高,肯定是过拟合了吧?不是。过拟合的典型特征是训练误差低、测试误差高,模型在训练集上已经"背下来"了,只是泛化不行。而退化问题的表现是训练误差本身就比浅层网络高,测试误差同样更高,两头都差。
这就很反常识了:一个56层的网络,理论上完全可以退化成一个20层的网络——把后面36层全部学成恒等映射(identity mapping,也就是输入什么就输出什么,y = x),那么它的表现至少应该和20层网络一样好,绝不会更差。既然理论上有解,为什么优化器找不到?答案藏在"让一堆卷积层学出恒等映射"这件事本身有多难上。
你以为恒等映射很简单,其实对一堆堆叠的非线性层来说,要把输出精确地还原成输入,需要权重和偏置配合得恰到好处,非线性激活会在中间不断扭曲信息。优化器在这么大的参数空间里去凑一个"什么都不做"的解,反而比凑一个"做点有用的事"的解更费劲。这就是残差块最妙的切入点:既然让网络学"什么都不做"很难,那就把"什么都不做"变成默认选项。
1.2 恒等映射其实很难学:一个思想实验
我更喜欢用下面这个思想实验来理解。假设有一个两层的小网络,输入x,输出要等于x。如果是普通的堆叠,网络需要满足:
f(x) = W2 * relu(W1 * x + b1) + b2 = x为了让这个式子成立,W1、W2、b1、b2要联合调解,而且relu是非线性的,中间一旦把某些维度的信息压没了,后面无论怎么调都补不回来。这不是不可能,而是优化器要走很远的路才能到达那个解,训练初期梯度信号又弱,很容易卡在半路。
而残差块的做法是把目标改成学"残差"F(x) = f(x) - x。当我们要的映射就是恒等时,F(x) = 0,也就是让那几层卷积的权重都趋向于0就行——把参数压到0比凑出一个精确的线性变换容易太多了,这也是权重衰减(weight decay)天然喜欢的方向。所以残差块本质上是在给网络一个"捷径":你搞不定复杂映射的时候,先把旁路打开,至少别把原来的信息弄丢。
| 对比项 | 普通堆叠层 | 残差块 |
|---|---|---|
| 待学目标 | 完整映射H(x) | 残差F(x) = H(x) - x |
| 恒等映射时的目标 | 逼近H(x)=x,难 | F(x)=0,参数趋零,易 |
| 深层退化 | 明显,层数越深越差 | 显著缓解 |
| 梯度传播 | 逐层连乘,易消失 | 带加法旁路,含恒等项 |
2. 残差块的数学骨架:那条直接通路到底做了什么
把退化现象讲清楚之后,就可以动真格看残差块的公式了。你会看到无数教程写y = F(x) + x,但大多数人只记住了这个等式,没真正理解F(x)和x各自扮演什么角色。我当年也是死记硬背,直到自己动手把参数改坏、观察loss怎么炸,才把这条路径的意义坐实。这一节我们把公式拆到骨头里,顺便用一个小数值例子把抽象符号落地。
2.1y = F(x) + x的逐项拆解
残差块的正向计算就一行:
y = F(x, {Wi}) + x这里x是这一块的输入,y是输出,F(x, {Wi})是那几层卷积(通常两层或三层)学出来的残差函数。注意几个细节:
第一,x是原封不动加到输出上的,中间不经过任何参数变换(除了维度不匹配的特殊情况,后面讲)。这条路径在文献里叫shortcut connection(捷径连接)或identity shortcut(恒等捷径),也有人叫skip connection。
第二,+是逐元素相加,不是拼接(concat)。这意味着F(x)的输出和x的通道数、空间尺寸必须完全一致,否则加不起来。这一点在写代码时是头号大坑,我会专门讲。
第三,加法之后通常还有一个ReLU,所以严格的输出是y = relu(F(x) + x)。这个位置很重要,让整个块保持非线性。
F(x)学的是"输入和期望输出之间的差",这也是"残差"这个词的来源。你可以把它理解成一种"修正项":网络先假设输出和输入差不多(直接抄过去),卷积层只负责调整那些不一样的部分。这种"先抄后改"的思路,比"从零开始构造"要省力得多。
2.2 为什么学残差比学原映射更容易:从优化地形看
我们换个角度,从优化的"地形"来理解。假设期望映射是H(x),普通网络直接拟合H(x),残差块拟合F(x) = H(x) - x,最后通过加法还原。两者在表达能力上是等价的(你总能从F还原H),但在优化难度上天差地别。
打个比方。你要画一条从A到B的路径。普通网络的做法是:给你一张白纸,你从头画整条路径。残差块的做法是:先把A到B的直线直接印在纸上,你只需要画那条线偏离多少。显然第二种更省事——直线大体上已经对了,你只需微调弯折的地方。
学术一点说,残差块把参数空间的搜索起点"挪"到了恒等映射附近。当网络刚初始化时,F(x)接近0,整个块接近恒等映射,深层网络初始状态就类似一个浅层网络,训练从"至少不差"的地方开始爬坡。普通网络初始状态是一堆随机映射,深层堆叠下信息被反复扭曲,起点就很糟。这两句话几乎解释了为什么残差块能让1000+层的网络也能训起来。
2.3 一个手算的小数值例子
光说概念容易飘,来点具体的。假设某个输入x = [2.0, -1.0],而这一块真正想实现的映射是H(x) = [2.2, -0.9](输入稍微被调整了一下)。
如果是普通层,它要直接输出[2.2, -0.9],那权重得凑出这个精确结果。如果这一层后面还接了别的层、还有BN缩放,凑起来更麻烦。
换成残差块,F(x)只需要输出H(x) - x = [0.2, 0.1]。网络的小卷积核只要产生一个很小的增量就行,初始权重接近0时,F(x)本身就接近0,输出接近x = [2.0, -1.0],误差只有[0.2, 0.1]。随着训练推进,F(x)慢慢逼近[0.2, 0.1],全程都是在恒等映射附近做小幅度修正。
这个例子虽然简单,但它精确传达了残差学习的精神:网络不需要从无到有地构造表示,只需要在已有信息上加一个小的、可控的扰动。数值稳定的起点、小的更新幅度,正是深层网络能训得下去的关键。我建议你在纸上真的算一遍,比看十遍公式有用。
3. 反向传播视角:梯度是怎么被加法项救回来的
正向的直觉讲完了,但残差块真正封神的地方在反向传播。前向只解释了"学起来容易",反向才解释了"为什么能学很深"。这一节我们要看链式法则里那个多出来的1,它是整个ResNet能堆到上百层的数学根源。这部分稍微有点数学味,但我会尽量用人话讲透,你别跳。
3.1 链式法则里的那个1 + F'
设损失为L,我们要算损失对输入x的梯度。前向是:
y = F(x) + x对x求导,用链式法则:
∂y/∂x = ∂F(x)/∂x + 1别小看这个+1。它来自x这条捷径——因为x直接加到输出上,对它求导就是1。整个块的梯度里就强行混进了一个单位项。
把这个结论往深了推。假设有L个残差块串联,第l个块的输入是x_l,输出是x_{l+1} = F(x_l) + x_l。展开递推,最终第L层的输出可以写成:
x_L = x_l + Σ_{i=l}^{L-1} F(x_i)也就是说,深层特征等于浅层特征加上中间所有残差之和。对x_l求导:
∂x_L/∂x_l = 1 + ∂/∂x_l Σ F(x_i)这个1就是关键。它意味着梯度从深层往浅层传的时候,至少有一条"高速公路"可以原样通过,不会被中间任何一层的权重连乘给吃掉。
3.2 连乘变连加:梯度消失的根被挖了
回忆一下普通深层网络为什么梯度消失。反向传播时,第l层的梯度要乘上后面每一层的雅可比矩阵:
∂L/∂x_l = ∂L/∂x_L * Π_{i=l}^{L-1} ∂x_{i+1}/∂x_i这一长串是连乘。只要每层的导数平均小于1,乘个几十次就趋近于0,浅层收不到有效梯度,学不动;如果每层都略大于1,又会爆炸成无穷。这就是深层网络训练的经典两难。
残差块把每一层的∂x_{i+1}/∂x_i变成了I + ∂F/∂x_i(I是单位矩阵)。连乘之后,展开式里总有一个全I的项——也就是那个1,它不衰减、不爆炸,直接把梯度从最后一层无损地送回第一层。剩下的项虽然还是会衰减,但整体上梯度有了下限保证。
用一句话概括:普通网络靠"连乘"传梯度,残差网络靠"连加"传梯度。连乘脆弱,连加鲁棒。这就是为什么残差块能在1000层的网络上依然稳定训练。
3.3 但它不是万能药:几个边界条件
我不喜欢把任何东西吹成神器,残差块也一样,有几个边界条件必须提醒:
- 捷径分支本身如果带参数缩放,比如乘一个0到1之间的系数,那个
1就会被削弱,梯度还是会衰减。所以原始ResNet的恒等捷径不加任何缩放,这是刻意的设计。后来有些人试过加门控(gating),效果并不总好。 F(x)本身发散时,I + ∂F/∂x里的第二项还是会爆炸,所以BN、权重初始化这些稳定手段依然不能省。残差块解决的是"梯度传播路径"问题,不是"数值稳定"的全部问题。- 维度不匹配时走的是带参数的捷径(1x1卷积),那条路径上不再有纯粹的
1,梯度高速路的成色会打折。这是有代价的,只是实验表明代价小到可以接受。
提示:如果你在做消融实验,想看残差块到底贡献了多少,可以把恒等捷径换成卷积捷径,固定其他配置,对比收敛速度和最终精度。你大概率会看到收敛慢一截,这正好反证了那条无参数捷径的价值。
4. 工程实现里的三个关键决定:维度、BN位置、瓶颈
理论捋顺了,落到代码还有几个绕不开的工程选择,每一个都直接影响模型能不能跑通、跑多快。我见过太多人抄了ResNet代码却不知道自己在抄什么,一旦换个输入尺寸就报错。这一节我把三个最关键的实现决定讲清楚,让你改代码时心里有底。
4.1 尺寸不匹配:F(x)和x怎么对齐
残差块里y = F(x) + x要求两边形状完全一致。但现实中,卷积经常带stride=2做下采样,通道数也会翻倍,这时候F(x)的输出是[N, 2C, H/2, W/2],而输入x还是[N, C, H, W],加不起来。解决方案有三种:
| 方案 | 做法 | 特点 |
|---|---|---|
| 零填充 + 下采样 | 对x做步长2的采样,补零对齐通道 | 无参数,但丢信息、通道对不齐 |
| 投影捷径(1x1卷积) | 用1x1, stride=2的卷积把x变换到目标形状 | 带参数,ResNet系列默认方案 |
| 直接加(无下采样) | 只在通道数不变时用恒等捷径 | 最理想,但只在同阶段内可行 |
主流实现用的是第二种。以ResNet的经典结构为例,每个stage第一次遇到通道翻倍时,捷径从纯恒等切换成Conv2d(1x1, stride=2) + BN。剩下的块通道和尺寸都不变,走纯恒等。判断条件通常是:
if stride != 1 or in_planes != planes * expansion: # 需要投影捷径这里的expansion是块内通道的放大倍数,后面讲瓶颈结构时会用到。常见错误是只判断了stride忘了判断通道数,结果在通道翻倍的块上报维度错误,或者反过来多加了一层没必要的1x1卷积。这个条件一定要写全。
4.2 BN和ReLU的摆放顺序:原始版和被吐槽的预激活
原始ResNet(v1)的顺序是:Conv → BN → ReLU → Conv → BN → 加法 → ReLU。注意最后一个ReLU是加完之后才做的。这样一来,捷径传过来的x在相加之前没有被激活,信息是"干净"的。
后来有人提出**预激活(pre-activation)**版本(ResNet v2),把顺序改成:BN → ReLU → Conv → BN → ReLU → Conv → 加法,加法后不再激活。作者发现这种摆法在极深网络(比如1000层)上更容易训练,收敛更稳。
两种摆法各有拥趸。我的经验是:
- 做常规深度(18到152层),两种差别不大,原始版够用,而且和大多数预训练权重对齐。
- 要往几百上千层堆,预激活版更稳。
- 最关键的是:加载别人预训练权重时,BN和ReLU的位置必须和权重训练时一致,否则精度会莫名其妙掉。
4.3 瓶颈结构与参数预算:为什么50层反而比34层轻
ResNet族里有个很有意思的现象:ResNet-50比ResNet-34深得多,但参数量反而更少。秘密在瓶颈(bottleneck)结构。34层及以下用的是BasicBlock,两块3x3卷积;50层及以上用的是Bottleneck,结构是1x1降维 → 3x3 → 1x1升维。
举个具体数字。假设输入输出通道都是256:
- BasicBlock:两个
3x3, 256→256卷积,参数量约3*3*256*256*2 ≈ 118万。 - Bottleneck:
1x1: 256→64,3x3: 64→64,1x1: 64→256,参数量约1*1*256*64 + 3*3*64*64 + 1*1*64*256 ≈ 7.4万。
同样输出256通道,Bottleneck的参数只有BasicBlock的十六分之一左右。它先用1x1把通道压到1/4,在低维空间做3x3卷积,再用1x1升回来。这种"两头大中间小"的形状像沙漏,所以叫瓶颈。代价是表达能力的结构变了,但实测精度更好,参数量和计算量都可控,所以成为主流。
| 结构 | 组成 | 用在 | 参数量(近似) |
|---|---|---|---|
| BasicBlock | 3x3 → 3x3 | ResNet-18/34 | 较高 |
| Bottleneck | 1x1 → 3x3 → 1x1 | ResNet-50/101/152 | 较低 |
5. 用PyTorch把一个残差块写扎实
看再多的公式,不如自己敲一遍。这一节我给两个可以直接跑的块实现,BasicBlock和Bottleneck,附上关键注释和我在调试时踩过的坑。你把这俩块拼起来,就是完整的ResNet。
5.1 BasicBlock:18层和34层的积木
import torch import torch.nn as nn class BasicBlock(nn.Module): expansion = 1 # 输出通道不放大 def __init__(self, in_planes, planes, stride=1): super().__init__() # 第一层3x3,可能带stride做下采样 self.conv1 = nn.Conv2d(in_planes, planes, kernel_size=3, stride=stride, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(planes) # 第二层3x3,stride固定为1,尺寸不变 self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=1, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(planes) self.relu = nn.ReLU(inplace=True) # 捷径分支:默认是恒等(什么都不做) self.shortcut = nn.Sequential() # 只有当stride或通道数变化时,才用1x1卷积投影对齐 if stride != 1 or in_planes != planes * self.expansion: self.shortcut = nn.Sequential( nn.Conv2d(in_planes, planes * self.expansion, kernel_size=1, stride=stride, bias=False), nn.BatchNorm2d(planes * self.expansion), ) def forward(self, x): out = self.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) out = out + self.shortcut(x) # 逐元素相加,形状必须一致 return self.relu(out)几个细节值得说:卷积全部设bias=False,因为后面紧跟BN,BN自带偏置,再加卷积偏置是冗余参数,还容易让训练初期不稳定。ReLU(inplace=True)省显存,但在某些需要保留中间激活的场景要小心。加法前不加激活,加法后统一ReLU,这是原始版摆法。
5.2 Bottleneck:50层以上的沙漏结构
class Bottleneck(nn.Module): expansion = 4 # 输出通道放大4倍 def __init__(self, in_planes, planes, stride=1): super().__init__() # 1x1降维 self.conv1 = nn.Conv2d(in_planes, planes, kernel_size=1, bias=False) self.bn1 = nn.BatchNorm2d(planes) # 3x3在低维空间卷积 self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(planes) # 1x1升维,放大expansion倍 self.conv3 = nn.Conv2d(planes, planes * self.expansion, kernel_size=1, bias=False) self.bn3 = nn.BatchNorm2d(planes * self.expansion) self.relu = nn.ReLU(inplace=True) self.shortcut = nn.Sequential() if stride != 1 or in_planes != planes * self.expansion: self.shortcut = nn.Sequential( nn.Conv2d(in_planes, planes * self.expansion, kernel_size=1, stride=stride, bias=False), nn.BatchNorm2d(planes * self.expansion), ) def forward(self, x): out = self.relu(self.bn1(self.conv1(x))) out = self.relu(self.bn2(self.conv2(out))) out = self.bn3(self.conv3(out)) out = out + self.shortcut(x) return self.relu(out)注意stride放在中间那个3x3上,两端的1x1都是stride=1。这个位置很重要:如果把stride放在第一个1x1上,降维和下采样同时发生,会丢更多信息,精度通常会掉。这也是很多人自己魔改时最容易放错的地方。
5.3 我实测踩过的几个坑
第一个坑是通道计算。捷径分支里planes * self.expansion这个表达式,在BasicBlock里expansion=1,在Bottleneck里是4,写错了不会立刻报错,而是在拼到第三、四个stage时才炸,排查起来很烦。建议用一个完整的ResNet类把每个stage的通道数打印出来核对一遍。
第二个坑是**inplace=True和残差相加的冲突**。理论上加法前的激活是bn2的输出、加法是新的张量,不冲突。但如果你在别处复用了中间变量,inplace会就地改掉原值,导致梯度算错。稳妥起见,调试阶段可以先把inplace关掉。
第三个坑是BN的running stats和batch size。残差块里BN很多,如果训练时batch size太小(比如小于8),BN的统计量估计不准,预训练模型微调时尤其明显。我一般会把BN的momentum调小一点,或者在极端小批量下考虑冻结BN层。
注意:加载官方预训练权重时,别自己重写block后直接
load_state_dict。哪怕结构只差一个BN的位置,键名和形状就对不上了,会报missing/unexpected keys。先用model.named_parameters()看一遍键名再对齐。
6. 预训练模型里的残差块:加载权重后你该关注什么
现在几乎没人从零训练大模型了,ResNet的预训练权重到处都是,拿来做迁移学习是常态。但预训练模型里的残差块有一些"文档里不写、踩了才知道"的细节,尤其在你改结构、换输入、做微调的时候。这一节聊聊这些实战注意事项。
6.1 结构对齐:为什么换个输入尺寸就报错
预训练权重是和特定结构严格绑定的。常见的一个需求是把ResNet的输入从224x224改成更大或更小的尺寸。这里有几个连锁反应:
- 下采样次数不变,因为
stride是写死在结构里的,改输入尺寸只改变特征图的绝对大小,不改变下采样倍率。 - 全连接层的输入维度要跟着改。
ResNet最后有个全局平均池化(AdaptiveAvgPool2d(1)),所以理论上输入尺寸变化后,全连接层维度不变——这也是为什么ResNet能直接接受非224输入。 - 如果换掉了分类头的类别数,
fc层要重建,其他层权重照样能加载,用strict=False加载时只会缺fc的键。
我一般会这么处理:先load_state_dict(weights, strict=False),然后打印缺了哪些键、多了哪些键,确认缺的只有分类头,多的只是新分类头,就可以放心微调。
6.2 微调时残差块该冻还是该放
迁移学习的经典问题是:到底冻结多少层。我的经验是分场景:
- 数据量小、和目标域高度相似:冻结到只剩最后几个stage和分类头。残差块前面的层学的是通用边缘、纹理特征,冻结能防过拟合,训练也快。
- 数据量中等、域差异大:全放开,但用较小的学习率(比如
1e-4),让残差块的权重缓慢适应新域。 - 数据量大:直接全放开,甚至可以把预训练权重当初始化,从头微调。
有个细节容易被忽略:冻结BN层。如果你只冻结了卷积权重但没冻结BN,BN的running_mean和running_var还在更新,等于结构又变了。正确做法是把BN设成eval()模式,或者显式冻结它的参数和统计量更新。
6.3 三个常见误解澄清
误解一:残差块越多越好。不是。超过某个深度后收益递减,甚至因为参数过多在小数据集上过拟合。选ResNet-18还是ResNet-50,看数据量和算力,不要盲目堆深。
误解二:捷径上的1x1卷积可以随便加。加在通道数不变、尺寸不变的块上,纯属浪费参数,还会削弱梯度高速路。只在必要时加。
误解三:残差块解决了所有深层训练问题。它主要解决梯度传播和退化,但不能替代好的初始化、归一化和学习率调度。一个没做初始化的残差网络照样训崩。我见过不少人把残差块当万能钥匙,忽略了训练技巧,最后怪结构不行,这就冤枉它了。
关于预训练模型还有一个隐含价值:它内部那些残差块学到的是层层递进的特征——浅层是边缘、颜色,中层是纹理、局部形状,深层是语义部件。你在做特征提取、风格迁移、目标检测骨干网络时,直接复用这些块,往往比自己训一个轻量模型效果更好。这也是残差块从2015年火到今天、依然是各类视觉任务默认骨干的根本原因。
最后分享一个我自己排查残差网络的小习惯:训练前先用一个假输入torch.randn(2, 3, 224, 224)跑一遍前向,把每个stage的输出形状print出来。你会发现很多维度错误在这一步就暴露了,比等到训练跑了几轮才报错省事得多。残差块的坑大多集中在形状对齐和捷径分支判断上,把这两点盯死,剩下的就是熟练度问题了。