Polars中over表达式的低效问题及相关技术疑问
Polars中窗口函数
over与group_by/agg+explode的性能差异分析 针对测试中发现的「over表达式运行速度比group_by/agg+explode组合慢2~3倍,且结果完全一致」的现象,逐一解答你的疑问:
1. 该性能表现是否符合预期?是否应始终用group_by/agg+explode替代over?
这种性能差异在当前测试场景下符合预期,但不能始终用后者替代over,原因如下:
- 底层逻辑差异:
- 窗口函数
over需要为每个聚合操作(比如mean().over、std().over)单独遍历分组,测试中每个v{i}列都调用了两次over(计算均值和标准差),相当于重复执行20次分组遍历,开销自然更高。 group_by/agg仅需执行一次分组,在分组内完成所有列的标准化计算,最后通过explode还原到原行结构,避免了重复的分组开销。
- 窗口函数
- 适用场景限制:
group_by/agg+explode仅适合「全分组聚合后需还原原行数量」的场景,且会改变原表的行顺序(需额外排序对齐)。而over的适用场景更广:比如需要保留原表行顺序、使用滑动/滚动窗口、或仅对部分列做窗口聚合并直接与原表其他列拼接时,over是更简洁且唯一的选择。
2. over表达式是否存在优化空间?
是的,存在明显的优化空间:
当前Polars对同一分组键的多个窗口聚合操作(比如测试中同一("id","id2")分组下的mean和std),没有做合并优化,会重复执行分组的哈希构建、数据遍历等步骤。未来可以通过以下方向优化:
- 共享同一分组键的中间结果:对相同分组的多个聚合操作,仅执行一次分组,复用分组后的数据集完成所有聚合计算。
- 批量窗口操作优化:针对多列执行相同窗口聚合的场景,提供批量处理逻辑,减少重复的函数调用和分组开销。
3. 两种方案的性能是否取决于具体场景,需用户自行测试选择?
完全取决于业务场景,建议根据以下情况选择:
- 优先选
over的场景:需要保留原表行顺序、使用滑动/滚动窗口、仅对部分列做窗口聚合且需与原表其他列直接拼接时。 - 优先选
group_by/agg+explode的场景:需要对全列做分组聚合后还原原行结构、对行顺序不敏感(或可接受后续排序)、分组基数适中(如本测试中的50*500=25000个分组)时,性能优势明显。 - 特殊情况:如果分组基数极大(比如每个分组仅1-2行),
explode的开销会显著增加,此时两种方案的性能差异会缩小,需实际测试验证。
测试代码
import time import numpy as np import polars as pl from polars.testing import assert_frame_equal ## setup rng = np.random.default_rng(1) nrows = 20_000_000 df = pl.DataFrame( dict( id=rng.integers(1, 50, nrows), id2=rng.integers(1, 500, nrows), v=rng.normal(0, 1, nrows), v1=rng.normal(0, 1, nrows), v2=rng.normal(0, 1, nrows), v3=rng.normal(0, 1, nrows), v4=rng.normal(0, 1, nrows), v5=rng.normal(0, 1, nrows), v6=rng.normal(0, 1, nrows), v7=rng.normal(0, 1, nrows), v8=rng.normal(0, 1, nrows), v9=rng.normal(0, 1, nrows), v10=rng.normal(0, 1, nrows), ) ) ## over start = time.perf_counter() res = ( df.lazy() .select( "id", "id2", *[ (pl.col(f"v{i}") - pl.col(f"v{i}").mean().over("id", "id2")) / pl.col(f"v{i}").std().over("id", "id2") for i in range(1, 11) ], ) .collect() ) print( time.perf_counter() - start ) # 8.541702497983351 ## groupby/agg + explode start = time.perf_counter() res2 = ( df.lazy() .group_by("id", "id2") .agg( (pl.col(f"v{i}") - pl.col(f"v{i}").mean()) / pl.col(f"v{i}").std() for i in range(1, 11) ) .explode(pl.exclude("id", "id2")) .collect() ) print( time.perf_counter() - start ) # 3.1841439900454134 ## compare results assert_frame_equal(res.sort(pl.all()), res2.sort(pl.all()))
内容的提问来源于stack exchange,提问作者lebesgue
相关产品推荐
相关产品推荐

