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

模型微调策略对比:

策略 训练参数 数据需求 训练速度 适用场景
全冻结 仅最后一层 较少 最快 小数据集,与预训练任务相似
部分微调 后几层+分类器 中等 中等 中等规模数据
全微调 所有参数 大量 最慢 大数据集,任务差异大

在本项目中,我们采用分阶段微调策略:

  1. 先冻结卷积层,只训练分类器(3个epoch)
  2. 解冻所有层,整体微调(10个epoch)
  3. 使用更小的学习率精细调整(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()

性能提升技巧:

  1. 标签平滑(Label Smoothing)

    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)
    
  2. 混合精度训练

    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()
    
  3. 模型集成

    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迁移学习的技术要点,还建立了一套可复用的图像分类流程。这套方法可以轻松扩展到其他细粒度分类任务,如鸟类识别、车辆型号识别等。关键在于根据具体问题调整数据增强策略和微调方法,同时持续监控模型在实际场景中的表现。

Logo

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

更多推荐