手把手教你用YOLOv5s训练自己的水果识别模型(附2611张标注数据集)
手把手教你用YOLOv5s训练自己的水果识别模型(附2611张标注数据集)
在计算机视觉领域,目标检测技术正以前所未有的速度改变着我们与数字世界的交互方式。对于想要快速入门AI应用开发的实践者来说,YOLOv5无疑是最友好的选择之一——它既保留了YOLO系列"一次看全"的实时检测优势,又大幅降低了模型训练门槛。本文将带你完整走通从数据集准备到模型部署的全流程,特别针对水果识别这个经典场景,使用包含6大类、2611张精细标注图片的现成数据集(已转换为YOLO格式),让你避开数据收集和标注的"深坑",直接进入模型训练的核心环节。
1. 开发环境配置与数据准备
1.1 基础环境搭建
推荐使用Python 3.8+和PyTorch 1.7+的组合,这是经过验证与YOLOv5兼容性最好的版本搭配。以下是一键配置命令:
# 创建并激活虚拟环境 conda create -n yolov5_fruit python=3.8 -y conda activate yolov5_fruit # 安装PyTorch(根据CUDA版本选择) pip install torch==1.8.1+cu111 torchvision==0.9.1+cu111 -f https://download.pytorch.org/whl/torch_stable.html # 克隆YOLOv5官方仓库 git clone https://github.com/ultralytics/yolov5 cd yolov5 pip install -r requirements.txt提示:如果使用Colab等云端环境,只需执行最后两个命令即可。建议选择T4或V100显卡以获得最佳训练速度。
1.2 数据集结构解析
我们提供的2611张水果数据集已按标准YOLO格式组织,包含以下6个类别:
- 黄冠苹果(golden delicious)
- 青苹果(granny smith)
- 梨(pear)
- 红富士苹果(red delicious)
- 红油桃(red nectarine)
- 黄桃(yellow peach)
数据集目录结构如下:
fruit_dataset_yolo/ ├── train/ │ ├── images/ # 存放训练集图片 │ └── labels/ # 存放对应的YOLO格式标签 ├── val/ │ ├── images/ # 验证集图片 │ └── labels/ └── test/ ├── images/ # 测试集图片 └── labels/每个.txt标签文件的格式示例:
0 0.543201 0.501212 0.123456 0.234567 # 类别索引 中心x 中心y 宽度 高度1.3 数据增强策略
在data/augmentation.yaml中添加以下增强配置可显著提升模型鲁棒性:
# 增强参数配置 hsv_h: 0.015 # 色调增强幅度 hsv_s: 0.7 # 饱和度增强范围 hsv_v: 0.4 # 明度增强范围 degrees: 15 # 旋转角度范围 translate: 0.1 # 平移比例 scale: 0.5 # 缩放幅度 shear: 0.0 # 剪切幅度 perspective: 0.0005 # 透视变换系数 flipud: 0.0 # 上下翻转概率 fliplr: 0.5 # 左右翻转概率 mosaic: 1.0 # mosaic增强概率 mixup: 0.1 # mixup增强概率2. 模型训练与调优实战
2.1 配置文件定制
创建data/fruits.yaml定义数据集路径和类别信息:
# 训练/验证集路径 train: ../fruit_dataset_yolo/train/images val: ../fruit_dataset_yolo/val/images # 类别数量 nc: 6 # 类别名称列表 names: ['golden delicious', 'granny smith', 'pear', 'red delicious', 'red nectarine', 'yellow peach']2.2 启动模型训练
使用以下命令开始训练YOLOv5s模型(小型版本适合快速迭代):
python train.py --img 640 --batch 16 --epochs 100 --data data/fruits.yaml \ --cfg models/yolov5s.yaml --weights yolov5s.pt --name fruit_detection \ --cache ram --device 0 --optimizer AdamW --patience 15关键参数说明:
| 参数 | 作用 | 推荐值 |
|---|---|---|
| --img | 输入图像尺寸 | 640 |
| --batch | 批次大小 | 根据显存调整(8-32) |
| --epochs | 训练轮次 | 50-300 |
| --weights | 预训练权重 | yolov5s.pt |
| --cache | 数据缓存方式 | ram/disk |
| --device | 训练设备 | 0(第一块GPU) |
| --optimizer | 优化器选择 | SGD/AdamW |
| --patience | 早停耐心值 | 10-20 |
2.3 训练监控与调优
训练过程中可以通过TensorBoard实时监控指标:
tensorboard --logdir runs/train常见问题解决方案:
- 过拟合:增加--patience值,添加--label-smoothing 0.1参数
- 显存不足:减小--batch-size,启用--multi-scale训练
- 类别不平衡:使用--class-weights参数自动计算权重
- 收敛慢:尝试--optimizer AdamW --lr0 0.001组合
3. 模型评估与性能分析
3.1 关键指标解读
训练完成后,在runs/train/exp*/目录下会生成结果文件:
results.png # 训练过程指标曲线 confusion_matrix.png # 混淆矩阵 val_batchX_labels.jpg # 验证集预测示例重点关注以下指标:
- mAP@0.5:IoU阈值为0.5时的平均精度
- mAP@0.5:0.95:IoU从0.5到0.95的平均精度
- Precision:预测为正样本中真实正样本比例
- Recall:真实正样本中被正确预测的比例
3.2 测试集验证
使用训练好的最佳模型进行最终测试:
python val.py --data data/fruits.yaml --weights runs/train/fruit_detection/weights/best.pt \ --task test --imgsz 640 --device 0 --save-json --save-conf输出示例:
Class Images Instances P R mAP@.5 mAP@.5:.95 all 328 891 0.92 0.89 0.91 0.68 golden 328 142 0.94 0.93 0.95 0.72 granny 328 138 0.91 0.88 0.89 0.65 ...3.3 可视化检测效果
使用detect.py脚本快速验证模型:
python detect.py --weights runs/train/fruit_detection/weights/best.pt \ --source ../fruit_dataset_yolo/test/images --save-txt --save-conf对于难例分析,可以添加--augment参数启用测试时增强:
python detect.py --weights best.pt --source test.jpg --augment4. 模型部署与应用开发
4.1 Flask API服务搭建
创建app.py实现简单的推理API:
from flask import Flask, request, jsonify import torch from PIL import Image import io app = Flask(__name__) model = torch.hub.load('ultralytics/yolov5', 'custom', path='runs/train/fruit_detection/weights/best.pt') @app.route('/predict', methods=['POST']) def predict(): if 'file' not in request.files: return jsonify({'error': 'No file uploaded'}), 400 file = request.files['file'] img_bytes = file.read() img = Image.open(io.BytesIO(img_bytes)) results = model(img, size=640) return jsonify(results.pandas().xyxy[0].to_dict(orient='records')) if __name__ == '__main__': app.run(host='0.0.0.0', port=5000)启动服务:
python app.py4.2 移动端集成方案
使用ONNX格式转换实现跨平台部署:
python export.py --weights runs/train/fruit_detection/weights/best.pt \ --include onnx --img 640 --device 0 --simplify在Android端可通过以下代码加载模型:
// 初始化ONNX运行时 val env = OrtEnvironment.getEnvironment() val sessionOptions = new OrtSession.SessionOptions() val session = env.createSession("fruit_detection.onnx", sessionOptions) // 准备输入数据 float[][][][] inputData = preprocessImage(bitmap) // 图像预处理 OnnxTensor inputTensor = OnnxTensor.createTensor(env, inputData) // 执行推理 try (OrtSession.Result results = session.run(Collections.singletonMap("images", inputTensor))) { float[][][] output = (float[][][]) results.get(0).getValue() // 解析检测结果 processOutput(output, confidenceThreshold=0.5) }4.3 性能优化技巧
量化压缩:将FP32模型转为INT8提升推理速度
python export.py --weights best.pt --include onnx --int8 --device 0TensorRT加速:
python export.py --weights best.pt --include engine --device 0多线程处理:
from threading import Lock model_lock = Lock() def threaded_predict(img): with model_lock: return model(img)