
使用Python语言 深度学习框架YOLOv8模型训练地铁安检数据集 通过训练的地铁铁路X光射线安检检测数据集的权重推理识别刀枪锤子打火机充电宝等文章目录使用Python语言 深度学习框架YOLOv8模型训练地铁安检数据集 通过训练的地铁铁路X光射线安检检测数据集的权重推理识别刀枪锤子打火机充电宝等✅ 一、环境准备Python YOLOv81. 安装 Anaconda可选但推荐2. 创建虚拟环境并安装依赖 二、数据集结构data.yaml 文件内容 三、训练模型train.py 四、推理检测detect.py 五、模型评估evaluate.py 六、如何调用训练好的模型封装函数 七、训练结果可视化可选数据集描述地铁铁路X光安检检测数据集12类地铁铁路X光安检检测数据集的表格类别英文名称训练集图片数验证集图片数测试集图片数警棍Baton---子弹Bullet---枪Gun---刀Knife---锤子Hammer---手铐HandCuffs---打火机Lighter---喷雾器Sprayer---钳子Pliers---充电宝Powerbank---剪刀Scissors---扳手Wrench---总计7838980980由9798张图片组成并按8:1:1的比例分配为训练集7838张、验证集980张和测试集980张。 若要获取每个类别的详细分布情况需要进一步的数据集说明或分析。1X-ray安检YOLO数据集12分类的完整、详细、适合初学者的训练、推理、评估及模型调用代码使用YOLOv8实现。✅ 一、环境准备Python YOLOv81. 安装 Anaconda可选但推荐下载地址https://www.anaconda.com/products/distribution2. 创建虚拟环境并安装依赖# 创建虚拟环境conda create-nxray_detectionpython3.8conda activate xray_detection# 安装 PyTorch根据 CUDA 版本选择# 例如CUDA 11.8pipinstalltorch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118# 安装 YOLOv8Ultralyticspipinstallultralytics opencv-python numpy matplotlib tqdm 验证安装成功importtorchprint(torch.cuda.is_available())# 应输出 True如果有GPU 二、数据集结构确保数据集按如下结构组织xray_dataset/ ├── images/ │ ├── train/ # 7838 张 │ ├── val/ # 980 张 │ └── test/ # 980 张 ├── labels/ │ ├── train/ │ ├── val/ │ └── test/ └── data.yamldata.yaml文件内容train:./xray_dataset/images/trainval:./xray_dataset/images/valtest:./xray_dataset/images/testnc:12names:-Baton-Bullet-Gun-Knife-Hammer-HandCuffs-Lighter-Sprayer-Pliers-Powerbank-Scissors-Wrench 三、训练模型train.py# train.pyfromultralyticsimportYOLOimporttorchdeftrain_model():# 检查 GPUdevice0iftorch.cuda.is_available()elsecpuprint(fUsing device:{device})# 加载预训练模型推荐使用 yolov8n / yolov8smodelYOLO(yolov8n.pt)# 也可以用 yolov8s.pt 获取更高精度# 开始训练resultsmodel.train(dataxray_dataset/data.yaml,# 数据配置文件路径epochs100,# 训练轮数imgsz640,# 输入图像大小可调整batch16,# 批次大小根据显存调整namexray_train,# 实验名称projectruns/detect,# 保存路径saveTrue,# 保存模型save_period10,# 每10轮保存一次devicedevice,# 使用 GPUworkers4,# 数据加载线程数patience20,# 早停20轮无提升则停止optimizerAdamW,# 优化器可选 SGD, AdamWlr00.01,# 初始学习率augmentTrue# 启用数据增强旋转、翻转等)print(训练完成)returnresultsif__name____main__:train_model()✅ 训练结束后模型权重将保存在runs/detect/xray_train/weights/best.pt # 最佳模型 runs/detect/xray_train/weights/last.pt # 最后一轮模型 四、推理检测detect.py# detect.pyfromultralyticsimportYOLOimportcv2defdetect_image(image_path,output_pathoutput.jpg):# 加载训练好的模型modelYOLO(runs/detect/xray_train/weights/best.pt)# 进行预测resultsmodel.predict(sourceimage_path,saveTrue,# 保存带框的图像projectruns/detect/predict,# 保存路径nameresult,conf0.5,# 置信度阈值showFalse# 是否显示窗口)# 可选手动绘制并保存forrinresults:im_arrayr.plot()# 绘制边界框和标签imr.orig_imgifhasattr(r,orig_img)elseim_array cv2.imwrite(output_path,im)print(f检测完成结果保存在{output_path})returnoutput_pathif__name____main__:detect_image(test_image.jpg) 五、模型评估evaluate.py# evaluate.pyfromultralyticsimportYOLOdefevaluate_model():# 加载训练好的模型modelYOLO(runs/detect/xray_train/weights/best.pt)# 在验证集上评估metricsmodel.val(dataxray_dataset/data.yaml,splitval,# 可选 val 或 testbatch16,imgsz640,device0iftorch.cuda.is_available()elsecpu)# 打印关键指标print(fmAP0.5:{metrics.box.map50:.4f})print(fmAP0.5:0.95:{metrics.box.map:.4f})print(fPrecision:{metrics.box.p:.4f})print(fRecall:{metrics.box.r:.4f})# 打印各类别 APfori,nameinenumerate(metrics.names):print(f{name}: AP0.5 {metrics.box.ap[i]:.4f})if__name____main__:evaluate_model() 六、如何调用训练好的模型封装函数# model_utils.pyfromultralyticsimportYOLOclassXRayDetector:def__init__(self,model_pathruns/detect/xray_train/weights/best.pt):self.modelYOLO(model_path)defpredict(self,image_path,conf0.5):resultsself.model.predict(sourceimage_path,confconf)detections[]forrinresults:forboxinr.boxes:cls_idint(box.cls)conf_scorefloat(box.conf)labelself.model.names[cls_id]detections.append({class:label,confidence:conf_score,bbox:box.xyxy[0].cpu().numpy().tolist()# [x1, y1, x2, y2]})returndetections# 使用示例if__name____main__:detectorXRayDetector()resultsdetector.predict(test_image.jpg)forresinresults:print(res) 七、训练结果可视化可选YOLOv8 会自动在runs/detect/xray_train/生成results.png训练损失、mAP 曲线confusion_matrix.png混淆矩阵PR_curve.png各类别 Precision-Recall 曲线你也可以用 Pandas 查看训练日志importpandasaspd dfpd.read_csv(runs/detect/xray_train/results.csv)print(df.tail())以上文字及代码仅供参考学习使用。