You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何提升Matplotlib批量生成400万张含4子图图像的速度?

提速建议:400万张Matplotlib子图生成优化

哇,400万张图确实是个超大的任务量,10小时才出4万张确实效率太低了,给你分享几个亲测有效的优化方向,从易到难一步步来:

1. 切换Matplotlib后端(立竿见影)

Matplotlib默认后端带交互式渲染开销,直接切到纯后端渲染模式,砍掉GUI相关的无用计算:

import matplotlib
matplotlib.use('Agg')  # 必须放在import plt之前
import matplotlib.pyplot as plt

或者用plt.switch_backend('Agg'),这一步能直接提升30%以上的渲染速度。

2. 预创建图形资源,避免重复初始化(最大性能瓶颈)

你现在每次循环都新建fig、gs、ax1-ax4,这个初始化的开销其实非常大!把图形和轴的创建移到循环外面,每次只更新数据、重绘、保存:

# 预初始化一次所有图形元素
plt.switch_backend('Agg')
fig = plt.figure(figsize=(6,6))
gs = gridspec.GridSpec(2, 2)
gs.update(left=0, wspace=0)
axes = [plt.subplot(gs[i]) for i in range(4)]
# 提前设置轴状态,避免重复操作
for ax in axes:
    ax.axis('off')
    ax.set_xticks([])
    ax.set_yticks([])

# 循环处理样本
for index, sample in enumerate(your_numpy_array):
    data_1, data_2, data_3, data_4 = sample.reshape(4, 40)
    
    # 清空轴内容并更新数据
    axes[0].clear()
    axes[0].plot(data_1, 's')
    axes[0].axis('off')  # clear后需重新关闭轴显示
    
    axes[1].clear()
    axes[1].plot(data_2, 's')
    axes[1].axis('off')
    
    axes[2].clear()
    axes[2].plot(data_3, 's')
    axes[2].axis('off')
    
    axes[3].clear()
    axes[3].plot(data_4, 's')
    axes[3].axis('off')
    
    # 保存图片,去掉不必要的参数
    plt.savefig(
        os.path.join(cwd, 'data/%d.png' % index),
        pad_inches=0,
        dpi=20,
        optimize=True  # 优化PNG压缩,加快保存速度
    )
# 循环结束后再关闭图形
plt.close()

这个改动能把单张图的生成时间砍掉至少一半,因为避免了重复创建图形对象的昂贵开销。

3. 砍掉不必要的绘图参数

  • 去掉bbox_inches='tight':你已经设置了left=0和pad_inches=0,这个参数会额外计算图形边界,完全没必要,反而增加耗时。
  • 简化标记样式:如果方形标记不是必须的,换成更简单的'o'或者直接用线条(去掉标记);如果必须用s,可以试试markerfacecolor='none'减少填充开销。
  • 提前设置轴的ticks为空:避免每次clear后Matplotlib自动生成刻度。

4. 调整进程数,避免CPU过载

36核机器跑36个进程容易导致上下文切换频繁,反而降低效率。试试把进程数降到24-32,留几个核心给系统IO调度和其他进程,实际CPU利用率会更高。

5. 终极优化:跳过Matplotlib,用Pillow直接绘制(速度提升一个数量级)

如果你的需求只是简单的线+方形标记,完全可以跳过Matplotlib的渲染层,用Pillow直接操作像素数组,绕开所有复杂的抽象层:

from PIL import Image, ImageDraw
import numpy as np

# 预定义参数
img_size = (120, 120)  # 6*20dpi
subplot_size = (60, 60)
marker_size = 2

for index, sample in enumerate(your_numpy_array):
    # 创建空白图像
    img = Image.new('L', img_size, color=255)
    draw = ImageDraw.Draw(img)
    
    data_list = sample.reshape(4, 40)
    # 四个子图的左上角坐标
    positions = [(0,0), (60,0), (0,60), (60,60)]
    
    for data, (x0, y0) in zip(data_list, positions):
        # 归一化数据到子图的y范围(留5像素边距)
        norm_data = np.interp(data, (data.min(), data.max()), (y0+5, y0+subplot_size[1]-5))
        # 生成x坐标(留5像素边距)
        x_coords = np.linspace(x0+5, x0+subplot_size[0]-5, 40)
        
        # 绘制线条
        for i in range(39):
            draw.line(
                [(x_coords[i], norm_data[i]), (x_coords[i+1], norm_data[i+1])],
                fill=0
            )
        # 绘制方形标记
        for x, y in zip(x_coords, norm_data):
            draw.rectangle(
                [(x-marker_size, y-marker_size), (x+marker_size, y+marker_size)],
                fill=0
            )
    
    # 保存图像
    img.save(os.path.join(cwd, 'data/%d.png' % index), optimize=True)

这个代码的执行速度会比Matplotlib版本快5-10倍,尤其适合大规模批量处理。

6. 文件系统小优化

  • 提前创建分级目录:如果index是连续的,可以按千/万级分目录(比如data/0-9999/、data/10000-19999/),避免单目录下文件过多导致的IO性能下降。
  • 用pathlib代替os.path:路径处理效率略高,代码也更简洁。

内容的提问来源于stack exchange,提问作者amba88

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.15 04:43:39