如何在Polars中更高效地执行条件连接?
优化Polars中带时间条件的自连接性能
问题描述
我手头有一个规模不小的DataFrame,对其进行自连接需要花费一定时间。我希望通过添加条件来执行连接,这样得到的结果DataFrame会小很多。但目前先做全连接再过滤的方式和全连接耗时差不多,想知道如何利用这些条件让条件连接比普通全连接更快?
示例代码如下:
import time import numpy as np import polars as pl # example dataframe rng = np.random.default_rng(1) nrows = 3_000_000 df = pl.DataFrame( dict( day=rng.integers(1, 300, nrows), id=rng.integers(1, 5_000, nrows), id2=rng.integers(1, 5, nrows), value=rng.normal(0, 1, nrows), ) ) # 普通自连接耗时约10-15秒(32核机器) start = time.perf_counter() df.join(df, on=["id", "id2"], how="left") time.perf_counter() - start # 先连接再过滤,耗时和普通自连接差不多 start = time.perf_counter() df.join(df, on=["id", "id2"], how="left").filter( (pl.col("day") < pl.col("day_right")) & (pl.col("day_right") - pl.col("day") <= 30) ) time.perf_counter() - start
注:过滤后的结果行数比全连接少10倍,但性能没有明显提升。
优化方案
核心思路是在连接阶段就过滤掉不符合条件的行,避免生成全连接的大中间结果。以下是两种高效实现方式:
方案1:使用join的condition参数(Polars 0.19.0+支持)
Polars从0.19.0版本开始支持在join中直接指定额外连接条件,查询优化器会将过滤逻辑融入连接操作,大幅减少中间数据量。
修改后的代码:
start = time.perf_counter() df.join( df, on=["id", "id2"], how="left", condition=(pl.col("day") < pl.col("day_right")) & (pl.col("day_right") - pl.col("day") <= 30) ) time.perf_counter() - start
这种方式无需先生成全连接结果,性能提升最为显著,通常能将耗时压缩至原全连接的1/5以内。
方案2:分组后用窗口函数+explode(兼容旧版本Polars)
如果你的Polars版本较低,不支持condition参数,可以先按id和id2分组,在组内为每行筛选符合时间条件的记录,再展开结果:
start = time.perf_counter() ( df.with_row_index("idx") .group_by(["id", "id2"], maintain_order=True) .agg( pl.struct(["day", "value", "idx"]).alias("left"), pl.struct(["day", "value", "idx"]).alias("right") ) .with_columns( pl.col("right") .map_elements( lambda right_rows: [ r for r in right_rows if (r["day"] > left_row["day"]) and (r["day"] - left_row["day"] <=30) ], return_dtype=pl.List(pl.Struct) ) .over("left") .alias("filtered_right") ) .explode("filtered_right") .filter(pl.col("filtered_right").is_not_null()) .select( pl.col("left").struct.field("idx").alias("idx_left"), pl.col("left").struct.field("day").alias("day"), pl.col("left").struct.field("value").alias("value"), pl.col("filtered_right").struct.field("idx").alias("idx_right"), pl.col("filtered_right").struct.field("day").alias("day_right"), pl.col("filtered_right").struct.field("value").alias("value_right"), pl.col("id"), pl.col("id2") ) ) time.perf_counter() - start
通过分组缩小了配对范围,组内仅处理同id+id2的数据,避免了全连接的无效计算。
内容的提问来源于stack exchange,提问作者lebesgue
相关产品推荐
相关产品推荐

