别再死记ResNet了!用PyTorch从零实现DenseNet-121,理解它的‘密集连接’到底好在哪

在深度学习领域,模型架构的创新往往伴随着对信息流动方式的重新思考。当ResNet通过残差连接(skip connection)解决了深层网络梯度消失问题时,另一种更为激进的设计——DenseNet的密集连接(dense connection)悄然登场。本文将带您用PyTorch从零构建DenseNet-121,通过代码实践揭示其核心设计哲学。

1. 密集连接:重新定义层间信息流动

传统卷积神经网络的层间连接像接力赛跑,每一层只能从前一层接收信息。ResNet引入了"捷径",允许信息跳过某些层。而DenseNet则更进一步——每一层都与之前所有层直接相连,形成全连接的信息网络。

这种设计带来三个关键优势:

  • 梯度流动更顺畅 :深层网络训练时,梯度可以直达浅层
  • 特征复用更高效 :所有中间特征都能被后续层直接利用
  • 参数效率更高 :每层只需学习新增特征,无需重复学习

用PyTorch定义一个基础的Dense Block:

import torch
import torch.nn as nn

class DenseLayer(nn.Module):
    def __init__(self, in_channels, growth_rate):
        super().__init__()
        self.bn = nn.BatchNorm2d(in_channels)
        self.conv = nn.Conv2d(in_channels, growth_rate, kernel_size=3, padding=1)
        
    def forward(self, x):
        out = self.conv(F.relu(self.bn(x)))
        return torch.cat([x, out], 1)  # 沿通道维度拼接

2. 构建DenseNet-121的核心模块

完整的DenseNet由多个Dense Block和Transition Layer交替组成。让我们分解实现每个组件:

2.1 带瓶颈层的Dense Block改进版

原始DenseNet的改进版本引入了瓶颈结构,显著减少了计算量:

class BottleneckDenseLayer(nn.Module):
    def __init__(self, in_channels, growth_rate, bn_size=4):
        super().__init__()
        inner_channels = bn_size * growth_rate
        self.bottleneck = nn.Sequential(
            nn.BatchNorm2d(in_channels),
            nn.ReLU(),
            nn.Conv2d(in_channels, inner_channels, kernel_size=1, bias=False)
        )
        self.conv = nn.Sequential(
            nn.BatchNorm2d(inner_channels),
            nn.ReLU(),
            nn.Conv2d(inner_channels, growth_rate, kernel_size=3, padding=1, bias=False)
        )
        
    def forward(self, x):
        out = self.bottleneck(x)
        out = self.conv(out)
        return torch.cat([x, out], 1)

2.2 过渡层的压缩设计

过渡层不仅降低特征图尺寸,还能通过压缩率θ减少通道数:

class TransitionLayer(nn.Module):
    def __init__(self, in_channels, compression=0.5):
        super().__init__()
        out_channels = int(in_channels * compression)
        self.transition = nn.Sequential(
            nn.BatchNorm2d(in_channels),
            nn.ReLU(),
            nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False),
            nn.AvgPool2d(2, stride=2)
        )
        
    def forward(self, x):
        return self.transition(x)

3. 完整DenseNet-121架构实现

结合上述模块,我们可以构建完整的DenseNet-121:

class DenseNet(nn.Module):
    def __init__(self, growth_rate=32, block_config=(6,12,24,16), 
                 num_classes=1000, compression=0.5):
        super().__init__()
        # 初始卷积层
        self.features = nn.Sequential(
            nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False),
            nn.BatchNorm2d(64),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
        )
        
        # 构建Dense Block和Transition Layer
        num_channels = 64
        for i, num_layers in enumerate(block_config):
            block = nn.Sequential()
            for j in range(num_layers):
                layer = BottleneckDenseLayer(num_channels + j*growth_rate, growth_rate)
                block.add_module(f'denselayer{i+1}_{j+1}', layer)
            num_channels += num_layers * growth_rate
            self.features.add_module(f'denseblock{i+1}', block)
            
            if i != len(block_config)-1:  # 最后一个block后不加transition
                trans = TransitionLayer(num_channels, compression)
                self.features.add_module(f'transition{i+1}', trans)
                num_channels = int(num_channels * compression)
        
        # 分类器
        self.classifier = nn.Linear(num_channels, num_classes)
        
    def forward(self, x):
        features = self.features(x)
        out = F.adaptive_avg_pool2d(features, (1,1))
        out = torch.flatten(out, 1)
        out = self.classifier(out)
        return out

4. 可视化分析与性能对比

4.1 特征重用可视化

我们可以通过hook机制捕获各层的特征图:

def visualize_features(model, input_tensor):
    features = {}
    def get_feature(name):
        def hook(model, input, output):
            features[name] = output.detach()
        return hook
    
    handles = []
    for name, layer in model.named_modules():
        if isinstance(layer, nn.Conv2d):
            handle = layer.register_forward_hook(get_feature(name))
            handles.append(handle)
    
    with torch.no_grad():
        model(input_tensor)
    
    for handle in handles:
        handle.remove()
    
    return features

4.2 参数量对比实验

与传统网络相比,DenseNet的参数效率显著提升:

模型 参数量(M) Top-1准确率(%)
ResNet-50 25.5 76.0
DenseNet-121 8.0 74.7
DenseNet-169 14.2 76.2

尽管DenseNet-121的参数量只有ResNet-50的31%,但准确率仅低1.3个百分点。

5. 训练技巧与实战建议

  1. 学习率调度 :使用余弦退火策略

    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
    
  2. 数据增强 :适合DenseNet的增强组合

    transform = transforms.Compose([
        transforms.RandomResizedCrop(224),
        transforms.RandomHorizontalFlip(),
        transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ])
    
  3. 内存优化 :当GPU内存不足时

    • 减小batch size
    • 使用梯度检查点技术
    • 尝试混合精度训练

在实际项目中,DenseNet特别适合以下场景:

  • 需要轻量级模型部署的移动端应用
  • 数据量相对较小的专业领域图像识别
  • 需要特征复用和多尺度特征融合的任务
Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐