社区所有版块导航
Python
python开源   Django   Python   DjangoApp   pycharm  
DATA
docker   Elasticsearch  
aigc
aigc   chatgpt  
WEB开发
linux   MongoDB   Redis   DATABASE   NGINX   其他Web框架   web工具   zookeeper   tornado   NoSql   Bootstrap   js   peewee   Git   bottle   IE   MQ   Jquery  
机器学习
机器学习算法  
Python88.com
反馈   公告   社区推广  
产品
短视频  
印度
印度  
Py学习  »  Python

基于YOLOv8旋转目标检测的工业齿轮多类型故障识别与定位(Python)

新机器视觉 • 9 月前 • 259 次点击  
   文章来源于:高斯的手稿
  链接:https://zhuanlan.zhihu.com/p/1946346577628206809

本文仅用于学术分享,如有侵权,请联系台作删文处理


基于YOLOv8旋转目标检测框架,针对齿轮图像中任意方向排列的缺陷进行识别和定位。首先配置并加载齿轮缺陷数据集,使用预训练的YOLOv8旋转边界框模型进行迁移学习,通过100个epoch的训练优化模型参数使其适应齿轮缺陷检测任务。训练过程中实时监控精确度、召回率、边界框损失、分类损失等关键指标确保模型收敛。训练完成后使用独立测试集进行模型验证,评估其在真实场景下的泛化能力。在预测阶段,系统对输入齿轮图像进行旋转边界框检测,识别出三种不同类型的齿轮缺陷,并使用不同颜色的半透明填充区域和边界框进行可视化标注,最后生成详细的性能分析报告和训练过程曲线,为工业质量控制提供可靠的自动化检测解决方案
算法流程如下
┌─────────────────────────────────────────────────────────────┐│                   齿轮故障检测系统流程图                       │└─────────────────────────────────────────────────────────────┘                               │                               ▼┌─────────────────┐    ┌─────────────────┐    ┌─────────────────┐│   数据准备阶段     │    │   模型训练阶段     │    │   模型验证阶段     ││                 │    │                 │    │                 ││ • 加载数据集配置  │───▶│ • 初始化YOLOv8   │───▶│ • 计算性能指标    ││ • 验证数据完整性  │    │   旋转检测模型    │    │ • 评估模型泛化    ││ • 设置数据路径    │    │ • 设置超参数     │    │   能力          │└─────────────────┘    │ • 开始模型训练    │    └─────────────────┘                       └─────────────────┘             │                               │                       ▼                               │              ┌─────────────────┐                               │              │   预测推理阶段    │                               │              │                 │                               └─────────────▶│ • 加载训练好的    │                                              │   模型          │                                              │ • 处理测试图像    │                                              │ • 生成旋转边界框   │                                              └─────────────────┘                                                       │                                                       ▼┌─────────────────┐    ┌─────────────────┐    ┌─────────────────┐│   结果可视化阶段   │    │   性能分析阶段    │    │   输出报告阶段    ││                 │    │                 │    │                 ││ • 绘制检测框     │◀──│ • 解析训练日志    │◀──│ • 统计最终指标    ││ • 添加类别标签    │    │ • 生成指标曲线    │    │ • 格式化输出     ││ • 保存可视化结果  │    │ • 分析训练过程    │    │ • 生成检测报告    │└─────────────────┘    └─────────────────┘    └─────────────────┘

详细算法步骤:

第一步进行数据集配置与验证,检查齿轮图像数据集的文件结构和标注格式,确保YAML配置文件正确指向训练集、验证集和测试集路径,同时验证图像文件和标注文件的对应关系。

第二步初始化YOLOv8旋转目标检测模型,加载预训练权重作为迁移学习的基础,设置训练超参数包括训练轮次、批次大小和输入图像尺寸,启动模型训练过程并实时保存训练状态。

第三步在独立验证集上评估训练完成的模型性能,计算精确度、召回率、平均精度均值等关键检测指标,验证模型在未见数据上的泛化能力和鲁棒性。

第四步使用训练好的模型对测试集图像进行批量预测,生成旋转边界框来精确框出各种齿轮缺陷的位置和方向,同时输出每个检测结果的类别置信度。

第五步对预测结果进行后处理和可视化,为每个检测到的缺陷绘制带有颜色编码的旋转边界框,添加类别标签说明,使用半透明填充增强视觉效果,并将标注结果保存为新的图像文件。

第六步从训练日志中提取历史训练数据,绘制多维度性能曲线包括精确度变化趋势、召回率进展、损失函数下降过程等,通过图表分析模型训练的动态过程和收敛情况。

第七步统计和计算最终模型性能指标,基于训练结束时的最佳权重计算综合评估分数,生成格式化的性能报告供质量评估和决策参考,完成整个齿轮故障检测系统的开发流程。

代码可适当参考:

# 齿轮故障检测系统 - 基于YOLOv8旋转目标检测的工业齿轮缺陷识别# ===================================================================# 导入所需的库和模块from ultralytics import YOLO  # YOLOv8目标检测框架import os  # 操作系统接口,用于文件路径操作import cv2  # OpenCV计算机视觉库,用于图像处理import matplotlib.pyplot as plt  # 数据可视化库import numpy as np  # 数值计算库import seaborn as sns  # 基于matplotlib的统计图形库import pandas as pd  # 数据分析库,用于处理训练日志数据
# ===================================================================# 1. 数据集配置与准备阶段# ===================================================================# 定义数据集在Kaggle环境中的存储路径dataset_path = "gears-dataset"# 构建数据集配置文件YAML的完整路径yaml_path = os.path.join(dataset_path, "gear.yaml")
# 打印YAML配置文件内容,用于验证数据集配置是否正确print("YAML配置文件内容:\n")with open(yaml_path, "r"as f:    print(f.read())
# ===================================================================# 2. 模型加载与训练阶段# ===================================================================# 加载预训练的YOLOv8中等尺寸旋转边界框检测模型# yolov8m-obb.pt支持旋转边界框检测,适用于任意方向的齿轮目标model = YOLO("yolov8m-obb.pt")
# 开始模型训练过程results = model.train(    data=yaml_path,  # 指定数据集配置文件路径    epochs=100,      # 设置训练总轮数为100轮    imgsz=640,       # 设置输入图像尺寸为640x640像素    batch=16,        # 设置每个批次的样本数量为16    name="gear-yolo-obb"  # 为本次训练运行命名,便于后续识别)
# ===================================================================# 3. 模型验证与性能评估阶段# ===================================================================# 使用验证集评估训练好的模型性能# 计算精确度、召回率、mAP等关键指标metrics = model.val()print("验证指标结果:", metrics)
# ===================================================================# 4. 模型预测与结果生成阶段# ===================================================================# 使用训练好的模型对测试集图像进行预测pred_results = model.predict(    source=os.path.join(dataset_path, "test/images"),  # 测试图像目录    save=False,    # 不自动保存预测结果,后续手动处理可视化    imgsz=640,     # 预测时使用的图像尺寸,与训练时保持一致    conf=0.25      # 设置检测置信度阈值为0.25,过滤低置信度检测)
# 定义不同齿轮缺陷类别的可视化颜色# 使用不同颜色区分不同类型的齿轮故障class_colors = {    "hp_cd": (0255255),  # 青色 - 表示某种齿轮缺陷    "hp_cm": (2550255),  # 洋红色 - 表示另一种齿轮缺陷    "kp": (1280255)      # 紫色 - 表示第三种齿轮缺陷}
# ===================================================================# 5. 检测结果可视化与输出保存阶段# ===================================================================# 创建输出目录,用于保存带检测框的可视化结果os.makedirs("overlay_outputs", exist_ok=True)
# 遍历所有预测结果,为每张图像添加检测框和标签for r in pred_results:    # 复制原始图像,避免修改原图    img = r.orig_img.copy()
    # 遍历当前图像中的所有旋转边界框检测结果    for box in r.obb:        # 提取目标类别ID并转换为整数        cls_id = int(box.cls[0])        # 根据类别ID获取对应的类别名称        label = r.names[cls_id]        # 提取旋转边界框的四个角点坐标,并转换为numpy数组        pts = box.xyxyxyxy[0].cpu().numpy().astype(int).reshape(-12)
        # 创建图像叠加层,用于半透明填充效果        overlay = img.copy()        # 使用对应类别的颜色填充旋转边界框区域        cv2.fillPoly(overlay, [pts], class_colors[label])        # 将填充层与原始图像以40%透明度混合        img = cv2.addWeighted(overlay, 0.4, img, 0.60)        # 在混合后的图像上绘制旋转边界框轮廓        cv2.polylines(img, [pts], isClosed=True, color=class_colors[label], thickness=2)
        # 提取边界框左上角坐标用于标签放置        x, y = pts[0]        # 在边界框上方添加类别标签文本        cv2.putText(img, label, (x, y - 5), cv2.FONT_HERSHEY_SIMPLEX, 0.7, class_colors[label], 2)
    # 构建输出文件路径,保持原文件名    out_path = os.path.join("overlay_outputs", os.path.basename(r.path))    # 保存带检测结果的可视化图像    cv2.imwrite(out_path, img)
print("✅ 叠加可视化结果已保存至 'overlay_outputs/' 目录")
# ===================================================================# 6. 样本检测结果展示阶段# ===================================================================# 获取所有输出图像文件列表sample_imgs = os.listdir("overlay_outputs")sample_imgs.sort()  # 对文件名进行排序,确保显示顺序一致# 确定要显示的图像数量(最多4张)num_display = min(4len(sample_imgs))
# 创建显示画布plt.figure(figsize=(1212))# 循环显示样本图像for i in range(num_display):    img_name = sample_imgs[i]  # 获取当前图像文件名    img_path = os.path.join("overlay_outputs", img_name)  # 构建完整路径    img = cv2.imread(img_path)  # 读取图像    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)  # 将BGR格式转换为RGB格式用于matplotlib显示
    # 在2x2网格中创建子图    plt.subplot(22, i + 1)    plt.imshow(img)  # 显示图像    plt.axis("off")  # 隐藏坐标轴    plt.title(img_name)  # 设置子图标题为文件名
# 调整子图布局,避免重叠plt.tight_layout()plt.show()
# ===================================================================# 7. 训练过程指标可视化分析阶段# ===================================================================# 定义训练结果存储目录run_dir = "runs/obb/gear-yolo-obb"# 构建训练日志CSV文件路径results_csv = os.path.join(run_dir, "results.csv")
# 检查训练日志文件是否存在if os.path.exists(results_csv):    # 读取训练日志数据到pandas DataFrame    df = pd.read_csv(results_csv)    print(f"\n✅ 成功从以下路径加载训练日志: {results_csv}")    print(f"可用数据列:\n{list(df.columns)}\n")
    # 设置图形样式和字体大小    sns.set(style="whitegrid", font_scale=1.2)    # 创建大尺寸图形画布    plt.figure(figsize=(18 12))
    # 1. 精确度曲线子图    plt.subplot(231)    plt.plot(df["epoch"], df["metrics/precision(B)"], marker="o", linewidth=2)    plt.title("Precision Curve")  # 图题保持英文    plt.xlabel("Epoch")  # 坐标轴标签保持英文    plt.ylabel("Precision")
    # 2. 召回率曲线子图    plt.subplot(232)    plt.plot(df["epoch"], df["metrics/recall(B)"], marker="o", color="orange", linewidth=2)    plt.title("Recall Curve")    plt.xlabel("Epoch")    plt.ylabel("Recall")
    # 3. F1分数曲线子图(手动计算)    plt.subplot(233)    # 根据精确度和召回率计算F1分数    f1 = 2 * (df["metrics/precision(B)"] * df["metrics/recall(B)"]) / (        df["metrics/precision(B)"] + df["metrics/recall(B)"] + 1e-9  # 添加小值避免除零错误    )    plt.plot(df["epoch"], f1, marker="o", color="green", linewidth=2)    plt.title("F1 Score Curve")    plt.xlabel("Epoch")    plt.ylabel("F1 Score")
    # 4. 边界框损失曲线子图    plt.subplot(234)    plt.plot(df["epoch"], df["train/box_loss"], marker="o", color="red", linewidth=2)    plt.title("Box Loss Curve")    plt.xlabel("Epoch")    plt.ylabel("Box Loss")
    # 5. 分类损失曲线子图    plt.subplot(235)    plt.plot(df["epoch"], df["train/cls_loss"], marker="o", color="purple", linewidth=2)    plt.title("Classification Loss Curve")    plt.xlabel("Epoch")    plt.ylabel("Cls Loss")
    # 6. DFL损失曲线子图    plt.subplot(236)    plt.plot(df["epoch"], df["train/dfl_loss"], marker="o", color="brown", linewidth=2)    plt.title("DFL Loss Curve")    plt.xlabel("Epoch")    plt.ylabel("DFL Loss")
    # 调整子图间距    plt.tight_layout()    plt.show()    print("\n✅ 所有训练指标曲线已成功显示!")
else:    print("⚠️ 训练结果文件results.csv未找到。请检查训练运行目录。")
# ===================================================================# 8. 最终模型性能指标统计与输出阶段# ===================================================================# 重新读取训练日志以获取最终性能指标run_dir = "runs/obb/gear-yolo-obb"results_csv = os.path.join(run_dir, "results.csv")
if os.path.exists(results_csv):    df = pd.read_csv(results_csv)    last_row = df.iloc[-1]  # 获取最后一行数据,即最终训练结果
    # 提取关键性能指标    precision = last_row["metrics/precision(B)"]  # 精确度    recall = last_row["metrics/recall(B)"]        # 召回率    # 计算F1分数(精确度和召回率的调和平均数)    f1_score = 2 * (precision * recall) / (precision + recall + 1e-6)    map50 = last_row["metrics/mAP50(B)"]          # mAP@0.5指标    map5095 = last_row["metrics/mAP50-95(B)"]     # mAP@0.5:0.95指标
    # 基于精确度和召回率计算整体准确度近似值    overall_accuracy = (2 * precision * recall) / (precision + recall + 1e-6)
    # 格式化输出最终模型性能指标    print( "========== 最终模型性能指标 ==========")    print(f"✅ 精确度(Precision):        {precision:.4f} ({precision*100:.2f}%)")    print(f"✅ 召回率(Recall):           {recall:.4f} ({recall*100:.2f}%)")    print(f"✅ F1分数:                   {f1_score:.4f} ({f1_score*100:.2f}%)")    print(f"✅ mAP@0.5:                  {map50:.4f} ({map50*100:.2f}%)")    print(f"✅ mAP@0.5:0.95:             {map5095:.4f} ({map5095*100:.2f}%)")    print(f"✅ 整体准确度:               {overall_accuracy:.4f} ({overall_accuracy*100:.2f}%)")    print("=====================================")else:    print("⚠️ 'results.csv' 文件未找到。请确保训练过程已成功完成。")


知乎学术咨询

https://www.zhihu.com/consult/people/792359672131756032?isMe=1

担任《Mechanical System and Signal Processing》《中国电机工程学报》《宇航学报》《控制与决策》等期刊审稿专家,擅长领域:信号滤波/降噪,机器学习/深度学习,时间序列预分析/预测,设备故障诊断/缺陷检测/异常检测

Python社区是高质量的Python/Django开发社区
本文地址:http://www.python88.com/topic/188686