手把手复现ResNet-34:用PyTorch从零搭建并可视化训练过程(附代码)

在计算机视觉领域,ResNet-34作为残差网络家族的经典成员,以其优雅的架构设计和突破性的性能表现,成为深度学习从业者必须掌握的里程碑式模型。本文将带您从零开始实现一个完整的ResNet-34模型,不仅包含网络结构的逐层构建,还会深入探讨残差连接的核心机制,并通过可视化技术让训练过程变得透明可观察。

1. 环境准备与基础架构

1.1 PyTorch环境配置

首先确保已安装最新版PyTorch和必要的可视化工具:

pip install torch torchvision tensorboard matplotlib

1.2 BasicBlock实现

ResNet-34的核心构建单元是BasicBlock,它包含两个3×3卷积层和跨层连接:

import torch
import torch.nn as nn

class BasicBlock(nn.Module):
    expansion = 1
    
    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)
        
        self.shortcut = nn.Sequential()
        if stride != 1 or in_channels != self.expansion * out_channels:
            self.shortcut = nn.Sequential(
                nn.Conv2d(
                    in_channels, self.expansion * out_channels,
                    kernel_size=1, stride=stride, bias=False
                ),
                nn.BatchNorm2d(self.expansion * 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)

关键点解析:

  • 残差连接:通过out += self.shortcut(x)实现特征图相加
  • 维度匹配:当输入输出通道数不一致时,使用1×1卷积调整维度
  • 下采样:通过设置stride=2实现特征图尺寸减半

2. 完整ResNet-34实现

2.1 网络主体结构

基于BasicBlock构建完整的ResNet-34:

class ResNet(nn.Module):
    def __init__(self, block, num_blocks, num_classes=1000):
        super().__init__()
        self.in_channels = 64
        
        self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)
        self.bn1 = nn.BatchNorm2d(64)
        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
        
        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.avgpool = nn.AdaptiveAvgPool2d((1, 1))
        self.fc = nn.Linear(512 * block.expansion, 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 * block.expansion
        return nn.Sequential(*layers)
    
    def forward(self, x):
        x = torch.relu(self.bn1(self.conv1(x)))
        x = self.maxpool(x)
        
        x = self.layer1(x)
        x = self.layer2(x)
        x = self.layer3(x)
        x = self.layer4(x)
        
        x = self.avgpool(x)
        x = torch.flatten(x, 1)
        x = self.fc(x)
        return x

2.2 模型实例化

创建ResNet-34实例:

def ResNet34():
    return ResNet(BasicBlock, [3, 4, 6, 3])

model = ResNet34()
print(model)

网络结构要点:

  • 初始卷积层:7×7卷积,stride=2,快速下采样
  • 四个阶段:分别包含3、4、6、3个BasicBlock
  • 全局平均池化:替代全连接层,减少参数量
  • 输出层:1000维对应ImageNet类别

3. 训练流程实现

3.1 数据准备与增强

使用PyTorch的ImageFolder加载数据并应用增强:

from torchvision import datasets, transforms

train_transform = transforms.Compose([
    transforms.RandomResizedCrop(224),
    transforms.RandomHorizontalFlip(),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], 
                         std=[0.229, 0.224, 0.225])
])

val_transform = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], 
                         std=[0.229, 0.224, 0.225])
])

train_set = datasets.ImageFolder('path/to/train', train_transform)
val_set = datasets.ImageFolder('path/to/val', val_transform)

train_loader = torch.utils.data.DataLoader(
    train_set, batch_size=64, shuffle=True, num_workers=4)
val_loader = torch.utils.data.DataLoader(
    val_set, batch_size=64, shuffle=False, num_workers=4)

3.2 训练循环实现

完整的训练流程包含损失函数、优化器和学习率调度:

import torch.optim as optim
from torch.utils.tensorboard import SummaryWriter

writer = SummaryWriter('runs/resnet34')

criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.1, 
                      momentum=0.9, weight_decay=1e-4)
scheduler = optim.lr_scheduler.StepLR(optimizer, 
                                     step_size=30, gamma=0.1)

def train(epoch):
    model.train()
    for batch_idx, (data, target) in enumerate(train_loader):
        optimizer.zero_grad()
        output = model(data)
        loss = criterion(output, target)
        loss.backward()
        optimizer.step()
        
        if batch_idx % 100 == 0:
            print(f'Train Epoch: {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)}]'
                  f'\tLoss: {loss.item():.6f}')
            writer.add_scalar('training loss', 
                            loss.item(), 
                            epoch * len(train_loader) + batch_idx)

def validate():
    model.eval()
    val_loss = 0
    correct = 0
    with torch.no_grad():
        for data, target in val_loader:
            output = model(data)
            val_loss += criterion(output, target).item()
            pred = output.argmax(dim=1, keepdim=True)
            correct += pred.eq(target.view_as(pred)).sum().item()
    
    val_loss /= len(val_loader.dataset)
    acc = 100. * correct / len(val_loader.dataset)
    print(f'\nValidation set: Average loss: {val_loss:.4f}, '
          f'Accuracy: {correct}/{len(val_loader.dataset)} ({acc:.2f}%)\n')
    writer.add_scalar('val accuracy', acc, epoch)
    return acc

for epoch in range(1, 100):
    train(epoch)
    val_acc = validate()
    scheduler.step()
    
    # 保存最佳模型
    if val_acc > best_acc:
        torch.save(model.state_dict(), 'resnet34_best.pth')
        best_acc = val_acc

writer.close()

4. 训练可视化与分析

4.1 TensorBoard监控

启动TensorBoard查看训练过程:

tensorboard --logdir=runs

关键监控指标:

  • 训练损失曲线
  • 验证集准确率
  • 权重分布直方图
  • 梯度流动情况

4.2 特征图可视化

实现中间层特征可视化:

import matplotlib.pyplot as plt

def visualize_feature_maps(model, img_tensor, layer_name):
    activations = {}
    
    def get_activation(name):
        def hook(model, input, output):
            activations[name] = output.detach()
        return hook
    
    # 注册hook
    for name, layer in model.named_modules():
        if name == layer_name:
            layer.register_forward_hook(get_activation(name))
    
    # 前向传播
    with torch.no_grad():
        model(img_tensor.unsqueeze(0))
    
    # 可视化
    act = activations[layer_name].squeeze()
    fig, axarr = plt.subplots(act.size(0)//8, 8, figsize=(20,20))
    for idx in range(act.size(0)):
        axarr[idx//8, idx%8].imshow(act[idx].cpu())
        axarr[idx//8, idx%8].axis('off')
    plt.show()

# 示例:可视化第一个残差块的输出
sample_img, _ = next(iter(val_loader))
visualize_feature_maps(model, sample_img[0], 'layer1.0')

4.3 梯度流动分析

检查梯度传播情况有助于诊断训练问题:

def plot_grad_flow(named_parameters):
    ave_grads = []
    layers = []
    for n, p in named_parameters:
        if(p.requires_grad) and ("bias" not in n):
            layers.append(n)
            ave_grads.append(p.grad.abs().mean().cpu())
    
    plt.figure(figsize=(10,6))
    plt.bar(np.arange(len(ave_grads)), ave_grads, alpha=0.5, lw=1)
    plt.hlines(0, 0, len(ave_grads)+1, lw=2, color="k") 
    plt.xticks(np.arange(len(ave_grads)), layers, rotation="vertical")
    plt.xlim(left=0, right=len(ave_grads))
    plt.ylim(bottom=-0.001, top=0.02)  # 自定义范围
    plt.xlabel("Layers")
    plt.ylabel("Average gradient")
    plt.title("Gradient flow")
    plt.grid(True)
    plt.tight_layout()
    plt.show()

# 在训练循环中添加
optimizer.step()
plot_grad_flow(model.named_parameters())

5. 调试技巧与性能优化

5.1 常见问题排查

维度不匹配错误

  • 检查每个残差块的输入输出通道数
  • 验证shortcut连接是否正确处理了stride>1的情况
  • 使用print(x.shape)在关键位置输出张量形状

训练不收敛

  • 检查初始学习率是否合适(ResNet通常使用0.1)
  • 验证数据归一化参数是否正确
  • 确保BatchNorm层处于训练模式

5.2 混合精度训练

利用NVIDIA的Apex库加速训练:

from apex import amp

model = ResNet34().cuda()
optimizer = optim.SGD(model.parameters(), lr=0.1)

model, optimizer = amp.initialize(model, optimizer, opt_level="O1")

with amp.scale_loss(loss, optimizer) as scaled_loss:
    scaled_loss.backward()

5.3 模型量化

训练后量化减小模型体积:

quantized_model = torch.quantization.quantize_dynamic(
    model, {nn.Linear, nn.Conv2d}, dtype=torch.qint8
)
torch.save(quantized_model.state_dict(), 'resnet34_quantized.pth')
Logo

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

更多推荐