Matplotlib 3.8.4 实战:3步从.txt文件绘制Loss/Acc曲线并导出高清PNG

在深度学习模型训练过程中,准确率和损失值的变化曲线是最直观反映模型学习状态的指标。许多初学者虽然能够获取这些数据,却苦于无法快速生成专业级的可视化图表用于报告或论文。本文将介绍一套基于Matplotlib 3.8.4的三步标准化流程,只需一个封装函数即可完成从数据读取到高质量图片导出的全过程。

1. 数据准备与预处理

任何可视化工作的起点都是规范化的数据源。在模型训练脚本中,我们需要将每个epoch的loss和accuracy值记录到文本文件中。以下是改进后的数据记录代码片段:

def save_training_log(train_loss, train_acc, loss_path='train_loss.txt', acc_path='train_acc.txt'):
    """保存训练日志到文本文件
    
    参数:
        train_loss (list): 训练损失值列表
        train_acc (list): 训练准确率列表
        loss_path (str): loss文件保存路径
        acc_path (str): accuracy文件保存路径
    """
    # 确保目录存在
    os.makedirs(os.path.dirname(loss_path), exist_ok=True)
    
    # 写入时保留6位小数
    with open(loss_path, 'w') as f:
        f.write(','.join(f'{x:.6f}' for x in train_loss))
    
    with open(acc_path, 'w') as f:
        f.write(','.join(f'{x:.6f}' for x in train_acc))

这种保存方式相比简单转储列表字符串有三大优势:

  1. 固定小数位数避免精度不一致
  2. 纯逗号分隔更易解析
  3. 自动创建不存在的目录

对应的数据读取函数也需要相应调整:

def load_training_data(file_path):
    """从文本文件加载训练数据
    
    参数:
        file_path (str): 数据文件路径
        
    返回:
        np.ndarray: 加载的数据数组
    """
    with open(file_path, 'r') as f:
        content = f.read().strip()
        if not content:  # 处理空文件
            return np.array([])
        return np.array([float(x) for x in content.split(',') if x])

注意:实际应用中建议添加try-except块处理文件不存在或格式错误的情况,这里为简洁省略

2. 核心绘图函数封装

我们将创建一个高度可配置的绘图函数,支持loss和accuracy曲线同图或分图显示。以下是函数的核心参数说明:

参数名 类型 默认值 说明
loss_data list/array 必填 损失值数据
acc_data list/array None 准确率数据
output_path str None 图片输出路径
dpi int 300 输出图片分辨率
figsize tuple (10,6) 图像尺寸(宽,高)
combined bool True 是否在同一图表显示双曲线

完整实现代码如下:

def plot_training_curve(loss_data, acc_data=None, output_path=None, 
                       dpi=300, figsize=(10,6), combined=True):
    """绘制训练曲线并保存
    
    参数:
        loss_data: 损失值数据
        acc_data: 准确率数据(可选)
        output_path: 输出路径(可选)
        dpi: 输出分辨率
        figsize: 图像尺寸
        combined: 是否合并显示
    """
    plt.figure(figsize=figsize)
    ax1 = plt.gca()
    
    # 绘制loss曲线
    color = 'tab:blue'
    ax1.set_xlabel('Epochs')
    ax1.set_ylabel('Loss', color=color)
    loss_line, = ax1.plot(loss_data, color=color, linewidth=2, label='Loss')
    ax1.tick_params(axis='y', labelcolor=color)
    
    # 自动调整Y轴范围,留出10%余量
    loss_min, loss_max = np.min(loss_data), np.max(loss_data)
    ax1.set_ylim(loss_min - 0.1*(loss_max-loss_min), 
                 loss_max + 0.1*(loss_max-loss_min))
    
    if acc_data is not None:
        if combined:
            # 双Y轴模式
            ax2 = ax1.twinx()
            color = 'tab:red'
            ax2.set_ylabel('Accuracy', color=color)
            acc_line, = ax2.plot(acc_data, color=color, linestyle='--', 
                                linewidth=2, label='Accuracy')
            ax2.tick_params(axis='y', labelcolor=color)
            
            # 合并图例
            lines = [loss_line, acc_line]
            ax1.legend(lines, [l.get_label() for l in lines], 
                      loc='upper right')
        else:
            # 分图模式
            plt.figure(figsize=(figsize[0], figsize[1]*0.6))
            plt.plot(acc_data, color='red', linestyle='--', 
                    linewidth=2, label='Accuracy')
            plt.xlabel('Epochs')
            plt.ylabel('Accuracy')
            plt.legend()
    
    # 美化样式
    ax1.grid(True, linestyle='--', alpha=0.6)
    plt.title('Training Metrics', pad=20)
    
    if output_path:
        # 确保目录存在
        os.makedirs(os.path.dirname(output_path), exist_ok=True)
        plt.savefig(output_path, dpi=dpi, bbox_inches='tight', 
                   facecolor='white', transparent=False)
    plt.close()

这个函数实现了几个关键优化:

  1. 智能调整Y轴范围,避免曲线贴边
  2. 支持双Y轴显示不同量纲的指标
  3. 自动处理输出目录创建
  4. 提供多种保存格式支持

3. 高质量输出与格式转换

Matplotlib支持多种输出格式,不同格式适用于不同场景:

格式 适用场景 优势 推荐DPI
PNG 网页/报告 无损压缩,通用性强 300-600
SVG 论文/印刷 矢量图,无限缩放 N/A
PDF 学术出版 矢量+文本可搜索 N/A

以下是批量导出多格式的示例代码:

def export_multiformat(base_name, formats=('png','svg','pdf'), dpi=300):
    """多格式导出训练曲线
    
    参数:
        base_name: 基础文件名(不含扩展名)
        formats: 要导出的格式列表
        dpi: 栅格图分辨率
    """
    loss = load_training_data('train_loss.txt')
    acc = load_training_data('train_acc.txt')
    
    for fmt in formats:
        output_path = f"{base_name}.{fmt}"
        plot_training_curve(loss, acc, output_path=output_path, 
                          dpi=dpi if fmt == 'png' else None)

关键参数设置建议:

  • DPI设置 :屏幕展示用300DPI,印刷品用600DPI
  • 尺寸选择
    • 单栏论文插图:3.5英寸宽
    • 双栏论文插图:7英寸宽
    • 演示文稿:10英寸宽
# 学术论文专用设置示例
plt.style.use('seaborn-paper')  # 学术风格
plt.rcParams.update({
    'font.family': 'serif',
    'font.serif': ['Times New Roman'],
    'font.size': 10,
    'axes.titlesize': 10,
    'axes.labelsize': 9,
    'xtick.labelsize': 8,
    'ytick.labelsize': 8,
    'legend.fontsize': 8,
    'figure.dpi': 600
})

4. 高级技巧与问题排查

实际应用中常会遇到几个典型问题:

曲线锯齿问题

  • 原因:数据点太少或波动太大
  • 解决方案:
    # 使用滑动平均平滑曲线
    window_size = 5
    smooth_loss = np.convolve(loss_data, np.ones(window_size)/window_size, mode='valid')
    

双曲线量纲差异

  • 解决方案:标准化处理
    from sklearn.preprocessing import MinMaxScaler
    
    scaler = MinMaxScaler()
    norm_loss = scaler.fit_transform(loss_data.reshape(-1,1)).flatten()
    norm_acc = scaler.fit_transform(acc_data.reshape(-1,1)).flatten()
    

常见错误处理

  1. 数据长度不一致:
min_length = min(len(loss_data), len(acc_data))
trunc_loss = loss_data[:min_length]
trunc_acc = acc_data[:min_length]
  1. 数据包含无效值:
clean_loss = loss_data[~np.isnan(loss_data)]
clean_acc = acc_data[~np.isinf(acc_data)]
  1. 图片导出空白:
# 在savefig前加上
plt.tight_layout()
# 或者调整bbox_inches
plt.savefig(..., bbox_inches='tight', pad_inches=0.1)

通过这套方法,即使是Matplotlib初学者也能快速生成可用于学术发表的精美曲线图。将上述代码封装成模块后,只需三行代码即可完成从数据到成图的转换:

from visualization import plot_training_curve

loss = load_training_data('train_loss.txt')
acc = load_training_data('train_acc.txt')
plot_training_curve(loss, acc, output_path='metrics.png')
Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐