从AlexNet到ResNeXt:用PyTorch复现7大经典图像分类网络实战指南

在计算机视觉领域,图像分类任务一直是衡量深度学习模型性能的重要基准。本文将带您从零开始,使用PyTorch框架完整复现七个里程碑式的神经网络架构。不同于单纯的理论讲解,我们将聚焦于工程实现细节实际编码技巧,帮助您真正掌握这些经典模型的精髓。

1. 环境准备与基础工具链搭建

1.1 PyTorch环境配置

首先需要确保您的开发环境已正确配置。推荐使用Python 3.8+和PyTorch 1.10+版本:

conda create -n torch-classify python=3.8
conda activate torch-classify
pip install torch torchvision torchaudio
pip install matplotlib tqdm tensorboard

对于GPU加速,需要额外安装CUDA工具包。可以通过以下命令验证GPU是否可用:

import torch
print(torch.cuda.is_available())  # 应输出True
print(torch.__version__)  # 确认版本号

1.2 数据预处理标准化流程

所有经典网络都使用ImageNet数据集作为基准,但实际训练时我们可以从CIFAR-10/100开始。这里给出通用的数据增强方案:

from torchvision import transforms

train_transform = transforms.Compose([
    transforms.RandomResizedCrop(224),
    transforms.RandomHorizontalFlip(),
    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),
    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])
])

2. 网络架构实现详解

2.1 AlexNet:深度学习的开山之作

AlexNet的PyTorch实现需要注意几个关键点:

import torch.nn as nn

class AlexNet(nn.Module):
    def __init__(self, num_classes=1000):
        super().__init__()
        self.features = nn.Sequential(
            nn.Conv2d(3, 64, kernel_size=11, stride=4, padding=2),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=3, stride=2),
            nn.Conv2d(64, 192, kernel_size=5, padding=2),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=3, stride=2),
            nn.Conv2d(192, 384, kernel_size=3, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(384, 256, kernel_size=3, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(256, 256, kernel_size=3, padding=1),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=3, stride=2),
        )
        self.avgpool = nn.AdaptiveAvgPool2d((6, 6))
        self.classifier = nn.Sequential(
            nn.Dropout(),
            nn.Linear(256 * 6 * 6, 4096),
            nn.ReLU(inplace=True),
            nn.Dropout(),
            nn.Linear(4096, 4096),
            nn.ReLU(inplace=True),
            nn.Linear(4096, num_classes),
        )

    def forward(self, x):
        x = self.features(x)
        x = self.avgpool(x)
        x = torch.flatten(x, 1)
        x = self.classifier(x)
        return x

实现要点

  • 原始论文使用LRN层,但现代实现通常省略
  • 使用inplace=True的ReLU可以节省内存
  • 全连接层前的自适应池化使网络适应不同输入尺寸

2.2 VGG:3×3卷积的胜利

VGG的核心在于堆叠小型卷积核。以下是VGG-16的模块化实现:

def make_layers(cfg, batch_norm=False):
    layers = []
    in_channels = 3
    for v in cfg:
        if v == 'M':
            layers += [nn.MaxPool2d(kernel_size=2, stride=2)]
        else:
            conv2d = nn.Conv2d(in_channels, v, kernel_size=3, padding=1)
            if batch_norm:
                layers += [conv2d, nn.BatchNorm2d(v), nn.ReLU(inplace=True)]
            else:
                layers += [conv2d, nn.ReLU(inplace=True)]
            in_channels = v
    return nn.Sequential(*layers)

cfgs = {
    'vgg16': [64, 64, 'M', 128, 128, 'M', 256, 256, 256, 'M', 
              512, 512, 512, 'M', 512, 512, 512, 'M'],
}

class VGG(nn.Module):
    def __init__(self, features, num_classes=1000):
        super().__init__()
        self.features = features
        self.avgpool = nn.AdaptiveAvgPool2d((7, 7))
        self.classifier = nn.Sequential(
            nn.Linear(512 * 7 * 7, 4096),
            nn.ReLU(inplace=True),
            nn.Dropout(),
            nn.Linear(4096, 4096),
            nn.ReLU(inplace=True),
            nn.Dropout(),
            nn.Linear(4096, num_classes),
        )

    def forward(self, x):
        x = self.features(x)
        x = self.avgpool(x)
        x = torch.flatten(x, 1)
        x = self.classifier(x)
        return x

工程技巧

  • 使用配置文件定义网络结构,便于扩展不同变体
  • 批量归一化可以显著提升训练稳定性
  • 预训练权重加载时需要匹配参数名称

2.3 ResNet:残差连接的革命

残差网络的核心是BasicBlock和Bottleneck设计。以下是关键实现:

class BasicBlock(nn.Module):
    expansion = 1

    def __init__(self, in_planes, planes, stride=1):
        super().__init__()
        self.conv1 = nn.Conv2d(
            in_planes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
        self.bn1 = nn.BatchNorm2d(planes)
        self.conv2 = nn.Conv2d(planes, planes, kernel_size=3,
                               stride=1, padding=1, bias=False)
        self.bn2 = nn.BatchNorm2d(planes)

        self.shortcut = nn.Sequential()
        if stride != 1 or in_planes != self.expansion*planes:
            self.shortcut = nn.Sequential(
                nn.Conv2d(in_planes, self.expansion*planes,
                          kernel_size=1, stride=stride, bias=False),
                nn.BatchNorm2d(self.expansion*planes)
            )

    def forward(self, x):
        out = F.relu(self.bn1(self.conv1(x)))
        out = self.bn2(self.conv2(out))
        out += self.shortcut(x)
        out = F.relu(out)
        return out

class ResNet(nn.Module):
    def __init__(self, block, num_blocks, num_classes=1000):
        super().__init__()
        self.in_planes = 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, planes, num_blocks, stride):
        strides = [stride] + [1]*(num_blocks-1)
        layers = []
        for stride in strides:
            layers.append(block(self.in_planes, planes, stride))
            self.in_planes = planes * block.expansion
        return nn.Sequential(*layers)

    def forward(self, x):
        x = F.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

关键细节

  • 残差连接需要处理维度不匹配情况
  • Bottleneck结构在深层网络中更高效
  • 预激活(BN-ReLU-Conv)变体通常表现更好

3. 训练技巧与调优策略

3.1 学习率调度策略比较

不同网络架构适合不同的学习率策略:

网络类型 推荐初始LR 调度策略 周期数 动量
AlexNet 1e-2 阶梯下降(每30epoch) 90 0.9
VGG 5e-3 余弦退火 120 0.9
ResNet 1e-1 预热+线性衰减 100 0.9
DenseNet 1e-1 OneCycleLR 300 0.9

实现OneCycleLR策略示例:

from torch.optim.lr_scheduler import OneCycleLR

optimizer = torch.optim.SGD(model.parameters(), lr=0.1, momentum=0.9)
scheduler = OneCycleLR(optimizer, max_lr=0.1, 
                      steps_per_epoch=len(train_loader), 
                      epochs=300)

3.2 混合精度训练实现

现代GPU支持混合精度训练,可大幅减少显存占用:

from torch.cuda.amp import GradScaler, autocast

scaler = GradScaler()

for inputs, targets in train_loader:
    optimizer.zero_grad()
    
    with autocast():
        outputs = model(inputs)
        loss = criterion(outputs, targets)
    
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()
    scheduler.step()

注意事项

  • 需要确保模型中没有不兼容FP16的操作
  • 损失缩放(loss scaling)对稳定性至关重要
  • 批量归一化层应保持FP32精度

4. 模型部署与性能优化

4.1 TorchScript导出与优化

将训练好的模型导出为可部署格式:

model.eval()
example_input = torch.rand(1, 3, 224, 224)
traced_script = torch.jit.trace(model, example_input)
traced_script.save("resnet50.pt")

# 进一步优化
optimized_model = torch.jit.optimize_for_inference(traced_script)

4.2 TensorRT加速实践

使用TensorRT可以获得显著的推理加速:

import tensorrt as trt

logger = trt.Logger(trt.Logger.WARNING)
builder = trt.Builder(logger)
network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH))
parser = trt.OnnxParser(network, logger)

with open("model.onnx", "rb") as f:
    parser.parse(f.read())

config = builder.create_builder_config()
config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 << 30)
serialized_engine = builder.build_serialized_network(network, config)

with open("engine.trt", "wb") as f:
    f.write(serialized_engine)

性能对比

模型 PyTorch(ms) TensorRT(ms) 加速比
ResNet-50 12.3 3.2 3.8x
DenseNet-121 15.7 4.1 3.8x
VGG-16 18.2 5.3 3.4x

在实际项目中,选择适合的模型架构需要平衡准确率、推理速度和资源消耗。ResNet系列因其优秀的准确率-速度权衡,仍然是许多工业应用的首选,而最新的EfficientNet等架构则在边缘设备上展现出更大优势。

Logo

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

更多推荐