Polars中高效对DataFrame子集应用自定义函数的最优方案咨询
问题描述
我想参照Pandas的思路优化Polars的map_elements用法,需求是:把无法修改的lat_lon_parse(来自lat_lon_parser包)仅应用于DataFrame的特定子集——列A包含°的行,其余行直接用.cast()转换即可。
现有两种实现方式:
from lat_lon_parser import parse as lat_lon_parse def test_lat_lon_parse(x): try: print(f"parsing {x}") return lat_lon_parse(x) except Exception: return None import polars as pl df = pl.DataFrame({'A':["1", "2", "5°N", "4°S"], "B":[1, 2, 3, 4]}) mask = pl.col('A').str.contains('°') # 方式1 def way_1(df): return df.with_columns( pl.when(mask) .then(pl.col('A').map_elements(test_lat_lon_parse)) .otherwise(pl.col('A')) .cast(pl.Float64) ) # 方式2 def way_2(df): return pl.concat([ df.filter(mask).with_columns(pl.col('A').map_elements(test_lat_lon_parse).alias('dummy')), df.filter(~mask).with_columns(pl.col('A').cast(pl.Float64).alias('dummy')) ])
目前发现:
- 方式1会遍历所有行,大数据量下效率很低;
- 方式2仅处理目标子集,
timeit测试大数据量下性能更优,但会打乱原DataFrame的行顺序。
想确认:我对这两种方式的理解是否正确?有没有更优的可扩展解决方案?
解答
对两种方式的理解确认
你的理解完全正确:
- 方式1:
when/then/otherwise搭配map_elements时,Polars会对整个列的每一行执行判断逻辑,哪怕不需要解析的行也会走一遍分支,加上map_elements本身是逐行处理的Python UDF,大数据量下确实会产生明显的性能损耗。 - 方式2:通过
filter拆分数据集,只对需要解析的子集调用map_elements,避免了无意义的计算,性能更优,但pl.concat会按传入的子集顺序合并,自然会打乱原有的行顺序。
更优的可扩展解决方案
要兼顾性能和保留原行顺序,可以借助Polars的with_row_index临时记录行号,处理完子集后再按行号排序恢复顺序,最后删除临时行号列。这种方法既保留了方式2的性能优势,又能保证数据顺序不变:
def optimal_way(df): # 添加临时行号列,记录原始顺序 df_with_idx = df.with_row_index("original_idx") # 拆分处理两个子集 parsed = df_with_idx.filter(mask).with_columns( pl.col('A').map_elements(test_lat_lon_parse).cast(pl.Float64).alias('A') ) casted = df_with_idx.filter(~mask).with_columns( pl.col('A').cast(pl.Float64).alias('A') ) # 合并后按原始行号排序,再删除临时列 return pl.concat([parsed, casted]).sort("original_idx").drop("original_idx")
额外性能优化建议
- 移除
test_lat_lon_parse中的print语句,避免IO操作拖慢处理速度; - 如果
lat_lon_parser支持批量处理(接受列表输入),可以用map_batches替代map_elements——map_batches是按批次处理数据,能大幅减少Python与Polars的交互开销,性能提升更明显:
def batch_lat_lon_parse(series): results = [] for x in series: try: results.append(lat_lon_parse(x)) except Exception: results.append(None) return pl.Series(results, dtype=pl.Float64) def optimal_way_with_batch(df): df_with_idx = df.with_row_index("original_idx") parsed = df_with_idx.filter(mask).with_columns( pl.col('A').map_batches(batch_lat_lon_parse).alias('A') ) casted = df_with_idx.filter(~mask).with_columns( pl.col('A').cast(pl.Float64).alias('A') ) return pl.concat([parsed, casted]).sort("original_idx").drop("original_idx")
内容的提问来源于stack exchange,提问作者tgrandje
相关产品推荐
相关产品推荐

