ResNet 迁移学习实战:PyTorch 预训练模型在5类花卉数据集上微调,准确率提升至95%
ResNet迁移学习实战:5类花卉分类的PyTorch完整解决方案
当面对特定领域的图像分类任务时,从头训练深度神经网络往往需要大量数据和计算资源。迁移学习技术让我们能够利用在大规模数据集(如ImageNet)上预训练的模型,通过微调快速适应新任务。本文将手把手带您实现一个基于PyTorch的ResNet迁移学习项目,在5类花卉数据集上达到95%的准确率。
1. 项目准备与环境配置
在开始之前,我们需要准备好开发环境和数据集。这个项目推荐使用Python 3.8+和PyTorch 1.10+版本,以下是环境配置的关键步骤:
# 创建并激活虚拟环境
python -m venv flower_cls
source flower_cls/bin/activate # Linux/Mac
flower_cls\Scripts\activate # Windows
# 安装核心依赖
pip install torch torchvision torchaudio
pip install matplotlib pillow pandas
花卉数据集可以从Kaggle或公开数据集平台获取,通常包含以下5个类别:
- 雏菊(daisy)
- 蒲公英(dandelion)
- 玫瑰(roses)
- 向日葵(sunflower)
- 郁金香(tulips)
数据集目录结构应如下:
flower_data/
train/
daisy/
dandelion/
roses/
sunflower/
tulips/
val/
daisy/
dandelion/
roses/
sunflower/
tulips/
2. 数据预处理与增强策略
图像数据的预处理和增强对模型性能至关重要。我们使用torchvision提供的工具来构建数据管道:
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], # ImageNet标准化
std=[0.229, 0.224, 0.225])
])
# 验证集转换(无需增强)
val_transform = transforms.Compose([
transforms.Resize(256), # 缩放至256x256
transforms.CenterCrop(224), # 中心裁剪
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406],
[0.229, 0.224, 0.225])
])
数据增强策略的选择需要平衡多样性和真实性。对于花卉分类,我们还考虑了以下增强方式:
- 随机旋转(±30度)
- 高斯模糊(模拟焦距变化)
- 随机灰度化(概率20%)
但要注意,过度增强可能导致模型难以学习有效特征。验证集应保持原始分布以准确评估模型性能。
3. ResNet模型加载与微调
PyTorch提供了预训练的ResNet模型,我们可以轻松加载并修改最后一层以适应我们的分类任务:
import torchvision.models as models
import torch.nn as nn
def initialize_model(num_classes):
# 加载预训练ResNet34
model = models.resnet34(pretrained=True)
# 冻结所有卷积层参数
for param in model.parameters():
param.requires_grad = False
# 替换最后的全连接层
num_ftrs = model.fc.in_features
model.fc = nn.Linear(num_ftrs, num_classes)
return model
model = initialize_model(num_classes=5)
model = model.to(device) # 移至GPU
模型微调策略对比:
| 策略 | 训练参数 | 数据需求 | 训练速度 | 适用场景 |
|---|---|---|---|---|
| 全冻结 | 仅最后一层 | 较少 | 最快 | 小数据集,与预训练任务相似 |
| 部分微调 | 后几层+分类器 | 中等 | 中等 | 中等规模数据 |
| 全微调 | 所有参数 | 大量 | 最慢 | 大数据集,任务差异大 |
在本项目中,我们采用分阶段微调策略:
- 先冻结卷积层,只训练分类器(3个epoch)
- 解冻所有层,整体微调(10个epoch)
- 使用更小的学习率精细调整(5个epoch)
4. 训练过程与超参数优化
训练过程中有几个关键因素需要特别注意:
损失函数与优化器选择:
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam([
{'params': model.conv1.parameters(), 'lr': 1e-5},
{'params': model.layer1.parameters(), 'lr': 1e-4},
{'params': model.layer2.parameters(), 'lr': 1e-4},
{'params': model.layer3.parameters(), 'lr': 1e-3},
{'params': model.layer4.parameters(), 'lr': 1e-3},
{'params': model.fc.parameters(), 'lr': 1e-3}
])
学习率调度策略:
scheduler = optim.lr_scheduler.ReduceLROnPlateau(
optimizer, mode='max', factor=0.1, patience=3, verbose=True)
训练过程中的关键指标监控:
for epoch in range(num_epochs):
model.train()
running_loss = 0.0
for inputs, labels in train_loader:
inputs = inputs.to(device)
labels = labels.to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
# 验证阶段
model.eval()
val_acc = 0.0
with torch.no_grad():
for inputs, labels in val_loader:
inputs = inputs.to(device)
labels = labels.to(device)
outputs = model(inputs)
_, preds = torch.max(outputs, 1)
val_acc += torch.sum(preds == labels.data)
val_acc = val_acc.double() / len(val_dataset)
scheduler.step(val_acc) # 根据验证准确率调整学习率
print(f'Epoch {epoch+1}/{num_epochs}')
print(f'Train Loss: {running_loss/len(train_loader):.4f}')
print(f'Val Acc: {val_acc:.4f}')
5. 模型评估与性能提升技巧
在完成训练后,我们需要全面评估模型性能。除了准确率外,还应考虑:
- 混淆矩阵分析
- 各类别的精确率、召回率和F1分数
- 推理速度(FPS)
混淆矩阵实现:
from sklearn.metrics import confusion_matrix
import seaborn as sns
def plot_confusion_matrix(model, data_loader):
model.eval()
all_preds = []
all_labels = []
with torch.no_grad():
for inputs, labels in data_loader:
inputs = inputs.to(device)
labels = labels.to(device)
outputs = model(inputs)
_, preds = torch.max(outputs, 1)
all_preds.extend(preds.cpu().numpy())
all_labels.extend(labels.cpu().numpy())
cm = confusion_matrix(all_labels, all_preds)
plt.figure(figsize=(10,8))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
xticklabels=class_names,
yticklabels=class_names)
plt.xlabel('Predicted')
plt.ylabel('Actual')
plt.show()
性能提升技巧:
-
标签平滑(Label Smoothing) :
criterion = nn.CrossEntropyLoss(label_smoothing=0.1) -
混合精度训练 :
from torch.cuda.amp import GradScaler, autocast scaler = GradScaler() with autocast(): outputs = model(inputs) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() -
模型集成 :
def ensemble_predict(models, input): with torch.no_grad(): outputs = [model(input) for model in models] avg_output = torch.mean(torch.stack(outputs), dim=0) _, pred = torch.max(avg_output, 1) return pred
6. 模型部署与推理优化
训练好的模型需要优化以便在实际应用中高效运行:
模型导出为ONNX格式:
dummy_input = torch.randn(1, 3, 224, 224).to(device)
torch.onnx.export(model, dummy_input, "flower_resnet34.onnx",
input_names=["input"], output_names=["output"],
dynamic_axes={"input": {0: "batch_size"},
"output": {0: "batch_size"}})
使用TorchScript优化:
traced_model = torch.jit.trace(model, dummy_input)
traced_model.save("flower_resnet34.pt")
推理代码示例:
from PIL import Image
def predict(image_path, model, transform):
img = Image.open(image_path).convert('RGB')
img_t = transform(img)
batch_t = torch.unsqueeze(img_t, 0).to(device)
model.eval()
with torch.no_grad():
output = model(batch_t)
prob = torch.nn.functional.softmax(output[0], dim=0)
_, pred = torch.max(output, 1)
return class_names[pred.item()], prob[pred.item()].item()
# 使用示例
class_name, confidence = predict("test_rose.jpg", model, val_transform)
print(f"预测结果: {class_name}, 置信度: {confidence:.2f}")
7. 实际应用中的挑战与解决方案
在实际部署花卉分类模型时,可能会遇到以下挑战及应对策略:
光照条件变化 :
- 解决方案:在数据增强中加入随机亮度调整
- 测试时使用直方图均衡化预处理
背景干扰 :
# 使用显著性检测减少背景干扰
from skimage.segmentation import quickshift
def salient_region_crop(image):
segments = quickshift(image, kernel_size=3, max_dist=6, ratio=0.5)
# 后续处理获取主要物体区域...
return cropped_image
类别不平衡 :
- 使用加权采样器
- 在损失函数中引入类别权重
class_weights = compute_class_weight('balanced', classes=np.unique(train_labels), y=train_labels) weights = torch.tensor(class_weights, dtype=torch.float).to(device) criterion = nn.CrossEntropyLoss(weight=weights)
模型轻量化 :
# 使用知识蒸馏训练小型模型
teacher_model = models.resnet34(pretrained=True)
student_model = models.resnet18()
# 蒸馏损失
def distillation_loss(y, labels, teacher_scores, temp=5.0, alpha=0.7):
return alpha * F.cross_entropy(y, labels) + \
(1-alpha) * F.kl_div(F.log_softmax(y/temp, dim=1),
F.softmax(teacher_scores/temp, dim=1))
通过本项目的完整实现,我们不仅掌握了ResNet迁移学习的技术要点,还建立了一套可复用的图像分类流程。这套方法可以轻松扩展到其他细粒度分类任务,如鸟类识别、车辆型号识别等。关键在于根据具体问题调整数据增强策略和微调方法,同时持续监控模型在实际场景中的表现。
更多推荐




所有评论(0)