如何使用Matplotlib高效绘制含100万数据点的散点图
百万级散点图绘制优化方案
核心性能瓶颈
当前代码耗时极高的核心原因是逐行循环调用plt.scatter(),每调用一次该函数就会创建独立的点集合对象,100万次调用产生的额外开销远高于渲染本身的耗时。同时代码中将x、y坐标转为字符串的写法属于错误用法,不仅会额外消耗性能,还会将数值坐标轴转为分类坐标轴,最终出图结果不符合预期。
优化方案
必改优化项
- 替换逐行读取CSV的逻辑,用批量读取工具一次性加载所有数据,避免循环开销
- 仅调用一次
plt.scatter(),传入数组格式的所有x、y、颜色参数,让Matplotlib内部批量渲染 - 删掉x、y坐标转字符串的错误逻辑,直接用数值类型传入
优化后代码(Pandas版本,性能最优)
import pandas as pd from matplotlib import pyplot as plt from matplotlib import style # 切换为无交互的Agg后端,减少不必要的渲染开销 plt.switch_backend('agg') style.use('ggplot') # 批量读取CSV,自动解析数值类型 df = pd.read_csv('total.csv', names=['点序号', 'x坐标', 'y坐标', '颜色值']) s = 0.5 # 单次传入所有数据渲染,超大数据量可开启rasterized=True提升大图保存速度 plt.scatter(df['x坐标'], df['y坐标'], color=df['颜色值'], s=s, rasterized=True) plt.savefig("graph.png", dpi=1000) # 手动释放内存,避免大数据量下内存泄漏 plt.close()
该版本处理100万数据点的耗时可从数十分钟降到数十秒,性能提升超过百倍。
无Pandas依赖版本(Numpy实现)
如果不想引入Pandas依赖,可以用Numpy批量读取数据:
import numpy as np from matplotlib import pyplot as plt from matplotlib import style plt.switch_backend('agg') style.use('ggplot') # 批量读取指定列,跳过第一列的点序号 data = np.genfromtxt( 'total.csv', delimiter=',', dtype=None, encoding='utf-8', usecols=(1, 2, 3) ) x = data['f0'].astype(float) y = data['f1'].astype(float) colors = data['f2'] s = 0.5 plt.scatter(x, y, color=colors, s=s, rasterized=True) plt.savefig("graph.png", dpi=1000) plt.close()
内容的提问来源于stack exchange,提问作者Bown Bleach
相关产品推荐
相关产品推荐

