AI残差网络:这次怎么落地的
先用最简单的例子说一下梯度消失是怎么发生的。
为什么深度网络会卡住
先用最简单的例子说一下梯度消失是怎么发生的。假设你有一个很普通的网络,每层就是一个线性变换加激活函数:
# 简单的深度网络示例(PyTorch)
import torch
import torch.nn as nn
class DeepNet(nn.Module):
def __init__(self, input_dim=784, hidden_dim=256, output_dim=10, num_layers=10):
super().__init__()
layers = []
layers.append(nn.Linear(input_dim, hidden_dim))
layers.append(nn.ReLU())
for _ in range(num_layers - 2):
layers.append(nn.Linear(hidden_dim, hidden_dim))
layers.append(nn.ReLU())
layers.append(nn.Linear(hidden_dim, output_dim))
self.net = nn.Sequential(*layers)
def forward(self, x):
return self.net(x)
这看起来很合理,但训练的时候你就会发现问题。假设隐藏层用的是 Sigmoid 而不是 ReLU,那梯度在反向传播时每次都要乘一个 Sigmoid 导数,这个导数最大也就 0.25,十层之后基本就没了。换成 ReLU 会好一些,因为正区间的导数是 1,但深层网络仍然会遇到"梯度退化":不同层的梯度差异巨大,深层神经元要么"死掉",要么训练极慢。
2015年的论文《Deep Residual Learning for Image Recognition》里给了一个直观的对比:普通 18 层网络和 34 层网络在训练集上的误差差不多,但测试集上 34 层反而更差。这说明深层网络不是学不动,而是难以找到一条好的优化路径。
ResNet 的核心想法:让信息直接跳跃
ResNet 的核心创新看起来简单到有点让人怀疑:在每个普通的卷积块旁边加一条"跳跃连接",让输入直接绕过卷积层,与输出相加。
# 基础残差块(Basic Block)
class BasicBlock(nn.Module):
def __init__(self, in_channels, out_channels, stride=1):
super().__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.relu = nn.ReLU(inplace=True)
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3,
stride=1, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(out_channels)
# 处理维度不匹配的情况
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):
out = self.conv1(x)
out = self.bn1(out)
out = self.relu(out)
out = self.conv2(out)
out = self.bn2(out)
# 残差连接
out += self.shortcut(x)
out = self.relu(out)
return out
这行代码里藏着两个关键设计:
第一个是 out += self.shortcut(x),也就是"残差连接"。它的直观理解是:与其让网络直接学习目标映射 H(x),不如让它学习残差 F(x) = H(x) - x,然后最终输出是 F(x) + x。如果 H(x) 本身就很简单(比如接近恒等映射),那网络只需要把残差学成零就行,比完全从头学容易太多。
第二个是 self.shortcut 处理维度不匹配。当 stride != 1(下采样)或者通道数变化时,不能直接相加,需要一个 1x1 卷积来调整形状。很多人第一眼看残差代码会觉得"为什么要有这个分支",因为它在反向传播时多了参数,实际上这是为了让梯度能顺利流下来。
用一个简单的图理解一下信息流动方向:
箭头越少的地方梯度就越容易消失。普通网络的梯度要一层层传下来,而残差网络可以通过捷径直接跳跃。
为什么跳跃连接解决了梯度消失
有个不太直观的事实:即使残差块里的卷积层训练得不好,跳跃连接也能保证"至少不比输入差"。假设残差 F(x) 学得很烂,那输出至少保留了原始输入 x,整个网络退化为更浅的网络。这种"安全网"性质让优化器敢于往更深的网络探索,因为最坏情况也不会比浅层网络差。
反向传播时,梯度可以通过捷径直接流过。用链式法则算一下:
$$\frac{\partial L}{\partial x} = \frac{\partial L}{\partial (F(x)+x)} \cdot \left(1 + \frac{\partial F(x)}{\partial x}\right)$$
注意那个 +1,这就是关键。普通网络的梯度是乘积相乘,越乘越小;残差网络多了一个恒等项 1,梯度至少能完整传下来,不会因为深度而消失。这不是理论推导,是训练时的真实感受:同样深度的网络,加了残差后收敛速度明显快,而且能稳定训练到更深的层数。
实践中的 ResNet 构建
完整 ResNet 的结构是把这些基础块叠起来,在一些地方下采样调整通道数:
class ResNet(nn.Module):
def __init__(self, block, num_blocks, num_classes=1000):
super().__init__()
self.in_channels = 64
# 初始卷积层
self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)
self.bn1 = nn.BatchNorm2d(64)
self.relu = nn.ReLU(inplace=True)
self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
# 四个残差阶段
self.layer1 = self._make_layer(block, 64, num_blocks[0], stride=1)
self.layer2 = self._make_layer(block, 128, num_blocks[1], stride=2)
self.layer3 = self._make_layer(block, 256, num_blocks[2], stride=2)
self.layer4 = self._make_layer(block, 512, num_blocks[3], stride=2)
# 分类头
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
self.fc = nn.Linear(512 * block.expansion, num_classes)
def _make_layer(self, block, out_channels, num_blocks, stride):
strides = [stride] + [1] * (num_blocks - 1)
layers = []
for stride in strides:
layers.append(block(self.in_channels, out_channels, stride))
self.in_channels = out_channels * block.expansion
return nn.Sequential(*layers)
def forward(self, self, x):
x = self.conv1(x)
x = self.bn1(x)
x = self.relu(x)
x = self.maxpool(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_classes=1000):
return ResNet(BasicBlock, [2, 2, 2, 2], num_classes)
def resnet34(num_classes=1000):
return ResNet(BasicBlock, [3, 4, 6, 3], num_classes)
这段代码跑起来很快,但实际用的时候有几个坑要小心。
踩过的坑和解决方案
第一个坑是通道数对不上。当初第一次手写 ResNet,在做 CIFAR-10 数据集时,输入是 3 通道 32x32,但照搬论文的 ImageNet 配置(初始 7x7 卷积、输出 64 通道)会在下采样后特征图太小,直接报错。解决办法是针对小数据集调整初始层:
# 适合 CIFAR-10 的轻量初始层
self.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(64)
self.relu = nn.ReLU(inplace=True)
# 去掉 maxpool,因为输入太小
第二个坑是预训练权重加载。很多人从 torchvision.models 加载预训练 ResNet,然后冻结部分层做迁移学习,结果发现准确率比随机初始化还差。原因是预训练模型是用 ImageNet 的标准化参数训练的(均值 [0.485, 0.456, 0.406],标准差 [0.229, 0.224, 0.225]),而迁移到其他数据集时忘了匹配这个预处理。
# 正确的迁移学习方式
import torchvision.models as models
# 加载预训练模型
pretrained = models.resnet18(pretrained=True)
# 替换最后的分类层
pretrained.fc = nn.Linear(512, num_classes)
# 数据预处理要匹配 ImageNet 的参数
transform = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)
])
第三个坑是瓶颈块(Bottleneck Block)用错。深层 ResNet(50层以上)用的是瓶颈块,结构是 1x1 -> 3x3 -> 1x1,目的是减少计算量。但如果你直接把 Basic Block 换成 Bottleneck 而不调整层数配置,通道数会爆炸,显存直接爆掉:
class Bottleneck(nn.Module):
expansion = 4 # 通道数膨胀倍数
def __init__(self, in_channels, out_channels, stride=1):
super().__init__()
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)
self.bn1 = nn.BatchNorm2d(out_channels)
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3,
stride=stride, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(out_channels)
self.conv3 = nn.Conv2d(out_channels, out_channels * self.expansion,
kernel_size=1, bias=False)
self.bn3 = nn.BatchNorm2d(out_channels * self.expansion)
self.relu = nn.ReLU(inplace=True)
self.shortcut = nn.Sequential()
if stride != 1 or in_channels != out_channels * self.expansion:
self.shortcut = nn.Sequential(
nn.Conv2d(in_channels, out_channels * self.expansion,
kernel_size=1, stride=stride, bias=False),
nn.BatchNorm2d(out_channels * self.expansion)
)
注意 expansion = 4,这意味着最后输出通道是输入的 4 倍。ResNet-50 的实际通道变化是 64 -> 256 -> 512 -> 1024 -> 2048,而不是表面上的 64 -> 128 -> 256 -> 512。第一次手动搭建 ResNet-50 时没注意到这个细节,算算显存需求才发现超了一倍。
第四个坑是初始化。残差网络对初始化还算宽容,但如果你用全零初始化,第一轮训练就会陷入困境。因为残差连接本质上是一个恒等映射,如果残差分支的参数全零,网络在训练初期很难有效学习。推荐做法是使用 Kaiming 初始化:
def initialize_weights(self):
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
elif isinstance(m, nn.BatchNorm2d):
nn.init.constant_(m.weight, 1)
nn.init.constant_(m.bias, 0)
ResNet 之后的演进
ResNet 之后残差思想被各种变体吸收,比如 Pre-activation ResNet(把 BN 和 ReLU 放在卷积前)、ResNeXt(用分组卷积增强容量)、DenseNet(直接把所有层连起来)。核心思路没变:让信息流动更顺畅。
实践中我常用的是 EfficientNet 和 Vision Transformer,但 ResNet 仍然是可靠的基线。它架构简单、实现成熟、文档丰富,对于大多数计算机视觉任务来说,先跑一个 ResNet-34 或 ResNet-50 看效果,再考虑要不要上更复杂的模型,这个流程很稳。
写在最后
残差网络的价值不在于它发明了什么复杂结构,而在于它用极简的方式回答了一个困扰多年的问题:深度网络为什么难训练。答案不是"深度本身的问题",而是"信息流动不畅"。这个观察后来影响了 Transformer、扩散模型等架构设计,甚至超出了计算机视觉的范畴。
做项目时,有时候遇到一个看似无法优化的瓶颈,不妨想想有没有"跳跃连接"的方案:绕过复杂路径,让信息直接传到需要的地方。不一定每次都有效,但至少值得一试。
可用性说明:本文发布于 2021 年 5 月,距今已超过五年。文中涉及的软件版本、接口、下载地址、命令参数和操作界面可能已经发生变化,部分方案在当前环境下可能失效。请结合官方最新文档核对后再操作,生产环境使用前务必先行验证。
版权声明: 本文首发于 指尖魔法屋-AI残差网络:这次怎么落地的(https://blog.thinkmoon.cn/post/242-deep-resnet-theory-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。