2D残差卷积网络入门指南:从理论到PyTorch实战

1次阅读
没有评论

共计 1883 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

为什么需要残差网络?

传统 CNN(如 VGG)随着层数加深会出现两大难题:

  1. 梯度消失 :反向传播时梯度连续相乘导致数值指数级减小,深层参数难以更新。数学表现为:
    $$\frac{\partial L}{\partial x} \approx \prod_{i=1}^{n}W_i \cdot \frac{\partial L}{\partial y} \rightarrow 0$$
  2. 网络退化 :实验表明 56 层 CNN 在 CIFAR-10 上的表现反而比 20 层更差,并非过拟合导致

残差学习(Residual Learning)的巧妙之处在于将目标函数改为:
$$H(x) = F(x) + x$$
其中 x 是恒等映射(identity mapping),F(x) 是残差部分。此时梯度变为:
$$\frac{\partial L}{\partial x} = \frac{\partial L}{\partial y} \cdot (1 + \frac{\partial F}{\partial x})$$
即使 $\frac{\partial F}{\partial x}$ 很小,梯度也不会完全消失。

主流架构对比

模型类型 参数量(百万) FLOPs(G) 适用场景
Plain CNN-34 21.3 3.6 浅层任务
ResNet-18 11.7 1.8 移动端 / 实时推理
ResNet-34 21.8 3.7 通用计算机视觉
ResNet-50 25.6 4.1 高精度需求场景

PyTorch 残差块实现

import torch
import torch.nn as nn

class BasicBlock(nn.Module):
    def __init__(self, in_channels, out_channels, stride=1):
        super().__init__()
        # 主分支:BN->ReLU->Conv 的标准顺序
        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):
        # [batch, channels, height, width]
        identity = x
        out = self.conv1(x)
        out = self.bn1(out)
        out = self.relu(out)
        out = self.conv2(out)
        out = self.bn2(out)
        out += self.shortcut(identity)  # 残差相加
        out = self.relu(out)
        return out

CIFAR-10 实验对比

我们使用相同超参数(学习率 0.1,batch size 128,训练 100 轮)测试:

  1. Plain CNN-34:最终准确率 72.3%,训练损失震荡明显
  2. ResNet-34:准确率提升至 89.7%,训练曲线平滑收敛

2D 残差卷积网络入门指南:从理论到 PyTorch 实战

常见问题解决方案

  1. 通道数不匹配 :当残差块进行下采样(stride=2)时,shortcut 分支需要用 1 ×1 卷积调整通道数和空间尺寸

  2. 梯度爆炸 :所有卷积层后必须加 BatchNorm,且初始化权重使用 He 初始化

  3. 性能下降 :避免在残差块内使用过大 kernel size(推荐 3 ×3),瓶颈结构(Bottleneck)中 1 ×1 卷积应先降维再升维

延伸思考方向

  1. 宽度调整 :能否根据输入图像分辨率动态调整通道数?比如 EfficientNet 的复合缩放方法

  2. 跨层连接 :DenseNet 的密集连接与 ResNet 的残差连接如何结合能进一步提升特征复用效率?

从实践角度看,ResNet 最吸引人的是其工程友好性——几乎不用调参就能获得不错的效果。建议初学者先复现基础版本,再逐步尝试改进方案。

正文完
 0
评论(没有评论)