📝 前言:为什么我们需要自己画图?

在使用 YOLO(无论是目标检测还是关键点姿态估计)训练完模型后,官方会自动生成图表。但对于发表论文的同学来说,官方生成的图往往存在以下痛点:

  1. 纯英文标签:很多国内期刊要求图表必须是全中文或中英文对照。

  2. 画幅拥挤:所有的 Loss 和 mAP 全部堆在同一张图里,重点不突出。

  3. 格式不符合学术规范:例如缺少 300 DPI 高清分辨率、边框未按学术要求处理、刻度线没有朝内等。

本文以 “关键点检测(Pose)” 为例,教大家如何利用 Python 读取 YOLO 导出的 results.csv 文件,使用面向对象的方式,绘制出高扩展性、符合期刊要求的精美数据图,并彻底解决 Matplotlib 中文显示为方块的。

部分示例图:

   


 一、 项目目录准备

为了保持工程的整洁,建议在你的 YOLO 项目目录下新建一个专门用于绘图的文件夹。目录结构如下:

项目目录/
│
├── runs/pose/train3/
│   └── results.csv       <-- YOLO训练生成的CSV数据文件
│
└── draw/             <-- 新建的绘图文件夹
    ├── draw_paper_figs.py   <-- 本文提供的Python运行脚本
    └── output_figs/         <-- 运行代码后,高清图片会自动保存在这里

 二、 核心 Python 代码(直接抄作业)

在新建的 draw_paper_figs.py 文件中,复制以下完整代码。

import pandas as pd
import matplotlib.pyplot as plt
import matplotlib as mpl
import os

class PaperFigureDrawer:
    def __init__(self, csv_path, save_dir):
        """
        初始化绘图器
        :param csv_path: YOLO训练结果results.csv的路径
        :param save_dir: 图片保存的目录
        """
        # 1. 读取数据并去除列名可能带有的多余空格(YOLO的CSV列名经常带有前置空格)
        self.df = pd.read_csv(csv_path)
        self.df.columns = self.df.columns.str.strip()
        
        self.save_dir = save_dir
        os.makedirs(self.save_dir, exist_ok=True)
        
        # 2. 配置学术论文级别的全图样式 (全局设置)
        self._setup_paper_style()

    def _setup_paper_style(self):
        """配置论文画图的字体和样式(彻底解决中文乱码问题)"""
        
        # --- 字体设置核心区 ---
        # 字体查找优先级:新罗马(英文) -> 宋体(中文Windows)  -> 黑体(兜底)
        mpl.rcParams['font.family'] = ['Times New Roman', 'SimSun', 'SimHei']
        mpl.rcParams['axes.unicode_minus'] = False  # 解决坐标轴负号显示为方块的问题
        
        # --- 学术排版级参数配置 ---
        config = {
            "figure.dpi": 300,           # 论文级高清分辨率 (300 dpi)
            "savefig.dpi": 300,
            "savefig.bbox": 'tight',     # 保存图表时去除多余的白边
            
            "axes.linewidth": 1.2,       # 边框粗细
            "xtick.direction": 'in',     # x轴刻度向内 (学术图表规范)
            "ytick.direction": 'in',     # y轴刻度向内
            "xtick.major.width": 1.2,
            "ytick.major.width": 1.2,
            "xtick.labelsize": 10,       # 刻度字体大小
            "ytick.labelsize": 10,
            "axes.labelsize": 12,        # 坐标轴标签字体大小
            "legend.fontsize": 10,       # 图例字体大小
            "legend.frameon": False      # 图例无边框 (显得更专业透气)
        }
        mpl.rcParams.update(config)

    def _clean_ax(self, ax):
        """格式化单个坐标轴:去掉右侧和顶部的线,添加半透明网格"""
        ax.spines['top'].set_visible(False)
        ax.spines['right'].set_visible(False)
        ax.grid(True, linestyle='--', alpha=0.5)

    def plot_map_comparison(self):
        """图1:绘制边界框与关键点的 mAP 对比图"""
        fig, ax = plt.subplots(figsize=(6, 4.5))
        epochs = self.df['epoch']
        
        # 提取数据
        map50_b = self.df['metrics/mAP50(B)']
        map50_p = self.df['metrics/mAP50(P)']
        
        # 绘图:使用学术规范色彩 (经典的蓝红配色),区分线型
        ax.plot(epochs, map50_b, color='#1f77b4', linestyle='-', linewidth=2, label='边界框 mAP@0.5 (Bbox)')
        ax.plot(epochs, map50_p, color='#d62728', linestyle='--', linewidth=2, label='关键点 mAP@0.5 (Pose)')
        
        ax.set_xlabel('训练轮次 (Epoch)')
        ax.set_ylabel('平均精度均值 (mAP)')
        ax.set_title('估计精度变化', fontsize=13, pad=15)
        
        ax.legend(loc='lower right')
        self._clean_ax(ax)
        
        save_path = os.path.join(self.save_dir, 'mAP_comparison.png')
        plt.savefig(save_path)
        plt.close()
        print(f"✅ 图表1已保存: {save_path}")

    def plot_loss_trends(self):
        """图2:绘制关键点损失 (Pose Loss) 的 Train/Val 收敛趋势图"""
        fig, ax = plt.subplots(figsize=(6, 4.5))
        epochs = self.df['epoch']
        
        train_pose_loss = self.df['train/pose_loss']
        val_pose_loss = self.df['val/pose_loss']
        
        ax.plot(epochs, train_pose_loss, color='#2ca02c', linestyle='-', linewidth=2, label='训练集损失 (Train)')
        ax.plot(epochs, val_pose_loss, color='#ff7f0e', linestyle='--', linewidth=2, label='验证集损失 (Val)')
        
        ax.set_xlabel('训练轮次 (Epoch)')
        ax.set_ylabel('关键点损失 (Pose Loss)')
        ax.set_title('损失收敛曲线', fontsize=13, pad=15)
        
        ax.legend(loc='upper right')
        self._clean_ax(ax)
        
        save_path = os.path.join(self.save_dir, 'Pose_Loss_trend.png')
        plt.savefig(save_path)
        plt.close()
        print(f"✅ 图表2已保存: {save_path}")

    def plot_comprehensive_metrics(self):
        """图3:在一个 1x2 的子图中展示综合数据 (适合占满论文双栏的宽度)"""
        fig, axes = plt.subplots(1, 2, figsize=(12, 4.5))
        epochs = self.df['epoch']

        # 左图:Loss对比
        axes[0].plot(epochs, self.df['val/box_loss'], color='#1f77b4', label='验证集边界框损失')
        axes[0].plot(epochs, self.df['val/pose_loss'], color='#d62728', linestyle='--', label='验证集关键点损失')
        axes[0].set_xlabel('训练轮次 (Epoch)')
        axes[0].set_ylabel('损失值 (Loss)')
        axes[0].set_title('(a) 验证集损失对比')
        axes[0].legend()
        self._clean_ax(axes[0])

        # 右图:P R 曲线
        axes[1].plot(epochs, self.df['metrics/precision(P)'], color='#2ca02c', label='精确率 (Precision)')
        axes[1].plot(epochs, self.df['metrics/recall(P)'], color='#9467bd', linestyle='-.', label='召回率 (Recall)')
        axes[1].set_xlabel('训练轮次 (Epoch)')
        axes[1].set_ylabel('百分比指标')
        axes[1].set_title('(b) 精确率与召回率')
        axes[1].legend(loc='lower right')
        self._clean_ax(axes[1])

        plt.tight_layout()
        save_path = os.path.join(self.save_dir, 'Comprehensive_metrics.png')
        plt.savefig(save_path)
        plt.close()
        print(f"✅ 图表3已保存: {save_path}")

if __name__ == "__main__":
    # --- 路径配置区 ---
    # 请根据实际位置修改,支持相对路径
    CSV_FILE_PATH = r"./runs/pose/train3/results.csv"  
    OUTPUT_DIRECTORY = r"./draw/output_figs"
    
    # 1. 实例化绘图器
    drawer = PaperFigureDrawer(csv_path=CSV_FILE_PATH, save_dir=OUTPUT_DIRECTORY)
    
    # 2. 调用具体的绘图方法
    drawer.plot_map_comparison()         # 单图:mAP 对比图 
    drawer.plot_loss_trends()            # 单图:Loss 收敛图
    drawer.plot_comprehensive_metrics()  # 拼图:综合数据双子图
    
    print("🎉 所有图片渲染完成!请前往 output_figs 文件夹查看。")

 三、 代码核心亮点

  1. 学术规范

    • 刻度线朝内:direction='in'。

    • 图例无边框:legend.frameon=False,让图表看起来不拥挤。

    • 右侧与顶部去边框:spines['top'].set_visible(False),这是顶级期刊十分偏爱的通透排版风格。

  2. 清除表头空格

    • YOLO 导出的 csv 格式不太标准,很多列名前面自带空格,直接读取会报 KeyError。代码中加入了 self.df.columns.str.strip() 自动清理,规避这个隐藏天坑。

  3. 可扩展性强

    • 以后想加图,直接在 PaperFigureDrawer 类里面新增一个 plot_xxx() 方法即可。


 四、 终极排雷:运行后中文全是方块怎么办?

上述代码的 _setup_paper_style 函数已经做好了字体降级策略(优先宋体,兜底黑体)。但如果你运行后图表的中文依然显示为方块,这是因为 Matplotlib 的字体缓存没有刷新!

解决办法(任选其一即可):

方案 A:强制使用系统黑体
把上面代码中的字体配置部分删掉,换成这最暴力的单行代码:

plt.rcParams['font.sans-serif'] = ['SimHei']  # 强制使用黑体

方案 B:清除字体缓存
去你的电脑 C 盘,找到 C:\Users\你的用户名\.matplotlib 文件夹,里面有一个 fontlist-vXXX.json 的文件。
👉 直接把这个 json 文件删掉!
然后重新运行上面的 Python 代码。Matplotlib 会重新扫描你电脑里所有的字体(包括自带的宋体、黑体),从此告别方块乱码!


结语

绘制精美的数据图是向审稿人展示科研态度的第一步。掌握这套代码后,无论后续是做目标检测、实例分割还是关键点检测,只需要稍微修改一下读取的列名,就能不断地产出精美的论文配图了!

以上内容均为个人记录和观点,仅供参考。

如果这篇文章对你有帮助,欢迎点赞、收藏,并在评论区交流问题!🎯

更多推荐