多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

数据集采集与微调:小样本训练专属抓取识别模型

数据集采集与微调:小样本训练专属抓取识别模型 数据集采集与微调小样本训练专属抓取识别模型通用 YOLO 模型能识别猫狗汽车但你的机械臂要抓的可能是某个特定型号的螺丝——这时候就得自己动手训练专属模型了。一、为什么通用模型不够用直接下个 YOLOv8 预训练模型跑起来确实能识别 80 类物体COCO 数据集但放到具身智能场景里立刻露馅特定物体不在类别里COCO 里有apple但没有红富士苹果有bottle但没有你家工厂产的 250ml 酱油瓶视角差异大COCO 的图多为平视视角机械臂相机是俯视/斜视背景特定工业场景背景是传送带/操作台和 COCO 的生活场景差太远小目标多螺丝、电子元件这种小目标通用模型召回率惨不忍睹所以必须训练专属模型。但问题是——标注数据哪来标注 10000 张图人工成本扛不住这就要靠小样本训练策略了。二、数据采集方案2.1 机械臂自动采集推荐最省事的方式让机械臂搭载相机自动绕着物体多角度拍摄。写个脚本控制机械臂走到 N 个预设位姿每个位姿拍一张importcv2importnumpyasnp CAPTURE_POSES[# [x, y, z, rx, ry, rz] 末端位姿[0.20,0.00,0.35,0,45,0],[0.20,0.10,0.35,0,45,0],[0.20,-0.10,0.35,0,45,0],[0.15,0.00,0.25,0,60,0],# ... 多角度、多距离、多俯仰]defauto_capture(arm,camera,save_dir):fori,poseinenumerate(CAPTURE_POSES):arm.move_to_pose(pose)arm.wait_stop()imgcamera.capture()cv2.imwrite(f{save_dir}/img_{i:04d}.jpg,img)这种方式的优点是位姿精确可控、可重复而且能自动记录每张图对应的相机位姿后续做标定也能用上。2.2 手动采集没有机械臂的话用手机/相机绕物体拍一圈。注意几点背景尽量贴近真实使用场景别在白墙前拍完再去仓库用光照要覆盖明暗两种条件物体姿态要多样正放/侧放/倒放2.3 数据增强扩充样本50 张原图通过数据增强可以扩充到几百张对小样本训练帮助极大fromalbumentationsimport(Compose,HorizontalFlip,RandomBrightnessContrast,Rotate,GaussianBlur,HueSaturationValue)augmentCompose([HorizontalFlip(p0.5),Rotate(limit15,p0.5),RandomBrightnessContrast(brightness_limit0.2,contrast_limit0.2,p0.5),GaussianBlur(blur_limit(3,5),p0.3),HueSaturationValue(hue_shift_limit10,sat_shift_limit20,p0.3),])defaugment_dataset(img,bboxes):augmentedaugment(imageimg,bboxesbboxes)returnaugmented[image],augmented[bboxes]注意旋转和翻转时bbox 坐标要同步变换Albumentations 库会自动处理但用 OpenCV 手写的话容易翻车。三、标注工具与格式3.1 标注工具对比工具平台优点缺点LabelImg桌面简单轻量、离线可用界面古老、无自动标注Roboflow网页在线协作、自动增强免费版有限制、数据要上传CVAT网页开源、功能强大、支持视频部署稍麻烦Label Studio网页多模态标注、开源配置复杂新手单人开发推荐 LabelImg团队协作选 CVAT 或 Roboflow。3.2 YOLO 标注格式YOLO 格式每张图对应一个.txt文件每行一个目标class x_center y_center width height坐标都是归一化到[0, 1]的相对值除以图片宽高。例如0 0.452 0.318 0.123 0.087 0 0.621 0.502 0.098 0.102 1 0.234 0.711 0.156 0.143class是类别 ID从 0 开始x_center y_center是目标框中心点坐标width height是框的宽高3.3 数据集目录结构YOLOv8 训练要求的数据集结构my_dataset/ ├── images/ │ ├── train/ ← 训练集图片 │ │ ├── img_0001.jpg │ │ └── img_0002.jpg │ └── val/ ← 验证集图片 │ └── img_0003.jpg ├── labels/ │ ├── train/ ← 训练集标注文件名和图片一一对应 │ │ ├── img_0001.txt │ │ └── img_0002.txt │ └── val/ │ └── img_0003.txt └── data.yaml ← 数据集配置data.yaml内容path:./my_datasettrain:images/trainval:images/valnc:2# 类别数names:[screw,nut]# 类别名四、小样本训练策略4.1 迁移学习这是小样本训练的核心。YOLOv8 预训练模型在 COCO 上学到的特征边缘、纹理、形状具有通用性我们只需要在新数据上微调一下让它认识新类别即可。Ultralytics 默认就是迁移学习——加载yolov8n.pt预训练权重开始训练yolo detect train\datadata.yaml\modelyolov8n.pt\# 预训练权重n/s/m/l/x 选型epochs100\imgsz640\batch16\lr00.01\freeze10# 冻结前 10 层4.2 冻结骨干网络freeze10表示冻结骨干网络backbone的前 10 层只训练检测头和部分颈部。骨干网络参数不更新训练参数量大幅减少50 张图也能跑出可用模型。冻结合数要根据样本数调整50 张以下冻结全部骨干freeze2250~200 张冻结前 10 层200~1000 张冻结前 5 层或不冻结4.3 数据增强加倍小样本训练最怕过拟合——模型把训练集背下来了测试集一塌糊涂。解决方法是训练时在线增强# hyp.yaml 训练超参hsv_h:0.015# 色调增强hsv_s:0.7# 饱和度增强hsv_v:0.4# 明度增强degrees:10.0# 旋转角度translate:0.1# 平移比例scale:0.5# 缩放比例fliplr:0.5# 水平翻转概率mosaic:1.0# Mosaic 增强4 图拼接超有用mixup:0.1# Mixup 增强mosaic增强对小目标特别有效把 4 张图拼成 1 张小目标比例被放大模型学得更准。五、训练环境搭建PC 端 GPU 训练环境# 1. 装 CUDA按显卡型号选版本# 2. 装 PyTorch带 CUDA 支持pipinstalltorch torchvision --index-url https://download.pytorch.org/whl/cu118# 3. 装 Ultralyticspipinstallultralytics# 4. 验证 GPU 可用python-cimport torch; print(torch.cuda.is_available())# 应输出 True没有 NVIDIA 显卡的话用 CPU 也能训练但速度慢 10 倍以上建议租云 GPUAutoDL、Colab 都行。六、训练参数配置详解yolo detect train\datadata.yaml\modelyolov8n.pt\epochs100\# 训练轮数小样本 100~300 足够batch16\# 批大小显存不够就降最小到 4imgsz640\# 输入尺寸640 够用显存富裕可上 1024lr00.01\# 初始学习率迁移学习用默认即可lrf0.01\# 最终学习率衰减系数patience20\# 早停20 轮无提升就停防过拟合weight_decay0.0005\# 权重衰减正则化防过拟合freeze10\# 冻结层数projectruns/train\# 输出目录nameexp1参数调优建议学习率loss 降不下去就调大0.02loss 抖动剧烈就调小0.005batch size显存允许越大越好但小样本别超 32否则一个 epoch 步数太少epochs看val/loss不降了就停别硬训过拟合反而掉点七、训练过程监控训练开始后关注这几个指标Epoch GPU_mem box_loss cls_loss dfl_loss Instances Size 1/100 3.52G 1.234 0.892 0.956 42 640 2/100 3.55G 0.845 0.621 0.812 38 640 ...box_loss边界框回归损失应该持续下降cls_loss分类损失应该持续下降dfl_loss分布焦点损失波动正常训练结束后看runs/train/exp1/results.png里面是所有指标的曲线图。重点看train/loss和val/loss都在下降 → 正常train/loss下降但val/loss上升 →过拟合该早停或加增强八、模型评估指标解读训练完会输出一堆指标新手容易看晕指标含义解读Precision精确率预测为正的样本里多少是对的查准率Recall召回率真实为正的样本里多少被找出来了查全率mAP50IoU0.5 时的平均精度偏宽松主要看这个mAP50-95IoU 从 0.5 到 0.95 的平均偏严格反映定位精度举个机械臂要抓螺丝Precision 高 Recall 低 → 漏抓多有些螺丝没识别到Precision 低 Recall 高 → 误抓多把螺母也当螺丝抓了。抓取场景一般优先保证 Recall漏抓比误抓更影响流程。九、模型导出与部署训练完的.pt模型要在 RK3588 上跑需要经过三步转换9.1 PyTorch → ONNXyoloexportmodelruns/train/exp1/weights/best.ptformatonnxsimplifyTrueopset12simplifyTrue会用 onnx-simplifier 简化图结构减少算子数量对后续 RKNN 转换更友好。9.2 ONNX → RKNN在 PC 上用 RKNN-Toolkit2 转换需要安装rknn-toolkit2fromrknn.apiimportRKNN rknnRKNN()# 配置量化参数rknn.config(mean_values[[0,0,0]],std_values[[255,255,255]],target_platformrk3588,quantized_dtypew8a8,# INT8 量化quantized_methodchannel,# 逐通道量化optimization_level3)# 加载 ONNX 模型rknn.load_onnx(modelbest.onnx)# 构建模型需要校准数据集做量化rknn.build(do_quantizationTrue,dataset./dataset.txt)# 导出 RKNN 模型rknn.export_rknn(best.rknn)dataset.txt是量化校准图片列表放 100~500 张代表性图片路径即可。9.3 RK3588 上推理把best.rknn拷到开发板用 rknn-lite API 推理fromrknnlite.apiimportRKNNLite rknnRKNNLite()rknn.load_rknn(best.rknn)rknn.init_runtime(core_maskRKNNLite.NPU_CORE_0_1_2)# 三核并行outputsrknn.infer(inputs[img_array])NPU 三核并行YOLOv8n 单帧推理能压到 15ms 以内30 FPS 没压力。十、Python 训练脚本整合不想敲命令行的话用 Python 脚本更灵活fromultralyticsimportYOLOdeftrain_custom_model():modelYOLO(yolov8n.pt)# 加载预训练模型resultsmodel.train(datadata.yaml,epochs100,batch16,imgsz640,lr00.01,freeze10,patience20,projectruns/train,namescrew_detect,device0,# GPU)# 验证metricsmodel.val()print(fmAP50:{metrics.box.map50:.4f})print(fmAP50-95:{metrics.box.map:.4f})# 导出 ONNXmodel.export(formatonnx,simplifyTrue)if__name____main__:train_custom_model()十一、结语小样本训练专属模型的核心就三招迁移学习 冻结骨干 数据增强。50 张图能训出可用模型500 张能上生产关键不在数据量而在数据质量和训练策略。标注别偷懒增强别过头指标别迷信实测最靠谱——拿机械臂实际跑 100 次抓取成功率才是硬道理。
返回列表