如何在Polars中高效利用Matplotlib cmap生成[r,g,b,a]颜色列
问题
我想通过Matplotlib的colormap(cmap),从一个浮点列生成[r,g,b,a]格式的颜色列表列。目前用下面的方式实现,但想知道有没有更快的方法:
data.with_columns( (pl.col("floatCol")/100).map_elements(cmap1) )
以下是最小可运行示例代码:
import matplotlib as mpl import polars as pl cmap1 = mpl.colors.LinearSegmentedColormap.from_list("GreenBlue", ["limegreen", "blue"]) data = pl.DataFrame( { "floatCol": [12,135.8, 1235.263,15.236], "boolCol": [True, True, False, False] } ) data = data.with_columns( pl.when(pl.col("boolCol").not_()) .then(mpl.colors.to_rgba("r")) .otherwise((pl.col("floatCol")/100).map_elements(cmap1)) .alias("c1") )
优化方案
你当前用的map_elements是逐元素处理,数据量大时效率很低。可以利用Matplotlib colormap的向量化特性(cmap.__call__直接支持数组输入),结合Polars的map_batches实现批量处理,速度会快很多:
简洁版实现
import matplotlib as mpl import polars as pl cmap1 = mpl.colors.LinearSegmentedColormap.from_list("GreenBlue", ["limegreen", "blue"]) red_rgba = mpl.colors.to_rgba("r") data = pl.DataFrame( { "floatCol": [12,135.8, 1235.263,15.236], "boolCol": [True, True, False, False] } ) data = data.with_columns( pl.when(pl.col("boolCol").not_()) .then(pl.lit(red_rgba)) .otherwise( pl.col("floatCol").map_batches(lambda s: cmap1(s/100).tolist()) ) .alias("c1") )
性能进阶版(更适合大数据集)
如果数据集非常大,可以手动处理掩码逻辑,减少不必要的计算:
import numpy as np import matplotlib as mpl import polars as pl cmap1 = mpl.colors.LinearSegmentedColormap.from_list("GreenBlue", ["limegreen", "blue"]) red_rgba = mpl.colors.to_rgba("r") data = pl.DataFrame( { "floatCol": [12,135.8, 1235.263,15.236], "boolCol": [True, True, False, False] } ) def batch_colorize(s: pl.Series) -> pl.Series: norm_vals = s.to_numpy() / 100 mask = data["boolCol"].to_numpy() # 先初始化全为红色 colors = np.full((len(s), 4), red_rgba) # 批量给符合条件的元素应用colormap colors[mask] = cmap1(norm_vals[mask]) return pl.Series(colors.tolist()) data = data.with_columns( pl.col("floatCol").map_batches(batch_colorize).alias("c1") )
为什么更快?
map_elements是逐元素遍历,每次调用cmap都有额外开销;map_batches是批量处理整列(或分块),利用Matplotlib的向量化计算,大幅减少函数调用次数,数据量越大,性能提升越明显。
额外优化:归一化浮点值
如果你的floatCol数值范围不是刚好适配cmap的[0,1]区间,可以先做归一化处理:
norm = mpl.colors.Normalize(vmin=data["floatCol"].min(), vmax=data["floatCol"].max()) # 在批量处理中替换为:cmap1(norm(s.to_numpy()))
内容的提问来源于stack exchange,提问作者Galedon
相关产品推荐
相关产品推荐

