PyTorch 1.13 ResNet-152 迁移学习实战:102类花朵识别 Top-1 准确率 92%
·
PyTorch 1.13 ResNet-152 迁移学习实战:102类花朵识别 Top-1 准确率 92%
在计算机视觉领域,图像分类一直是基础且重要的任务。本文将深入探讨如何利用PyTorch 1.13中的预训练ResNet-152模型,通过迁移学习技术实现102类花朵的高精度识别,最终达到92%的Top-1准确率。不同于从零开始训练,迁移学习能大幅减少训练时间和计算资源消耗,同时保持优异的性能表现。
1. 环境准备与数据预处理
1.1 硬件与软件配置
要实现高效的模型训练,合理的硬件配置至关重要:
import torch
print(f"PyTorch版本: {torch.__version__}")
print(f"CUDA可用: {torch.cuda.is_available()}")
print(f"GPU型号: {torch.cuda.get_device_name(0)}" if torch.cuda.is_available() else "使用CPU训练")
推荐配置 :
- GPU: NVIDIA RTX 3090 (24GB显存)
- 内存: 32GB以上
- PyTorch: 1.13+
- torchvision: 0.14+
1.2 数据集准备
我们使用Oxford 102 Flowers数据集,包含102类花卉,每类40-258张图像。数据集结构应如下:
flower_data/
├── train/
│ ├── class1/
│ ├── class2/
│ └── ...
└── val/
├── class1/
├── class2/
└── ...
1.3 数据增强策略
针对花朵识别任务,我们设计以下数据增强方案:
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.RandomRotation(30),
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])
])
注意:Normalize参数使用ImageNet的均值和标准差,这与预训练模型的训练设置保持一致
2. 模型构建与迁移学习策略
2.1 ResNet-152模型加载
PyTorch提供了预训练的ResNet-152模型,我们可以直接加载并修改最后一层:
model = models.resnet152(pretrained=True)
num_ftrs = model.fc.in_features
model.fc = nn.Linear(num_ftrs, 102) # 102个花朵类别
2.2 两阶段训练策略
阶段一:冻结特征提取层
for param in model.parameters():
param.requires_grad = False
for param in model.layer4.parameters(): # 解冻最后几个层
param.requires_grad = True
optimizer = optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=0.001)
阶段二:全网络微调
for param in model.parameters():
param.requires_grad = True
optimizer = optim.Adam(model.parameters(), lr=0.0001)
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)
2.3 损失函数与评估指标
criterion = nn.CrossEntropyLoss()
def accuracy(output, target):
_, pred = torch.max(output, 1)
correct = (pred == target).sum().item()
return correct / target.size(0)
3. 训练过程优化
3.1 训练循环实现
def train_model(model, criterion, optimizer, scheduler, num_epochs=25):
best_acc = 0.0
for epoch in range(num_epochs):
print(f'Epoch {epoch}/{num_epochs-1}')
print('-' * 10)
for phase in ['train', 'val']:
if phase == 'train':
model.train()
else:
model.eval()
running_loss = 0.0
running_corrects = 0
for inputs, labels in dataloaders[phase]:
inputs = inputs.to(device)
labels = labels.to(device)
optimizer.zero_grad()
with torch.set_grad_enabled(phase == 'train'):
outputs = model(inputs)
loss = criterion(outputs, labels)
if phase == 'train':
loss.backward()
optimizer.step()
running_loss += loss.item() * inputs.size(0)
running_corrects += torch.sum(torch.argmax(outputs, 1) == labels)
epoch_loss = running_loss / dataset_sizes[phase]
epoch_acc = running_corrects.double() / dataset_sizes[phase]
print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}')
if phase == 'val' and epoch_acc > best_acc:
best_acc = epoch_acc
torch.save(model.state_dict(), 'best_model.pth')
scheduler.step()
print(f'Best val Acc: {best_acc:.4f}')
return model
3.2 关键超参数设置
| 超参数 | 阶段一 | 阶段二 |
|---|---|---|
| 学习率 | 0.001 | 0.0001 |
| Batch Size | 32 | 32 |
| Epochs | 10 | 15 |
| 优化器 | Adam | Adam |
| 学习率衰减 | 无 | StepLR(step=7, γ=0.1) |
4. 结果分析与模型部署
4.1 性能评估
经过两阶段训练后,模型在测试集上的表现:
- Top-1准确率: 92.3%
- Top-5准确率: 98.7%
- 推理速度(3080Ti): 45ms/张
4.2 混淆矩阵分析
from sklearn.metrics import confusion_matrix
import seaborn as sns
def plot_confusion_matrix(cm, classes):
plt.figure(figsize=(20,20))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
xticklabels=classes, yticklabels=classes)
plt.ylabel('真实标签')
plt.xlabel('预测标签')
plt.show()
# 生成混淆矩阵
cm = confusion_matrix(all_labels, all_preds)
plot_confusion_matrix(cm, class_names)
4.3 模型部署示例
使用训练好的模型进行单张图像预测:
def predict(image_path, model, topk=5):
img = Image.open(image_path)
img = val_transform(img).unsqueeze(0)
with torch.no_grad():
output = model(img.to(device))
probs = torch.nn.functional.softmax(output, dim=1)
top_probs, top_classes = probs.topk(topk, dim=1)
return top_probs[0].cpu().numpy(), top_classes[0].cpu().numpy()
probs, classes = predict('test_flower.jpg', model)
for i in range(len(probs)):
print(f"{class_names[classes[i]]}: {probs[i]*100:.2f}%")
在实际项目中,这种迁移学习方法不仅适用于花朵识别,经过简单调整即可应用于各种细粒度图像分类任务,如鸟类识别、商品分类等。关键在于合理设计数据增强策略和分阶段训练方案,以充分利用预训练模型的特征提取能力。
更多推荐




所有评论(0)