别再死记ResNet结构了!用PyTorch手撕一个可运行的残差块,彻底搞懂‘短路连接’
从零实现ResNet残差块:用PyTorch解剖深度学习中的"高速公路"
当你第一次看到ResNet的网络结构图时,那些跨越多个层的"弯弯绕绕"的连接线是否让你感到困惑?为什么简单的"抄近道"设计就能让神经网络突破千层大关?今天我们不谈空洞的理论,直接打开PyTorch的代码编辑器,用最直观的方式拆解这个改变了计算机视觉领域的神奇结构。
1. 残差连接的本质:当神经网络遇上高速公路
想象你正在驾驶一辆车从A点到B点。传统神经网络就像必须严格按照导航路线行驶,即使遇到堵车也不能变道;而残差网络则像拥有多条应急车道的高速公路,当主路拥堵时,车辆可以随时切换到更畅通的捷径。这种设计理念在2015年由微软研究院提出后,立刻让神经网络的深度突破了以往难以想象的限制。
在代码层面,一个基础的残差块(BasicBlock)只需要实现一个核心思想:输出 = 输入 + 变换后的输入。用数学表达就是:
y = x + F(x)
其中x是输入,F(x)是经过几层神经网络变换后的结果。这个看似简单的加法操作,实际上解决了深度神经网络训练中的两大难题:
- 梯度消失问题:在反向传播时,梯度可以直接通过加法操作回传到浅层
- 特征复用问题:网络可以自主决定哪些特征需要进一步加工,哪些直接传递
让我们用PyTorch实现一个最简单的残差块:
import torch
import torch.nn as nn
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.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):
out = torch.relu(self.bn1(self.conv1(x)))
out = self.bn2(self.conv2(out))
out += self.shortcut(x) # 关键残差连接
return torch.relu(out)
2. 维度匹配的艺术:残差连接的实现细节
观察上面的代码,你会发现一个关键组件self.shortcut。这是因为在实际应用中,输入x和变换后的F(x)必须保持相同的维度才能相加。当出现维度不匹配时,我们通常有三种处理方案:
| 方案 | 实现方式 | 优点 | 缺点 |
|---|---|---|---|
| 零填充 | 在不足的维度补零 | 不增加参数 | 可能损失信息 |
| 1x1卷积 | 使用1x1卷积调整通道数 | 灵活调整维度 | 引入少量参数 |
| 池化 | 使用池化调整空间尺寸 | 计算量小 | 可能丢失细节 |
在我们的实现中选择了最常用的1x1卷积方案,这也是原论文中推荐的"方案B"。下面是一个维度不匹配时的处理示例:
# 输入通道64,输出通道128,特征图尺寸减半
block = BasicBlock(in_channels=64, out_channels=128, stride=2)
# 前向传播过程
x = torch.randn(1, 64, 32, 32) # (batch, channels, height, width)
output = block(x)
print(output.shape) # torch.Size([1, 128, 16, 16])
3. 从BasicBlock到Bottleneck:深度网络的进化
当网络深度增加到50层以上时,研究者们发现了一种更高效的残差块设计——Bottleneck结构。它通过1x1卷积先压缩再扩展通道数,形成了"瓶颈"形状,大幅减少了计算量:
class Bottleneck(nn.Module):
def __init__(self, in_channels, out_channels, stride=1, expansion=4):
super().__init__()
mid_channels = out_channels // expansion
self.conv1 = nn.Conv2d(in_channels, mid_channels, kernel_size=1, bias=False)
self.bn1 = nn.BatchNorm2d(mid_channels)
self.conv2 = nn.Conv2d(mid_channels, mid_channels, kernel_size=3, stride=stride, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(mid_channels)
self.conv3 = nn.Conv2d(mid_channels, out_channels, kernel_size=1, bias=False)
self.bn3 = 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 = torch.relu(self.bn1(self.conv1(x)))
out = torch.relu(self.bn2(self.conv2(out)))
out = self.bn3(self.conv3(out))
out += self.shortcut(x)
return torch.relu(out)
两种结构的对比如下:
BasicBlock:
- 两个3x3卷积层
- 适用于浅层网络(如ResNet18/34)
- 计算量相对较大
Bottleneck:
- 1x1 → 3x3 → 1x1的结构
- 适用于深层网络(如ResNet50/101/152)
- 计算量减少约40%
4. 实战演练:在CIFAR-10上训练残差网络
现在我们将完整的残差块组装成一个可以运行的网络,并在CIFAR-10数据集上进行测试。以下是完整的训练代码框架:
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
# 数据准备
transform = transforms.Compose([
transforms.RandomHorizontalFlip(),
transforms.RandomCrop(32, padding=4),
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])
trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
trainloader = torch.DataLoader(trainset, batch_size=128, shuffle=True, num_workers=2)
# 网络定义
class ResNet(nn.Module):
def __init__(self, block, num_blocks, num_classes=10):
super().__init__()
self.in_channels = 64
self.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(64)
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.linear = nn.Linear(512, 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
return nn.Sequential(*layers)
def forward(self, x):
out = torch.relu(self.bn1(self.conv1(x)))
out = self.layer1(out)
out = self.layer2(out)
out = self.layer3(out)
out = self.layer4(out)
out = torch.avg_pool2d(out, 4)
out = out.view(out.size(0), -1)
return self.linear(out)
# 创建ResNet18
def ResNet18():
return ResNet(BasicBlock, [2,2,2,2])
# 训练循环
net = ResNet18().cuda()
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.1, momentum=0.9, weight_decay=5e-4)
for epoch in range(200):
running_loss = 0.0
for i, data in enumerate(trainloader, 0):
inputs, labels = data
inputs, labels = inputs.cuda(), labels.cuda()
optimizer.zero_grad()
outputs = net(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
print(f'Epoch {epoch+1}, Loss: {running_loss/len(trainloader):.3f}')
提示:在实际训练中,学习率应该按照计划进行调整。通常在训练到1/3和2/3进度时将学习率降低10倍
残差网络与传统网络在训练过程中的对比非常明显:
- 收敛速度:残差网络通常在前几个epoch就能达到不错的准确率
- 训练稳定性:即使网络很深,训练曲线也很平滑
- 最终性能:深层残差网络可以突破传统网络的性能瓶颈
5. 残差连接的变体与最新发展
自ResNet提出以来,研究者们不断改进残差连接的形式。以下是几种重要的变体:
- Pre-activation ResNet:将BN和ReLU移到卷积之前,进一步改善梯度流动
- Wide ResNet:增加每层的通道数,减少深度
- ResNeXt:在残差块中引入分组卷积
- Attention Residual Learning:结合注意力机制
最新的研究趋势表明,残差连接的思想已经超越了计算机视觉领域,被广泛应用于:
- 自然语言处理(Transformer中的残差连接)
- 生成对抗网络
- 图神经网络
- 强化学习
在实现这些变体时,PyTorch的灵活性让我们可以轻松修改基础残差块。例如,实现一个Pre-activation版本的残差块只需要调整前向传播顺序:
class PreActBlock(nn.Module):
def __init__(self, in_channels, out_channels, stride=1):
super().__init__()
self.bn1 = nn.BatchNorm2d(in_channels)
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(out_channels)
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)
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)
)
def forward(self, x):
out = torch.relu(self.bn1(x))
shortcut = self.shortcut(out) # 注意这里也应用了预处理
out = self.conv1(out)
out = self.conv2(torch.relu(self.bn2(out)))
return out + shortcut
6. 调试残差网络的实用技巧
在实际项目中应用残差网络时,有几个关键点需要注意:
常见问题排查清单:
-
梯度爆炸/消失:
- 检查初始化是否正确
- 确保使用了Batch Normalization
- 尝试减小学习率
-
维度不匹配错误:
- 打印每一层的输入输出维度
- 确保shortcut分支正确处理了stride和通道变化
-
性能不如预期:
- 尝试不同的学习率调度策略
- 检查数据增强是否足够
- 尝试不同的残差块变体
性能优化技巧:
- 使用混合精度训练可以显著减少显存占用
- 对于小图像(如CIFAR),可以去掉第一个池化层
- 使用学习率warmup有助于训练极深网络
# 混合精度训练示例
from torch.cuda.amp import GradScaler, autocast
scaler = GradScaler()
for inputs, labels in trainloader:
inputs, labels = inputs.cuda(), labels.cuda()
optimizer.zero_grad()
with autocast():
outputs = net(inputs)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
在CIFAR-10数据集上,一个正确实现的ResNet18应该能够达到:
- 训练准确率:>95%
- 测试准确率:>90%
- 训练时间(单卡1080Ti):约30分钟
如果结果明显低于这些指标,很可能实现中存在某些问题。最常见的问题是残差连接没有正确工作,导致网络实际上退化成了普通CNN。可以通过检查梯度流动或可视化特征图来诊断这类问题。
理解残差网络的最好方式就是亲手实现它。建议从最简单的BasicBlock开始,逐步添加更复杂的组件,并在每个阶段验证网络的行为是否符合预期。当你看到自己实现的残差网络成功训练到100层以上时,那种成就感是阅读任何论文都无法替代的。
更多推荐


所有评论(0)