九、项目B · 案例三:模型验证

案例目标:在验证集上输出 mAP、混淆矩阵(存 matrix.csv)、各类 F1/AUC(存 val.csv)、ROC 与 PR 曲线(存 roc.jpg / pr.jpg)。先把四个概念搞明白,代码才不会写错方向:

示意图
图 9-1 mAP / 混淆矩阵 / F1 / AUC 一图看懂

9.1 快捷通道:ultralytics 一行验证 + 导出混淆矩阵

from ultralytics import YOLO
import numpy as np

model = YOLO("runs/sign/exp1/weights/best.pt")
metrics = model.val(data="data.yaml", split="val")
print("mAP50:", metrics.box.map50, "mAP50-95:", metrics.box.map)

# 混淆矩阵 → matrix.csv(ultralytics 自带混淆矩阵对象)
cm = metrics.confusion_matrix.matrix          # numpy 数组
np.savetxt("matrix.csv", cm, delimiter=",", fmt="%.0f")

9.2 通用写法:sklearn 计算 F1 / AUC / ROC / PR(通用方案)

# validate.py —— 用 sklearn/matplotlib 输出全部验证成果
import numpy as np, cv2, matplotlib
matplotlib.use("Agg")                          # 无显示器环境必须加这句!
import matplotlib.pyplot as plt
from sklearn.metrics import (confusion_matrix, f1_score, roc_curve, auc,
                             precision_recall_curve, classification_report)

# 前提:你已经把验证集跑完推理,得到每个样本的 真实类别 y_true 和 预测类别 y_pred、
# 以及各类别的置信度得分 y_score(用分类/检测的输出分数即可)
# 下面的 yolov8 批量推理示例演示如何收集:
results = model.predict(source="dataset/images/val", conf=0.25, save=False)
y_true, y_pred, y_score = [], [], []
for r in results:
    for b in r.boxes:
        y_pred.append(int(b.cls)); y_score.append(float(b.conf))
        y_true.append(真实类别)   # 从 val 标签 txt 对应读入

# ① 混淆矩阵 → matrix.csv
cm = confusion_matrix(y_true, y_pred)
np.savetxt("matrix.csv", cm, delimiter=",", fmt="%d")

# ② 各类 F1 / AUC → val.csv
f1 = f1_score(y_true, y_pred, average=None)
aucs = []
for c in range(3):                             # 每类做 one-vs-rest 算 AUC
    y_bin = (np.array(y_true) == c).astype(int)
    fpr, tpr, _ = roc_curve(y_bin, np.array(y_score))
    aucs.append(auc(fpr, tpr))
np.savetxt("val.csv", np.vstack([f1, aucs]), delimiter=",",
           header="F1,AUC", comments="")

# ③ ROC 曲线 → roc.jpg
plt.figure()
for c in range(3):
    y_bin = (np.array(y_true) == c).astype(int)
    fpr, tpr, _ = roc_curve(y_bin, np.array(y_score))
    plt.plot(fpr, tpr, label=f"class {c} (AUC={auc(fpr,tpr):.3f})")
plt.plot([0, 1], [0, 1], "--", color="gray")     # 50%参考线
plt.xlabel("False Positive Rate"); plt.ylabel("True Positive Rate")
plt.title("ROC Curve"); plt.legend(); plt.savefig("roc.jpg", dpi=150)

# ④ PR 曲线 → pr.jpg
plt.figure()
for c in range(3):
    y_bin = (np.array(y_true) == c).astype(int)
    prec, rec, _ = precision_recall_curve(y_bin, np.array(y_score))
    plt.plot(rec, prec, label=f"class {c}")
plt.xlabel("Recall"); plt.ylabel("Precision")
plt.title("PR Curve"); plt.legend(); plt.savefig("pr.jpg", dpi=150)
⚠️ 验证任务交付物对照(缺一个文件都不完整)

成果文件夹里应包含:matrix.csv(混淆矩阵)、val.csv(mAP/F1/AUC 指标)、roc.jpgpr.jpg,全部 FTP 上传。画图务必加 matplotlib.use("Agg"),否则云端没有图形界面会直接报错。

← 上一页下一页 →

— AI训练师技术交流教程 · 仅供学习交流 —