ResNet风格2D卷积残差网络:从原理到新手友好实现

1次阅读
没有评论

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

image.webp

为什么需要残差网络?

传统深度 CNN 随着层数增加会出现两大难题:

ResNet 风格 2D 卷积残差网络:从原理到新手友好实现

  • 梯度消失 :反向传播时链式法则导致浅层权重更新幅度指数级衰减
  • 网络退化 :56 层网络的训练误差反而比 20 层更高(非过拟合导致)

残差学习通过引入跨层连接(skip connection)将原始映射转化为 $F(x) = H(x) – x$ 的残差形式,其数学优势在于:

  1. 极端情况下可使 $F(x) \rightarrow 0$ 退化为恒等映射
  2. 反向传播时梯度多了一条无损通路:$\frac{\partial loss}{\partial x} = \frac{\partial loss}{\partial F(x)} \cdot \frac{\partial F(x)}{\partial x} + 1$

架构对比实验

使用 CIFAR-10 数据集对比 VGG-16 与 ResNet-18 的表现:

指标 VGG-16 ResNet-18
训练准确率 72.3% 94.8%
测试准确率 68.5% 92.1%
收敛周期 120 60

训练曲线显示 ResNet 的损失值下降更快且更稳定,验证了残差连接的有效性。

PyTorch 实现详解

BasicBlock 核心组件

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)

        # 处理维度不匹配的 shortcut
        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 = F.relu(self.bn1(self.conv1(x)))
        out = self.bn2(self.conv2(out))
        out += self.shortcut(x)  # 残差连接
        return F.relu(out)

完整网络搭建

def ResNet18(num_classes=10):
    model = nn.Sequential(nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False),
        nn.BatchNorm2d(64),
        nn.ReLU(),

        # 4 个残差阶段
        make_layer(64, 64, stride=1, num_blocks=2),
        make_layer(64, 128, stride=2, num_blocks=2),
        make_layer(128, 256, stride=2, num_blocks=2),
        make_layer(256, 512, stride=2, num_blocks=2),

        nn.AdaptiveAvgPool2d((1,1)),
        nn.Flatten(),
        nn.Linear(512, num_classes)
    )
    return model

关键实践技巧

  1. BatchNorm 放置顺序
  2. 坚持 conv -> bn -> relu 标准顺序
  3. 切勿在残差相加后遗漏 BN 层

  4. 初始化策略

  5. 残差分支最后一层 BN 的 γ 初始化为 0
  6. 使网络初始状态接近恒等映射

  7. 优化器配置

    optimizer = torch.optim.SGD(model.parameters(),
        lr=0.1,
        weight_decay=5e-4,
        momentum=0.9
    )
    scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=[30, 60], gamma=0.1
    )

性能优化实战

  • FLOPs 计算

    from ptflops import get_model_complexity_info
    flops, params = get_model_complexity_info(model, (3,32,32), as_strings=True)
    print(f"FLOPs: {flops}, Params: {params}")

  • 内存优化

  • 使用梯度检查点(checkpointing)
  • 混合精度训练(AMP)

延伸思考方向

  1. 跨层连接设计
  2. DenseNet 的密集连接模式
  3. 随机深度(Stochastic Depth)

  4. 语义分割改造

  5. 全卷积形式的跳跃连接
  6. 多尺度特征融合

通过这个实现案例,我们可以体会到残差网络 ” 大道至简 ” 的设计哲学。建议读者尝试在自定义数据集上微调网络深度和宽度,观察模型性能的变化规律。

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