PyTorch实战:预训练、微调与迁移学习的代码级解析

在深度学习领域,掌握预训练模型的高效使用已成为开发者必备技能。本文将通过PyTorch代码实战,带你深入理解如何在实际图像分类任务中应用这些技术。我们将以花卉分类为例,使用ResNet50模型,对比不同策略的效果差异。

1. 环境准备与数据加载

首先确保已安装最新版PyTorch和TorchVision:

pip install torch torchvision torchaudio
pip install matplotlib pandas

花卉数据集可采用Oxford 102 Flowers数据集,其包含102类花卉图像。以下是数据加载与预处理的标准流程:

from torchvision import transforms, datasets

# 数据增强与归一化
train_transform = transforms.Compose([
    transforms.RandomResizedCrop(224),
    transforms.RandomHorizontalFlip(),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])

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

# 加载数据集
train_set = datasets.Flowers102(root='./data', split='train', transform=train_transform, download=True)
val_set = datasets.Flowers102(root='./data', split='val', transform=val_transform)

提示:ImageNet的均值和标准差([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])是预训练模型的标配参数,保持一致性对模型性能至关重要

2. 预训练模型加载策略

PyTorch提供了便捷的预训练模型加载方式。以下是三种典型加载方法的对比:

加载方式 代码示例 适用场景 内存占用
完整加载 model = torchvision.models.resnet50(pretrained=True) 需要全部层参数
部分加载 model.load_state_dict(torch.load(path), strict=False) 自定义修改网络结构
特征提取 features = list(model.children())[:-1] 仅需卷积特征

实际项目中推荐使用完整加载+自定义修改的方式:

import torchvision.models as models

def build_model(num_classes=102):
    # 加载预训练ResNet50
    model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)
    
    # 替换最后一层全连接
    in_features = model.fc.in_features
    model.fc = torch.nn.Linear(in_features, num_classes)
    
    return model

3. 微调技术实战详解

微调的核心在于参数更新策略的控制。我们通过冻结层来实现不同程度的微调:

3.1 基础微调方法

model = build_model()

# 方案1:仅训练最后一层(特征提取器模式)
for param in model.parameters():
    param.requires_grad = False
for param in model.fc.parameters():
    param.requires_grad = True

# 方案2:微调最后两个阶段(适用于中等规模数据集)
layers_to_train = ['layer4', 'fc']
for name, param in model.named_parameters():
    if any(layer in name for layer in layers_to_train):
        param.requires_grad = True
    else:
        param.requires_grad = False

# 方案3:全网络微调(大数据场景)
for param in model.parameters():
    param.requires_grad = True

3.2 分层学习率设置

不同层通常需要不同的学习率,这可以通过参数分组实现:

optimizer = torch.optim.SGD([
    {'params': model.conv1.parameters(), 'lr': 1e-4},
    {'params': model.layer1.parameters(), 'lr': 5e-4},
    {'params': model.layer2.parameters(), 'lr': 1e-3},
    {'params': model.fc.parameters(), 'lr': 5e-3}
], momentum=0.9)

注意:浅层网络通常提取基础特征,学习率应设置较小;深层网络和分类层需要更大学习率以适应新任务

4. 迁移学习的进阶技巧

4.1 特征提取与模型蒸馏

除了直接微调,我们还可以提取中间特征用于其他模型:

# 创建特征提取器
feature_extractor = torch.nn.Sequential(*list(model.children())[:-1])

# 提取特征
with torch.no_grad():
    features = feature_extractor(images)
    features = features.view(features.size(0), -1)

4.2 多任务学习框架

迁移学习可与多任务学习结合,共享底层特征:

class MultiTaskModel(nn.Module):
    def __init__(self, base_model):
        super().__init__()
        self.base = nn.Sequential(*list(base_model.children())[:-2])
        self.task1_head = nn.Linear(2048, 102)  # 花卉分类
        self.task2_head = nn.Linear(2048, 10)   # 附加任务
        
    def forward(self, x):
        features = self.base(x)
        features = features.mean([2, 3])  # 全局平均池化
        return self.task1_head(features), self.task2_head(features)

5. 实验对比与结果分析

我们对比了三种训练策略在花卉数据集上的表现:

策略 训练时间 Top-1准确率 Top-5准确率 过拟合风险
从头训练 4h 58.2% 82.1%
仅微调最后一层 1.5h 76.5% 92.3%
全网络微调 3h 89.7% 98.1%

从训练曲线可以看出,使用预训练模型的策略收敛更快:

# 绘制准确率曲线
plt.figure(figsize=(12, 4))
plt.subplot(121)
plt.plot(scratch_train_acc, label='Scratch Train')
plt.plot(finetune_train_acc, label='Finetune Train')
plt.title('Training Accuracy')
plt.legend()

plt.subplot(122)
plt.plot(scratch_val_acc, label='Scratch Val')
plt.plot(finetune_val_acc, label='Finetune Val')
plt.title('Validation Accuracy')
plt.legend()

在实际项目中,当遇到类似花卉分类这样的中等规模数据集(数千张图像)时,我通常会采用分阶段微调策略:先冻结所有层训练分类头,然后逐步解冻深层网络,最后微调全部参数。这种方法在保证模型性能的同时,能有效控制训练成本。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐