从零实现ResNet残差块:用PyTorch解剖深度学习中的"高速公路"

当你第一次看到ResNet的网络结构图时,那些跨越多个层的"弯弯绕绕"的连接线是否让你感到困惑?为什么简单的"抄近道"设计就能让神经网络突破千层大关?今天我们不谈空洞的理论,直接打开PyTorch的代码编辑器,用最直观的方式拆解这个改变了计算机视觉领域的神奇结构。

1. 残差连接的本质:当神经网络遇上高速公路

想象你正在驾驶一辆车从A点到B点。传统神经网络就像必须严格按照导航路线行驶,即使遇到堵车也不能变道;而残差网络则像拥有多条应急车道的高速公路,当主路拥堵时,车辆可以随时切换到更畅通的捷径。这种设计理念在2015年由微软研究院提出后,立刻让神经网络的深度突破了以往难以想象的限制。

在代码层面,一个基础的残差块(BasicBlock)只需要实现一个核心思想:输出 = 输入 + 变换后的输入。用数学表达就是:

y = x + F(x)

其中x是输入,F(x)是经过几层神经网络变换后的结果。这个看似简单的加法操作,实际上解决了深度神经网络训练中的两大难题:

  1. 梯度消失问题:在反向传播时,梯度可以直接通过加法操作回传到浅层
  2. 特征复用问题:网络可以自主决定哪些特征需要进一步加工,哪些直接传递

让我们用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

  1. 两个3x3卷积层
  2. 适用于浅层网络(如ResNet18/34)
  3. 计算量相对较大

Bottleneck

  1. 1x1 → 3x3 → 1x1的结构
  2. 适用于深层网络(如ResNet50/101/152)
  3. 计算量减少约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倍

残差网络与传统网络在训练过程中的对比非常明显:

  1. 收敛速度:残差网络通常在前几个epoch就能达到不错的准确率
  2. 训练稳定性:即使网络很深,训练曲线也很平滑
  3. 最终性能:深层残差网络可以突破传统网络的性能瓶颈

5. 残差连接的变体与最新发展

自ResNet提出以来,研究者们不断改进残差连接的形式。以下是几种重要的变体:

  1. Pre-activation ResNet:将BN和ReLU移到卷积之前,进一步改善梯度流动
  2. Wide ResNet:增加每层的通道数,减少深度
  3. ResNeXt:在残差块中引入分组卷积
  4. 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. 调试残差网络的实用技巧

在实际项目中应用残差网络时,有几个关键点需要注意:

常见问题排查清单

  1. 梯度爆炸/消失

    • 检查初始化是否正确
    • 确保使用了Batch Normalization
    • 尝试减小学习率
  2. 维度不匹配错误

    • 打印每一层的输入输出维度
    • 确保shortcut分支正确处理了stride和通道变化
  3. 性能不如预期

    • 尝试不同的学习率调度策略
    • 检查数据增强是否足够
    • 尝试不同的残差块变体

性能优化技巧

  • 使用混合精度训练可以显著减少显存占用
  • 对于小图像(如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层以上时,那种成就感是阅读任何论文都无法替代的。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐