在 Ultralytics YOLO 中训练 Fashion-MNIST 图像分类模型:数据集解析与完整实践指南
【免费下载链接】yolov10YOLOv10: Real-Time End-to-End Object Detection [NeurIPS 2024]项目地址: https://gitcode.com/GitHub_Trending/yo/yolov10
导读
Fashion-MNIST 是 Zalando Research 发布的时尚商品图像分类数据集,被设计为经典 MNIST 手写数字数据集的"即插即用"替代品,广泛用于卷积神经网络(CNN)等图像分类模型的训练与基准测试。本指南以 docs/en/datasets/classify/fashion-mnist.md 为骨架,结合当前仓库中ultralytics分类任务的源码实现(数据集加载、训练器、验证器、预测器)与默认配置,系统讲解该数据集的结构、标签体系,以及如何使用 YOLOv8n-cls 分类模型完成从训练、验证到推理的完整闭环,帮助你掌握在 Ultralytics YOLO 框架中跑通一个标准分类数据集的全部技术细节。
Fashion-MNIST 数据集概览
Fashion-MNIST 是 Zalando 商品图像的数据库,由 Zalando Research 发布。它包含60,000 张训练图像和10,000 张测试图像,每张图像都是28×28 像素的灰度图,并关联一个从 10 个类别中选出的标签。
该数据集的设计初衷是作为原始 MNIST 手写数字数据集的**直接替代品(drop-in replacement)**用于机器学习算法基准测试:两者的样本规模、图像尺寸、灰度格式、类别数量(10 类)与任务类型完全一致,但 Fashion-MNIST 的样本内容从手写数字变成了复杂的服装商品图像,从而为算法评估提供了更具区分度的挑战。在 docs/en/datasets/classify/index.md 支持自动下载的分类数据集中,Fashion-MNIST 与 CIFAR-10、ImageNet、MNIST 等并列,可直接通过data='fashion-mnist'触发自动下载。
关键特性
- 共 70,000 张 Zalando 商品图像,其中 60,000 张用于训练,10,000 张用于测试。
- 全部为 28×28 像素的灰度图像。
- 每个像素携带单一像素值,表示该点的明暗程度,数值越大越暗,取值范围为 0 到 255 的整数。
- 是机器学习和图像分类领域广泛使用的训练与测试基准。
数据集结构
Fashion-MNIST 划分为两个子集:
- 训练集(Training Set):包含 60,000 张图像,用于训练机器学习模型。
- 测试集(Testing Set):包含 10,000 张图像,用于测试和基准评估训练完成的模型。
标签体系
每个训练与测试样本都被分配为以下 10 个类别标签之一:
| 标签值 | 类别(英文) | 类别(中文参考) |
|---|---|---|
| 0 | T-shirt/top | T 恤/上衣 |
| 1 | Trouser | 裤子 |
| 2 | Pullover | 套头衫 |
| 3 | Dress | 连衣裙 |
| 4 | Coat | 外套 |
| 5 | Sandal | 凉鞋 |
| 6 | Shirt | 衬衫 |
| 7 | Sneaker | 运动鞋 |
| 8 | Bag | 包 |
| 9 | Ankle boot | 踝靴 |
应用场景
Fashion-MNIST 被广泛用于训练和评估图像分类任务中的深度学习模型,例如卷积神经网络(CNN)、支持向量机(SVM)以及其他各类机器学习算法。得益于其简单而结构规整的格式,该数据集是机器学习和计算机视觉领域研究人员与实践者的基础资源,常被用作:
- 新模型的快速基线测试与算法对比基准;
- 教学与入门实验中替代 MNIST 的进阶选项;
- 数据增强、正则化、超参数调优等训练技巧的快速验证平台。
使用 Ultralytics YOLO 训练分类模型
在 Ultralytics YOLO 中,分类任务的入口是yolov8n-cls等以-cls后缀标识的分类模型。以下代码与命令基于 docs/en/datasets/classify/fashion-mnist.md 中的训练示例,并补充了完整可运行细节。
Python 方式
from ultralytics import YOLO # 加载模型(推荐加载预训练模型作为起点) model = YOLO('yolov8n-cls.pt') # 在 Fashion-MNIST 上训练 100 个 epoch,图像尺寸 28x28 results = model.train(data='fashion-mnist', epochs=100, imgsz=28)CLI 方式
# 从预训练 *.pt 模型开始训练 yolo classify train data=fashion-mnist model=yolov8n-cls.pt epochs=100 imgsz=28注意:原始文档中 CLI 示例使用的是yolo detect train,而在当前仓库的分类任务中,正确的任务关键字是classify(见 ultralytics/models/yolo/classify/train.py 中ClassificationTrainer将task强制设为"classify")。关于全部可用的训练参数,请参阅 训练模式文档。
数据自动下载与目录组织
当data='fashion-mnist'时,Ultralytics 会自动下载数据集,并将其组织为 torchvisionImageFolder风格的标准分类目录结构。按 docs/en/datasets/classify/index.md 所述,分类数据集的目录格式为:
root/ |-- class1/ | |-- img1.jpg | |-- img2.jpg | |-- ... |-- class2/ | |-- img1.jpg | |-- img2.jpg | |-- ... |-- class3/ | |-- img1.jpg | |-- img2.jpg | |-- ... |-- ...即root目录下为每个类别建立一个以类别名命名的子目录,子目录内是该类别的全部图像。如果你有自定义数据集,只需按此格式组织目录,再将data参数指向数据集目录即可无缝复用同一套训练流程。
源码级解析:分类数据流与训练闭环
当前仓库的ultralytics代码库完整实现了分类任务的数据加载、训练、验证与预测链路。理解这些实现有助于你更好地调参和排查问题。
数据加载:基于 torchvision ImageFolder 的扩展
分类数据集由 ultralytics/data/dataset.py 中的ClassificationDataset类承载,它直接继承自torchvision.datasets.ImageFolder,因此天然支持按类目子目录组织的数据集。该类的核心职责包括:
- 图像校验与缓存:
verify_images()会扫描全部图像,过滤损坏样本,并将校验结果写入.cache文件,后续加载时通过文件哈希判断缓存是否有效,从而加速重复训练。 - 缓存加速:支持
cache=True/ram(图像缓存进内存)与cache='disk'(图像以无压缩.npy文件缓存到磁盘)两种方式,减少训练时 IO 开销。 - 数据增强:训练模式下通过
classify_augmentations()应用缩放、水平/垂直翻转、随机擦除(erasing)、HSV 扰动与可选的auto_augment(randaugment / augmix / autoaugment);验证/测试模式则通过classify_transforms()仅做缩放与中心裁剪。 - 类别采样:当
args.fraction < 1.0时可截取训练样本的前 N% 用于快速实验(见 ultralytics/data/dataset.py)。
在 ultralytics/data/utils.py 中,分类数据集的统计信息同样通过torchvision.datasets.ImageFolder获取每个 split(train/val/test)的类别分布,用于输出数据集统计报告。
训练器:ClassificationTrainer
ultralytics/models/yolo/classify/train.py 中的ClassificationTrainer继承自BaseTrainer,承担分类任务的完整训练编排,值得关注的设计点包括:
- 默认输入尺寸:构造函数中若未显式指定
imgsz,默认设为 224(而非检测任务的 640),因此 Fashion-MNIST 的 28×28 图像在使用imgsz=28时会被直接缩放送入模型。 - 模型来源的三种路径:
setup_model()支持从本地*.pt权重加载、从*.yaml配置构建、或直接传入 torchvision 内置模型名(如resnet18)加载 ImageNet 预训练权重;最后统一通过ClassificationModel.reshape_outputs()将输出头调整为当前数据集的类别数nc。 - Dropout 正则化:当配置了
dropout参数时,会为模型中的所有torch.nn.Dropout层设置丢弃概率(见 ultralytics/models/yolo/classify/train.py),这对在 Fashion-MNIST 这种规模较小的数据集上抑制过拟合很有帮助。 - 类别名同步:
set_model_attributes()将数据集加载出的类别名写入模型names,保证预测与验证输出的标签可读。
验证器:ClassificationValidator 与评估指标
ultralytics/models/yolo/classify/val.py 中的ClassificationValidator定义了分类任务的评估流程,其核心指标是top-1 与 top-5 准确率(见get_desc()中打印的top1_acc/top5_acc列)。验证阶段会:
- 对每个 batch 取每个样本的 top-5 预测(
n5 = min(len(self.names), 5))累积预测与目标; - 构建分类混淆矩阵
ConfusionMatrix,并在plots=True时输出归一化与未归一化两种混淆矩阵图(ultralytics/models/yolo/classify/val.py),方便分析 Fashion-MNIST 中容易混淆的类别(如 T-shirt/top 与 Shirt、Pullover 与 Coat); - 训练结束后自动绘制训练样本批次图、验证标签图与预测图,保存到训练输出目录。
推理:ClassificationPredictor
训练完成后,可使用 ultralytics/models/yolo/classify/predict.py 中的ClassificationPredictor对任意图像执行分类推理,结果以Results对象携带类别概率(probs)返回。一个典型的 Python 推理示例如下:
from ultralytics import YOLO model = YOLO('runs/classify/train/weights/best.pt') # 加载训练好的模型 results = model('path/to/image.jpg') # 对图像分类 print(results[0].probs.top1, results[0].names[results[0].probs.top1]) # top-1 类别CLI 对应命令为:
yolo classify predict model=runs/classify/train/weights/best.pt source='path/to/image.jpg'关键训练参数速查
以下参数均可在 ultralytics/cfg/default.yaml 中找到默认值,并可通过 Python 关键字参数或 CLI 覆盖:
| 参数 | 默认值 | 说明 |
|---|---|---|
epochs | 100 | 训练轮数 |
imgsz | 224(分类默认) | 输入图像尺寸,Fashion-MNIST 建议 28 |
batch | 16 | 每批图像数(-1 启用 AutoBatch) |
cache | False | True/ram/disk,缓存图像加速训练 |
device | 空 | 运行设备,如0、0,1,2,3、cpu |
workers | 8 | 数据加载线程数 |
pretrained | True | 是否使用预训练权重 |
dropout | 0.0 | 分类任务专用的 Dropout 概率 |
optimizer | auto | 优化器,可选 SGD、Adam、AdamW 等或 auto 自动选择 |
fraction | 1.0 | 训练集使用比例,可用于快速试跑 |
cos_lr | False | 是否使用余弦学习率调度 |
resume | False | 是否从最近 checkpoint 恢复训练 |
amp | True | 是否启用自动混合精度训练 |
plots | True | 是否保存训练/验证过程图 |
在 Fashion-MNIST 这类 70K 小图上,合理组合dropout(如 0.1)、cache=True(数据量小可直接缓存进内存)以及fraction快速试跑,可以显著提升实验迭代效率。
总结
Fashion-MNIST 凭借与 MNIST 完全一致的数据形态和更具语义复杂度的服装图像,成为评估图像分类算法的理想基准。在 Ultralytics YOLO 框架下,通过data='fashion-mnist'一条指令即可完成数据集自动下载,配合yolov8n-cls预训练分类模型,仅需几行代码就能跑通训练、验证与推理全流程。结合ClassificationDataset的缓存与校验机制、ClassificationTrainer的多源模型加载与 Dropout 支持,以及ClassificationValidator的 top-1/top-5 与混淆矩阵评估,你可以快速在标准基准上验证模型设计,并将同一套流程平滑迁移到自己的自定义分类数据集上。
致谢
如果你在研究或开发工作中使用了 Fashion-MNIST 数据集,请通过其 GitHub 仓库(由 Zalando Research 提供)标注数据来源。该数据集由 Zalando Research 制作发布,感谢其为机器学习社区提供的宝贵资源。
【免费下载链接】yolov10YOLOv10: Real-Time End-to-End Object Detection [NeurIPS 2024]项目地址: https://gitcode.com/GitHub_Trending/yo/yolov10
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考