【YOLO实战】提取results.csv训练数据,Python绘制高颜值学术论文图(附解决中文乱码方案)
📝 前言:为什么我们需要自己画图?
在使用 YOLO(无论是目标检测还是关键点姿态估计)训练完模型后,官方会自动生成图表。但对于发表论文的同学来说,官方生成的图往往存在以下痛点:
纯英文标签:很多国内期刊要求图表必须是全中文或中英文对照。
画幅拥挤:所有的 Loss 和 mAP 全部堆在同一张图里,重点不突出。
格式不符合学术规范:例如缺少 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 文件夹查看。")
三、 代码核心亮点
-
学术规范
-
刻度线朝内:direction='in'。
-
图例无边框:legend.frameon=False,让图表看起来不拥挤。
-
右侧与顶部去边框:spines['top'].set_visible(False),这是顶级期刊十分偏爱的通透排版风格。
-
-
清除表头空格
-
YOLO 导出的 csv 格式不太标准,很多列名前面自带空格,直接读取会报 KeyError。代码中加入了 self.df.columns.str.strip() 自动清理,规避这个隐藏天坑。
-
-
可扩展性强
-
以后想加图,直接在 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 会重新扫描你电脑里所有的字体(包括自带的宋体、黑体),从此告别方块乱码!
结语
绘制精美的数据图是向审稿人展示科研态度的第一步。掌握这套代码后,无论后续是做目标检测、实例分割还是关键点检测,只需要稍微修改一下读取的列名,就能不断地产出精美的论文配图了!
以上内容均为个人记录和观点,仅供参考。
如果这篇文章对你有帮助,欢迎点赞、收藏,并在评论区交流问题!🎯
更多推荐


所有评论(0)