102花卉数据集实战指南:从数据加载到模型训练的全流程解析

1次阅读
没有评论

共计 2358 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

数据集背景

102 花卉数据集(Oxford 102 Flowers)是一个经典的细粒度图像分类数据集,包含 102 种英国常见花卉类别。每个类别有 40 到 258 张图像,总计 8189 张。图像尺寸不固定,但多在 500×500 像素左右。该数据集常用于验证图像分类算法的性能,特别是在迁移学习和细粒度分类任务中。

102 花卉数据集实战指南:从数据加载到模型训练的全流程解析

痛点分析

新手在处理这个数据集时经常会遇到以下问题:

  • 数据加载慢:原始图像尺寸较大,直接加载会消耗大量内存
  • 类别不均衡:各类别样本数量差异明显(最多 / 最少相差 6 倍)
  • 预处理复杂:需要统一图像尺寸并做标准化
  • 数据增强选择困难:不清楚哪些变换对花卉图像有效
  • 标签处理混乱:数据集提供的 matlab 格式标签需要特殊处理

技术方案

1. 构建高效数据管道

使用 PyTorch 的 Dataset 和 DataLoader 可以很好地解决数据加载问题:

from torch.utils.data import Dataset
from PIL import Image
import numpy as np

class Flowers102Dataset(Dataset):
    def __init__(self, img_paths, labels, transform=None):
        self.img_paths = img_paths
        self.labels = labels
        self.transform = transform

    def __len__(self):
        return len(self.img_paths)

    def __getitem__(self, idx):
        img = Image.open(self.img_paths[idx]).convert('RGB')
        label = self.labels[idx]

        if self.transform:
            img = self.transform(img)

        return img, label

2. 完整预处理代码

典型预处理包含 resize、裁剪、归一化等操作:

from torchvision import transforms

train_transform = transforms.Compose([transforms.RandomResizedCrop(224),  # 随机裁剪并 resize 到 224x224
    transforms.RandomHorizontalFlip(),  # 随机水平翻转
    transforms.ColorJitter(0.2, 0.2, 0.2),  # 颜色扰动
    transforms.ToTensor(),  # 转为 Tensor
    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])  # ImageNet 均值方差
])

val_transform = transforms.Compose([transforms.Resize(256),  # 先缩放到 256x256
    transforms.CenterCrop(224),  # 中心裁剪 224x224
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])

3. 迁移学习实现

使用预训练 ResNet18 模型进行微调:

import torch.nn as nn
import torchvision.models as models

# 加载预训练模型
model = models.resnet18(pretrained=True)

# 冻结除最后一层外的所有参数
for param in model.parameters():
    param.requires_grad = False

# 替换最后的全连接层
num_features = model.fc.in_features
model.fc = nn.Linear(num_features, 102)  # 102 个类别 

避坑指南

处理类别不均衡

两种实用方法:

  1. 过采样少数类:使用 WeightedRandomSampler
  2. 损失函数加权:nn.CrossEntropyLoss(weight=class_weights)

内存优化

如果内存不足,可以考虑:

  • 使用 LMDB 等高效存储格式
  • 降低 batch size(如从 64 降到 32)
  • 使用梯度累积技术

数据增强参数

对于花卉图像,建议:

  • 旋转角度不超过 30 度(避免花蕊变形)
  • 亮度调节范围 0.7-1.3
  • 谨慎使用垂直翻转(花朵通常朝上)

效果验证

评估指标

经过 20 个 epoch 训练后,典型结果:

  • Top- 1 准确率:约 85%
  • Top- 5 准确率:约 95%

可视化训练过程

使用 TensorBoard 记录训练曲线:

from torch.utils.tensorboard import SummaryWriter

writer = SummaryWriter()

for epoch in range(epochs):
    # ... 训练代码...
    writer.add_scalar('Loss/train', train_loss, epoch)
    writer.add_scalar('Accuracy/train', train_acc, epoch)
    writer.add_scalar('Loss/val', val_loss, epoch)
    writer.add_scalar('Accuracy/val', val_acc, epoch)

延伸思考

  1. 尝试不同的预训练模型:
  2. EfficientNet 通常表现更好
  3. Vision Transformer 也可以尝试

  4. 数据增强策略优化:

  5. 添加 MixUp 或 CutMix
  6. 尝试 AutoAugment 策略

  7. 更精细的微调策略:

  8. 分层解冻(先解冻最后几层)
  9. 使用不同的学习率

经过这套流程,你应该能够快速搭建一个 baseline 模型。记住在实际项目中,数据质量往往比模型选择更重要,所以要多花时间在数据分析和预处理上。

正文完
 0
评论(没有评论)