如何提升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
相关产品推荐
相关产品推荐

