Polars滚动模式优化咨询:1亿行数据错误ID高效修正
优化Polars中滚动模式窗口的性能问题
问题背景
我有一个约1亿行的DataFrame,按group分组,id列中存在以99标识的错误值。尝试用rolling_map()结合scipy.stats.mode修正这些错误,但运行速度极慢,求更优实现方式。
原代码及结果
import polars as pl from scipy import stats def dummy(input): return stats.mode(input)[0] df = pl.DataFrame({'group': [10, 10, 10, 10, 10, 10, 10, 20, 20, 20, 20], 'id': [1, 1, 99, 1, 1, 2, 2, 3, 3, 99, 3]}) df.with_columns(pl.col('id') .rolling_map(function=dummy, window_size=3, min_periods=1, center=True) .over('group') .alias('id_mode'))
运行结果:
shape: (11, 3) ╭───────┬─────┬─────────╮ │ group ┆ id ┆ id_mode │ │ i64 ┆ i64 ┆ i64 │ ╞═══════╪═════╪═════════╡ │ 10 ┆ 1 ┆ 1 │ │ 10 ┆ 1 ┆ 1 │ │ 10 ┆ 99 ┆ 1 │ │ 10 ┆ 1 ┆ 1 │ │ 10 ┆ 1 ┆ 1 │ │ 10 ┆ 2 ┆ 2 │ │ 10 ┆ 2 ┆ 2 │ │ 20 ┆ 3 ┆ 3 │ │ 20 ┆ 3 ┆ 3 │ │ 20 ┆ 99 ┆ 3 │ │ 20 ┆ 3 ┆ 3 │ ╰───────┴─────┴─────────╯
性能瓶颈分析
rolling_map()会对每个窗口调用Python自定义函数,存在大量Python与底层引擎的交互开销。对于1亿行的大规模数据,这种逐窗口的Python调用会导致性能急剧下降。
优化方案
改用Polars内置的向量化操作,完全避免Python UDF的开销,所有计算在Rust层面完成,性能提升显著。
方案1:移位列+内置模式计算
针对中心窗口大小为3的场景,通过生成前一行、当前行、后一行的id列,合并为列表后逐行计算模式:
import polars as pl df = pl.DataFrame({'group': [10, 10, 10, 10, 10, 10, 10, 20, 20, 20, 20], 'id': [1, 1, 99, 1, 1, 2, 2, 3, 3, 99, 3]}) result = df.with_columns( # 按分组生成前一行和后一行的id prev_id=pl.col('id').shift(1).over('group'), next_id=pl.col('id').shift(-1).over('group') ).with_columns( # 合并三列成列表,计算每行的模式 id_mode=pl.concat_list([pl.col('prev_id'), pl.col('id'), pl.col('next_id')]) .list.eval(pl.element().mode()) .list.first() ).drop(['prev_id', 'next_id']) print(result)
方案2:排除错误值99的优化
如果99是明确的错误值,计算模式时可先过滤掉99,避免错误值干扰结果:
import polars as pl df = pl.DataFrame({'group': [10, 10, 10, 10, 10, 10, 10, 20, 20, 20, 20], 'id': [1, 1, 99, 1, 1, 2, 2, 3, 3, 99, 3]}) result = df.with_columns( prev_id=pl.col('id').shift(1).over('group'), next_id=pl.col('id').shift(-1).over('group') ).with_columns( id_mode=pl.concat_list([pl.col('prev_id'), pl.col('id'), pl.col('next_id')]) .list.filter(pl.element() != 99) # 过滤错误值99 .list.eval(pl.element().mode()) .list.first() ).drop(['prev_id', 'next_id']) print(result)
大规模数据优化:使用LazyFrame
对于1亿行的超大数据集,建议使用Polars的LazyFrame延迟执行,优化查询计划并减少内存占用:
import polars as pl # 从文件读取数据时直接使用LazyFrame lf = pl.scan_csv('large_data.csv') # 替换为你的数据源 result = lf.with_columns( prev_id=pl.col('id').shift(1).over('group'), next_id=pl.col('id').shift(-1).over('group') ).with_columns( id_mode=pl.concat_list([pl.col('prev_id'), pl.col('id'), pl.col('next_id')]) .list.filter(pl.element() != 99) .list.eval(pl.element().mode()) .list.first() ).drop(['prev_id', 'next_id']).collect()
性能对比
对于1亿行数据,上述向量化方案的速度会比rolling_map()快数十倍甚至上百倍,完全规避了Python UDF的调用开销。
内容的提问来源于stack exchange,提问作者usdn
相关产品推荐
相关产品推荐

