别再用目标检测的YOLOv5了!手把手教你用它的分类模块(v6.2+)搞定图片分类

当提到YOLOv5时,大多数开发者脑海中浮现的往往是它在目标检测任务中的卓越表现。但鲜为人知的是,从v6.2版本开始,这个强大的框架已经悄然内置了完整的图像分类功能。对于已经熟悉YOLOv5生态的开发者来说,这无疑是一条通往图像分类任务的捷径——无需切换技术栈,就能复用已有的知识和工具链。

1. 为什么选择YOLOv5做图像分类?

在深度学习领域,ResNet、EfficientNet等经典分类网络早已深入人心。但YOLOv5的分类模块却有几个独特的优势:

  • 技术栈统一:如果你已经在使用YOLOv5进行目标检测,那么分类任务可以共享同一套环境配置、数据预处理流程和部署工具
  • 性能平衡:YOLOv5的主干网络经过目标检测任务的锤炼,在速度和精度之间取得了很好的平衡
  • 迁移成本低:从检测转向分类时,90%的工程经验可以直接复用

实测对比(ImageNet-1k验证集):

模型参数量(M)Top-1 Acc(%)推理速度(ms)
ResNet5025.576.18.2
EfficientNet-B05.377.16.8
YOLOv5s-cls7.275.35.4

提示:YOLOv5的分类模型在保持相当精度的同时,展现了更优的推理速度,这对边缘设备部署尤为重要。

2. YOLOv5分类模块架构解析

YOLOv5的分类网络实际上只使用了其目标检测网络的主干部分(backbone),去掉了检测专用的头部结构。这种设计带来了几个关键特性:

2.1 核心组件拆解

  1. CBS模块:Conv-BatchNorm-SiLU的基本组合,构成了网络的基础单元
  2. C3模块:借鉴了Cross Stage Partial网络的思想,通过残差连接增强特征复用
  3. SPPF层:空间金字塔池化结构,能够处理不同尺寸的输入
# 典型的YOLOv5分类网络结构(简化版)
Model(
  (model): Sequential(
    (0): Conv(3, 64, kernel_size=6, stride=2, padding=2)
    (1): Conv(64, 128, kernel_size=3, stride=2)
    (2): C3(128, 128, n=3)
    (3): Conv(128, 256, kernel_size=3, stride=2)
    (4): C3(256, 256, n=6)
    (5): Conv(256, 512, kernel_size=3, stride=2)
    (6): C3(512, 512, n=9)
    (7): SPPF(512, 512, k=5)
    (8): Linear(512, num_classes)
  )
)

2.2 与检测模型的区别

  • 去除了PANet特征金字塔结构
  • 移除了锚框相关的预测头
  • 最终输出层改为全连接分类器
  • 默认输入尺寸从640x640调整为224x224

3. 从零开始训练分类模型

3.1 数据准备规范

YOLOv5分类模块要求数据按照特定结构组织:

dataset_root/
├── train/
│   ├── class1/
│   │   ├── img1.jpg
│   │   └── img2.jpg
│   └── class2/
│       ├── img1.jpg
│       └── img2.jpg
└── val/
    ├── class1/
    └── class2/

关键参数配置

# data.yaml示例
train: ../datasets/mydata/train
val: ../datasets/mydata/val
nc: 10  # 类别数
names: ['cat', 'dog', ...]  # 类别名称

3.2 训练流程详解

启动训练的基本命令:

python classify/train.py \
  --model yolov5s-cls.pt \
  --data mydata \
  --epochs 100 \
  --img 224 \
  --batch-size 64 \
  --pretrained

常见调参技巧

  • 学习率策略:使用--lr0设置初始学习率,--lrf设置最终学习率
  • 数据增强:通过--augment启用,可自定义增强策略
  • 混合精度:--amp可加速训练并减少显存占用
  • 多GPU训练:添加--device 0,1参数

注意:首次运行时会自动下载预训练权重,若网络环境受限可手动下载后放置到指定位置。

4. 模型评估与部署实战

4.1 性能评估方法

YOLOv5提供了完整的评估工具:

# 验证集评估
python classify/val.py \
  --weights runs/train-cls/exp/weights/best.pt \
  --data mydata
  
# 单张图片测试
python classify/predict.py \
  --weights runs/train-cls/exp/weights/best.pt \
  --source test.jpg

关键评估指标

  • Top-1 Accuracy:预测最可能类别正确的比例
  • Top-5 Accuracy:预测前五个可能类别中包含正确答案的比例
  • 混淆矩阵:可视化各类别间的误判情况

4.2 生产环境部署方案

  1. ONNX导出
import torch
model = torch.load('best.pt', map_location='cpu')['model'].float()
model.eval()
torch.onnx.export(model, torch.randn(1,3,224,224), 'model.onnx')
  1. TensorRT加速
trtexec --onnx=model.onnx \
  --saveEngine=model.engine \
  --fp16 \
  --workspace=2048
  1. 移动端部署
    • 使用LibTorch在Android/iOS上运行
    • 转换为CoreML格式适配Apple生态
    • 转换为TFLite格式部署到边缘设备

5. 进阶技巧与性能优化

5.1 数据增强策略

YOLOv5分类模块内置了两种增强方式:

  1. Albumentations增强

    • 色彩抖动
    • 随机旋转
    • 网格畸变
    • 随机光照
  2. Torchvision增强

    • 随机水平翻转
    • 标准化
    • 随机裁剪

自定义增强示例

# 在classify/train.py中添加
self.album_transforms = A.Compose([
    A.RandomBrightnessContrast(p=0.5),
    A.RGBShift(p=0.3),
    A.GaussNoise(p=0.1),
])

5.2 模型压缩技巧

  1. 知识蒸馏

    • 使用更大的YOLOv5-cls模型作为教师模型
    • 通过KL散度损失传递知识
  2. 量化部署

    • 动态量化(8bit整数)
    • 静态量化(校准后量化)
# 动态量化示例
quantized_model = torch.quantization.quantize_dynamic(
    model, {torch.nn.Linear}, dtype=torch.qint8
)

在实际项目中,我发现YOLOv5的分类模块特别适合需要同时处理检测和分类任务的场景。比如在一个工业质检系统中,我们可以用检测模型定位缺陷位置,再用分类模型判断缺陷类型,两者共享大部分基础设施,显著降低了部署复杂度。

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐